code wiki / (root) / nx_f32_llm_read_dims.nx

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}