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}