nx_f32_llm_read_dims.nx source
↩ module page · 236 lines · 10140 B
1// nx_f32_llm_read_dims.nx -- read model dims from GGUF metadata.
2//
3// Walks the GGUF metadata section to derive model dimensions and
4// writes them into NxF32LlamaModel. Closes the "caller pre-fills
5// dims" gap from nx_f32_llm_load_weights_from_gguf.
6//
7// Architecture awareness:
8// Reads "general.architecture" string; uses it as the key prefix
9// for "<arch>.block_count" etc. Supports qwen2, llama (Llama-2),
10// llama (Llama-3 same key namespace), and any other arch that
11// follows the standard <arch>.{block_count,embedding_length,...}
12// convention.
13//
14// Keys read:
15// general.architecture (string)
16// <arch>.block_count (u32 -> n_layers)
17// <arch>.embedding_length (u32 -> hidden_dim)
18// <arch>.attention.head_count (u32 -> n_heads)
19// <arch>.attention.head_count_kv (u32 -> n_kv_heads, optional;
20// defaults to n_heads if missing
21// -- the no-GQA case)
22// <arch>.feed_forward_length (u32 -> ffn_dim)
23//
24// vocab_size derived from token_embd.weight's dim_0.
25// head_dim = hidden_dim / n_heads.
26//
27// genealogy_id: gguf_v3_metadata_spec_gerganov_2024 + llama_qwen_arch_keys
28// lineage_id: substrate_f32_llm_read_dims_v1
29
30import "nx_syscalls.nx"
31import "nx_tier.nx"
32import "nx_gguf.nx"
33import "nx_gguf_load.nx"
34import "nx_gguf_meta.nx"
35import "nx_f32_llm.nx"
36
37const NX_FLD_OK: nx_int = 0
38const NX_FLD_ERR_NULL: nx_int = 1
39const NX_FLD_ERR_NO_ARCH: nx_int = 2
40const NX_FLD_ERR_NO_KEY: nx_int = 3
41const NX_FLD_ERR_BAD_TYPE: nx_int = 4
42const NX_FLD_ERR_NO_EMBED: nx_int = 5
43const NX_FLD_ERR_BAD_DIM: nx_int = 6
44const NX_FLD_N_VERDICTS: nx_int = 7
45
46const NX_GGUF_TYPE_U32: i64 = 4
47const NX_GGUF_TYPE_STR: i64 = 8
48
49func nx_fld_verdict_is_valid(v: nx_int) -> nx_int {
50 if v < 0 { return 0 }
51 if v >= NX_FLD_N_VERDICTS { return 0 }
52 return 1
53}
54
55// Build "<arch>.<suffix>" key in out_buf. Returns total length.
56
57func _fld_concat_key(arch: *u8, arch_len: nx_int,
58 suffix: *u8, suffix_len: nx_int,
59 out_buf: *u8) -> nx_int {
60 var i: nx_int = 0
61 while i < arch_len {
62 out_buf[i] = arch[i]
63 i = i + 1
64 }
65 out_buf[arch_len] = 0x2e as u8 // '.'
66 var j: nx_int = 0
67 while j < suffix_len {
68 out_buf[arch_len + 1 + j] = suffix[j]
69 j = j + 1
70 }
71 return arch_len + 1 + suffix_len
72}
73
74// Read a u32 metadata value for "<arch>.<suffix>". Writes value to
75// out_val. Returns OK / NO_KEY / BAD_TYPE.
76
77func _fld_read_arch_u32(buf: *u8, len: i64, hdr: *NxGgufHeader,
78 arch: *u8, arch_len: nx_int,
79 suffix: *u8, suffix_len: nx_int,
80 out_val: *i64) -> nx_int {
81 let key_buf: *u8 = sys_mmap(64)
82 let n: nx_int = _fld_concat_key(arch, arch_len, suffix, suffix_len, key_buf)
83
84 let off_out: *i64 = sys_mmap(8) as *i64
85 let typ_out: *i64 = sys_mmap(8) as *i64
86 let v: nx_int = nx_gguf_meta_find(buf, len, hdr, key_buf, n, off_out, typ_out)
87 if v == NX_GMETA_NOT_FOUND { return NX_FLD_ERR_NO_KEY }
88 if v != NX_GMETA_OK { return NX_FLD_ERR_BAD_TYPE }
89 if typ_out[0] != NX_GGUF_TYPE_U32 { return NX_FLD_ERR_BAD_TYPE }
90 out_val[0] = nx_gguf_meta_read_u32(buf, off_out[0])
91 return NX_FLD_OK
92}
93
94// Find token_embd.weight + return its dim_0 (vocab_size).
95
96func _fld_read_vocab_size_from_embed(hdr: *NxGgufHeader) -> i64 {
97 let n_te: *u8 = sys_mmap(17)
98 n_te[0]=0x74 as u8; n_te[1]=0x6f as u8; n_te[2]=0x6b as u8; n_te[3]=0x65 as u8
99 n_te[4]=0x6e as u8; n_te[5]=0x5f as u8; n_te[6]=0x65 as u8; n_te[7]=0x6d as u8
100 n_te[8]=0x62 as u8; n_te[9]=0x64 as u8; n_te[10]=0x2e as u8; n_te[11]=0x77 as u8
101 n_te[12]=0x65 as u8; n_te[13]=0x69 as u8; n_te[14]=0x67 as u8; n_te[15]=0x68 as u8
102 n_te[16]=0x74 as u8
103 let idx: nx_int = nx_gguf_find_tensor(hdr, n_te, 17)
104 if idx < 0 { return -1 }
105 let ti: *NxGgufTensorInfo = nx_gguf_tensor_at(hdr, idx)
106 // GGUF/ggml convention: token_embd.weight is stored with shape
107 // [hidden_dim, vocab_size] where dim_0 = hidden_dim (fastest-varying
108 // in memory) and dim_1 = vocab_size. This matches the forward's
109 // embed lookup pattern: memory[tok * hidden_dim + d] = embed[tok][d].
110 return ti.dim_1
111}
112
113// Read all dims from metadata + embed shape into the model struct.
114
115func nx_f32_llm_read_dims_from_gguf(buf: *u8, len: i64, hdr: *NxGgufHeader,
116 model: *NxF32LlamaModel,
117 out_err: *i64) -> nx_int {
118 if model == (0 as *NxF32LlamaModel) {
119 out_err[0] = NX_FLD_ERR_NULL
120 return NX_FLD_ERR_NULL
121 }
122
123 // ===== general.architecture =====
124 let k_arch: *u8 = sys_mmap(20)
125 k_arch[0]=0x67 as u8; k_arch[1]=0x65 as u8; k_arch[2]=0x6e as u8; k_arch[3]=0x65 as u8
126 k_arch[4]=0x72 as u8; k_arch[5]=0x61 as u8; k_arch[6]=0x6c as u8; k_arch[7]=0x2e as u8
127 k_arch[8]=0x61 as u8; k_arch[9]=0x72 as u8; k_arch[10]=0x63 as u8; k_arch[11]=0x68 as u8
128 k_arch[12]=0x69 as u8; k_arch[13]=0x74 as u8; k_arch[14]=0x65 as u8; k_arch[15]=0x63 as u8
129 k_arch[16]=0x74 as u8; k_arch[17]=0x75 as u8; k_arch[18]=0x72 as u8; k_arch[19]=0x65 as u8
130
131 let off_out: *i64 = sys_mmap(8) as *i64
132 let typ_out: *i64 = sys_mmap(8) as *i64
133 let v_arch: nx_int = nx_gguf_meta_find(buf, len, hdr, k_arch, 20, off_out, typ_out)
134 if v_arch != NX_GMETA_OK {
135 out_err[0] = NX_FLD_ERR_NO_ARCH
136 return NX_FLD_ERR_NO_ARCH
137 }
138 if typ_out[0] != NX_GGUF_TYPE_STR {
139 out_err[0] = NX_FLD_ERR_BAD_TYPE
140 return NX_FLD_ERR_BAD_TYPE
141 }
142 let arch_len: nx_int = nx_gguf_meta_read_string_len(buf, off_out[0]) as nx_int
143 let arch_str: *u8 = nx_gguf_meta_read_string_ptr(buf, off_out[0])
144
145 // ===== Build each dim key with the architecture prefix =====
146
147 // ".block_count" (12 chars including leading dot; we strip leading dot since
148 // _fld_concat_key adds it.)
149 let s_bc: *u8 = sys_mmap(11)
150 s_bc[0]=0x62 as u8; s_bc[1]=0x6c as u8; s_bc[2]=0x6f as u8; s_bc[3]=0x63 as u8
151 s_bc[4]=0x6b as u8; s_bc[5]=0x5f as u8; s_bc[6]=0x63 as u8; s_bc[7]=0x6f as u8
152 s_bc[8]=0x75 as u8; s_bc[9]=0x6e as u8; s_bc[10]=0x74 as u8
153
154 // "embedding_length"
155 let s_el: *u8 = sys_mmap(16)
156 s_el[0]=0x65 as u8; s_el[1]=0x6d as u8; s_el[2]=0x62 as u8; s_el[3]=0x65 as u8
157 s_el[4]=0x64 as u8; s_el[5]=0x64 as u8; s_el[6]=0x69 as u8; s_el[7]=0x6e as u8
158 s_el[8]=0x67 as u8; s_el[9]=0x5f as u8; s_el[10]=0x6c as u8; s_el[11]=0x65 as u8
159 s_el[12]=0x6e as u8; s_el[13]=0x67 as u8; s_el[14]=0x74 as u8; s_el[15]=0x68 as u8
160
161 // "attention.head_count"
162 let s_hc: *u8 = sys_mmap(20)
163 s_hc[0]=0x61 as u8; s_hc[1]=0x74 as u8; s_hc[2]=0x74 as u8; s_hc[3]=0x65 as u8
164 s_hc[4]=0x6e as u8; s_hc[5]=0x74 as u8; s_hc[6]=0x69 as u8; s_hc[7]=0x6f as u8
165 s_hc[8]=0x6e as u8; s_hc[9]=0x2e as u8; s_hc[10]=0x68 as u8; s_hc[11]=0x65 as u8
166 s_hc[12]=0x61 as u8; s_hc[13]=0x64 as u8; s_hc[14]=0x5f as u8; s_hc[15]=0x63 as u8
167 s_hc[16]=0x6f as u8; s_hc[17]=0x75 as u8; s_hc[18]=0x6e as u8; s_hc[19]=0x74 as u8
168
169 // "attention.head_count_kv"
170 let s_hckv: *u8 = sys_mmap(23)
171 s_hckv[0]=0x61 as u8; s_hckv[1]=0x74 as u8; s_hckv[2]=0x74 as u8; s_hckv[3]=0x65 as u8
172 s_hckv[4]=0x6e as u8; s_hckv[5]=0x74 as u8; s_hckv[6]=0x69 as u8; s_hckv[7]=0x6f as u8
173 s_hckv[8]=0x6e as u8; s_hckv[9]=0x2e as u8; s_hckv[10]=0x68 as u8; s_hckv[11]=0x65 as u8
174 s_hckv[12]=0x61 as u8; s_hckv[13]=0x64 as u8; s_hckv[14]=0x5f as u8; s_hckv[15]=0x63 as u8
175 s_hckv[16]=0x6f as u8; s_hckv[17]=0x75 as u8; s_hckv[18]=0x6e as u8; s_hckv[19]=0x74 as u8
176 s_hckv[20]=0x5f as u8; s_hckv[21]=0x6b as u8; s_hckv[22]=0x76 as u8
177
178 // "feed_forward_length"
179 let s_ffl: *u8 = sys_mmap(19)
180 s_ffl[0]=0x66 as u8; s_ffl[1]=0x65 as u8; s_ffl[2]=0x65 as u8; s_ffl[3]=0x64 as u8
181 s_ffl[4]=0x5f as u8; s_ffl[5]=0x66 as u8; s_ffl[6]=0x6f as u8; s_ffl[7]=0x72 as u8
182 s_ffl[8]=0x77 as u8; s_ffl[9]=0x61 as u8; s_ffl[10]=0x72 as u8; s_ffl[11]=0x64 as u8
183 s_ffl[12]=0x5f as u8; s_ffl[13]=0x6c as u8; s_ffl[14]=0x65 as u8; s_ffl[15]=0x6e as u8
184 s_ffl[16]=0x67 as u8; s_ffl[17]=0x74 as u8; s_ffl[18]=0x68 as u8
185
186 // ===== Read scalars =====
187 let val_out: *i64 = sys_mmap(8) as *i64
188
189 let v_nl: nx_int = _fld_read_arch_u32(buf, len, hdr, arch_str, arch_len,
190 s_bc, 11, val_out)
191 if v_nl != NX_FLD_OK { out_err[0] = v_nl; return v_nl }
192 model.n_layers = val_out[0] as nx_int
193
194 let v_hd: nx_int = _fld_read_arch_u32(buf, len, hdr, arch_str, arch_len,
195 s_el, 16, val_out)
196 if v_hd != NX_FLD_OK { out_err[0] = v_hd; return v_hd }
197 model.hidden_dim = val_out[0] as nx_int
198
199 let v_h: nx_int = _fld_read_arch_u32(buf, len, hdr, arch_str, arch_len,
200 s_hc, 20, val_out)
201 if v_h != NX_FLD_OK { out_err[0] = v_h; return v_h }
202 model.n_heads = val_out[0] as nx_int
203
204 let v_kv: nx_int = _fld_read_arch_u32(buf, len, hdr, arch_str, arch_len,
205 s_hckv, 23, val_out)
206 if v_kv == NX_FLD_ERR_NO_KEY {
207 // No GQA: default to n_heads.
208 model.n_kv_heads = model.n_heads
209 } else {
210 if v_kv != NX_FLD_OK { out_err[0] = v_kv; return v_kv }
211 model.n_kv_heads = val_out[0] as nx_int
212 }
213
214 let v_ff: nx_int = _fld_read_arch_u32(buf, len, hdr, arch_str, arch_len,
215 s_ffl, 19, val_out)
216 if v_ff != NX_FLD_OK { out_err[0] = v_ff; return v_ff }
217 model.ffn_dim = val_out[0] as nx_int
218
219 // head_dim derived.
220 if model.n_heads <= 0 {
221 out_err[0] = NX_FLD_ERR_BAD_DIM
222 return NX_FLD_ERR_BAD_DIM
223 }
224 model.head_dim = model.hidden_dim / model.n_heads
225
226 // vocab_size from embed tensor shape.
227 let vs: i64 = _fld_read_vocab_size_from_embed(hdr)
228 if vs <= 0 {
229 out_err[0] = NX_FLD_ERR_NO_EMBED
230 return NX_FLD_ERR_NO_EMBED
231 }
232 model.vocab_size = vs as nx_int
233
234 out_err[0] = NX_FLD_OK
235 return NX_FLD_OK
236}