code wiki / (root) / nx_f32_llama_v4b.nx

nx_f32_llama_v4b.nx source

↩ module page · 344 lines · 13322 B

1// nx_f32_llama_v4b.nx -- BATCHED MULTI-SEQUENCE decode twins: one forward 2// step advances M INDEPENDENT paged sequences by one token each (the 3// continuous-batching primitive, vLLM-class serving). 4// 5// WHY IT PAYS: QKV / FFN / lm_head rows are sequence-INDEPENDENT, so M 6// forks batch through the same matmuls -- every weight byte is read once 7// for M tokens instead of once per token (the same amortization that makes 8// chunked prefill fast). Attention is per-sequence: each row r attends 9// its own seqs[r] via nx_f32_attn_with_paged with n_q=1 (100%% reuse). 10// Per-cell matmul math is independent of m, so each row's logits are 11// BIT-IDENTICAL to the m=1 sequential path (gated by nx_batched_gate). 12// 13// Faithful derivative of nx_f32_llama_v4p with EXACTLY these deltas: 14// * pseq -> seqs: *i64 (M pointers to NxPagedSeq), one NEW token per seq 15// * RoPE position per row = seqs[r].seq_len (each row is ITS seq's next) 16// * attention: per-row loop over nx_f32_attn_with_paged(n_q=1) 17// * stack: ensure_append(seq_r, 1) each BEFORE layers; advance each AFTER 18// 19// genealogy_id: kwon_2023_pagedattention + yu_2022_orca_continuous_batching 20// lineage_id: substrate_f32_llama_v4b_v1 21 22import "nx_syscalls.nx" 23import "nx_tier.nx" 24import "nx_f32.nx" 25import "nx_f32_rmsnorm.nx" 26import "nx_f32_matmul.nx" 27import "nx_f32_activations.nx" 28import "nx_f32_rope.nx" 29import "nx_f32_attn_multi.nx" 30import "nx_f32_lazy_weight.nx" 31import "nx_thread_pool.nx" 32import "nx_f32_llama_block.nx" 33import "nx_f32_llama_block_v4.nx" 34import "nx_f32_llama_stack_v4.nx" 35import "nx_f32_llm.nx" 36import "nx_f32_llm_v4.nx" 37import "nx_kvcache.nx" 38import "nx_f32_attn_paged.nx" 39import "nx_f32_llama_v4p.nx" 40 41// ===== Block: M rows, one per sequence ============================== 42 43func nx_f32_llama_block_forward_v4b( 44 x: *i64, 45 M: nx_int, 46 hidden_dim: nx_int, 47 n_heads: nx_int, 48 n_kv_heads: nx_int, 49 head_dim: nx_int, 50 ffn_dim: nx_int, 51 layer: *NxF32LlamaLayerLazy, 52 seqs: *i64, 53 layer_idx: nx_int, 54 eps: i64, 55 attn_scale: i64, 56 rope_log_base: i64, 57 apply_rope: nx_int, 58 out: *i64) -> nx_int { 59 60 if M <= 0 { return NX_BLK4_ERR_BAD_DIM } 61 if hidden_dim <= 0 { return NX_BLK4_ERR_BAD_DIM } 62 if n_heads * head_dim != hidden_dim { return NX_BLK4_ERR_BAD_DIM } 63 if layer == (0 as *NxF32LlamaLayerLazy) { return NX_BLK4_ERR_NULL } 64 if seqs == (0 as *i64) { return NX_BLK4_ERR_NULL } 65 if x == (0 as *i64) { return NX_BLK4_ERR_NULL } 66 if out == (0 as *i64) { return NX_BLK4_ERR_NULL } 67 68 let q_dim: nx_int = n_heads * head_dim 69 let kv_dim: nx_int = n_kv_heads * head_dim 70 71 let attn_in: *i64 = sys_mmap(M * hidden_dim * 8) as *i64 72 let Q: *i64 = sys_mmap(M * q_dim * 8) as *i64 73 let K_new: *i64 = sys_mmap(M * kv_dim * 8) as *i64 74 let V_new: *i64 = sys_mmap(M * kv_dim * 8) as *i64 75 let attn_concat: *i64 = sys_mmap(M * q_dim * 8) as *i64 76 let attn_proj: *i64 = sys_mmap(M * hidden_dim * 8) as *i64 77 let x_mid: *i64 = sys_mmap(M * hidden_dim * 8) as *i64 78 let ffn_in: *i64 = sys_mmap(M * hidden_dim * 8) as *i64 79 let gate_raw: *i64 = sys_mmap(M * ffn_dim * 8) as *i64 80 let up_buf: *i64 = sys_mmap(M * ffn_dim * 8) as *i64 81 let hidden_buf: *i64 = sys_mmap(M * ffn_dim * 8) as *i64 82 let ffn_proj: *i64 = sys_mmap(M * hidden_dim * 8) as *i64 83 84 // ===== Attention sublayer ===== 85 var t: nx_int = 0 86 while t < M { 87 let x_row: *i64 = (((x as i64) + t * hidden_dim * 8)) as *i64 88 let n_row: *i64 = (((attn_in as i64) + t * hidden_dim * 8)) as *i64 89 nx_f32_rmsnorm(x_row, layer.gamma_attn, hidden_dim, eps, n_row) 90 t = t + 1 91 } 92 93 // batched QKV: M rows through each weight (bytes read once for all M). 94 nx_f32_lazy_matmul(attn_in, layer.W_q, Q, M, hidden_dim, q_dim) 95 nx_f32_lazy_matmul(attn_in, layer.W_k, K_new, M, hidden_dim, kv_dim) 96 nx_f32_lazy_matmul(attn_in, layer.W_v, V_new, M, hidden_dim, kv_dim) 97 98 if (layer.bias_q as i64) != 0 { 99 var tb: nx_int = 0 100 while tb < M { 101 var iq: nx_int = 0 102 while iq < q_dim { Q[tb*q_dim+iq] = nx_f32_add(Q[tb*q_dim+iq], layer.bias_q[iq]); iq = iq + 1 } 103 var ik: nx_int = 0 104 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 } 105 var iv: nx_int = 0 106 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 } 107 tb = tb + 1 108 } 109 } 110 111 // RoPE: position per ROW = its sequence's current length. 112 if apply_rope != 0 { 113 let cs: *i64 = sys_mmap(head_dim * 8) as *i64 114 var t_r: nx_int = 0 115 while t_r < M { 116 let sq: *NxPagedSeq = seqs[t_r] as *NxPagedSeq 117 let pos: nx_int = sq.seq_len 118 nx_f32_rope_build_cs(cs, head_dim, pos, rope_log_base) 119 var h_q: nx_int = 0 120 while h_q < n_heads { 121 let qv: *i64 = (((Q as i64) + (t_r * q_dim + h_q * head_dim) * 8)) as *i64 122 nx_f32_rope_apply_cs_neox(qv, head_dim, cs) 123 h_q = h_q + 1 124 } 125 var h_kv: nx_int = 0 126 while h_kv < n_kv_heads { 127 let kv: *i64 = (((K_new as i64) + (t_r * kv_dim + h_kv * head_dim) * 8)) as *i64 128 nx_f32_rope_apply_cs_neox(kv, head_dim, cs) 129 h_kv = h_kv + 1 130 } 131 t_r = t_r + 1 132 } 133 sys_munmap(cs, head_dim * 8) 134 } 135 136 // attention: PER SEQUENCE (n_q = 1 row each; appends into its own seq). 137 var r: nx_int = 0 138 while r < M { 139 let sq2: *NxPagedSeq = seqs[r] as *NxPagedSeq 140 let Qr: *i64 = (((Q as i64) + r * q_dim * 8)) as *i64 141 let Kr: *i64 = (((K_new as i64) + r * kv_dim * 8)) as *i64 142 let Vr: *i64 = (((V_new as i64) + r * kv_dim * 8)) as *i64 143 let Cr: *i64 = (((attn_concat as i64) + r * q_dim * 8)) as *i64 144 let v_attn: nx_int = nx_f32_attn_with_paged(Qr, Kr, Vr, 1, 145 n_heads, n_kv_heads, head_dim, 146 sq2, layer_idx, 1, attn_scale, Cr) 147 if v_attn != NX_F32_AP_OK { return NX_BLK4_ERR_CACHE } 148 r = r + 1 149 } 150 151 nx_f32_lazy_matmul(attn_concat, layer.W_o, attn_proj, M, q_dim, hidden_dim) 152 153 var i: nx_int = 0 154 while i < M * hidden_dim { 155 x_mid[i] = nx_f32_add(x[i], attn_proj[i]) 156 i = i + 1 157 } 158 159 // ===== FFN (SwiGLU) sublayer ===== 160 var t2: nx_int = 0 161 while t2 < M { 162 let xm_row: *i64 = (((x_mid as i64) + t2 * hidden_dim * 8)) as *i64 163 let fi_row: *i64 = (((ffn_in as i64) + t2 * hidden_dim * 8)) as *i64 164 nx_f32_rmsnorm(xm_row, layer.gamma_ffn, hidden_dim, eps, fi_row) 165 t2 = t2 + 1 166 } 167 168 nx_f32_lazy_matmul(ffn_in, layer.W_gate, gate_raw, M, hidden_dim, ffn_dim) 169 nx_f32_lazy_matmul(ffn_in, layer.W_up, up_buf, M, hidden_dim, ffn_dim) 170 nx_blk4_swiglu_pool(gate_raw, up_buf, hidden_buf, M * ffn_dim) 171 nx_f32_lazy_matmul(hidden_buf, layer.W_down, ffn_proj, M, ffn_dim, hidden_dim) 172 173 var k: nx_int = 0 174 while k < M * hidden_dim { 175 out[k] = nx_f32_add(x_mid[k], ffn_proj[k]) 176 k = k + 1 177 } 178 179 sys_munmap(attn_in, M * hidden_dim * 8) 180 sys_munmap(Q, M * q_dim * 8) 181 sys_munmap(K_new, M * kv_dim * 8) 182 sys_munmap(V_new, M * kv_dim * 8) 183 sys_munmap(attn_concat, M * q_dim * 8) 184 sys_munmap(attn_proj, M * hidden_dim * 8) 185 sys_munmap(x_mid, M * hidden_dim * 8) 186 sys_munmap(ffn_in, M * hidden_dim * 8) 187 sys_munmap(gate_raw, M * ffn_dim * 8) 188 sys_munmap(up_buf, M * ffn_dim * 8) 189 sys_munmap(hidden_buf, M * ffn_dim * 8) 190 sys_munmap(ffn_proj, M * hidden_dim * 8) 191 return NX_BLK4_OK 192} 193 194// ===== Stack ========================================================== 195 196func nx_f32_llama_stack_forward_v4b( 197 x_in: *i64, 198 x_out: *i64, 199 M: nx_int, 200 n_layers: nx_int, 201 hidden_dim: nx_int, 202 n_heads: nx_int, 203 n_kv_heads: nx_int, 204 head_dim: nx_int, 205 ffn_dim: nx_int, 206 layers: *i64, 207 seqs: *i64, 208 eps: i64, 209 attn_scale: i64, 210 rope_log_base: i64, 211 apply_rope: nx_int) -> nx_int { 212 213 if M <= 0 { return NX_STK4_ERR_BAD_DIM } 214 if n_layers <= 0 { return NX_STK4_ERR_BAD_DIM } 215 if x_in == (0 as *i64) { return NX_STK4_ERR_NULL } 216 if x_out == (0 as *i64) { return NX_STK4_ERR_NULL } 217 if layers == (0 as *i64) { return NX_STK4_ERR_NULL } 218 if seqs == (0 as *i64) { return NX_STK4_ERR_NULL } 219 220 // capacity + COW privatization for EVERY seq before any layer appends. 221 var e: nx_int = 0 222 while e < M { 223 let sq: *NxPagedSeq = seqs[e] as *NxPagedSeq 224 let ve: nx_int = nx_pkv_ensure_append(sq, 1) 225 if ve != NX_PKV_OK { return NX_STK4_ERR_LAYER } 226 e = e + 1 227 } 228 229 let n_elem: nx_int = M * hidden_dim 230 let buf_a: *i64 = sys_mmap(n_elem * 8) as *i64 231 let buf_b: *i64 = sys_mmap(n_elem * 8) as *i64 232 233 var i: nx_int = 0 234 while i < n_elem { buf_a[i] = x_in[i]; i = i + 1 } 235 236 var L: nx_int = 0 237 var use_a_as_src: nx_int = 1 238 while L < n_layers { 239 let layer_ptr_i64: i64 = layers[L] 240 let layer: *NxF32LlamaLayerLazy = layer_ptr_i64 as *NxF32LlamaLayerLazy 241 if layer == (0 as *NxF32LlamaLayerLazy) { return NX_STK4_ERR_LAYER } 242 243 if use_a_as_src != 0 { 244 let v: nx_int = nx_f32_llama_block_forward_v4b( 245 buf_a, M, hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 246 layer, seqs, L, eps, attn_scale, rope_log_base, apply_rope, buf_b) 247 if v != NX_BLK4_OK { return NX_STK4_ERR_LAYER } 248 use_a_as_src = 0 249 } else { 250 let v: nx_int = nx_f32_llama_block_forward_v4b( 251 buf_b, M, hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 252 layer, seqs, L, eps, attn_scale, rope_log_base, apply_rope, buf_a) 253 if v != NX_BLK4_OK { return NX_STK4_ERR_LAYER } 254 use_a_as_src = 1 255 } 256 L = L + 1 257 } 258 259 var a2: nx_int = 0 260 while a2 < M { 261 let sq2: *NxPagedSeq = seqs[a2] as *NxPagedSeq 262 nx_pkv_advance(sq2, 1) 263 a2 = a2 + 1 264 } 265 266 if use_a_as_src != 0 { 267 var k: nx_int = 0 268 while k < n_elem { x_out[k] = buf_a[k]; k = k + 1 } 269 } else { 270 var k2: nx_int = 0 271 while k2 < n_elem { x_out[k2] = buf_b[k2]; k2 = k2 + 1 } 272 } 273 274 sys_munmap(buf_a, n_elem * 8) 275 sys_munmap(buf_b, n_elem * 8) 276 return NX_STK4_OK 277} 278 279// ===== Forward: one token per sequence -> M logit rows =============== 280 281func nx_f32_llm_forward_v4b( 282 model: *NxF32LlamaModel, 283 token_ids: *i64, 284 M: nx_int, 285 seqs: *i64, 286 eps: i64, 287 attn_scale: i64, 288 rope_log_base: i64, 289 apply_rope: nx_int, 290 logits: *i64) -> nx_int { 291 292 if model == (0 as *NxF32LlamaModel) { return NX_FLV4_ERR_NULL } 293 if token_ids == (0 as *i64) { return NX_FLV4_ERR_NULL } 294 if seqs == (0 as *i64) { return NX_FLV4_ERR_NULL } 295 if logits == (0 as *i64) { return NX_FLV4_ERR_NULL } 296 if M <= 0 { return NX_FLV4_ERR_BAD_DIM } 297 298 let hidden_dim: nx_int = model.hidden_dim 299 let vocab_size: nx_int = model.vocab_size 300 301 let x_embed: *i64 = sys_mmap(M * hidden_dim * 8) as *i64 302 let x_stacked: *i64 = sys_mmap(M * hidden_dim * 8) as *i64 303 let x_normed: *i64 = sys_mmap(M * hidden_dim * 8) as *i64 304 305 var t: nx_int = 0 306 while t < M { 307 let tok: nx_int = token_ids[t] as nx_int 308 if tok < 0 { return NX_FLV4_ERR_TOKEN } 309 if tok >= vocab_size { return NX_FLV4_ERR_TOKEN } 310 var d: nx_int = 0 311 while d < hidden_dim { 312 x_embed[t * hidden_dim + d] = model.embed_weights[tok * hidden_dim + d] 313 d = d + 1 314 } 315 t = t + 1 316 } 317 318 let v_stk: nx_int = nx_f32_llama_stack_forward_v4b( 319 x_embed, x_stacked, M, model.n_layers, 320 hidden_dim, model.n_heads, model.n_kv_heads, model.head_dim, model.ffn_dim, 321 model.layers, seqs, eps, attn_scale, rope_log_base, apply_rope) 322 if v_stk != NX_STK4_OK { return NX_FLV4_ERR_STACK } 323 324 var t2: nx_int = 0 325 while t2 < M { 326 let src: *i64 = (((x_stacked as i64) + t2 * hidden_dim * 8)) as *i64 327 let dst: *i64 = (((x_normed as i64) + t2 * hidden_dim * 8)) as *i64 328 nx_f32_rmsnorm(src, model.gamma_out, hidden_dim, eps, dst) 329 t2 = t2 + 1 330 } 331 332 if model.lm_head_q8 != 0 { 333 nx_f32_lazy_matmul(x_normed, model.lm_head_q8 as *NxF32LazyWeight, logits, 334 M, hidden_dim, vocab_size) 335 } else { 336 nx_f32_matmul_t_pool(nx_lw_shared_pool(), x_normed, model.lm_head, logits, 337 M, hidden_dim, vocab_size) 338 } 339 340 sys_munmap(x_embed, M * hidden_dim * 8) 341 sys_munmap(x_stacked, M * hidden_dim * 8) 342 sys_munmap(x_normed, M * hidden_dim * 8) 343 return NX_FLV4_OK 344}