code wiki / (root) / nx_simd_i32x8_test.nx

nx_simd_i32x8_test.nx source

↩ module page · 128 lines · 4493 B

1// nx_simd_i32x8_test.nx -- exercise the i32x8 SIMD intrinsics. 2// 3// Same scalar-vs-SIMD-bit-exact gate as nx_simd_test.nx but with 4// 8 i32 lanes per vector (double the throughput of i64x4 for 5// 32-bit workloads -- ML token IDs, image pixels, hashes). 6 7import "nx_kernel_v2.nx" 8import "nx_log.nx" 9 10const NX_MO_SEQ_CST: i64 = 5 11 12func t1_round_trip_load_store() -> i64 { 13 let a_raw: *u8 = sys_mmap(64) 14 let r_raw: *u8 = sys_mmap(64) 15 let a: *i64 = a_raw as *i64 16 let r: *i64 = r_raw as *i64 17 // Fill 8 i32 lanes as 4 i64 words (lane order LE: low-i32 first). 18 // Encode (10, 20, 30, 40, 50, 60, 70, 80) -> pack into i64s. 19 a[0] = (20 << 32) | 10 20 a[1] = (40 << 32) | 30 21 a[2] = (60 << 32) | 50 22 a[3] = (80 << 32) | 70 23 let v: i64 = __simd_vload_i32_x8(a as *i64) 24 __simd_vstore_i32_x8(v, r as *i64) 25 if r[0] != ((20 << 32) | 10) { return 1 } 26 if r[1] != ((40 << 32) | 30) { return 2 } 27 if r[2] != ((60 << 32) | 50) { return 3 } 28 if r[3] != ((80 << 32) | 70) { return 4 } 29 return 0 30} 31 32func t2_vadd() -> i64 { 33 let a_raw: *u8 = sys_mmap(64) 34 let b_raw: *u8 = sys_mmap(64) 35 let r_raw: *u8 = sys_mmap(64) 36 let a: *i64 = a_raw as *i64 37 let b: *i64 = b_raw as *i64 38 let r: *i64 = r_raw as *i64 39 a[0] = (20 << 32) | 10 40 a[1] = (40 << 32) | 30 41 a[2] = (60 << 32) | 50 42 a[3] = (80 << 32) | 70 43 b[0] = (2 << 32) | 1 44 b[1] = (4 << 32) | 3 45 b[2] = (6 << 32) | 5 46 b[3] = (8 << 32) | 7 47 let va: i64 = __simd_vload_i32_x8(a as *i64) 48 let vb: i64 = __simd_vload_i32_x8(b as *i64) 49 let vsum: i64 = __simd_vadd_i32_x8(va, vb) 50 __simd_vstore_i32_x8(vsum, r as *i64) 51 // (10+1, 20+2, 30+3, 40+4, 50+5, 60+6, 70+7, 80+8) = 52 // (11, 22, 33, 44, 55, 66, 77, 88) 53 if r[0] != ((22 << 32) | 11) { return 10 } 54 if r[1] != ((44 << 32) | 33) { return 11 } 55 if r[2] != ((66 << 32) | 55) { return 12 } 56 if r[3] != ((88 << 32) | 77) { return 13 } 57 return 0 58} 59 60func t3_vmul() -> i64 { 61 let a_raw: *u8 = sys_mmap(64) 62 let b_raw: *u8 = sys_mmap(64) 63 let r_raw: *u8 = sys_mmap(64) 64 let a: *i64 = a_raw as *i64 65 let b: *i64 = b_raw as *i64 66 let r: *i64 = r_raw as *i64 67 a[0] = (20 << 32) | 10 68 a[1] = (40 << 32) | 30 69 a[2] = (60 << 32) | 50 70 a[3] = (80 << 32) | 70 71 b[0] = (2 << 32) | 1 72 b[1] = (4 << 32) | 3 73 b[2] = (6 << 32) | 5 74 b[3] = (8 << 32) | 7 75 let va: i64 = __simd_vload_i32_x8(a as *i64) 76 let vb: i64 = __simd_vload_i32_x8(b as *i64) 77 let vp: i64 = __simd_vmul_i32_x8(va, vb) 78 __simd_vstore_i32_x8(vp, r as *i64) 79 // (10, 40, 90, 160, 250, 360, 490, 640) 80 if r[0] != ((40 << 32) | 10) { return 20 } 81 if r[1] != ((160 << 32) | 90) { return 21 } 82 if r[2] != ((360 << 32) | 250) { return 22 } 83 if r[3] != ((640 << 32) | 490) { return 23 } 84 return 0 85} 86 87func t4_vbroadcast() -> i64 { 88 let r_raw: *u8 = sys_mmap(64) 89 let r: *i64 = r_raw as *i64 90 let v: i64 = __simd_vbroadcast_i32_x8(7) 91 __simd_vstore_i32_x8(v, r as *i64) 92 if r[0] != ((7 << 32) | 7) { return 30 } 93 if r[1] != ((7 << 32) | 7) { return 31 } 94 if r[2] != ((7 << 32) | 7) { return 32 } 95 if r[3] != ((7 << 32) | 7) { return 33 } 96 return 0 97} 98 99func t5_vreduce_sum() -> i64 { 100 let a_raw: *u8 = sys_mmap(64) 101 let a: *i64 = a_raw as *i64 102 a[0] = (20 << 32) | 10 103 a[1] = (40 << 32) | 30 104 a[2] = (60 << 32) | 50 105 a[3] = (80 << 32) | 70 106 let v: i64 = __simd_vload_i32_x8(a as *i64) 107 let s: i64 = __simd_vreduce_sum_i32_x8(v) 108 // Sum (10+20+30+40+50+60+70+80) = 360 109 if s != 360 { return 40 } 110 return 0 111} 112 113func main() -> nx_exit { 114 println("=== nx_simd_i32x8 smoke (8 lanes of 32-bit) ===" as *u8) 115 if t1_round_trip_load_store() != 0 { println("FAIL T1" as *u8); return 1 } 116 println("T1 vload+vstore PASS 8-lane i32 round-trip" as *u8) 117 if t2_vadd() != 0 { println("FAIL T2" as *u8); return 2 } 118 println("T2 vadd PASS lanes (11..88) bit-exact" as *u8) 119 if t3_vmul() != 0 { println("FAIL T3" as *u8); return 3 } 120 println("T3 vmul PASS lanes (10,40,90,160,250,360,490,640)" as *u8) 121 if t4_vbroadcast() != 0 { println("FAIL T4" as *u8); return 4 } 122 println("T4 vbroadcast PASS all 8 lanes = 7" as *u8) 123 if t5_vreduce_sum() != 0 { println("FAIL T5" as *u8); return 5 } 124 println("T5 vreduce_sum PASS 10+20+...+80 = 360" as *u8) 125 println("" as *u8) 126 println("All i32x8 SIMD ops bit-exact vs scalar; 8-lane throughput proven." as *u8) 127 return 0 128}