yum-mirror/slang

Making it easier to work with shaders

git clone https://git.yummers.dev/yum-mirror/slang

Darren WihandiAdd full support for SPV_NV_shader_subgroup_partitioned (#7103)0476b57fa

master
5.6 KiB146 linesraw
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