nx_batched_tput.nx source
↩ module page · 149 lines · 6615 B
1// nx_batched_tput.nx -- AGGREGATE THROUGHPUT of batched multi-sequence
2// decode (the server axis). Single-stream decode is at its CPU floor
3// (~15 tok/s); batching reads the 209MB weights ONCE per round for M
4// sequences, so aggregate tok/s should scale ~M until compute catches
5// memory. Prefill ONE prompt, fork M ways (prefix-shared), run T batched
6// v4b rounds, report aggregate M*T tokens / wall. Sweep M=1,2,4,8. This
7// is the honest "how fast can the sovereign engine actually serve" number.
8// expect_exit: 0 license_tier: ORIGINAL
9import "nx_syscalls.nx"
10import "nx_tier.nx"
11import "nx_le.nx"
12import "nx_bpe.nx"
13import "nx_gguf.nx"
14import "nx_gguf_load.nx"
15import "nx_gguf_meta.nx"
16import "nx_f32.nx"
17import "nx_f32_kv_cache.nx"
18import "nx_f32_lazy_weight.nx"
19import "nx_f32_llama_block.nx"
20import "nx_f32_llama_block_v4.nx"
21import "nx_f32_llama_stack_v4.nx"
22import "nx_f32_llama_layer_lazy_load.nx"
23import "nx_f32_llm.nx"
24import "nx_f32_llm_v4.nx"
25import "nx_f32_llm_read_dims.nx"
26import "nx_f32_bpe_load.nx"
27import "nx_f32_llm_special_tokens.nx"
28import "nx_f32_sampler.nx"
29import "nx_reasoning.nx"
30import "nx_kvcache.nx"
31import "nx_f32_attn_paged.nx"
32import "nx_f32_llama_v4p.nx"
33import "nx_f32_llama_v4b.nx"
34const BT_MAGIC_1000000000: i64 = 1000000000
35const BT_MAGIC_67108864: i64 = 67108864
36const BT_MAGIC_262144: i64 = 262144
37const BT_MAGIC_524288: i64 = 524288
38const BT_MAGIC_151644: i64 = 151644
39const BT_MAGIC_151645: i64 = 151645
40
41const BT_T: nx_int = 16 // decode rounds (tokens/seq) per measurement
42
43func bt_w(s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } sys_write(1, s, n); return 0 }
44func bt_wn(v: i64) -> i64 {
45 var m: i64 = v
46 if m < 0 { bt_w("-" as *u8); m = 0 - m }
47 let t: *u8 = sys_mmap(28)
48 var k: i64 = 0
49 if m == 0 { t[0] = 48 as u8; k = 1 }
50 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 }
51 let o: *u8 = sys_mmap(28)
52 var i: i64 = 0
53 while i < k { o[i] = t[k - 1 - i]; i = i + 1 }
54 sys_write(1, o, k)
55 return 0
56}
57
58// run M forks x BT_T batched rounds; return aggregate tokens/sec (milli).
59func bt_measure(rc: *NxReasonCfg, pool: *NxPagedPool, seq0: *NxPagedSeq,
60 base_row: *i64, M: nx_int) -> i64 {
61 let model: *NxF32LlamaModel = rc.model
62 let vs: nx_int = model.vocab_size
63 let seqs: *i64 = sys_mmap(M * 8) as *i64
64 let cur: *i64 = sys_mmap(M * 8) as *i64
65 var f: nx_int = 0
66 while f < M {
67 let fs: *NxPagedSeq = nx_pkv_seq_fork(seq0)
68 seqs[f] = fs as i64
69 cur[f] = nx_f32_sampler_argmax(base_row, vs) as i64 // deterministic
70 f = f + 1
71 }
72 let blg: *i64 = sys_mmap(M * vs * 8) as *i64
73 let ids: *i64 = sys_mmap(M * 8) as *i64
74 let t0: i64 = sys_now_us()
75 var r: nx_int = 0
76 while r < BT_T {
77 var i: nx_int = 0
78 while i < M { ids[i] = cur[i]; i = i + 1 }
79 if nx_f32_llm_forward_v4b(model, ids, M, seqs, rc.eps, rc.attn_scale,
80 rc.rope_log_base, 1, blg) != NX_FLV4_OK { return 0 - 1 }
81 var j: nx_int = 0
82 while j < M {
83 let row: *i64 = ((blg as i64) + j * vs * 8) as *i64
84 cur[j] = nx_f32_sampler_argmax(row, vs) as i64
85 j = j + 1
86 }
87 r = r + 1
88 }
89 let us: i64 = sys_now_us() - t0
90 var f2: nx_int = 0
91 while f2 < M { nx_pkv_seq_free(seqs[f2] as *NxPagedSeq); f2 = f2 + 1 }
92 var u: i64 = us
93 if u < 1 { u = 1 }
94 // aggregate tokens = M * BT_T; tok/s milli = tokens*1e6/us... report tok/s*1000
95 return (M as i64) * (BT_T as i64) * BT_MAGIC_1000000000 / u
96}
97
98func main() -> i64 {
99 let path: *u8 = "/tmp/nx_real_model.gguf" as *u8
100 let len_out: *i64 = sys_mmap(8) as *i64
101 let buf: *u8 = sys_read_file(path, len_out)
102 if buf == (0 as *u8) { return 10 }
103 let hdr: *NxGgufHeader = sys_mmap(NX_GGUF_HDR_BYTES) as *NxGgufHeader
104 if nx_gguf_parse(buf, len_out[0], hdr) != NX_GGUF_OK { return 20 }
105 let model: *NxF32LlamaModel = nx_f32_llama_model_alloc()
106 let oe: *i64 = sys_mmap(8) as *i64
107 if nx_f32_llm_read_dims_from_gguf(buf, len_out[0], hdr, model, oe) != NX_FLD_OK { return 30 }
108 if nx_f32_llm_load_weights_v4_from_gguf(buf, hdr, model, oe) != NX_FLV4_OK { return 40 }
109 let vocab: *NxBpeVocab = nx_bpe_vocab_new(BT_MAGIC_67108864, BT_MAGIC_262144, BT_MAGIC_524288)
110 let nt: *i64 = sys_mmap(8) as *i64
111 let nm: *i64 = sys_mmap(8) as *i64
112 if nx_f32_bpe_load_from_gguf(buf, len_out[0], hdr, vocab, nt, nm, oe) != NX_FBL_OK { return 50 }
113 let eos: nx_int = nx_f32_llm_read_eos(buf, len_out[0], hdr)
114
115 let rc: *NxReasonCfg = nx_reason_cfg_alloc()
116 rc.model = model; rc.vocab = vocab; rc.cache = 0 as *NxF32KVCache
117 rc.max_new = BT_T; rc.inv_temp_f32 = 0x3F800000; rc.top_k = 0
118 rc.eps = 0x358637BD; rc.attn_scale = 0x3E000000; rc.rope_log_base = 0x415D0EAB
119 rc.eos = eos; rc.im_start = BT_MAGIC_151644; rc.im_end = BT_MAGIC_151645
120
121 let kv_dim: nx_int = model.n_kv_heads * model.head_dim
122 let pool: *NxPagedPool = nx_pkv_pool_new(256, model.n_layers, kv_dim)
123
124 // one shared prefill (excluded from timing).
125 let q: *u8 = "The history of the world in brief:" as *u8
126 let toks: *i64 = sys_mmap(512 * 8) as *i64
127 let ntq: nx_int = nx_reason_build_chat_toks(rc, q, 34, toks)
128 let seq0: *NxPagedSeq = nx_pkv_seq_new(pool, 64)
129 let vs: nx_int = model.vocab_size
130 let lg: *i64 = sys_mmap((ntq + 2) * vs * 8) as *i64
131 if nx_f32_llm_forward_v4p(model, toks, ntq, seq0, rc.eps, rc.attn_scale,
132 rc.rope_log_base, 1, lg) != NX_FLV4_OK { return 60 }
133 let base_row: *i64 = sys_mmap(vs * 8) as *i64
134 let src: *i64 = ((lg as i64) + (ntq - 1) * vs * 8) as *i64
135 var bc: nx_int = 0
136 while bc < vs { base_row[bc] = src[bc]; bc = bc + 1 }
137
138 bt_w("=== BATCHED DECODE AGGREGATE THROUGHPUT (futex pool + hw f32) ===\n" as *u8)
139 let one: i64 = bt_measure(rc, pool, seq0, base_row, 1)
140 bt_w("M=1 agg_tok_s_milli="); bt_wn(one); bt_w("\n" as *u8)
141 let two: i64 = bt_measure(rc, pool, seq0, base_row, 2)
142 bt_w("M=2 agg_tok_s_milli="); bt_wn(two); bt_w(" scale_x100_vs_M1="); bt_wn(two * 100 / (one + 1)); bt_w("\n" as *u8)
143 let four: i64 = bt_measure(rc, pool, seq0, base_row, 4)
144 bt_w("M=4 agg_tok_s_milli="); bt_wn(four); bt_w(" scale_x100_vs_M1="); bt_wn(four * 100 / (one + 1)); bt_w("\n" as *u8)
145 let eight: i64 = bt_measure(rc, pool, seq0, base_row, 8)
146 bt_w("M=8 agg_tok_s_milli="); bt_wn(eight); bt_w(" scale_x100_vs_M1="); bt_wn(eight * 100 / (one + 1)); bt_w("\n" as *u8)
147 bt_w("BATCHED_TPUT DONE\n" as *u8)
148 return 0
149}