code wiki / (root) / nx_f32x8_dot_bench.nx

nx_f32x8_dot_bench.nx source

↩ module page · 100 lines · 3617 B

1// nx_f32x8_dot_bench.nx -- x86 f32 SIMD dot (__f32x8_dot, AVX vmulps) vs scalar: verify + speedup. 2// 3// Complements the integer vpmaddwd win: the f32 side (DiT projections + linear attention) accelerates via 4// __f32x8_dot (8 f32 mul-adds per call). A 4096-wide dot = 512 __f32x8_dot chunk calls. Verified vs scalar 5// f32 + timed. f32 packed as 4-byte f32 in memory (the intrinsic's layout). 6// license_tier: ORIGINAL 7import "nx_syscalls.nx" 8import "nx_tier.nx" 9import "nx_le.nx" 10import "nx_strconv.nx" 11import "nx_f32.nx" 12import "nx_f32_div.nx" 13import "nx_f32_cvt.nx" 14import "nx_clock.nx" 15const K_MAGIC_4096: i64 = 4096 16const K_MAGIC_20000: i64 = 20000 17 18func fb_emit(fd: i64, key: *u8, kl: i64, v: i64) -> i64 { 19 let line: *u8 = sys_mmap(64) 20 var lo: i64 = 0 21 var i: i64 = 0 22 while i < kl { line[lo] = key[i]; lo = lo + 1; i = i + 1 } 23 line[lo] = 0x3D; lo = lo + 1 24 let dec: *u8 = sys_mmap(32) 25 let nd: i64 = nx_strconv_format_i64(v, dec) 26 var k: i64 = 0 27 while k < nd { line[lo] = dec[k]; lo = lo + 1; k = k + 1 } 28 line[lo] = 0x0A; lo = lo + 1 29 return sys_write(fd, line, lo) 30} 31 32func main() -> i64 { 33 let N: i64 = K_MAGIC_4096 34 let a: *u8 = sys_mmap(N * 4) // packed 4-byte f32 35 let b: *u8 = sys_mmap(N * 4) 36 let seven: i64 = nx_i32_to_f32(7) 37 let five: i64 = nx_i32_to_f32(5) 38 var i: i64 = 0 39 while i < N { 40 nx_le_write_u32(a, i * 4, nx_f32_div(nx_i32_to_f32((i - (i / 7) * 7) + 1), seven)) 41 nx_le_write_u32(b, i * 4, nx_f32_div(nx_i32_to_f32((i - (i / 5) * 5) + 1), five)) 42 i = i + 1 43 } 44 45 // scalar dot (read packed f32) 46 var scalar: i64 = 0 47 i = 0 48 while i < N { scalar = nx_f32_add(scalar, nx_f32_mul(nx_le_read_u32(a, i * 4), nx_le_read_u32(b, i * 4))); i = i + 1 } 49 50 // SIMD dot: sum of per-chunk __f32x8_dot (8 f32 each) 51 var simd: i64 = 0 52 var c: i64 = 0 53 while c < N / 8 { 54 simd = nx_f32_add(simd, __f32x8_dot(((a as i64) + c * 32) as *i64, ((b as i64) + c * 32) as *i64)) 55 c = c + 1 56 } 57 // relative check (f32 accumulation order differs): |simd-scalar| < 1% |scalar| 58 let tol: i64 = nx_f32_mul(nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(100)), scalar & 0x7FFFFFFF) 59 if (nx_f32_sub(simd, scalar) & 0x7FFFFFFF) >= tol { return 80 } 60 61 let IT: i64 = K_MAGIC_20000 62 let t0: i64 = nx_clock_monotonic_ns() 63 var s1: i64 = 0 64 var it: i64 = 0 65 while it < IT { 66 var d: i64 = 0 67 i = 0 68 while i < N { d = nx_f32_add(d, nx_f32_mul(nx_le_read_u32(a, i * 4), nx_le_read_u32(b, i * 4))); i = i + 1 } 69 s1 = nx_f32_add(s1, d) 70 it = it + 1 71 } 72 let t1: i64 = nx_clock_monotonic_ns() 73 let scalar_ns: i64 = t1 - t0 74 75 let t2: i64 = nx_clock_monotonic_ns() 76 var s2: i64 = 0 77 it = 0 78 while it < IT { 79 var d: i64 = 0 80 c = 0 81 while c < N / 8 { d = nx_f32_add(d, __f32x8_dot(((a as i64) + c * 32) as *i64, ((b as i64) + c * 32) as *i64)); c = c + 1 } 82 s2 = nx_f32_add(s2, d) 83 it = it + 1 84 } 85 let t3: i64 = nx_clock_monotonic_ns() 86 let simd_ns: i64 = t3 - t2 87 88 let ofd: i64 = sys_openat_wr("/tmp/f32x8_dot.txt" as *u8, 0x1a4) 89 if ofd >= 0 { 90 fb_emit(ofd, "N" as *u8, 1, N) 91 fb_emit(ofd, "scalar_bits" as *u8, 11, scalar) 92 fb_emit(ofd, "simd_bits" as *u8, 9, simd) 93 fb_emit(ofd, "scalar_ns" as *u8, 9, scalar_ns) 94 fb_emit(ofd, "simd_ns" as *u8, 7, simd_ns) 95 if simd_ns > 0 { fb_emit(ofd, "speedup_x100" as *u8, 12, scalar_ns * 100 / simd_ns) } 96 sys_close(ofd) 97 } 98 if simd_ns >= scalar_ns { return 81 } 99 return 0 100}