nx_simd_dot_bench.nx source
↩ module page · 125 lines · 4247 B
1// nx_simd_dot_bench.nx -- x86 integer-SIMD dot (vpmaddwd via __i16x16_madd) vs scalar: verify + speedup.
2//
3// Unblocks the perf-exceed: __i16x16_madd(acc_i32x8, a_i16x16, b_i16x16) does 16 int16 mul-adds per call
4// via one vpmaddwd. A 4096-wide dot = 256 madd calls + one horizontal sum. Verified bit-exact vs scalar +
5// timed. This is the reusable SIMD integer-dot primitive for the Q4_K GEMM (q*col per sub-block).
6// license_tier: ORIGINAL
7import "nx_syscalls.nx"
8import "nx_tier.nx"
9import "nx_strconv.nx"
10import "nx_clock.nx"
11
12func db_pack4(a: i64, b: i64, c: i64, d: i64) -> i64 {
13 return (a & 0xFFFF) | ((b & 0xFFFF) << 16) | ((c & 0xFFFF) << 32) | ((d & 0xFFFF) << 48)
14}
15
16// horizontal-sum the 8 signed i32 lanes packed in acc[0..3]
17func db_hsum(acc: *i64) -> i64 {
18 var sum: i64 = 0
19 var i: i64 = 0
20 while i < 4 {
21 let v: i64 = acc[i]
22 var lo: i64 = v & 0xFFFFFFFF
23 if lo >= 0x80000000 { lo = lo - 0x100000000 }
24 var hi: i64 = (v >> 32) & 0xFFFFFFFF
25 if hi >= 0x80000000 { hi = hi - 0x100000000 }
26 sum = sum + lo + hi
27 i = i + 1
28 }
29 return sum
30}
31
32func db_emit(fd: i64, key: *u8, kl: i64, v: i64) -> i64 {
33 let line: *u8 = sys_mmap(64)
34 var lo: i64 = 0
35 var i: i64 = 0
36 while i < kl { line[lo] = key[i]; lo = lo + 1; i = i + 1 }
37 line[lo] = 0x3D; lo = lo + 1
38 let dec: *u8 = sys_mmap(32)
39 let nd: i64 = nx_strconv_format_i64(v, dec)
40 var k: i64 = 0
41 while k < nd { line[lo] = dec[k]; lo = lo + 1; k = k + 1 }
42 line[lo] = 0x0A; lo = lo + 1
43 return sys_write(fd, line, lo)
44}
45
46func main() -> i64 {
47 let N: i64 = 4096
48 let va: *i64 = sys_mmap(N * 8) as *i64 // plain i64-per-value (scalar ref)
49 let vb: *i64 = sys_mmap(N * 8) as *i64
50 let ap: *i64 = sys_mmap(N * 2) as *i64 // packed 4 i16 per i64
51 let bp: *i64 = sys_mmap(N * 2) as *i64
52 var i: i64 = 0
53 while i < N {
54 va[i] = (i - (i / 31) * 31) + 1 // 1..31
55 vb[i] = (i - (i / 17) * 17) + 1 // 1..17
56 i = i + 1
57 }
58 var j: i64 = 0
59 while j < N / 4 {
60 ap[j] = db_pack4(va[j * 4], va[j * 4 + 1], va[j * 4 + 2], va[j * 4 + 3])
61 bp[j] = db_pack4(vb[j * 4], vb[j * 4 + 1], vb[j * 4 + 2], vb[j * 4 + 3])
62 j = j + 1
63 }
64
65 // scalar dot
66 var scalar: i64 = 0
67 i = 0
68 while i < N { scalar = scalar + va[i] * vb[i]; i = i + 1 }
69
70 // SIMD dot: acc(8 i32)=0; madd over N/16 chunks (each 16 i16 = 4 i64); hsum.
71 let acc: *i64 = sys_mmap(32) as *i64
72 acc[0] = 0; acc[1] = 0; acc[2] = 0; acc[3] = 0
73 let chunks: i64 = N / 16
74 var c: i64 = 0
75 while c < chunks {
76 __i16x16_madd(acc as *i64, ((ap as i64) + c * 32) as *i64, ((bp as i64) + c * 32) as *i64)
77 c = c + 1
78 }
79 let simd: i64 = db_hsum(acc)
80 if simd != scalar { return 80 } // must be bit-exact
81
82 // ---- timing ----
83 let IT: i64 = 20000
84 let t0: i64 = nx_clock_monotonic_ns()
85 var s_acc: i64 = 0
86 var it: i64 = 0
87 while it < IT {
88 var d: i64 = 0
89 i = 0
90 while i < N { d = d + va[i] * vb[i]; i = i + 1 }
91 s_acc = s_acc + d
92 it = it + 1
93 }
94 let t1: i64 = nx_clock_monotonic_ns()
95 let scalar_ns: i64 = t1 - t0
96
97 let t2: i64 = nx_clock_monotonic_ns()
98 var m_acc: i64 = 0
99 it = 0
100 while it < IT {
101 acc[0] = 0; acc[1] = 0; acc[2] = 0; acc[3] = 0
102 c = 0
103 while c < chunks {
104 __i16x16_madd(acc as *i64, ((ap as i64) + c * 32) as *i64, ((bp as i64) + c * 32) as *i64)
105 c = c + 1
106 }
107 m_acc = m_acc + db_hsum(acc)
108 it = it + 1
109 }
110 let t3: i64 = nx_clock_monotonic_ns()
111 let simd_ns: i64 = t3 - t2
112
113 let ofd: i64 = sys_openat_wr("/tmp/simd_dot.txt" as *u8, 0x1a4)
114 if ofd >= 0 {
115 db_emit(ofd, "N" as *u8, 1, N)
116 db_emit(ofd, "dot_value" as *u8, 9, simd)
117 db_emit(ofd, "scalar_ns" as *u8, 9, scalar_ns)
118 db_emit(ofd, "simd_ns" as *u8, 7, simd_ns)
119 if simd_ns > 0 { db_emit(ofd, "speedup_x100" as *u8, 12, scalar_ns * 100 / simd_ns) }
120 db_emit(ofd, "sink" as *u8, 4, (s_acc & 1) + (m_acc & 1))
121 sys_close(ofd)
122 }
123 if simd_ns >= scalar_ns { return 81 }
124 return 0
125}