code wiki / (root) / nx_f32_ffn_path_gate.nx

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}