code wiki / (root) / nx_f32x8_range_gate.nx

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}