code wiki / (root) / nx_llm_sched_gate.nx

nx_llm_sched_gate.nx source

↩ module page · 228 lines · 8978 B

1// nx_llm_sched_gate.nx -- MEASURED gate for the CONTINUOUS-BATCHING 2// scheduler (nx_llm_sched) on the REAL model. 3// 4// STAGGER 4 requests admitted at different times (2 up front, 1 after 5// round 2, 1 after round 4) -- requests JOIN the running batch 6// mid-flight (the continuous-batching property) 7// INVARIANT each request's bytes BIT-IDENTICAL to its SOLO sequential 8// run (same seed) -- output independent of co-tenants and 9// admission timing (the fundamental serving correctness) 10// HYGIENE pool all-free after releases 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" 40import "nx_llm_sched.nx" 41 42const SG_MAXNEW: nx_int = 6 43 44func sg2_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 sg2_wn(v: i64) -> i64 { 46 var m: i64 = v 47 if m < 0 { sg2_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 sg2_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} 63 64// solo reference: fresh seq, sequential m=1 decode, same structure as the 65// scheduler's sample accounting (prefill-sample, then per-step fwd+sample). 66func sg2_solo(rc: *NxReasonCfg, pool: *NxPagedPool, toks: *i64, nt: nx_int, 67 seed: i64, out: *u8, out_cap: nx_int) -> i64 { 68 let model: *NxF32LlamaModel = rc.model 69 let vs: nx_int = model.vocab_size 70 let seq: *NxPagedSeq = nx_pkv_seq_new(pool, nt + SG_MAXNEW + 32) 71 let lg: *i64 = sys_mmap((nt + 2) * vs * 8) as *i64 72 if nx_f32_llm_forward_v4p(model, toks, nt, seq, rc.eps, rc.attn_scale, 73 rc.rope_log_base, 1, lg) != NX_FLV4_OK { return 0 - 1 } 74 let prng: *i64 = sys_mmap(8) as *i64 75 nx_prng_init(prng, seed) 76 let one: *i64 = sys_mmap(8) as *i64 77 let db: *u8 = sys_mmap(64) 78 let lrow: *i64 = ((lg as i64) + (nt - 1) * vs * 8) as *i64 79 var cur: nx_int = nx_f32_sampler_sample_top_k(lrow, vs, rc.top_k, 80 rc.inv_temp_f32, prng) 81 var no: i64 = 0 82 var emitted: i64 = 0 83 var alive: nx_int = 1 84 if cur == rc.im_end { alive = 0 } 85 if rc.eos >= 0 { if cur == rc.eos { alive = 0 } } 86 while alive == 1 { 87 one[0] = cur as i64 88 let dn: nx_int = nx_bpe_decode_bytelevel(rc.vocab, one, 1, db) 89 var bi: nx_int = 0 90 while bi < dn { 91 if no < out_cap { out[no] = db[bi]; no = no + 1 } 92 bi = bi + 1 93 } 94 emitted = emitted + 1 95 if emitted >= (SG_MAXNEW as i64) { alive = 0 } else { 96 if nx_f32_llm_forward_v4p(model, one, 1, seq, rc.eps, rc.attn_scale, 97 rc.rope_log_base, 1, lg) != NX_FLV4_OK { return 0 - 1 } 98 let nid: nx_int = nx_f32_sampler_sample_top_k(lg, vs, rc.top_k, 99 rc.inv_temp_f32, prng) 100 if nid == rc.im_end { alive = 0 } else { 101 var st2: nx_int = 0 102 if rc.eos >= 0 { if nid == rc.eos { st2 = 1 } } 103 if st2 == 1 { alive = 0 } else { cur = nid } 104 } 105 } 106 } 107 nx_pkv_seq_free(seq) 108 sys_munmap(lg, (nt + 2) * vs * 8) 109 return no 110} 111 112func main() -> i64 { 113 let path: *u8 = "/tmp/nx_real_model.gguf" as *u8 114 let len_out: *i64 = sys_mmap(8) as *i64 115 let buf: *u8 = sys_read_file(path, len_out) 116 if buf == (0 as *u8) { return 10 } 117 let hdr: *NxGgufHeader = sys_mmap(NX_GGUF_HDR_BYTES) as *NxGgufHeader 118 if nx_gguf_parse(buf, len_out[0], hdr) != NX_GGUF_OK { return 20 } 119 let model: *NxF32LlamaModel = nx_f32_llama_model_alloc() 120 let out_err: *i64 = sys_mmap(8) as *i64 121 if nx_f32_llm_read_dims_from_gguf(buf, len_out[0], hdr, model, out_err) != NX_FLD_OK { return 30 } 122 if nx_f32_llm_load_weights_v4_from_gguf(buf, hdr, model, out_err) != NX_FLV4_OK { return 40 } 123 let vocab: *NxBpeVocab = nx_bpe_vocab_new(67108864, 262144, 524288) 124 let nt2: *i64 = sys_mmap(8) as *i64 125 let nm: *i64 = sys_mmap(8) as *i64 126 if nx_f32_bpe_load_from_gguf(buf, len_out[0], hdr, vocab, nt2, nm, out_err) != NX_FBL_OK { return 50 } 127 let eos: nx_int = nx_f32_llm_read_eos(buf, len_out[0], hdr) 128 129 let rc: *NxReasonCfg = nx_reason_cfg_alloc() 130 rc.model = model 131 rc.vocab = vocab 132 rc.cache = 0 as *NxF32KVCache 133 rc.max_new = SG_MAXNEW 134 rc.inv_temp_f32 = 0x3FA00000 135 rc.top_k = 40 136 rc.eps = 0x358637BD 137 rc.attn_scale = 0x3E000000 138 rc.rope_log_base = 0x415D0EAB 139 rc.eos = eos 140 rc.im_start = 151644 141 rc.im_end = 151645 142 143 let kv_dim: nx_int = model.n_kv_heads * model.head_dim 144 let pool: *NxPagedPool = nx_pkv_pool_new(32, model.n_layers, kv_dim) 145 146 // two distinct prompts, four requests, four seeds. 147 let qA: *u8 = sys_mmap(96) 148 let nqA0: nx_int = nx_reason_build_q_short(qA, 47, 0, 38) 149 let tA: *i64 = sys_mmap(256 * 8) as *i64 150 let ntA: nx_int = nx_reason_build_chat_toks(rc, qA, nqA0, tA) 151 let qB: *u8 = sys_mmap(96) 152 let nqB0: nx_int = nx_reason_build_q_short(qB, 83, 1, 47) 153 let tB: *i64 = sys_mmap(256 * 8) as *i64 154 let ntB: nx_int = nx_reason_build_chat_toks(rc, qB, nqB0, tB) 155 156 // ---- scheduler with staggered admission -------------------------- 157 let S: *NxLlmSched = nx_sched_new(rc, pool) 158 let o0: *u8 = sys_mmap(96) 159 let o1: *u8 = sys_mmap(96) 160 let o2: *u8 = sys_mmap(96) 161 let o3: *u8 = sys_mmap(96) 162 let s0: nx_int = nx_sched_admit(S, tA, ntA, 111, SG_MAXNEW, o0, 96) 163 let s1: nx_int = nx_sched_admit(S, tB, ntB, 222, SG_MAXNEW, o1, 96) 164 if s0 < 0 { return 60 } 165 if s1 < 0 { return 60 } 166 var rounds: i64 = 0 167 nx_sched_round(S); rounds = rounds + 1 168 nx_sched_round(S); rounds = rounds + 1 169 let s2: nx_int = nx_sched_admit(S, tA, ntA, 333, SG_MAXNEW, o2, 96) 170 if s2 < 0 { return 60 } 171 nx_sched_round(S); rounds = rounds + 1 172 nx_sched_round(S); rounds = rounds + 1 173 let s3: nx_int = nx_sched_admit(S, tB, ntB, 444, SG_MAXNEW, o3, 96) 174 if s3 < 0 { return 60 } 175 var act: nx_int = 1 176 var guard: i64 = 0 177 while act > 0 { 178 act = nx_sched_round(S) 179 if act < 0 { return 61 } 180 rounds = rounds + 1 181 guard = guard + 1 182 if guard > 32 { return 62 } 183 } 184 sg2_w("SCHED rounds=" as *u8); sg2_wn(rounds) 185 sg2_w(" lens:" as *u8) 186 let R0: *NxLlmReq = nx_sched_req(S, s0) 187 let R1: *NxLlmReq = nx_sched_req(S, s1) 188 let R2: *NxLlmReq = nx_sched_req(S, s2) 189 let R3: *NxLlmReq = nx_sched_req(S, s3) 190 sg2_w(" " as *u8); sg2_wn(R0.out_len) 191 sg2_w(" " as *u8); sg2_wn(R1.out_len) 192 sg2_w(" " as *u8); sg2_wn(R2.out_len) 193 sg2_w(" " as *u8); sg2_wn(R3.out_len) 194 sg2_w("\n" as *u8) 195 196 // ---- solo references (same seeds) --------------------------------- 197 let r0: *u8 = sys_mmap(96) 198 let r1: *u8 = sys_mmap(96) 199 let r2: *u8 = sys_mmap(96) 200 let r3: *u8 = sys_mmap(96) 201 let n0: i64 = sg2_solo(rc, pool, tA, ntA, 111, r0, 96) 202 let n1: i64 = sg2_solo(rc, pool, tB, ntB, 222, r1, 96) 203 let n2: i64 = sg2_solo(rc, pool, tA, ntA, 333, r2, 96) 204 let n3: i64 = sg2_solo(rc, pool, tB, ntB, 444, r3, 96) 205 206 if n0 != R0.out_len { return 70 } 207 if n1 != R1.out_len { return 71 } 208 if n2 != R2.out_len { return 72 } 209 if n3 != R3.out_len { return 73 } 210 if sg2_beq(o0, r0, n0) != 1 { return 74 } 211 if sg2_beq(o1, r1, n1) != 1 { return 75 } 212 if sg2_beq(o2, r2, n2) != 1 { return 76 } 213 if sg2_beq(o3, r3, n3) != 1 { return 77 } 214 sg2_w("SCHED INVARIANT all 4 staggered outputs == solo runs BIT-IDENTICAL OK\n" as *u8) 215 216 nx_sched_release(S, s0) 217 nx_sched_release(S, s1) 218 nx_sched_release(S, s2) 219 nx_sched_release(S, s3) 220 if pool.n_free != 32 { 221 sg2_w("HYGIENE n_free=" as *u8); sg2_wn(pool.n_free); sg2_w(" want 32\n" as *u8) 222 return 80 223 } 224 sg2_w("SCHED HYGIENE pool all-free OK\n" as *u8) 225 sg2_w("LIAR-KILL stagger-admission=1 batch-invariance=1 hygiene=1\n" as *u8) 226 sg2_w("LLM_SCHED_GATE DONE\n" as *u8) 227 return 0 228}