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}