yum-mirror/slang
Making it easier to work with shaders
git clone https://git.yummers.dev/yum-mirror/slang
0476b57fa
master
1//TEST_CATEGORY(wave, compute) 2//TEST:COMPARE_COMPUTE_EX(filecheck-buffer=CHECK):-vk -compute -shaderobj -emit-spirv-directly 3//TEST:COMPARE_COMPUTE_EX(filecheck-buffer=CHECK):-vk -compute -shaderobj -emit-spirv-via-glsl 4//TEST:COMPARE_COMPUTE_EX(filecheck-buffer=CHECK):-cuda -compute -shaderobj -xslang -DCUDA 5 6//TEST:COMPARE_COMPUTE_EX(filecheck-buffer=CHECK):-vk -compute -shaderobj -emit-spirv-directly -xslang -DUSE_GLSL_SYNTAX -allow-glsl 7 8//TEST_INPUT:ubuffer(data=[0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0], stride=4):out,name outputBuffer 9RWStructuredBuffer<uint> outputBuffer; 10 11#if defined(USE_GLSL_SYNTAX) 12#define __partitionedAnd subgroupPartitionedAndNV 13#define __partitionedOr subgroupPartitionedOrNV 14#define __partitionedXor subgroupPartitionedXorNV 15#else 16#define __partitionedAnd WaveMultiBitAnd 17#define __partitionedOr WaveMultiBitOr 18#define __partitionedXor WaveMultiBitXor 19#endif 20 21static uint gAndValue = 0; 22static uint gOrValue = 0; 23static uint gOrResult = 0; 24static uint gXorValue = 0; 25static uint gXorResult = 0; 26 27__generic<T : __BuiltinLogicalType> 28bool test1Bitwise(uint4 mask) 29{ 30 let andValue = T(gAndValue); 31 let orValue = T(gOrValue); 32 let orResult = T(gOrResult); 33 let xorValue = T(gXorValue); 34 let xorResult = T(gXorResult); 35 36 return true 37 & (__partitionedAnd(andValue, mask) == andValue) 38 & (__partitionedOr(orValue, mask) == orResult) 39 & (__partitionedXor(xorValue, mask) == xorResult) 40 ; 41} 42 43__generic<T : __BuiltinLogicalType, let N : int> 44bool testVBitwise(uint4 mask) { 45 typealias GVec = vector<T, N>; 46 47 let andValue = GVec(T(gAndValue)); 48 let orValue = GVec(T(gOrValue)); 49 let orResult = GVec(T(gOrResult)); 50 let xorValue = GVec(T(gXorValue)); 51 let xorResult = GVec(T(gXorResult)); 52 53 return true 54 & all(__partitionedAnd(andValue, mask) == andValue) 55 & all(__partitionedOr(orValue, mask) == orResult) 56 & all(__partitionedXor(xorValue, mask) == xorResult) 57 ; 58} 59 60bool testBitwise(uint4 mask) 61{ 62 return true 63 & test1Bitwise<int>(mask) 64 & testVBitwise<int, 2>(mask) 65 & testVBitwise<int, 3>(mask) 66 & testVBitwise<int, 4>(mask) 67 & test1Bitwise<uint>(mask) 68 & testVBitwise<uint, 2>(mask) 69 & testVBitwise<uint, 3>(mask) 70 & testVBitwise<uint, 4>(mask) 71 72 // TODO: these are failing SPIRV validation and should be fixed. 73 // SPIRV's ops do not directly accept/return bool. 74 // & test1Bitwise<bool>(mask) 75 // & testVBitwise<bool, 2>(mask) 76 // & testVBitwise<bool, 3>(mask) 77 // & testVBitwise<bool, 4>(mask) 78 79#if !defined(CUDA) 80 & test1Bitwise<int8_t>(mask) 81 & testVBitwise<int8_t, 2>(mask) 82 & testVBitwise<int8_t, 3>(mask) 83 & testVBitwise<int8_t, 4>(mask) 84 & test1Bitwise<int16_t>(mask) 85 & testVBitwise<int16_t, 2>(mask) 86 & testVBitwise<int16_t, 3>(mask) 87 & testVBitwise<int16_t, 4>(mask) 88 & test1Bitwise<int64_t>(mask) 89 & testVBitwise<int64_t, 2>(mask) 90 & testVBitwise<int64_t, 3>(mask) 91 & testVBitwise<int64_t, 4>(mask) 92 & test1Bitwise<uint8_t>(mask) 93 & testVBitwise<uint8_t, 2>(mask) 94 & testVBitwise<uint8_t, 3>(mask) 95 & testVBitwise<uint8_t, 4>(mask) 96 & test1Bitwise<uint16_t>(mask) 97 & testVBitwise<uint16_t, 2>(mask) 98 & testVBitwise<uint16_t, 3>(mask) 99 & testVBitwise<uint16_t, 4>(mask) 100 & test1Bitwise<uint64_t>(mask) 101 & testVBitwise<uint64_t, 2>(mask) 102 & testVBitwise<uint64_t, 3>(mask) 103 & testVBitwise<uint64_t, 4>(mask) 104#endif 105 ; 106} 107 108[numthreads(32, 1, 1)] 109[shader("compute")] 110void computeMain(uint3 dispatchThreadID : SV_DispatchThreadID) 111{ 112 let index = dispatchThreadID.x; 113 114 let isSecondGroup = index >= 15; 115 let mask = isSecondGroup ? uint4(0xFFFF8000, 0, 0, 0) : uint4(0x0007FFF, 0, 0, 0); 116 117 // One invocation in second group is different from others to test or and xor operations. 118 let isOrSet = (index == 15); 119 120 gAndValue = isSecondGroup ? uint(1) : uint(0); 121 gOrValue = isOrSet ? uint(1) : uint(0); 122 gOrResult = isSecondGroup ? uint(1) : uint(0); 123 124 // Alternate 0s and 1s for xor. 125 gXorValue = (index % 2 == 0) ? uint(0) : uint(1); 126 if (isOrSet) 127 { 128 // This is in second group - disrupt the alternating sequence. 129 gXorValue = uint(0); 130 } 131 gXorResult = isSecondGroup ? uint(0) : uint(1); 132 133 bool result = true 134 & testBitwise(mask) 135 ; 136 137 // CHECK-COUNT-32: 1 138 outputBuffer[index] = uint(result); 139}