nx_f32_llama_layer_load.nx source
↩ module page · 192 lines · 8085 B
1// nx_f32_llama_layer_load.nx -- bind one layer's 9 GGUF tensors into NxF32LlamaLayer.
2//
3// Composes nx_gguf_load_tensor_to_f32 nine times with the canonical
4// Llama tensor names ("blk.{N}.attn_norm.weight" etc.) and stores the
5// resulting f32 raw-bit pointers into the layer struct.
6//
7// Per ggml docs (Gerganov 2024) -- clean-room from public spec; no
8// copied code. Tensor naming:
9// blk.{N}.attn_norm.weight -> gamma_attn
10// blk.{N}.attn_q.weight -> W_q
11// blk.{N}.attn_k.weight -> W_k
12// blk.{N}.attn_v.weight -> W_v
13// blk.{N}.attn_output.weight -> W_o
14// blk.{N}.ffn_norm.weight -> gamma_ffn
15// blk.{N}.ffn_gate.weight -> W_gate
16// blk.{N}.ffn_up.weight -> W_up
17// blk.{N}.ffn_down.weight -> W_down
18//
19// genealogy_id: llama_gguf_naming_gerganov_2024 + standard_layer_binding
20// lineage_id: substrate_f32_llama_layer_load_v1
21
22import "nx_syscalls.nx"
23import "nx_tier.nx"
24import "nx_dec_emit.nx"
25import "nx_gguf.nx"
26import "nx_gguf_load_f32.nx"
27import "nx_f32_llama_block.nx"
28
29const NX_FLL_OK: nx_int = 0
30const NX_FLL_ERR_BAD_LAYER: nx_int = 1
31const NX_FLL_ERR_NOT_FOUND: nx_int = 2
32const NX_FLL_ERR_NULL: nx_int = 3
33const NX_FLL_N_VERDICTS: nx_int = 4
34
35func nx_fll_verdict_is_valid(v: nx_int) -> nx_int {
36 if v < 0 { return 0 }
37 if v >= NX_FLL_N_VERDICTS { return 0 }
38 return 1
39}
40
41// Build "blk.{layer_idx}.{suffix}" into name_out. Returns total length.
42// name_out must have >= 4 + 5 (max digits) + 1 + suffix_len bytes.
43
44func _fll_fmt_name(layer_idx: nx_int, suffix: *u8, suffix_len: nx_int,
45 name_out: *u8) -> nx_int {
46 name_out[0] = 0x62 as u8 // 'b'
47 name_out[1] = 0x6c as u8 // 'l'
48 name_out[2] = 0x6b as u8 // 'k'
49 name_out[3] = 0x2e as u8 // '.'
50 let li_len: i64 = nx_dec_emit_u63(name_out, 4, layer_idx as i64)
51 let after_idx: nx_int = 4 + (li_len as nx_int)
52 name_out[after_idx] = 0x2e as u8 // '.'
53 var i: nx_int = 0
54 while i < suffix_len {
55 name_out[after_idx + 1 + i] = suffix[i]
56 i = i + 1
57 }
58 return after_idx + 1 + suffix_len
59}
60
61// One-shot helper: format name, load tensor, return *i64 storage.
62
63func _fll_load_one(buf: *u8, hdr: *NxGgufHeader,
64 layer_idx: nx_int,
65 suffix: *u8, suffix_len: nx_int,
66 out_err: *i64) -> *i64 {
67 let name_buf: *u8 = sys_mmap(64)
68 let n_len: nx_int = _fll_fmt_name(layer_idx, suffix, suffix_len, name_buf)
69 let nv_out: *i64 = sys_mmap(8) as *i64
70 let inner_err: *i64 = sys_mmap(8) as *i64
71 let storage: *i64 = nx_gguf_load_tensor_to_f32(buf, hdr, name_buf, n_len,
72 nv_out, inner_err)
73 if inner_err[0] != NX_GLF_OK {
74 out_err[0] = NX_FLL_ERR_NOT_FOUND
75 return 0 as *i64
76 }
77 return storage
78}
79
80// Populate all 9 weight pointers in the layer struct from GGUF.
81
82func nx_f32_llama_layer_load_from_gguf(buf: *u8, hdr: *NxGgufHeader,
83 layer_idx: nx_int,
84 layer_out: *NxF32LlamaLayer,
85 out_err: *i64) -> nx_int {
86 if layer_idx < 0 {
87 out_err[0] = NX_FLL_ERR_BAD_LAYER
88 return NX_FLL_ERR_BAD_LAYER
89 }
90 if layer_out == (0 as *NxF32LlamaLayer) {
91 out_err[0] = NX_FLL_ERR_NULL
92 return NX_FLL_ERR_NULL
93 }
94
95 // ===== Suffix byte arrays =====
96
97 // "attn_norm.weight" (16)
98 let s_an: *u8 = sys_mmap(16)
99 s_an[0]=0x61 as u8; s_an[1]=0x74 as u8; s_an[2]=0x74 as u8; s_an[3]=0x6e as u8
100 s_an[4]=0x5f as u8; s_an[5]=0x6e as u8; s_an[6]=0x6f as u8; s_an[7]=0x72 as u8
101 s_an[8]=0x6d as u8; s_an[9]=0x2e as u8; s_an[10]=0x77 as u8; s_an[11]=0x65 as u8
102 s_an[12]=0x69 as u8; s_an[13]=0x67 as u8; s_an[14]=0x68 as u8; s_an[15]=0x74 as u8
103
104 // "attn_q.weight" (13)
105 let s_q: *u8 = sys_mmap(13)
106 s_q[0]=0x61 as u8; s_q[1]=0x74 as u8; s_q[2]=0x74 as u8; s_q[3]=0x6e as u8
107 s_q[4]=0x5f as u8; s_q[5]=0x71 as u8; s_q[6]=0x2e as u8; s_q[7]=0x77 as u8
108 s_q[8]=0x65 as u8; s_q[9]=0x69 as u8; s_q[10]=0x67 as u8; s_q[11]=0x68 as u8
109 s_q[12]=0x74 as u8
110
111 // "attn_k.weight" (13)
112 let s_k: *u8 = sys_mmap(13)
113 s_k[0]=0x61 as u8; s_k[1]=0x74 as u8; s_k[2]=0x74 as u8; s_k[3]=0x6e as u8
114 s_k[4]=0x5f as u8; s_k[5]=0x6b as u8; s_k[6]=0x2e as u8; s_k[7]=0x77 as u8
115 s_k[8]=0x65 as u8; s_k[9]=0x69 as u8; s_k[10]=0x67 as u8; s_k[11]=0x68 as u8
116 s_k[12]=0x74 as u8
117
118 // "attn_v.weight" (13)
119 let s_v: *u8 = sys_mmap(13)
120 s_v[0]=0x61 as u8; s_v[1]=0x74 as u8; s_v[2]=0x74 as u8; s_v[3]=0x6e as u8
121 s_v[4]=0x5f as u8; s_v[5]=0x76 as u8; s_v[6]=0x2e as u8; s_v[7]=0x77 as u8
122 s_v[8]=0x65 as u8; s_v[9]=0x69 as u8; s_v[10]=0x67 as u8; s_v[11]=0x68 as u8
123 s_v[12]=0x74 as u8
124
125 // "attn_output.weight" (18)
126 let s_o: *u8 = sys_mmap(18)
127 s_o[0]=0x61 as u8; s_o[1]=0x74 as u8; s_o[2]=0x74 as u8; s_o[3]=0x6e as u8
128 s_o[4]=0x5f as u8; s_o[5]=0x6f as u8; s_o[6]=0x75 as u8; s_o[7]=0x74 as u8
129 s_o[8]=0x70 as u8; s_o[9]=0x75 as u8; s_o[10]=0x74 as u8; s_o[11]=0x2e as u8
130 s_o[12]=0x77 as u8; s_o[13]=0x65 as u8; s_o[14]=0x69 as u8; s_o[15]=0x67 as u8
131 s_o[16]=0x68 as u8; s_o[17]=0x74 as u8
132
133 // "ffn_norm.weight" (15)
134 let s_fn: *u8 = sys_mmap(15)
135 s_fn[0]=0x66 as u8; s_fn[1]=0x66 as u8; s_fn[2]=0x6e as u8; s_fn[3]=0x5f as u8
136 s_fn[4]=0x6e as u8; s_fn[5]=0x6f as u8; s_fn[6]=0x72 as u8; s_fn[7]=0x6d as u8
137 s_fn[8]=0x2e as u8; s_fn[9]=0x77 as u8; s_fn[10]=0x65 as u8; s_fn[11]=0x69 as u8
138 s_fn[12]=0x67 as u8; s_fn[13]=0x68 as u8; s_fn[14]=0x74 as u8
139
140 // "ffn_gate.weight" (15)
141 let s_fg: *u8 = sys_mmap(15)
142 s_fg[0]=0x66 as u8; s_fg[1]=0x66 as u8; s_fg[2]=0x6e as u8; s_fg[3]=0x5f as u8
143 s_fg[4]=0x67 as u8; s_fg[5]=0x61 as u8; s_fg[6]=0x74 as u8; s_fg[7]=0x65 as u8
144 s_fg[8]=0x2e as u8; s_fg[9]=0x77 as u8; s_fg[10]=0x65 as u8; s_fg[11]=0x69 as u8
145 s_fg[12]=0x67 as u8; s_fg[13]=0x68 as u8; s_fg[14]=0x74 as u8
146
147 // "ffn_up.weight" (13)
148 let s_fu: *u8 = sys_mmap(13)
149 s_fu[0]=0x66 as u8; s_fu[1]=0x66 as u8; s_fu[2]=0x6e as u8; s_fu[3]=0x5f as u8
150 s_fu[4]=0x75 as u8; s_fu[5]=0x70 as u8; s_fu[6]=0x2e as u8; s_fu[7]=0x77 as u8
151 s_fu[8]=0x65 as u8; s_fu[9]=0x69 as u8; s_fu[10]=0x67 as u8; s_fu[11]=0x68 as u8
152 s_fu[12]=0x74 as u8
153
154 // "ffn_down.weight" (15)
155 let s_fd: *u8 = sys_mmap(15)
156 s_fd[0]=0x66 as u8; s_fd[1]=0x66 as u8; s_fd[2]=0x6e as u8; s_fd[3]=0x5f as u8
157 s_fd[4]=0x64 as u8; s_fd[5]=0x6f as u8; s_fd[6]=0x77 as u8; s_fd[7]=0x6e as u8
158 s_fd[8]=0x2e as u8; s_fd[9]=0x77 as u8; s_fd[10]=0x65 as u8; s_fd[11]=0x69 as u8
159 s_fd[12]=0x67 as u8; s_fd[13]=0x68 as u8; s_fd[14]=0x74 as u8
160
161 // ===== Load all 9 =====
162
163 layer_out.gamma_attn = _fll_load_one(buf, hdr, layer_idx, s_an, 16, out_err)
164 if out_err[0] != NX_FLL_OK { return out_err[0] as nx_int }
165
166 layer_out.W_q = _fll_load_one(buf, hdr, layer_idx, s_q, 13, out_err)
167 if out_err[0] != NX_FLL_OK { return out_err[0] as nx_int }
168
169 layer_out.W_k = _fll_load_one(buf, hdr, layer_idx, s_k, 13, out_err)
170 if out_err[0] != NX_FLL_OK { return out_err[0] as nx_int }
171
172 layer_out.W_v = _fll_load_one(buf, hdr, layer_idx, s_v, 13, out_err)
173 if out_err[0] != NX_FLL_OK { return out_err[0] as nx_int }
174
175 layer_out.W_o = _fll_load_one(buf, hdr, layer_idx, s_o, 18, out_err)
176 if out_err[0] != NX_FLL_OK { return out_err[0] as nx_int }
177
178 layer_out.gamma_ffn = _fll_load_one(buf, hdr, layer_idx, s_fn, 15, out_err)
179 if out_err[0] != NX_FLL_OK { return out_err[0] as nx_int }
180
181 layer_out.W_gate = _fll_load_one(buf, hdr, layer_idx, s_fg, 15, out_err)
182 if out_err[0] != NX_FLL_OK { return out_err[0] as nx_int }
183
184 layer_out.W_up = _fll_load_one(buf, hdr, layer_idx, s_fu, 13, out_err)
185 if out_err[0] != NX_FLL_OK { return out_err[0] as nx_int }
186
187 layer_out.W_down = _fll_load_one(buf, hdr, layer_idx, s_fd, 15, out_err)
188 if out_err[0] != NX_FLL_OK { return out_err[0] as nx_int }
189
190 out_err[0] = NX_FLL_OK
191 return NX_FLL_OK
192}