code wiki / (root) / nx_f32_llama_v4p.nx

nx_f32_llama_v4p.nx source

↩ module page · 348 lines · 13678 B

1// nx_f32_llama_v4p.nx -- PAGED-KV twins of the v4 forward chain: block / 2// stack / model-forward over *NxPagedSeq instead of the contiguous 3// *NxF32KVCache. FAITHFUL copies of nx_f32_llama_block_v4 / 4// nx_f32_llama_stack_v4 / nx_f32_llm_forward_v4 with EXACTLY these 5// substitutions (the codebase twin convention -- v2/v3/v4 precedent): 6// block: cache -> pseq; RoPE cache_before = pseq.seq_len; 7// nx_f32_attn_with_cache -> nx_f32_attn_with_paged 8// stack: nx_pkv_ensure_append ONCE before the layer loop (blocks + 9// copy-on-append privatization must precede ALL layers' 10// appends); nx_f32_kv_cache_advance -> nx_pkv_advance 11// forward: cache -> pseq; stack_v4 -> stack_v4p (embed/rmsnorm/lm_head 12// identical) 13// Debug dumps (blk4_dump8/ffnmax "log2mag" stdout spam) omitted -- 14// numerically irrelevant. bp_* phase profiling kept (parity with the 15// contiguous profiler). Bit-exactness vs the contiguous chain is gated 16// end-to-end by nx_paged_fwd_gate on the REAL model. 17// 18// This is what makes the paged pool PAY: prefill a prompt ONCE, fork N 19// sequences (refcounted, copy-on-append) -- the prefix-shared 20// self-consistency / best-of-N serving pattern. 21// 22// genealogy_id: kwon_2023_pagedattention + standard_transformer_stack 23// lineage_id: substrate_f32_llama_v4p_v1 24 25import "nx_syscalls.nx" 26import "nx_tier.nx" 27import "nx_f32.nx" 28import "nx_f32_rmsnorm.nx" 29import "nx_f32_matmul.nx" 30import "nx_f32_activations.nx" 31import "nx_f32_rope.nx" 32import "nx_f32_attn_multi.nx" 33import "nx_f32_lazy_weight.nx" 34import "nx_thread_pool.nx" 35import "nx_f32_llama_block.nx" 36import "nx_f32_llama_block_v4.nx" 37import "nx_f32_llama_stack_v4.nx" 38import "nx_f32_llm.nx" 39import "nx_f32_llm_v4.nx" 40import "nx_kvcache.nx" 41import "nx_f32_attn_paged.nx" 42 43// ===== Block twin ==================================================== 44 45func nx_f32_llama_block_forward_v4p( 46 x: *i64, 47 n_tokens: nx_int, 48 hidden_dim: nx_int, 49 n_heads: nx_int, 50 n_kv_heads: nx_int, 51 head_dim: nx_int, 52 ffn_dim: nx_int, 53 layer: *NxF32LlamaLayerLazy, 54 pseq: *NxPagedSeq, 55 layer_idx: nx_int, 56 eps: i64, 57 attn_scale: i64, 58 rope_log_base: i64, 59 apply_rope: nx_int, 60 out: *i64) -> nx_int { 61 62 if n_tokens <= 0 { return NX_BLK4_ERR_BAD_DIM } 63 if hidden_dim <= 0 { return NX_BLK4_ERR_BAD_DIM } 64 if n_heads <= 0 { return NX_BLK4_ERR_BAD_DIM } 65 if n_kv_heads <= 0 { return NX_BLK4_ERR_BAD_DIM } 66 if head_dim <= 0 { return NX_BLK4_ERR_BAD_DIM } 67 if ffn_dim <= 0 { return NX_BLK4_ERR_BAD_DIM } 68 if n_heads * head_dim != hidden_dim { return NX_BLK4_ERR_BAD_DIM } 69 if layer == (0 as *NxF32LlamaLayerLazy) { return NX_BLK4_ERR_NULL } 70 if pseq == (0 as *NxPagedSeq) { return NX_BLK4_ERR_NULL } 71 if x == (0 as *i64) { return NX_BLK4_ERR_NULL } 72 if out == (0 as *i64) { return NX_BLK4_ERR_NULL } 73 74 let q_dim: nx_int = n_heads * head_dim 75 let kv_dim: nx_int = n_kv_heads * head_dim 76 77 let _tp0: i64 = bp_t() 78 let attn_in: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 79 let Q: *i64 = sys_mmap(n_tokens * q_dim * 8) as *i64 80 let K_new: *i64 = sys_mmap(n_tokens * kv_dim * 8) as *i64 81 let V_new: *i64 = sys_mmap(n_tokens * kv_dim * 8) as *i64 82 let attn_concat: *i64 = sys_mmap(n_tokens * q_dim * 8) as *i64 83 let attn_proj: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 84 let x_mid: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 85 let ffn_in: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 86 let gate_raw: *i64 = sys_mmap(n_tokens * ffn_dim * 8) as *i64 87 let gate_act: *i64 = sys_mmap(n_tokens * ffn_dim * 8) as *i64 88 let up_buf: *i64 = sys_mmap(n_tokens * ffn_dim * 8) as *i64 89 let hidden_buf: *i64 = sys_mmap(n_tokens * ffn_dim * 8) as *i64 90 let ffn_proj: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 91 bp_add(BP_ALLOC, _tp0) 92 93 // ===== Attention sublayer ===== 94 let _tp1: i64 = bp_t() 95 var t: nx_int = 0 96 while t < n_tokens { 97 let x_row: *i64 = (((x as i64) + t * hidden_dim * 8)) as *i64 98 let n_row: *i64 = (((attn_in as i64) + t * hidden_dim * 8)) as *i64 99 nx_f32_rmsnorm(x_row, layer.gamma_attn, hidden_dim, eps, n_row) 100 t = t + 1 101 } 102 bp_add(BP_RMSNORM, _tp1) 103 104 let _tp2: i64 = bp_t() 105 nx_f32_lazy_matmul(attn_in, layer.W_q, Q, n_tokens, hidden_dim, q_dim) 106 nx_f32_lazy_matmul(attn_in, layer.W_k, K_new, n_tokens, hidden_dim, kv_dim) 107 nx_f32_lazy_matmul(attn_in, layer.W_v, V_new, n_tokens, hidden_dim, kv_dim) 108 bp_add(BP_MATMUL, _tp2) 109 110 // Qwen2 attention bias (attention_bias=true): Q/K/V += bias, per token, 111 // BEFORE RoPE. Guarded so no-bias models are unaffected. 112 let _tp3: i64 = bp_t() 113 if (layer.bias_q as i64) != 0 { 114 var tb: nx_int = 0 115 while tb < n_tokens { 116 var iq: nx_int = 0 117 while iq < q_dim { Q[tb*q_dim+iq] = nx_f32_add(Q[tb*q_dim+iq], layer.bias_q[iq]); iq = iq + 1 } 118 var ik: nx_int = 0 119 while ik < kv_dim { K_new[tb*kv_dim+ik] = nx_f32_add(K_new[tb*kv_dim+ik], layer.bias_k[ik]); ik = ik + 1 } 120 var iv: nx_int = 0 121 while iv < kv_dim { V_new[tb*kv_dim+iv] = nx_f32_add(V_new[tb*kv_dim+iv], layer.bias_v[iv]); iv = iv + 1 } 122 tb = tb + 1 123 } 124 } 125 126 // RoPE: (cos,sin) table once per ABSOLUTE position (= pseq.seq_len + t). 127 if apply_rope != 0 { 128 let cache_before: nx_int = pseq.seq_len 129 let cs: *i64 = sys_mmap(head_dim * 8) as *i64 130 var t_r: nx_int = 0 131 while t_r < n_tokens { 132 let pos: nx_int = cache_before + t_r 133 nx_f32_rope_build_cs(cs, head_dim, pos, rope_log_base) 134 var h_q: nx_int = 0 135 while h_q < n_heads { 136 let qv: *i64 = (((Q as i64) + (t_r * q_dim + h_q * head_dim) * 8)) as *i64 137 nx_f32_rope_apply_cs_neox(qv, head_dim, cs) 138 h_q = h_q + 1 139 } 140 var h_kv: nx_int = 0 141 while h_kv < n_kv_heads { 142 let kv: *i64 = (((K_new as i64) + (t_r * kv_dim + h_kv * head_dim) * 8)) as *i64 143 nx_f32_rope_apply_cs_neox(kv, head_dim, cs) 144 h_kv = h_kv + 1 145 } 146 t_r = t_r + 1 147 } 148 sys_munmap(cs, head_dim * 8) 149 } 150 bp_add(BP_ROPE, _tp3) 151 152 let _tp4: i64 = bp_t() 153 let v_attn: nx_int = nx_f32_attn_with_paged(Q, K_new, V_new, n_tokens, 154 n_heads, n_kv_heads, head_dim, 155 pseq, layer_idx, 1, attn_scale, 156 attn_concat) 157 if v_attn != NX_F32_AP_OK { return NX_BLK4_ERR_CACHE } 158 bp_add(BP_ATTN, _tp4) 159 160 let _tp5: i64 = bp_t() 161 nx_f32_lazy_matmul(attn_concat, layer.W_o, attn_proj, n_tokens, q_dim, hidden_dim) 162 bp_add(BP_MATMUL, _tp5) 163 164 let _tp6: i64 = bp_t() 165 var i: nx_int = 0 166 while i < n_tokens * hidden_dim { 167 x_mid[i] = nx_f32_add(x[i], attn_proj[i]) 168 i = i + 1 169 } 170 bp_add(BP_RESID, _tp6) 171 172 // ===== FFN (SwiGLU) sublayer ===== 173 let _tp7: i64 = bp_t() 174 var t2: nx_int = 0 175 while t2 < n_tokens { 176 let xm_row: *i64 = (((x_mid as i64) + t2 * hidden_dim * 8)) as *i64 177 let fi_row: *i64 = (((ffn_in as i64) + t2 * hidden_dim * 8)) as *i64 178 nx_f32_rmsnorm(xm_row, layer.gamma_ffn, hidden_dim, eps, fi_row) 179 t2 = t2 + 1 180 } 181 bp_add(BP_RMSNORM, _tp7) 182 183 let _tp8: i64 = bp_t() 184 nx_f32_lazy_matmul(ffn_in, layer.W_gate, gate_raw, n_tokens, hidden_dim, ffn_dim) 185 nx_f32_lazy_matmul(ffn_in, layer.W_up, up_buf, n_tokens, hidden_dim, ffn_dim) 186 bp_add(BP_MATMUL, _tp8) 187 188 let _tp9: i64 = bp_t() 189 nx_blk4_swiglu_pool(gate_raw, up_buf, hidden_buf, n_tokens * ffn_dim) 190 bp_add(BP_ACT, _tp9) 191 192 let _tp10: i64 = bp_t() 193 nx_f32_lazy_matmul(hidden_buf, layer.W_down, ffn_proj, n_tokens, ffn_dim, hidden_dim) 194 bp_add(BP_MATMUL, _tp10) 195 196 let _tp11: i64 = bp_t() 197 var k: nx_int = 0 198 while k < n_tokens * hidden_dim { 199 out[k] = nx_f32_add(x_mid[k], ffn_proj[k]) 200 k = k + 1 201 } 202 bp_add(BP_RESID, _tp11) 203 204 return NX_BLK4_OK 205} 206 207// ===== Stack twin ==================================================== 208 209func nx_f32_llama_stack_forward_v4p( 210 x_in: *i64, 211 x_out: *i64, 212 n_tokens: nx_int, 213 n_layers: nx_int, 214 hidden_dim: nx_int, 215 n_heads: nx_int, 216 n_kv_heads: nx_int, 217 head_dim: nx_int, 218 ffn_dim: nx_int, 219 layers: *i64, 220 pseq: *NxPagedSeq, 221 eps: i64, 222 attn_scale: i64, 223 rope_log_base: i64, 224 apply_rope: nx_int) -> nx_int { 225 226 if n_tokens <= 0 { return NX_STK4_ERR_BAD_DIM } 227 if n_layers <= 0 { return NX_STK4_ERR_BAD_DIM } 228 if hidden_dim <= 0 { return NX_STK4_ERR_BAD_DIM } 229 if x_in == (0 as *i64) { return NX_STK4_ERR_NULL } 230 if x_out == (0 as *i64) { return NX_STK4_ERR_NULL } 231 if layers == (0 as *i64) { return NX_STK4_ERR_NULL } 232 if pseq == (0 as *NxPagedSeq) { return NX_STK4_ERR_NULL } 233 234 // blocks + copy-on-append privatization ONCE, before any layer appends. 235 let ve: nx_int = nx_pkv_ensure_append(pseq, n_tokens) 236 if ve != NX_PKV_OK { return NX_STK4_ERR_LAYER } 237 238 let n_elem: nx_int = n_tokens * hidden_dim 239 let buf_a: *i64 = sys_mmap(n_elem * 8) as *i64 240 let buf_b: *i64 = sys_mmap(n_elem * 8) as *i64 241 242 var i: nx_int = 0 243 while i < n_elem { buf_a[i] = x_in[i]; i = i + 1 } 244 245 var L: nx_int = 0 246 var use_a_as_src: nx_int = 1 247 while L < n_layers { 248 let layer_ptr_i64: i64 = layers[L] 249 let layer: *NxF32LlamaLayerLazy = layer_ptr_i64 as *NxF32LlamaLayerLazy 250 if layer == (0 as *NxF32LlamaLayerLazy) { return NX_STK4_ERR_LAYER } 251 252 if use_a_as_src != 0 { 253 let v: nx_int = nx_f32_llama_block_forward_v4p( 254 buf_a, n_tokens, hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 255 layer, pseq, L, eps, attn_scale, rope_log_base, apply_rope, buf_b) 256 if v != NX_BLK4_OK { return NX_STK4_ERR_LAYER } 257 use_a_as_src = 0 258 } else { 259 let v: nx_int = nx_f32_llama_block_forward_v4p( 260 buf_b, n_tokens, hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 261 layer, pseq, L, eps, attn_scale, rope_log_base, apply_rope, buf_a) 262 if v != NX_BLK4_OK { return NX_STK4_ERR_LAYER } 263 use_a_as_src = 1 264 } 265 L = L + 1 266 } 267 268 nx_pkv_advance(pseq, n_tokens) 269 270 if use_a_as_src != 0 { 271 var k: nx_int = 0 272 while k < n_elem { x_out[k] = buf_a[k]; k = k + 1 } 273 } else { 274 var k2: nx_int = 0 275 while k2 < n_elem { x_out[k2] = buf_b[k2]; k2 = k2 + 1 } 276 } 277 278 sys_munmap(buf_a, n_elem * 8) 279 sys_munmap(buf_b, n_elem * 8) 280 return NX_STK4_OK 281} 282 283// ===== Forward twin ================================================== 284 285func nx_f32_llm_forward_v4p( 286 model: *NxF32LlamaModel, 287 token_ids: *i64, 288 n_tokens: nx_int, 289 pseq: *NxPagedSeq, 290 eps: i64, 291 attn_scale: i64, 292 rope_log_base: i64, 293 apply_rope: nx_int, 294 logits: *i64) -> nx_int { 295 296 if model == (0 as *NxF32LlamaModel) { return NX_FLV4_ERR_NULL } 297 if token_ids == (0 as *i64) { return NX_FLV4_ERR_NULL } 298 if pseq == (0 as *NxPagedSeq) { return NX_FLV4_ERR_NULL } 299 if logits == (0 as *i64) { return NX_FLV4_ERR_NULL } 300 if n_tokens <= 0 { return NX_FLV4_ERR_BAD_DIM } 301 302 let hidden_dim: nx_int = model.hidden_dim 303 let vocab_size: nx_int = model.vocab_size 304 305 let x_embed: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 306 let x_stacked: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 307 let x_normed: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 308 309 var t: nx_int = 0 310 while t < n_tokens { 311 let tok: nx_int = token_ids[t] as nx_int 312 if tok < 0 { return NX_FLV4_ERR_TOKEN } 313 if tok >= vocab_size { return NX_FLV4_ERR_TOKEN } 314 var d: nx_int = 0 315 while d < hidden_dim { 316 x_embed[t * hidden_dim + d] = model.embed_weights[tok * hidden_dim + d] 317 d = d + 1 318 } 319 t = t + 1 320 } 321 322 let v_stk: nx_int = nx_f32_llama_stack_forward_v4p( 323 x_embed, x_stacked, n_tokens, model.n_layers, 324 hidden_dim, model.n_heads, model.n_kv_heads, model.head_dim, model.ffn_dim, 325 model.layers, pseq, eps, attn_scale, rope_log_base, apply_rope) 326 if v_stk != NX_STK4_OK { return NX_FLV4_ERR_STACK } 327 328 var t2: nx_int = 0 329 while t2 < n_tokens { 330 let src: *i64 = (((x_stacked as i64) + t2 * hidden_dim * 8)) as *i64 331 let dst: *i64 = (((x_normed as i64) + t2 * hidden_dim * 8)) as *i64 332 nx_f32_rmsnorm(src, model.gamma_out, hidden_dim, eps, dst) 333 t2 = t2 + 1 334 } 335 336 if model.lm_head_q8 != 0 { 337 nx_f32_lazy_matmul(x_normed, model.lm_head_q8 as *NxF32LazyWeight, logits, 338 n_tokens, hidden_dim, vocab_size) 339 } else { 340 nx_f32_matmul_t_pool(nx_lw_shared_pool(), x_normed, model.lm_head, logits, 341 n_tokens, hidden_dim, vocab_size) 342 } 343 344 sys_munmap(x_embed, n_tokens * hidden_dim * 8) 345 sys_munmap(x_stacked, n_tokens * hidden_dim * 8) 346 sys_munmap(x_normed, n_tokens * hidden_dim * 8) 347 return NX_FLV4_OK 348}