nx_f32_ffn_path_gate.nx source
↩ module page · 118 lines · 4232 B
1// nx_f32_ffn_path_gate.nx -- ISOLATED, host-noise-immune proof that the
2// F32 packed range-dot cached path beats the scalar threaded matmul_t
3// at the real FFN decode shape (m=1, k=896, n=4864 = Qwen W_gate).
4//
5// Both paths run in the SAME loaded conditions, so the RATIO holds even
6// when the host is noisy (the trap that contaminated the full-forward
7// profiler). Warms the cache first so the timed cached calls are pure
8// SIMD dots (the fill is one-time, amortized across all tokens in the
9// real forward).
10//
11// Exact-int regime (dense small-int f32) so bit-exact compare across the
12// two different accumulation orders (scalar sequential vs 8-wide) is
13// legitimate.
14//
15// Checks:
16// 1 cached (range-dot) result == scalar matmul_t, bit-exact
17// 2 cached path >= floor x faster than scalar over N reps (both timed
18// threaded; ratio printed)
19//
20// lineage_id: f32_ffn_path_gate_v1
21
22import "nx_f32_lazy_weight.nx"
23import "nx_f32_matmul_t.nx"
24import "nx_f32_cvt.nx"
25import "nx_fmt.nx"
26
27const FK: i64 = 896 // hidden
28const FN: i64 = 4864 // ffn (W_gate n)
29const FREPS: i64 = 200
30const FFLOOR_X100: i64 = 150
31
32func fg_lcg(s: i64) -> i64 {
33 var v: i64 = s * 1103515245 + 12345
34 v = v & 2147483647
35 return v
36}
37// exact small-int f32 in i64 slots
38func fg_fill(p: *i64, count: i64, seed: i64, half: i64) -> i64 {
39 var s: i64 = seed
40 var i: i64 = 0
41 while i < count {
42 s = fg_lcg(s)
43 p[i] = nx_i32_to_f32((s % (half + half)) - half)
44 i = i + 1
45 }
46 return 0
47}
48func fg_same(a: *i64, b: *i64, count: i64) -> i64 {
49 var i: i64 = 0
50 while i < count { if a[i] != b[i] { return 0 } i = i + 1 }
51 return 1
52}
53func fg_poison(p: *i64, count: i64) -> i64 {
54 var i: i64 = 0
55 while i < count { p[i] = 0 - 777777; i = i + 1 }
56 return 0
57}
58func fg_nl() -> i64 { fmt_puts("\n" as *u8); return 0 }
59
60func main() -> i64 {
61 let A: *i64 = sys_mmap(FK * 8) as *i64
62 let W: *i64 = sys_mmap(FN * FK * 8) as *i64 // ggml layout: row j = k contiguous
63 let Cref: *i64 = sys_mmap(FN * 8) as *i64
64 let Cd: *i64 = sys_mmap(FN * 8) as *i64
65 fg_fill(A, FK, 20260708, 4)
66 fg_fill(W, FN * FK, 31337, 4)
67
68 let pool: *NxThreadPool = nx_lw_shared_pool()
69
70 var pass: i64 = 0
71
72 // scalar oracle (threaded matmul_t) -- W is ggml col-major B[kk + j*k]
73 nx_f32_matmul_t_pool(pool, A, W, Cref, 1, FK, FN)
74
75 // cached range-dot path via the dispatcher (first call fills the cache)
76 let Wl: *NxF32LazyWeight = nx_f32_lazy_weight_new_f32(W, FK, FN)
77 fg_poison(Cd, FN)
78 let v1: nx_int = nx_f32_lazy_matmul(A, Wl, Cd, 1, FK, FN)
79 var ok1: i64 = 0
80 if v1 == NX_LW_OK { if Wl.pk_state == 1 { ok1 = fg_same(Cd, Cref, FN) } }
81 if ok1 != 1 {
82 fmt_puts("FFN 1 EXACT FAIL v="); fmt_putn(v1); fmt_puts(" state="); fmt_putn(Wl.pk_state); fg_nl()
83 return 11
84 }
85 fmt_puts("FFN 1 CACHED==SCALAR EXACT OK (pk_state=1, used="); fmt_putn(nx_lw_cache_used()); fmt_puts(")"); fg_nl()
86 pass = pass + 1
87
88 // timing: scalar matmul_t vs cached range-dot, same shape, N reps
89 let t0: i64 = sys_now_us()
90 var r0: i64 = 0
91 while r0 < FREPS { nx_f32_matmul_t_pool(pool, A, W, Cref, 1, FK, FN); r0 = r0 + 1 }
92 let us_scalar: i64 = sys_now_us() - t0
93
94 let t1: i64 = sys_now_us()
95 var r1: i64 = 0
96 while r1 < FREPS {
97 let vr: nx_int = nx_f32_lazy_matmul(A, Wl, Cd, 1, FK, FN)
98 if vr != NX_LW_OK { return 20 }
99 r1 = r1 + 1
100 }
101 let us_cached: i64 = sys_now_us() - t1
102
103 var us_s: i64 = us_scalar
104 if us_s < 1 { us_s = 1 }
105 var us_c: i64 = us_cached
106 if us_c < 1 { us_c = 1 }
107 let macs: i64 = FK * FN * FREPS
108 fmt_puts("scalar_us="); fmt_putn(us_scalar); fmt_puts(" mflops="); fmt_putn(2 * macs / us_s); fg_nl()
109 fmt_puts("cached_us="); fmt_putn(us_cached); fmt_puts(" mflops="); fmt_putn(2 * macs / us_c); fg_nl()
110 let sx100: i64 = us_s * 100 / us_c
111 fmt_puts("cached_vs_scalar_x100="); fmt_putn(sx100); fg_nl()
112 if sx100 < FFLOOR_X100 { fmt_puts("FFN 2 SPEEDUP FAIL"); fg_nl(); return 12 }
113 fmt_puts("FFN 2 SPEEDUP OK"); fg_nl()
114 pass = pass + 1
115
116 fmt_puts("F32_FFN_PATH_GATE "); fmt_putn(pass); fmt_puts("/2 GREEN"); fg_nl()
117 return 0
118}