code wiki / (root) / nx_simd_madd_probe.nx

nx_simd_madd_probe.nx source

↩ module page · 65 lines · 2170 B

1// nx_simd_madd_probe.nx -- test the x86 integer-SIMD primitive __i16x16_madd (vpmaddwd) on nx_cc_known_good. 2// 3// __simd_vdot_i16_x16 is RISC-V-only codegen (nx_parse.nx:1436 says nx_riscv.nx). But __i16x16_madd 4// (nx_parse.nx:1497) is documented as x86 vpmaddwd: *acc(i32x8) += a(i16x16) . b(i16x16) pairwise. If THIS 5// works on x86, the integer GEMM can be SIMD-accelerated here -> the whole perf-exceed unblocks. 6// Golden: a=b=[1..16], vpmaddwd pairs -> 8 i32, sum = sum(i^2, 1..16) = 1496. 7// license_tier: ORIGINAL 8import "nx_syscalls.nx" 9import "nx_tier.nx" 10const K_MAGIC_1496: i64 = 1496 11const K_MAGIC_2992: i64 = 2992 12 13func sp3_pack(a: i64, b: i64, c: i64, d: i64) -> i64 { 14 return (a & 0xFFFF) | ((b & 0xFFFF) << 16) | ((c & 0xFFFF) << 32) | ((d & 0xFFFF) << 48) 15} 16 17func main() -> i64 { 18 let a: *i64 = sys_mmap(64) as *i64 // 16 i16 = 4 i64 19 let b: *i64 = sys_mmap(64) as *i64 20 let acc: *i64 = sys_mmap(32) as *i64 // 8 i32 = 4 i64 21 a[0] = sp3_pack(1, 2, 3, 4) 22 a[1] = sp3_pack(5, 6, 7, 8) 23 a[2] = sp3_pack(9, 10, 11, 12) 24 a[3] = sp3_pack(13, 14, 15, 16) 25 b[0] = a[0] 26 b[1] = a[1] 27 b[2] = a[2] 28 b[3] = a[3] 29 acc[0] = 0 30 acc[1] = 0 31 acc[2] = 0 32 acc[3] = 0 33 34 __i16x16_madd(acc as *i64, a as *i64, b as *i64) 35 36 // sum the 8 signed i32 lanes in acc 37 var sum: i64 = 0 38 var i: i64 = 0 39 while i < 4 { 40 let v: i64 = acc[i] 41 var lo: i64 = v & 0xFFFFFFFF 42 if lo >= 0x80000000 { lo = lo - 0x100000000 } 43 var hi: i64 = (v >> 32) & 0xFFFFFFFF 44 if hi >= 0x80000000 { hi = hi - 0x100000000 } 45 sum = sum + lo + hi 46 i = i + 1 47 } 48 if sum != K_MAGIC_1496 { return 1 } 49 50 // second call must ACCUMULATE (acc += again) -> 2992 51 __i16x16_madd(acc as *i64, a as *i64, b as *i64) 52 sum = 0 53 i = 0 54 while i < 4 { 55 let v: i64 = acc[i] 56 var lo: i64 = v & 0xFFFFFFFF 57 if lo >= 0x80000000 { lo = lo - 0x100000000 } 58 var hi: i64 = (v >> 32) & 0xFFFFFFFF 59 if hi >= 0x80000000 { hi = hi - 0x100000000 } 60 sum = sum + lo + hi 61 i = i + 1 62 } 63 if sum != K_MAGIC_2992 { return 2 } 64 return 0 65}