nx_f32x8_range_gate.nx source
↩ module page · 127 lines · 4211 B
1// nx_f32x8_range_gate.nx -- validate a VECTOR-ACCUMULATOR range dot
2// (__f32x8_fma across the whole k + ONE __f32x8_hsum) vs the current
3// __f32x4_dot loop (hsum every 4).
4//
5// The matmul cached path (_lw_dot_task) uses __f32x4_dot, which does a
6// horizontal sum per 4 lanes (~2 uops/MAC = no better than scalar --
7// measured 3x this arc). The proper SIMD keeps an 8-wide accumulator
8// in a register/memory across the reduction and hsums ONCE. Both
9// intrinsics already exist; this gate proves (a) they agree bit-exact
10// in the exact-int regime and (b) the range dot is faster.
11//
12// Checks:
13// 1 f32x8 range dot == __f32x4_dot loop, bit-exact (exact ints)
14// 2 range dot >= floor x faster over many reps at k=896
15//
16// lineage_id: f32x8_range_gate_v1
17
18import "nx_syscalls.nx"
19import "nx_f32.nx"
20import "nx_f32_cvt.nx"
21import "nx_fmt.nx"
22
23const RK: i64 = 896 // Qwen hidden = a real reduction dim (div by 8)
24const RREPS: i64 = 40000
25const RFLOOR_X100: i64 = 150
26
27func rg_st4(p: *u8, idx: i64, bits: i64) -> i64 {
28 p[idx * 4 + 0] = bits as u8
29 p[idx * 4 + 1] = (bits >> 8) as u8
30 p[idx * 4 + 2] = (bits >> 16) as u8
31 p[idx * 4 + 3] = (bits >> 24) as u8
32 return 0
33}
34func rg_ld4(p: *u8, off: i64) -> i64 {
35 return (p[off] as i64) | ((p[off+1] as i64) << 8) | ((p[off+2] as i64) << 16) | ((p[off+3] as i64) << 24)
36}
37
38// current cached-dot inner: __f32x4_dot, hsum every 4.
39func dot_x4(a: *u8, b: *u8, count: i64) -> i64 {
40 let ab: i64 = a as i64
41 let bb: i64 = b as i64
42 var s: i64 = 0
43 var l: i64 = 0
44 while l < count {
45 s = __f32_add(s, __f32x4_dot((ab + l * 4) as *u8, (bb + l * 4) as *u8))
46 l = l + 4
47 }
48 return s
49}
50
51// vector-accumulator range dot: __f32x8_fma across k, hsum ONCE.
52// acc = caller-owned 32-byte (8 f32) buffer, zeroed here each call.
53func dot_x8_range(a: *u8, b: *u8, count: i64, acc: *u8) -> i64 {
54 let az: *i64 = acc as *i64
55 az[0] = 0; az[1] = 0; az[2] = 0; az[3] = 0
56 let ab: i64 = a as i64
57 let bb: i64 = b as i64
58 let n8: i64 = (count / 8) * 8
59 var l: i64 = 0
60 while l < n8 {
61 __f32x8_fma(acc, (ab + l * 4) as *u8, (bb + l * 4) as *u8)
62 l = l + 8
63 }
64 var s: i64 = __f32x8_hsum(acc)
65 // tail (count % 8) scalar
66 while l < count {
67 s = __f32_add(s, __f32_mul(rg_ld4(a, l * 4), rg_ld4(b, l * 4)))
68 l = l + 1
69 }
70 return s
71}
72
73func rg_nl() -> i64 { fmt_puts("\n" as *u8); return 0 }
74
75func main() -> i64 {
76 let a: *u8 = sys_mmap(RK * 4)
77 let b: *u8 = sys_mmap(RK * 4)
78 var i: i64 = 0
79 while i < RK {
80 rg_st4(a, i, nx_i32_to_f32((i % 5) - 2))
81 rg_st4(b, i, nx_i32_to_f32((i % 7) - 3))
82 i = i + 1
83 }
84 let acc: *u8 = sys_mmap(32)
85
86 var pass: i64 = 0
87
88 // ---- 1: exactness ----
89 let s4: i64 = dot_x4(a, b, RK)
90 let s8: i64 = dot_x8_range(a, b, RK, acc)
91 if s4 != s8 {
92 fmt_puts("RX8 1 EXACT FAIL x4="); fmt_putn(s4); fmt_puts(" x8="); fmt_putn(s8); rg_nl()
93 return 11
94 }
95 fmt_puts("RX8 1 EXACT OK (bits="); fmt_putn(s4); fmt_puts(")"); rg_nl()
96 pass = pass + 1
97
98 // ---- 2: speed ----
99 let t0: i64 = sys_now_us()
100 var r0: i64 = 0
101 var sink4: i64 = 0
102 while r0 < RREPS { sink4 = dot_x4(a, b, RK); r0 = r0 + 1 }
103 let us4: i64 = sys_now_us() - t0
104
105 let t1: i64 = sys_now_us()
106 var r1: i64 = 0
107 var sink8: i64 = 0
108 while r1 < RREPS { sink8 = dot_x8_range(a, b, RK, acc); r1 = r1 + 1 }
109 let us8: i64 = sys_now_us() - t1
110
111 var u4: i64 = us4
112 if u4 < 1 { u4 = 1 }
113 var u8: i64 = us8
114 if u8 < 1 { u8 = 1 }
115 let macs: i64 = RK * RREPS
116 fmt_puts("x4_us="); fmt_putn(us4); fmt_puts(" mflops="); fmt_putn(2 * macs / u4); rg_nl()
117 fmt_puts("x8_us="); fmt_putn(us8); fmt_puts(" mflops="); fmt_putn(2 * macs / u8); rg_nl()
118 let sx100: i64 = u4 * 100 / u8
119 fmt_puts("x8_speedup_x100="); fmt_putn(sx100); rg_nl()
120 if sink4 != sink8 { fmt_puts("RX8 sink mismatch"); rg_nl(); return 20 }
121 if sx100 < RFLOOR_X100 { fmt_puts("RX8 2 SPEEDUP FAIL"); rg_nl(); return 12 }
122 fmt_puts("RX8 2 SPEEDUP OK"); rg_nl()
123 pass = pass + 1
124
125 fmt_puts("F32X8_RANGE_GATE "); fmt_putn(pass); fmt_puts("/2 GREEN"); rg_nl()
126 return 0
127}