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}