code wiki / (root) / nx_simd_vdot_test.nx

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}