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}