nx_simd_vdot_test.nx source
↩ module page · 149 lines · 5476 B
1// nx_simd_vdot_test.nx -- widening dot product i16x16 -> i64.
2//
3// The kernel of every ML inference inner loop, every convolution
4// stride, every audio/DSP filter on x86. vpmaddwd does (a*b)
5// pair-wise into 8 i32 lanes in ONE instruction; RV-V does it via
6// vwmul.vv + vwredsum.vs (widening at the codegen layer).
7//
8// We test the bit-exact behaviour vs a hand-written scalar reduce
9// because the widening prevents intermediate overflow that the
10// naive vmul+vreduce_sum would suffer with values > 181 (since
11// 181*181 = 32761 ~ INT16_MAX).
12
13import "nx_kernel_v2.nx"
14import "nx_log.nx"
15
16// Pack 4 i16 values (low 16 bits each) into one i64.
17func pack4_i16(a: i64, b: i64, c: i64, d: i64) -> i64 {
18 return (a & 0xFFFF)
19 | ((b & 0xFFFF) << 16)
20 | ((c & 0xFFFF) << 32)
21 | ((d & 0xFFFF) << 48)
22}
23
24// Sign-extend a 16-bit value from low bits of an i64.
25func sx_i16(x: i64) -> i64 {
26 let m: i64 = x & 0xFFFF
27 if m >= 0x8000 { return m - 0x10000 }
28 return m
29}
30
31// Scalar reference dot product: sum_i a[i] * b[i] over 16 i16 lanes.
32func ref_dot(a: *i64, b: *i64) -> i64 {
33 var acc: i64 = 0
34 var w: i64 = 0
35 while w < 4 {
36 let av: i64 = a[w]
37 let bv: i64 = b[w]
38 acc = acc + sx_i16(av) * sx_i16(bv)
39 acc = acc + sx_i16(av >> 16) * sx_i16(bv >> 16)
40 acc = acc + sx_i16(av >> 32) * sx_i16(bv >> 32)
41 acc = acc + sx_i16(av >> 48) * sx_i16(bv >> 48)
42 w = w + 1
43 }
44 return acc
45}
46
47func t1_identity_dot() -> i64 {
48 // dot([1,2,...,16], [1,2,...,16]) = sum(i^2) for i=1..16 = 1496
49 let a_raw: *u8 = sys_mmap(64)
50 let b_raw: *u8 = sys_mmap(64)
51 let a: *i64 = a_raw as *i64
52 let b: *i64 = b_raw as *i64
53 a[0] = pack4_i16(1, 2, 3, 4)
54 a[1] = pack4_i16(5, 6, 7, 8)
55 a[2] = pack4_i16(9, 10, 11, 12)
56 a[3] = pack4_i16(13, 14, 15, 16)
57 b[0] = pack4_i16(1, 2, 3, 4)
58 b[1] = pack4_i16(5, 6, 7, 8)
59 b[2] = pack4_i16(9, 10, 11, 12)
60 b[3] = pack4_i16(13, 14, 15, 16)
61 let va: i64 = __simd_vload_i16_x16(a as *i64)
62 let vb: i64 = __simd_vload_i16_x16(b as *i64)
63 let simd_sum: i64 = __simd_vdot_i16_x16(va, vb)
64 let ref_sum: i64 = ref_dot(a, b)
65 if simd_sum != 1496 { return 1 }
66 if ref_sum != 1496 { return 2 }
67 if simd_sum != ref_sum { return 3 }
68 return 0
69}
70
71func t2_large_values_no_overflow() -> i64 {
72 // Each lane = 200; dot = 200*200*16 = 640000.
73 // 200*200 = 40000 overflows int16 -- proves widening works.
74 let a_raw: *u8 = sys_mmap(64)
75 let b_raw: *u8 = sys_mmap(64)
76 let a: *i64 = a_raw as *i64
77 let b: *i64 = b_raw as *i64
78 a[0] = pack4_i16(200, 200, 200, 200)
79 a[1] = pack4_i16(200, 200, 200, 200)
80 a[2] = pack4_i16(200, 200, 200, 200)
81 a[3] = pack4_i16(200, 200, 200, 200)
82 b[0] = pack4_i16(200, 200, 200, 200)
83 b[1] = pack4_i16(200, 200, 200, 200)
84 b[2] = pack4_i16(200, 200, 200, 200)
85 b[3] = pack4_i16(200, 200, 200, 200)
86 let va: i64 = __simd_vload_i16_x16(a as *i64)
87 let vb: i64 = __simd_vload_i16_x16(b as *i64)
88 let simd_sum: i64 = __simd_vdot_i16_x16(va, vb)
89 if simd_sum != 640000 { return 10 }
90 return 0
91}
92
93func t3_mixed_sign() -> i64 {
94 // Pattern: a = [+1, -1, +2, -2, +3, -3, +4, -4, +5, -5, +6, -6, +7, -7, +8, -8]
95 // b = [1..16]
96 // dot = 1 - 2 + 6 - 8 + 15 - 18 + 28 - 32 + 45 - 50 + 66 - 72 + 91 - 98 + 120 - 128
97 // = -36
98 let a_raw: *u8 = sys_mmap(64)
99 let b_raw: *u8 = sys_mmap(64)
100 let a: *i64 = a_raw as *i64
101 let b: *i64 = b_raw as *i64
102 a[0] = pack4_i16(1, -1, 2, -2)
103 a[1] = pack4_i16(3, -3, 4, -4)
104 a[2] = pack4_i16(5, -5, 6, -6)
105 a[3] = pack4_i16(7, -7, 8, -8)
106 b[0] = pack4_i16(1, 2, 3, 4)
107 b[1] = pack4_i16(5, 6, 7, 8)
108 b[2] = pack4_i16(9, 10, 11, 12)
109 b[3] = pack4_i16(13, 14, 15, 16)
110 let va: i64 = __simd_vload_i16_x16(a as *i64)
111 let vb: i64 = __simd_vload_i16_x16(b as *i64)
112 let simd_sum: i64 = __simd_vdot_i16_x16(va, vb)
113 let ref_sum: i64 = ref_dot(a, b)
114 if simd_sum != ref_sum { return 20 }
115 if simd_sum != -36 { return 21 }
116 return 0
117}
118
119func t4_zero_vector() -> i64 {
120 let a_raw: *u8 = sys_mmap(64)
121 let b_raw: *u8 = sys_mmap(64)
122 let a: *i64 = a_raw as *i64
123 let b: *i64 = b_raw as *i64
124 a[0] = 0; a[1] = 0; a[2] = 0; a[3] = 0
125 b[0] = pack4_i16(100, 200, 300, 400)
126 b[1] = pack4_i16(500, 600, 700, 800)
127 b[2] = pack4_i16(900, 1000, 1100, 1200)
128 b[3] = pack4_i16(1300, 1400, 1500, 1600)
129 let va: i64 = __simd_vload_i16_x16(a as *i64)
130 let vb: i64 = __simd_vload_i16_x16(b as *i64)
131 let simd_sum: i64 = __simd_vdot_i16_x16(va, vb)
132 if simd_sum != 0 { return 30 }
133 return 0
134}
135
136func main() -> nx_exit {
137 println("=== nx_simd_vdot smoke (i16x16 widening dot product) ===" as *u8)
138 if t1_identity_dot() != 0 { println("FAIL T1" as *u8); return 1 }
139 println("T1 identity PASS sum(1..16)^2 = 1496" as *u8)
140 if t2_large_values_no_overflow() != 0 { println("FAIL T2" as *u8); return 2 }
141 println("T2 large vals PASS 16 * 200*200 = 640000 (no i16 overflow)" as *u8)
142 if t3_mixed_sign() != 0 { println("FAIL T3" as *u8); return 3 }
143 println("T3 mixed sign PASS alternating signs = -36" as *u8)
144 if t4_zero_vector() != 0 { println("FAIL T4" as *u8); return 4 }
145 println("T4 zero PASS zero . anything = 0" as *u8)
146 println("" as *u8)
147 println("Widening dot product OPTIMAL on RV-V (vwmul.vv + vwredsum.vs) and x86 (vpmaddwd)." as *u8)
148 return 0
149}