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
4.5 KiB139 linesraw
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}