code wiki / (root) / nx_simd_dot_bench.nx

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}