yum-mirror/slang
Making it easier to work with shaders
git clone https://git.yummers.dev/yum-mirror/slang
0476b57fa
master
1//TEST:SIMPLE(filecheck=CHECK_SPIRV): -stage compute -entry computeMain -target spirv -DNO_INTEGER_MATRIX 2//TEST:SIMPLE(filecheck=CHECK_GLSL): -stage compute -entry computeMain -target glsl -DNO_INTEGER_MATRIX 3//TEST:SIMPLE(filecheck=CHECK_CUDA): -stage compute -entry computeMain -target cuda 4//TEST:SIMPLE(filecheck=CHECK_HLSL): -stage compute -entry computeMain -target hlsl 5 6// 7// Tests all variants and overloads of WaveMultiPrefix* arithmetic intrinsics. 8// 9 10struct OutputData 11{ 12 int scalarSum; 13 int scalarProduct; 14 int scalarBitAnd; 15 int scalarBitOr; 16 int scalarBitXor; 17 int vectorSum; 18 int vectorProduct; 19 int vectorBitAnd; 20 int vectorBitOr; 21 int vectorBitXor; 22 int matrixSum; 23 int matrixProduct; 24 int matrixBitAnd; 25 int matrixBitOr; 26 int matrixBitXor; 27 float floatScalarSum; 28 float floatScalarProduct; 29 float floatVectorSum; 30 float floatVectorProduct; 31 float floatMatrixSum; 32 float floatMatrixProduct; 33}; 34 35RWStructuredBuffer<OutputData> outputBuffer; 36 37// CHECK_SPIRV: OpCapability GroupNonUniformPartitionedNV 38// CHECK_SPIRV: OpExtension "SPV_NV_shader_subgroup_partitioned" 39// CHECK_SPIRV: OpGroupNonUniformIAdd{{.*}}PartitionedExclusiveScanNV 40// CHECK_SPIRV: OpGroupNonUniformIMul{{.*}}PartitionedExclusiveScanNV 41// CHECK_SPIRV: OpGroupNonUniformBitwiseAnd{{.*}}PartitionedExclusiveScanNV 42// CHECK_SPIRV: OpGroupNonUniformBitwiseOr{{.*}}PartitionedExclusiveScanNV 43// CHECK_SPIRV: OpGroupNonUniformBitwiseXor{{.*}}PartitionedExclusiveScanNV 44// CHECK_SPIRV: OpGroupNonUniformFAdd{{.*}}PartitionedExclusiveScanNV 45 46// CHECK_GLSL: GL_NV_shader_subgroup_partitioned 47// CHECK_GLSL: subgroupPartitionedExclusiveAddNV 48// CHECK_GLSL: subgroupPartitionedExclusiveMulNV 49// CHECK_GLSL: subgroupPartitionedExclusiveAndNV 50// CHECK_GLSL: subgroupPartitionedExclusiveOrNV 51// CHECK_GLSL: subgroupPartitionedExclusiveXorNV 52 53// CHECK_CUDA: _wavePrefixSum 54// CHECK_CUDA: _wavePrefixProduct 55// CHECK_CUDA: _wavePrefixAnd 56// CHECK_CUDA: _wavePrefixOr 57// CHECK_CUDA: _wavePrefixXor 58// CHECK_CUDA: _wavePrefixSumMultiple 59// CHECK_CUDA: _wavePrefixProductMultiple 60// CHECK_CUDA: _wavePrefixAndMultiple 61// CHECK_CUDA: _wavePrefixOrMultiple 62// CHECK_CUDA: _wavePrefixXorMultiple 63 64// CHECK_HLSL: WaveMultiPrefixSum 65// CHECK_HLSL: WaveMultiPrefixProduct 66// CHECK_HLSL: WaveMultiPrefixBitAnd 67// CHECK_HLSL: WaveMultiPrefixBitOr 68// CHECK_HLSL: WaveMultiPrefixBitXor 69 70 71[numthreads(1, 1, 1)] 72void computeMain(uint3 dTid : SV_DispatchThreadID) 73{ 74 int scalarVal = dTid.x; 75 uint4 mask = WaveMatch(scalarVal); 76 77 int scalarSum = WaveMultiPrefixSum(scalarVal, mask); 78 int scalarProduct = WaveMultiPrefixProduct(scalarVal, mask); 79 int scalarBitAnd = WaveMultiPrefixBitAnd(scalarVal, mask); 80 int scalarBitOr = WaveMultiPrefixBitOr(scalarVal, mask); 81 int scalarBitXor = WaveMultiPrefixBitXor(scalarVal, mask); 82 83 int3 vectorVal = int3(dTid.x, dTid.y, dTid.z); 84 int3 vectorSum = WaveMultiPrefixSum(vectorVal, mask); 85 int3 vectorProduct = WaveMultiPrefixProduct(vectorVal, mask); 86 int3 vectorBitAnd = WaveMultiPrefixBitAnd(vectorVal, mask); 87 int3 vectorBitOr = WaveMultiPrefixBitOr(vectorVal, mask); 88 int3 vectorBitXor = WaveMultiPrefixBitXor(vectorVal, mask); 89 90 float floatScalarVal = float(dTid.x) + 0.5f; // Example floating-point scalar value 91 uint4 floatMask = WaveMatch(floatScalarVal); // Create a mask for matching lanes 92 93 float floatScalarSum = WaveMultiPrefixSum(floatScalarVal, floatMask); 94 float floatScalarProduct = WaveMultiPrefixProduct(floatScalarVal, floatMask); 95 96 float3 floatVectorVal = float3(dTid.x, dTid.y, dTid.z) + 0.5f; // Example floating-point vector value 97 float3 floatVectorSum = WaveMultiPrefixSum(floatVectorVal, floatMask); 98 float3 floatVectorProduct = WaveMultiPrefixProduct(floatVectorVal, floatMask); 99 100 OutputData output; 101 output.scalarSum = scalarSum; 102 output.scalarProduct = scalarProduct; 103 output.scalarBitAnd = scalarBitAnd; 104 output.scalarBitOr = scalarBitOr; 105 output.scalarBitXor = scalarBitXor; 106 output.vectorSum = vectorSum.x; 107 output.vectorProduct = vectorProduct.x; 108 output.vectorBitAnd = vectorBitAnd.x; 109 output.vectorBitOr = vectorBitOr.x; 110 output.vectorBitXor = vectorBitXor.x; 111 output.floatScalarSum = floatScalarSum; 112 output.floatScalarProduct = floatScalarProduct; 113 output.floatVectorSum = floatVectorSum.x; 114 output.floatVectorProduct = floatVectorProduct.x; 115 116 float3x3 floatMatrixVal = float3x3( 117 float(dTid.x) + 0.5f, float(dTid.y) + 0.5f, float(dTid.z) + 0.5f, 118 float(dTid.z) + 0.5f, float(dTid.x) + 0.5f, float(dTid.y) + 0.5f, 119 float(dTid.y) + 0.5f, float(dTid.z) + 0.5f, float(dTid.x) + 0.5f 120 ); 121 float3x3 floatMatrixSum = WaveMultiPrefixSum(floatMatrixVal, floatMask); 122 float3x3 floatMatrixProduct = WaveMultiPrefixProduct(floatMatrixVal, floatMask); 123 output.floatMatrixSum = floatMatrixSum[0][0]; 124 output.floatMatrixProduct = floatMatrixProduct[0][0]; 125 126#if !defined(NO_INTEGER_MATRIX) 127 int3x3 matrixVal = int3x3( 128 dTid.x, dTid.y, dTid.z, 129 dTid.z, dTid.x, dTid.y, 130 dTid.y, dTid.z, dTid.x 131 ); 132 int3x3 matrixSum = WaveMultiPrefixSum(matrixVal, mask); 133 int3x3 matrixProduct = WaveMultiPrefixProduct(matrixVal, mask); 134 int3x3 matrixBitAnd = WaveMultiPrefixBitAnd(matrixVal, mask); 135 int3x3 matrixBitOr = WaveMultiPrefixBitOr(matrixVal, mask); 136 int3x3 matrixBitXor = WaveMultiPrefixBitXor(matrixVal, mask); 137 output.matrixSum = matrixSum[0][0]; 138 output.matrixProduct = matrixProduct[0][0]; 139 output.matrixBitAnd = matrixBitAnd[0][0]; 140 output.matrixBitOr = matrixBitOr[0][0]; 141 output.matrixBitXor = matrixBitXor[0][0]; 142#endif 143 144 outputBuffer[dTid.x] = output; 145} 146