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}