nx_simd_usat_shift_test.nx source
↩ module page · 159 lines · 6510 B
1// nx_simd_usat_shift_test.nx -- unsigned saturating arith + per-lane shifts (i16x16).
2//
3// Two primitive families exercised in one smoke:
4//
5// vsaddu / vssubu -- unsigned saturating add/sub. Required for
6// RGBA channel data (signed sat would clip at 32767 instead of
7// 65535, mangling bright channels). Maps to vpaddusw/vpsubusw
8// on x86 and vsaddu.vv/vssubu.vv on RV-V.
9//
10// vsll / vsrl / vsra -- per-lane immediate shifts. Image
11// processing channel shuffles (pack ARGB into 5-6-5), crypto
12// round shifts, fast pow-of-2 multiply/divide. Maps to
13// vpsllw/vpsrlw/vpsraw on x86 and vsll.vx/vsrl.vx/vsra.vx on RV-V.
14
15import "nx_kernel_v2.nx"
16import "nx_log.nx"
17
18func pack4_i16(a: i64, b: i64, c: i64, d: i64) -> i64 {
19 return (a & 0xFFFF)
20 | ((b & 0xFFFF) << 16)
21 | ((c & 0xFFFF) << 32)
22 | ((d & 0xFFFF) << 48)
23}
24
25func t1_vsaddu_clamp_high() -> i64 {
26 // a = [50000, 50000, ...] (u16; signed view = -15536)
27 // b = [20000, 20000, ...]
28 // u16 sat: 50000+20000 = 70000 -> 65535
29 // signed sat would do something completely wrong.
30 let a_raw: *u8 = sys_mmap(64)
31 let b_raw: *u8 = sys_mmap(64)
32 let r_raw: *u8 = sys_mmap(64)
33 let a: *i64 = a_raw as *i64
34 let b: *i64 = b_raw as *i64
35 let r: *i64 = r_raw as *i64
36 a[0] = pack4_i16(50000, 50000, 50000, 50000)
37 a[1] = pack4_i16(50000, 50000, 50000, 50000)
38 a[2] = pack4_i16(50000, 50000, 50000, 50000)
39 a[3] = pack4_i16(50000, 50000, 50000, 50000)
40 b[0] = pack4_i16(20000, 20000, 20000, 20000)
41 b[1] = pack4_i16(20000, 20000, 20000, 20000)
42 b[2] = pack4_i16(20000, 20000, 20000, 20000)
43 b[3] = pack4_i16(20000, 20000, 20000, 20000)
44 let va: i64 = __simd_vload_i16_x16(a as *i64)
45 let vb: i64 = __simd_vload_i16_x16(b as *i64)
46 let vr: i64 = __simd_vsaddu_i16_x16(va, vb)
47 __simd_vstore_i16_x16(vr, r as *i64)
48 // 65535 = 0xFFFF -- packed bits.
49 let exp: i64 = pack4_i16(65535, 65535, 65535, 65535)
50 if r[0] != exp { return 1 }
51 if r[1] != exp { return 2 }
52 return 0
53}
54
55func t2_vssubu_clamp_low() -> i64 {
56 // a = [100, 200, 300, 400], b = [1000, 1000, 1000, 1000]
57 // u16 sat: a-b underflows below 0 -> 0
58 let a_raw: *u8 = sys_mmap(64)
59 let b_raw: *u8 = sys_mmap(64)
60 let r_raw: *u8 = sys_mmap(64)
61 let a: *i64 = a_raw as *i64
62 let b: *i64 = b_raw as *i64
63 let r: *i64 = r_raw as *i64
64 a[0] = pack4_i16(100, 200, 300, 400)
65 a[1] = pack4_i16(500, 600, 700, 800)
66 a[2] = pack4_i16(900, 1000, 1100, 1200)
67 a[3] = pack4_i16(1300, 1400, 1500, 1600)
68 b[0] = pack4_i16(1000, 1000, 1000, 1000)
69 b[1] = pack4_i16(1000, 1000, 1000, 1000)
70 b[2] = pack4_i16(1000, 1000, 1000, 1000)
71 b[3] = pack4_i16(1000, 1000, 1000, 1000)
72 let va: i64 = __simd_vload_i16_x16(a as *i64)
73 let vb: i64 = __simd_vload_i16_x16(b as *i64)
74 let vr: i64 = __simd_vssubu_i16_x16(va, vb)
75 __simd_vstore_i16_x16(vr, r as *i64)
76 // r[0]: max(0,100-1000)=0, max(0,200-1000)=0, 0, 0
77 if r[0] != 0 { return 10 }
78 // r[1]: max(0,500-1000)=0, 0, 0, 0
79 if r[1] != 0 { return 11 }
80 // r[2]: max(0,900-1000)=0, 1000-1000=0, 1100-1000=100, 1200-1000=200
81 if r[2] != pack4_i16(0, 0, 100, 200) { return 12 }
82 // r[3]: 1300-1000=300, 400, 500, 600
83 if r[3] != pack4_i16(300, 400, 500, 600) { return 13 }
84 return 0
85}
86
87func t3_vsll_shift_left() -> i64 {
88 // Shift each lane left by 3 = multiply by 8.
89 let a_raw: *u8 = sys_mmap(64)
90 let r_raw: *u8 = sys_mmap(64)
91 let a: *i64 = a_raw as *i64
92 let r: *i64 = r_raw as *i64
93 a[0] = pack4_i16(1, 2, 3, 4)
94 a[1] = pack4_i16(5, 6, 7, 8)
95 a[2] = pack4_i16(9, 10, 11, 12)
96 a[3] = pack4_i16(13, 14, 15, 16)
97 let v: i64 = __simd_vload_i16_x16(a as *i64)
98 let vs: i64 = __simd_vsll_i16_x16(v, 3)
99 __simd_vstore_i16_x16(vs, r as *i64)
100 if r[0] != pack4_i16(8, 16, 24, 32) { return 20 }
101 if r[1] != pack4_i16(40, 48, 56, 64) { return 21 }
102 if r[2] != pack4_i16(72, 80, 88, 96) { return 22 }
103 if r[3] != pack4_i16(104, 112, 120, 128) { return 23 }
104 return 0
105}
106
107func t4_vsrl_shift_right_logical() -> i64 {
108 // Logical right shift -- top bits zero-filled even on negatives.
109 let a_raw: *u8 = sys_mmap(64)
110 let r_raw: *u8 = sys_mmap(64)
111 let a: *i64 = a_raw as *i64
112 let r: *i64 = r_raw as *i64
113 // Lane 0: -2 = 0xFFFE. vsrl by 1 = 0x7FFF = 32767 (logical fill).
114 // (vsra would give 0xFFFF = -1.)
115 a[0] = pack4_i16(-2, 64, 256, 1024)
116 a[1] = pack4_i16(0, 0, 0, 0)
117 a[2] = pack4_i16(0, 0, 0, 0)
118 a[3] = pack4_i16(0, 0, 0, 0)
119 let v: i64 = __simd_vload_i16_x16(a as *i64)
120 let vs: i64 = __simd_vsrl_i16_x16(v, 1)
121 __simd_vstore_i16_x16(vs, r as *i64)
122 if r[0] != pack4_i16(32767, 32, 128, 512) { return 30 }
123 return 0
124}
125
126func t5_vsra_shift_right_arithmetic() -> i64 {
127 // Arithmetic right shift -- sign-extends top bits.
128 let a_raw: *u8 = sys_mmap(64)
129 let r_raw: *u8 = sys_mmap(64)
130 let a: *i64 = a_raw as *i64
131 let r: *i64 = r_raw as *i64
132 a[0] = pack4_i16(-2, 64, -256, 1024)
133 a[1] = pack4_i16(0, 0, 0, 0)
134 a[2] = pack4_i16(0, 0, 0, 0)
135 a[3] = pack4_i16(0, 0, 0, 0)
136 let v: i64 = __simd_vload_i16_x16(a as *i64)
137 let vs: i64 = __simd_vsra_i16_x16(v, 1)
138 __simd_vstore_i16_x16(vs, r as *i64)
139 // -2 >> 1 = -1 (sign-extended). -256 >> 1 = -128.
140 if r[0] != pack4_i16(-1, 32, -128, 512) { return 40 }
141 return 0
142}
143
144func main() -> nx_exit {
145 println("=== nx_simd_usat_shift smoke (u-sat + per-lane shifts) ===" as *u8)
146 if t1_vsaddu_clamp_high() != 0 { println("FAIL T1" as *u8); return 1 }
147 println("T1 vsaddu clamp_h PASS 50000+20000 -> 65535 (u16 max, RGBA correct)" as *u8)
148 if t2_vssubu_clamp_low() != 0 { println("FAIL T2" as *u8); return 2 }
149 println("T2 vssubu clamp_l PASS underflow -> 0 (u16 min, RGBA correct)" as *u8)
150 if t3_vsll_shift_left() != 0 { println("FAIL T3" as *u8); return 3 }
151 println("T3 vsll<<3 PASS 16 lanes shifted (= mult by 8) bit-exact" as *u8)
152 if t4_vsrl_shift_right_logical() != 0 { println("FAIL T4" as *u8); return 4 }
153 println("T4 vsrl>>1 PASS logical fill: -2>>1 = 32767 (top zero)" as *u8)
154 if t5_vsra_shift_right_arithmetic() != 0 { println("FAIL T5" as *u8); return 5 }
155 println("T5 vsra>>1 PASS arith fill: -2>>1 = -1 (sign-extended)" as *u8)
156 println("" as *u8)
157 println("u-sat + shifts OPTIMAL on RV-V (vsaddu/vssubu + vsll/vsrl/vsra.vx) and x86 (vp*usw + vpsllw/vpsrlw/vpsraw)." as *u8)
158 return 0
159}