code wiki / (root) / nx_batched_gate.nx

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}