nx_batched_gate.nx source
↩ module page · 288 lines · 11217 B
1// nx_batched_gate.nx -- MEASURED gate for BATCHED MULTI-SEQUENCE decode
2// (nx_f32_llama_v4b) on the REAL model.
3//
4// EQUIV batched decode == sequential fork decode BIT-IDENTICAL:
5// prefill once, fork N=3 twice (same seeds); decode one set
6// sequentially (v4p, m=1 per fork) and one set BATCHED (v4b,
7// one M-row forward per round); texts must match byte-for-byte
8// SPEED decode phases timed separately from the shared prefill:
9// sequential N*T m=1 forwards vs T M-row forwards
10// HYGIENE pool returns to all-free
11//
12// license_tier: ORIGINAL expect_exit: 0
13
14import "nx_syscalls.nx"
15import "nx_tier.nx"
16import "nx_le.nx"
17import "nx_bpe.nx"
18import "nx_gguf.nx"
19import "nx_gguf_load.nx"
20import "nx_gguf_meta.nx"
21import "nx_f32.nx"
22import "nx_f32_kv_cache.nx"
23import "nx_f32_lazy_weight.nx"
24import "nx_f32_llama_block.nx"
25import "nx_f32_llama_block_v4.nx"
26import "nx_f32_llama_stack_v4.nx"
27import "nx_f32_llama_layer_lazy_load.nx"
28import "nx_f32_llm.nx"
29import "nx_f32_llm_v4.nx"
30import "nx_f32_llm_read_dims.nx"
31import "nx_f32_bpe_load.nx"
32import "nx_f32_llm_special_tokens.nx"
33import "nx_f32_sampler.nx"
34import "nx_prng.nx"
35import "nx_reasoning.nx"
36import "nx_kvcache.nx"
37import "nx_f32_attn_paged.nx"
38import "nx_f32_llama_v4p.nx"
39import "nx_f32_llama_v4b.nx"
40
41const BG_N: nx_int = 3
42const BG_MAXNEW: nx_int = 8
43
44func bg_w(s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } sys_write(1, s, n); return 0 }
45func bg_wn(v: i64) -> i64 {
46 var m: i64 = v
47 if m < 0 { bg_w("-" as *u8); m = 0 - m }
48 let t: *u8 = sys_mmap(28)
49 var k: i64 = 0
50 if m == 0 { t[0] = 48 as u8; k = 1 }
51 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 }
52 let o: *u8 = sys_mmap(28)
53 var i: i64 = 0
54 while i < k { o[i] = t[k - 1 - i]; i = i + 1 }
55 sys_write(1, o, k)
56 return 0
57}
58func bg_beq(a: *u8, b: *u8, n: i64) -> i64 {
59 var i: i64 = 0
60 while i < n { if a[i] != b[i] { return 0 } i = i + 1 }
61 return 1
62}
63func bg_stop(rc: *NxReasonCfg, t: nx_int) -> nx_int {
64 if t == rc.im_end { return 1 }
65 if rc.eos >= 0 { if t == rc.eos { return 1 } }
66 return 0
67}
68// append decoded bytes of token t to buf at *len (cap 96); returns 0.
69func bg_emit(rc: *NxReasonCfg, t: nx_int, buf: *u8, len: *i64, one: *i64, db: *u8) -> i64 {
70 one[0] = t as i64
71 let dn: nx_int = nx_bpe_decode_bytelevel(rc.vocab, one, 1, db)
72 var bi: nx_int = 0
73 var no: i64 = len[0]
74 while bi < dn {
75 if no < 96 { buf[no] = db[bi]; no = no + 1 }
76 bi = bi + 1
77 }
78 len[0] = no
79 return 0
80}
81
82func main() -> i64 {
83 // ---- load the real model once ----------------------------------
84 let path: *u8 = "/tmp/nx_real_model.gguf" as *u8
85 let len_out: *i64 = sys_mmap(8) as *i64
86 let buf: *u8 = sys_read_file(path, len_out)
87 if buf == (0 as *u8) { return 10 }
88 let hdr: *NxGgufHeader = sys_mmap(NX_GGUF_HDR_BYTES) as *NxGgufHeader
89 if nx_gguf_parse(buf, len_out[0], hdr) != NX_GGUF_OK { return 20 }
90 let model: *NxF32LlamaModel = nx_f32_llama_model_alloc()
91 let out_err: *i64 = sys_mmap(8) as *i64
92 if nx_f32_llm_read_dims_from_gguf(buf, len_out[0], hdr, model, out_err) != NX_FLD_OK { return 30 }
93 if nx_f32_llm_load_weights_v4_from_gguf(buf, hdr, model, out_err) != NX_FLV4_OK { return 40 }
94 let vocab: *NxBpeVocab = nx_bpe_vocab_new(67108864, 262144, 524288)
95 let nt2: *i64 = sys_mmap(8) as *i64
96 let nm: *i64 = sys_mmap(8) as *i64
97 if nx_f32_bpe_load_from_gguf(buf, len_out[0], hdr, vocab, nt2, nm, out_err) != NX_FBL_OK { return 50 }
98 let eos: nx_int = nx_f32_llm_read_eos(buf, len_out[0], hdr)
99
100 let rc: *NxReasonCfg = nx_reason_cfg_alloc()
101 rc.model = model
102 rc.vocab = vocab
103 rc.cache = 0 as *NxF32KVCache
104 rc.max_new = BG_MAXNEW
105 rc.inv_temp_f32 = 0x3FA00000
106 rc.top_k = 40
107 rc.eps = 0x358637BD
108 rc.attn_scale = 0x3E000000
109 rc.rope_log_base = 0x415D0EAB
110 rc.eos = eos
111 rc.im_start = 151644
112 rc.im_end = 151645
113
114 let kv_dim: nx_int = model.n_kv_heads * model.head_dim
115 let pool: *NxPagedPool = nx_pkv_pool_new(24, model.n_layers, kv_dim)
116 let vs: nx_int = model.vocab_size
117
118 let q1: *u8 = "What is 47 plus 38? Answer with only the number." as *u8
119 let toks: *i64 = sys_mmap(512 * 8) as *i64
120 let ntq: nx_int = nx_reason_build_chat_toks(rc, q1, 49, toks)
121
122 // ---- shared prefill once ----------------------------------------
123 let seq0: *NxPagedSeq = nx_pkv_seq_new(pool, 64)
124 let lg: *i64 = sys_mmap((ntq + 2) * vs * 8) as *i64
125 let tpf0: i64 = sys_now_us()
126 if nx_f32_llm_forward_v4p(model, toks, ntq, seq0, rc.eps, rc.attn_scale,
127 rc.rope_log_base, 1, lg) != NX_FLV4_OK { return 60 }
128 let prefill_us: i64 = sys_now_us() - tpf0
129 let base_row: *i64 = sys_mmap(vs * 8) as *i64
130 let srcrow: *i64 = ((lg as i64) + (ntq - 1) * vs * 8) as *i64
131 var bc: nx_int = 0
132 while bc < vs { base_row[bc] = srcrow[bc]; bc = bc + 1 }
133
134 let one: *i64 = sys_mmap(8) as *i64
135 let db: *u8 = sys_mmap(64)
136
137 // ---- SEQUENTIAL reference: fork + m=1 decode per fork ------------
138 let txS: *u8 = sys_mmap(BG_N * 96)
139 let lnS: *i64 = sys_mmap(BG_N * 8) as *i64
140 let ts0: i64 = sys_now_us()
141 var s: nx_int = 0
142 while s < BG_N {
143 let fseq: *NxPagedSeq = nx_pkv_seq_fork(seq0)
144 let prng: *i64 = sys_mmap(8) as *i64
145 nx_prng_init(prng, 20260709 + (s as i64) * 7919)
146 var cw: nx_int = 0
147 while cw < vs { lg[cw] = base_row[cw]; cw = cw + 1 }
148 let obuf: *u8 = ((txS as i64) + s * 96) as *u8
149 let olen: *i64 = ((lnS as i64) + s * 8) as *i64
150 olen[0] = 0
151 var step: nx_int = 0
152 while step < BG_MAXNEW {
153 let nid: nx_int = nx_f32_sampler_sample_top_k(lg, vs, rc.top_k,
154 rc.inv_temp_f32, prng)
155 if bg_stop(rc, nid) == 1 { step = BG_MAXNEW } else {
156 bg_emit(rc, nid, obuf, olen, one, db)
157 one[0] = nid as i64
158 if nx_f32_llm_forward_v4p(model, one, 1, fseq, rc.eps, rc.attn_scale,
159 rc.rope_log_base, 1, lg) != NX_FLV4_OK { return 61 }
160 step = step + 1
161 }
162 }
163 nx_pkv_seq_free(fseq)
164 s = s + 1
165 }
166 let seq_us: i64 = sys_now_us() - ts0
167
168 // ---- BATCHED: fork again (same seeds), decode via v4b rounds -----
169 let txB: *u8 = sys_mmap(BG_N * 96)
170 let lnB: *i64 = sys_mmap(BG_N * 8) as *i64
171 let fseqs: *i64 = sys_mmap(BG_N * 8) as *i64
172 let prngs: *i64 = sys_mmap(BG_N * 8) as *i64
173 let curs: *i64 = sys_mmap(BG_N * 8) as *i64
174 let done: *i64 = sys_mmap(BG_N * 8) as *i64
175 let emitted: *i64 = sys_mmap(BG_N * 8) as *i64
176 let tb0: i64 = sys_now_us()
177 var f: nx_int = 0
178 while f < BG_N {
179 let fs: *NxPagedSeq = nx_pkv_seq_fork(seq0)
180 fseqs[f] = fs as i64
181 let pr: *i64 = sys_mmap(8) as *i64
182 nx_prng_init(pr, 20260709 + (f as i64) * 7919)
183 prngs[f] = pr as i64
184 let ob: *u8 = ((txB as i64) + f * 96) as *u8
185 let ol: *i64 = ((lnB as i64) + f * 8) as *i64
186 ol[0] = 0
187 done[f] = 0
188 emitted[f] = 0
189 // first sample from the SHARED base row (non-mutating sampler).
190 let nid0: nx_int = nx_f32_sampler_sample_top_k(base_row, vs, rc.top_k,
191 rc.inv_temp_f32, pr)
192 if bg_stop(rc, nid0) == 1 { done[f] = 1 } else {
193 bg_emit(rc, nid0, ob, ol, one, db)
194 curs[f] = nid0 as i64
195 emitted[f] = 1
196 }
197 f = f + 1
198 }
199 // rounds: ONE batched forward advances every active fork.
200 let bat_ids: *i64 = sys_mmap(BG_N * 8) as *i64
201 let bat_seqs: *i64 = sys_mmap(BG_N * 8) as *i64
202 let bat_map: *i64 = sys_mmap(BG_N * 8) as *i64
203 let blg: *i64 = sys_mmap(BG_N * vs * 8) as *i64
204 var running: nx_int = 1
205 while running == 1 {
206 // collect active forks that still need a forward (emitted < max_new).
207 var Mb: nx_int = 0
208 var f2: nx_int = 0
209 while f2 < BG_N {
210 if done[f2] == 0 {
211 if emitted[f2] < BG_MAXNEW {
212 bat_ids[Mb] = curs[f2]
213 bat_seqs[Mb] = fseqs[f2]
214 bat_map[Mb] = f2
215 Mb = Mb + 1
216 } else {
217 done[f2] = 1
218 }
219 }
220 f2 = f2 + 1
221 }
222 if Mb == 0 { running = 0 } else {
223 if nx_f32_llm_forward_v4b(model, bat_ids, Mb, bat_seqs, rc.eps,
224 rc.attn_scale, rc.rope_log_base, 1,
225 blg) != NX_FLV4_OK { return 62 }
226 var r2: nx_int = 0
227 while r2 < Mb {
228 let fk: nx_int = bat_map[r2] as nx_int
229 let row: *i64 = ((blg as i64) + r2 * vs * 8) as *i64
230 let pr2: *i64 = prngs[fk] as *i64
231 let nid: nx_int = nx_f32_sampler_sample_top_k(row, vs, rc.top_k,
232 rc.inv_temp_f32, pr2)
233 if bg_stop(rc, nid) == 1 { done[fk] = 1 } else {
234 let ob2: *u8 = ((txB as i64) + fk * 96) as *u8
235 let ol2: *i64 = ((lnB as i64) + fk * 8) as *i64
236 bg_emit(rc, nid, ob2, ol2, one, db)
237 curs[fk] = nid as i64
238 emitted[fk] = emitted[fk] + 1
239 }
240 r2 = r2 + 1
241 }
242 }
243 }
244 let bat_us: i64 = sys_now_us() - tb0
245 var f3: nx_int = 0
246 while f3 < BG_N {
247 nx_pkv_seq_free(fseqs[f3] as *NxPagedSeq)
248 f3 = f3 + 1
249 }
250
251 // ---- EQUIV: texts bit-identical ----------------------------------
252 var c: nx_int = 0
253 while c < BG_N {
254 if lnS[c] != lnB[c] {
255 bg_w("BATCH EQUIV LEN MISMATCH f=" as *u8); bg_wn(c as i64)
256 bg_w(" seq=" as *u8); bg_wn(lnS[c])
257 bg_w(" bat=" as *u8); bg_wn(lnB[c]); bg_w("\n" as *u8)
258 return 70
259 }
260 let pa: *u8 = ((txS as i64) + c * 96) as *u8
261 let pb: *u8 = ((txB as i64) + c * 96) as *u8
262 if bg_beq(pa, pb, lnS[c]) != 1 {
263 bg_w("BATCH EQUIV BYTES MISMATCH f=" as *u8); bg_wn(c as i64); bg_w("\n" as *u8)
264 return 71
265 }
266 c = c + 1
267 }
268 bg_w("BG EQUIV batched decode == sequential fork decode BIT-IDENTICAL (N=3) OK\n" as *u8)
269
270 bg_w("SPEED prefill_us=" as *u8); bg_wn(prefill_us)
271 bg_w(" seq_decode_us=" as *u8); bg_wn(seq_us)
272 bg_w(" batched_decode_us=" as *u8); bg_wn(bat_us)
273 var b2: i64 = bat_us
274 if b2 < 1 { b2 = 1 }
275 bg_w(" seq_vs_batched_x100=" as *u8); bg_wn(seq_us * 100 / b2)
276 bg_w("\n" as *u8)
277
278 // ---- HYGIENE ------------------------------------------------------
279 nx_pkv_seq_free(seq0)
280 if pool.n_free != 24 {
281 bg_w("HYGIENE pool n_free=" as *u8); bg_wn(pool.n_free); bg_w(" want 24\n" as *u8)
282 return 80
283 }
284 bg_w("BG HYGIENE pool all-free OK\n" as *u8)
285 bg_w("LIAR-KILL equiv-bitident=1 hygiene=1\n" as *u8)
286 bg_w("BATCHED_GATE DONE\n" as *u8)
287 return 0
288}