nx_llm_sched.nx source
↩ module page · 193 lines · 6345 B
1// nx_llm_sched.nx -- sovereign CONTINUOUS-BATCHING scheduler (Orca / vLLM
2// serving core). Requests ADMIT at any time (chunked prefill onto their own
3// paged sequence, forked-nothing: independent prompts) and JOIN the decode
4// batch mid-flight; each ROUND advances every active request by one token
5// via ONE batched forward (nx_f32_llama_v4b); finished requests leave the
6// batch and free their pages. The socket seat wraps THIS; the correctness
7// property gated by nx_llm_sched_gate is BATCH-INVARIANCE: a request's
8// bytes are BIT-IDENTICAL to its solo run, regardless of co-tenants or
9// admission timing.
10//
11// genealogy_id: yu_2022_orca_continuous_batching + kwon_2023_pagedattention
12// lineage_id: substrate_llm_sched_v1
13
14import "nx_syscalls.nx"
15import "nx_tier.nx"
16import "nx_prng.nx"
17import "nx_bpe.nx"
18import "nx_f32_llm.nx"
19import "nx_f32_llm_v4.nx"
20import "nx_f32_sampler.nx"
21import "nx_reasoning.nx"
22import "nx_kvcache.nx"
23import "nx_f32_attn_paged.nx"
24import "nx_f32_llama_v4p.nx"
25import "nx_f32_llama_v4b.nx"
26
27const NX_SCHED_MAX_REQS: nx_int = 16
28
29// per-request slot (12 fields x 8 = 96 bytes)
30struct NxLlmReq {
31 seq: *NxPagedSeq,
32 prng: *i64,
33 cur: i64, // next input token
34 out: *u8,
35 out_cap: nx_int,
36 out_len: i64,
37 emitted: i64,
38 max_new: nx_int,
39 done: i64, // 0 active, 1 finished
40 in_use: i64,
41 seed: i64,
42 pad: i64
43}
44const NX_LLM_REQ_BYTES: nx_int = 96
45
46struct NxLlmSched {
47 rc: *NxReasonCfg,
48 pool: *NxPagedPool,
49 reqs: *u8, // NX_SCHED_MAX_REQS slots
50 n_reqs: nx_int
51}
52const NX_LLM_SCHED_BYTES: nx_int = 32
53
54func nx_sched_new(rc: *NxReasonCfg, pool: *NxPagedPool) -> *NxLlmSched {
55 let S: *NxLlmSched = sys_mmap(NX_LLM_SCHED_BYTES) as *NxLlmSched
56 S.rc = rc
57 S.pool = pool
58 S.reqs = sys_mmap(NX_SCHED_MAX_REQS * NX_LLM_REQ_BYTES)
59 S.n_reqs = 0
60 return S
61}
62
63func nx_sched_req(S: *NxLlmSched, i: nx_int) -> *NxLlmReq {
64 return ((S.reqs as i64) + i * NX_LLM_REQ_BYTES) as *NxLlmReq
65}
66
67func _sched_stop(rc: *NxReasonCfg, t: nx_int) -> nx_int {
68 if t == rc.im_end { return 1 }
69 if rc.eos >= 0 { if t == rc.eos { return 1 } }
70 return 0
71}
72
73func _sched_emit(rc: *NxReasonCfg, R: *NxLlmReq, t: nx_int) -> i64 {
74 let one: *i64 = sys_mmap(8) as *i64
75 let db: *u8 = sys_mmap(64)
76 one[0] = t as i64
77 let dn: nx_int = nx_bpe_decode_bytelevel(rc.vocab, one, 1, db)
78 var bi: nx_int = 0
79 var no: i64 = R.out_len
80 while bi < dn {
81 if no < R.out_cap { R.out[no] = db[bi]; no = no + 1 }
82 bi = bi + 1
83 }
84 R.out_len = no
85 return 0
86}
87
88// ADMIT: chunked prefill onto a fresh paged seq; sample the first token;
89// the request joins the next round. Returns slot id or -1 (full/error).
90
91func nx_sched_admit(S: *NxLlmSched, toks: *i64, nt: nx_int, seed: i64,
92 max_new: nx_int, out: *u8, out_cap: nx_int) -> nx_int {
93 if S.n_reqs >= NX_SCHED_MAX_REQS { return 0 - 1 }
94 let rc: *NxReasonCfg = S.rc
95 let model: *NxF32LlamaModel = rc.model
96 let vs: nx_int = model.vocab_size
97
98 let seq: *NxPagedSeq = nx_pkv_seq_new(S.pool, nt + max_new + 32)
99 let lg: *i64 = sys_mmap((nt + 2) * vs * 8) as *i64
100 if nx_f32_llm_forward_v4p(model, toks, nt, seq, rc.eps, rc.attn_scale,
101 rc.rope_log_base, 1, lg) != NX_FLV4_OK {
102 return 0 - 1
103 }
104 let slot: nx_int = S.n_reqs
105 let R: *NxLlmReq = nx_sched_req(S, slot)
106 R.seq = seq
107 R.prng = sys_mmap(8) as *i64
108 nx_prng_init(R.prng, seed)
109 R.out = out
110 R.out_cap = out_cap
111 R.out_len = 0
112 R.emitted = 0
113 R.max_new = max_new
114 R.done = 0
115 R.in_use = 1
116 R.seed = seed
117
118 let lrow: *i64 = ((lg as i64) + (nt - 1) * vs * 8) as *i64
119 let t0: nx_int = nx_f32_sampler_sample_top_k(lrow, vs, rc.top_k,
120 rc.inv_temp_f32, R.prng)
121 if _sched_stop(rc, t0) == 1 { R.done = 1 } else {
122 _sched_emit(rc, R, t0)
123 R.cur = t0 as i64
124 R.emitted = 1
125 }
126 sys_munmap(lg, (nt + 2) * vs * 8)
127 S.n_reqs = S.n_reqs + 1
128 return slot
129}
130
131// ONE ROUND: batched forward advances every active request by one token.
132// Returns the number of requests still active after the round.
133
134func nx_sched_round(S: *NxLlmSched) -> nx_int {
135 let rc: *NxReasonCfg = S.rc
136 let model: *NxF32LlamaModel = rc.model
137 let vs: nx_int = model.vocab_size
138
139 let bat_ids: *i64 = sys_mmap(NX_SCHED_MAX_REQS * 8) as *i64
140 let bat_seqs: *i64 = sys_mmap(NX_SCHED_MAX_REQS * 8) as *i64
141 let bat_map: *i64 = sys_mmap(NX_SCHED_MAX_REQS * 8) as *i64
142 var Mb: nx_int = 0
143 var i: nx_int = 0
144 while i < S.n_reqs {
145 let R: *NxLlmReq = nx_sched_req(S, i)
146 if R.in_use == 1 {
147 if R.done == 0 {
148 if R.emitted < (R.max_new as i64) {
149 bat_ids[Mb] = R.cur
150 bat_seqs[Mb] = R.seq as i64
151 bat_map[Mb] = i as i64
152 Mb = Mb + 1
153 } else {
154 R.done = 1
155 }
156 }
157 }
158 i = i + 1
159 }
160 if Mb == 0 { return 0 }
161
162 let blg: *i64 = sys_mmap(Mb * vs * 8) as *i64
163 if nx_f32_llm_forward_v4b(model, bat_ids, Mb, bat_seqs, rc.eps,
164 rc.attn_scale, rc.rope_log_base, 1,
165 blg) != NX_FLV4_OK { return 0 - 1 }
166 var r: nx_int = 0
167 var still: nx_int = 0
168 while r < Mb {
169 let ri: nx_int = bat_map[r] as nx_int
170 let R2: *NxLlmReq = nx_sched_req(S, ri)
171 let row: *i64 = ((blg as i64) + r * vs * 8) as *i64
172 let nid: nx_int = nx_f32_sampler_sample_top_k(row, vs, rc.top_k,
173 rc.inv_temp_f32, R2.prng)
174 if _sched_stop(rc, nid) == 1 { R2.done = 1 } else {
175 _sched_emit(rc, R2, nid)
176 R2.cur = nid as i64
177 R2.emitted = R2.emitted + 1
178 if R2.emitted >= (R2.max_new as i64) { R2.done = 1 } else { still = still + 1 }
179 }
180 r = r + 1
181 }
182 sys_munmap(blg, Mb * vs * 8)
183 return still
184}
185
186func nx_sched_release(S: *NxLlmSched, slot: nx_int) -> i64 {
187 let R: *NxLlmReq = nx_sched_req(S, slot)
188 if R.in_use == 1 {
189 nx_pkv_seq_free(R.seq)
190 R.in_use = 0
191 }
192 return 0
193}