code wiki / (root) / nx_simd_usat_shift_test.nx

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}