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}