code wiki / (root) / nx_llm_sched.nx

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}