nx_q8_mix_ab.nx source
↩ module page · 128 lines · 4906 B
1// nx_q8_mix_ab.nx -- same-process, same-load A/B of the REAL per-layer block
2// matmul mix (Qwen2.5-0.5B: W_q/W_k/W_v/W_o/W_gate/W_up/W_down), m=1 decode.
3// A = force ALL through the pool (old behavior). B = threshold-routed
4// (small -> single-thread, big -> pool = new behavior). Run back-to-back
5// under IDENTICAL host load so the RATIO is noise-immune (absolute forward
6// benches are not -- this box runs concurrent sessions). Reports total us
7// for one "layer mix" repeated NLAY times, both ways, + the speedup.
8// license_tier: ORIGINAL expect_exit: 0
9import "nx_syscalls.nx"
10import "nx_tier.nx"
11import "nx_le.nx"
12import "nx_f32.nx"
13import "nx_f32_cvt.nx"
14import "nx_thread_pool.nx"
15import "nx_f32_lazy_weight.nx"
16import "nx_fmt.nx"
17
18const HID: i64 = 896
19const FFN: i64 = 4864
20const KVD: i64 = 128
21const Q8B: i64 = 34
22const Q8V: i64 = 32
23const NLAY: i64 = 24
24
25func ab_nl() -> i64 { fmt_puts("\n" as *u8); return 0 }
26func ab_lcg(s: i64) -> i64 { var v: i64 = s * 1103515245 + 12345; v = v & 2147483647; return v }
27func ab_weight(k: i64, n: i64, seed: i64) -> *NxF32LazyWeight {
28 let bpr: i64 = (k / Q8V) * Q8B
29 let w: *u8 = sys_mmap(n * bpr)
30 var s: i64 = seed
31 var r: i64 = 0
32 while r < n {
33 var b: i64 = 0
34 while b < k / Q8V {
35 let off: i64 = r * bpr + b * Q8B
36 w[off + 1] = 0x2C as u8
37 var q: i64 = 0
38 while q < 32 { s = ab_lcg(s); w[off + 2 + q] = ((s % 17) - 8) as u8; q = q + 1 }
39 b = b + 1
40 }
41 r = r + 1
42 }
43 return nx_f32_lazy_weight_new_q8_0(w, 0, n, k)
44}
45
46func main() -> i64 {
47 // A vectors: hidden-dim input, and ffn-dim input (for W_down).
48 let Ah: *i64 = sys_mmap(HID * 8) as *i64
49 let Af: *i64 = sys_mmap(FFN * 8) as *i64
50 var s: i64 = 4242
51 var i: i64 = 0
52 while i < HID { s = ab_lcg(s); Ah[i] = nx_i32_to_f32((s % 9) - 4); i = i + 1 }
53 i = 0
54 while i < FFN { s = ab_lcg(s); Af[i] = nx_i32_to_f32((s % 9) - 4); i = i + 1 }
55
56 // one layer's 7 weights (reused across NLAY -- same warmth for A and B).
57 let Wq: *NxF32LazyWeight = ab_weight(HID, HID, 11) // 896x896
58 let Wk: *NxF32LazyWeight = ab_weight(HID, KVD, 12) // 896x128
59 let Wv: *NxF32LazyWeight = ab_weight(HID, KVD, 13)
60 let Wo: *NxF32LazyWeight = ab_weight(HID, HID, 14)
61 let Wg: *NxF32LazyWeight = ab_weight(HID, FFN, 15) // 896x4864
62 let Wu: *NxF32LazyWeight = ab_weight(HID, FFN, 16)
63 let Wd: *NxF32LazyWeight = ab_weight(FFN, HID, 17) // 4864x896
64
65 let Ch: *i64 = sys_mmap(HID * 8) as *i64
66 let Ck: *i64 = sys_mmap(KVD * 8) as *i64
67 let Cf: *i64 = sys_mmap(FFN * 8) as *i64
68
69 nx_lw_shared_pool() // warm pool
70
71 // ---- A: force ALL through the pool ----
72 let ta: i64 = sys_now_us()
73 var la: i64 = 0
74 while la < NLAY {
75 _lw_q8_0_matmul_pool_force(Wq, Ah, Ch, 1, HID, HID)
76 _lw_q8_0_matmul_pool_force(Wk, Ah, Ck, 1, HID, KVD)
77 _lw_q8_0_matmul_pool_force(Wv, Ah, Ck, 1, HID, KVD)
78 _lw_q8_0_matmul_pool_force(Wo, Ah, Ch, 1, HID, HID)
79 _lw_q8_0_matmul_pool_force(Wg, Ah, Cf, 1, HID, FFN)
80 _lw_q8_0_matmul_pool_force(Wu, Ah, Cf, 1, HID, FFN)
81 _lw_q8_0_matmul_pool_force(Wd, Af, Ch, 1, FFN, HID)
82 la = la + 1
83 }
84 let usa: i64 = sys_now_us() - ta
85
86 // ---- B: threshold-routed (the shipped nx_f32_lazy_matmul path) ----
87 let tb: i64 = sys_now_us()
88 var lb: i64 = 0
89 while lb < NLAY {
90 _lw_q8_0_matmul(Wq, Ah, Ch, 1, HID, HID)
91 _lw_q8_0_matmul(Wk, Ah, Ck, 1, HID, KVD)
92 _lw_q8_0_matmul(Wv, Ah, Ck, 1, HID, KVD)
93 _lw_q8_0_matmul(Wo, Ah, Ch, 1, HID, HID)
94 _lw_q8_0_matmul(Wg, Ah, Cf, 1, HID, FFN)
95 _lw_q8_0_matmul(Wu, Ah, Cf, 1, HID, FFN)
96 _lw_q8_0_matmul(Wd, Af, Ch, 1, FFN, HID)
97 lb = lb + 1
98 }
99 let usb: i64 = sys_now_us() - tb
100
101 // ---- C: force ALL through single-thread ----
102 let tc: i64 = sys_now_us()
103 var lc: i64 = 0
104 while lc < NLAY {
105 _lw_q8_0_matmul_st(Wq, Ah, Ch, 1, HID, HID)
106 _lw_q8_0_matmul_st(Wk, Ah, Ck, 1, HID, KVD)
107 _lw_q8_0_matmul_st(Wv, Ah, Ck, 1, HID, KVD)
108 _lw_q8_0_matmul_st(Wo, Ah, Ch, 1, HID, HID)
109 _lw_q8_0_matmul_st(Wg, Ah, Cf, 1, HID, FFN)
110 _lw_q8_0_matmul_st(Wu, Ah, Cf, 1, HID, FFN)
111 _lw_q8_0_matmul_st(Wd, Af, Ch, 1, FFN, HID)
112 lc = lc + 1
113 }
114 let usc: i64 = sys_now_us() - tc
115
116 fmt_puts("A(all-pool) us="); fmt_putn(usa); ab_nl()
117 fmt_puts("B(threshold) us="); fmt_putn(usb); ab_nl()
118 fmt_puts("C(all-ST) us="); fmt_putn(usc); ab_nl()
119 var ub: i64 = usb
120 if ub < 1 { ub = 1 }
121 var uc: i64 = usc
122 if uc < 1 { uc = 1 }
123 fmt_puts("A_vs_B_x100="); fmt_putn(usa * 100 / ub); ab_nl()
124 fmt_puts("A_vs_C_x100="); fmt_putn(usa * 100 / uc); ab_nl()
125 fmt_puts("B_vs_C_x100="); fmt_putn(usb * 100 / uc); ab_nl()
126 fmt_puts("Q8_MIX_AB DONE"); ab_nl()
127 return 0
128}