code wiki / (root) / nx_f32_llama_layer_load.nx

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}