yum-mirror/slang

Making it easier to work with shaders

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

Harsh Aggarwal (NVIDIA)Fix 7723 - Add autodiff tests (#7919)d10732742

master
7.5 KiB224 linesraw
1//TEST(compute):COMPARE_COMPUTE_EX:-cuda -compute -shaderobj -output-using-type
2//TEST(compute, vulkan):COMPARE_COMPUTE_EX:-vk -compute -shaderobj -output-using-type
3//TEST(compute):COMPARE_COMPUTE_EX:-slang -compute -shaderobj -output-using-type
4
5//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 0 0 0 0 0 0], stride=4):out,name=outputBuffer
6RWStructuredBuffer<float> outputBuffer;
7
8typedef DifferentialPair<float> dpfloat;
9typedef DifferentialPair<float2> dpfloat2;
10typedef DifferentialPair<float3> dpfloat3;
11
12[Differentiable]
13float _clamp(float x, float min, float max)
14{
15    return clamp(x, min, max);
16}
17
18[Differentiable]
19float3 _clamp3(float3 x, float3 min, float3 max)
20{
21    return clamp(x, min, max);
22}
23
24[Differentiable]
25float _clamp_equiv(float x, float _min, float _max)
26{
27    return max(_min, min(_max, x));
28}
29
30[Differentiable]
31float3 _clamp_equiv(float3 x, float3 _min, float3 _max)
32{
33    return max(_min, min(_max, x));
34}
35
36[numthreads(1, 1, 1)]
37void computeMain(uint3 dispatchThreadID: SV_DispatchThreadID)
38{
39    // x in between max and min
40    {
41        dpfloat dpx = dpfloat(2.0, 0.1);
42        dpfloat dpmax = dpfloat(3.0, 0.2);
43        dpfloat dpmin = dpfloat(1.0, 0.3);
44
45        dpfloat res = fwd_diff(_clamp)(dpx, dpmin, dpmax);
46        outputBuffer[0] = res.d; // Expected: 0.1
47    }
48
49    // x less than min
50    {
51        dpfloat dpx = dpfloat(0.5, 0.1);
52        dpfloat dpmax = dpfloat(3.0, 0.2);
53        dpfloat dpmin = dpfloat(1.0, 0.3);
54
55        dpfloat res = fwd_diff(_clamp)(dpx, dpmin, dpmax);
56        outputBuffer[1] = res.d; // Expected: 0.3
57    }
58
59    // x greater than max
60    {
61        dpfloat dpx = dpfloat(4.0, 0.1);
62        dpfloat dpmax = dpfloat(3.0, 0.2);
63        dpfloat dpmin = dpfloat(1.0, 0.3);
64
65        dpfloat res = fwd_diff(_clamp)(dpx, dpmin, dpmax);
66        outputBuffer[2] = res.d; // Expected: 0.2
67    }
68
69    // float3 version with one in between, one below min and one above max.
70    {
71        dpfloat3 dpx = dpfloat3(float3(2.0, 0.5, 4.0), float3(0.1, 0.1, 0.1));
72        dpfloat3 dpmax = dpfloat3(float3(3.0, 3.0, 3.0), float3(0.2, 0.2, 0.2));
73        dpfloat3 dpmin = dpfloat3(float3(1.0, 1.0, 1.0), float3(0.3, 0.3, 0.3));
74
75        dpfloat3 res = fwd_diff(_clamp3)(dpx, dpmin, dpmax);
76        outputBuffer[3] = res.d.x; // Expected: 0.1
77        outputBuffer[4] = res.d.y; // Expected: 0.3
78        outputBuffer[5] = res.d.z; // Expected: 0.2
79    }
80
81    // Equivalent to the first test, but with a different implementation of clamp
82    {
83        dpfloat dpx = dpfloat(2.0, 0.1);
84        dpfloat dpmax = dpfloat(3.0, 0.2);
85        dpfloat dpmin = dpfloat(1.0, 0.3);
86
87        dpfloat res = fwd_diff(_clamp_equiv)(dpx, dpmin, dpmax);
88        outputBuffer[6] = res.d; // Expected: 0.1
89    }
90
91    // Equivalent to the second test, but with a different implementation of clamp
92    {
93        dpfloat dpx = dpfloat(0.5, 0.1);
94        dpfloat dpmax = dpfloat(3.0, 0.2);
95        dpfloat dpmin = dpfloat(1.0, 0.3);
96
97        dpfloat res = fwd_diff(_clamp_equiv)(dpx, dpmin, dpmax);
98        outputBuffer[7] = res.d; // Expected: 0.3
99    }
100
101    // Equivalent to the third test, but with a different implementation of clamp
102    {
103        dpfloat dpx = dpfloat(4.0, 0.1);
104        dpfloat dpmax = dpfloat(3.0, 0.2);
105        dpfloat dpmin = dpfloat(1.0, 0.3);
106
107        dpfloat res = fwd_diff(_clamp_equiv)(dpx, dpmin, dpmax);
108        outputBuffer[8] = res.d; // Expected: 0.2
109    }
110
111    // Equivalent to the fourth test, but with a different implementation of clamp
112    {
113        dpfloat3 dpx = dpfloat3(float3(2.0, 0.5, 4.0), float3(0.1, 0.1, 0.1));
114        dpfloat3 dpmax = dpfloat3(float3(3.0, 3.0, 3.0), float3(0.2, 0.2, 0.2));
115        dpfloat3 dpmin = dpfloat3(float3(1.0, 1.0, 1.0), float3(0.3, 0.3, 0.3));
116
117        dpfloat3 res = fwd_diff(_clamp_equiv)(dpx, dpmin, dpmax);
118        outputBuffer[9] = res.d.x; // Expected: 0.1
119        outputBuffer[10] = res.d.y; // Expected: 0.3
120        outputBuffer[11] = res.d.z; // Expected: 0.2
121    }
122
123    // Reverse-mode tests.
124    
125    // x in between max and min
126    {
127        dpfloat dpx = dpfloat(2.0, 0.0);
128        dpfloat dpmax = dpfloat(3.0, 0.0);
129        dpfloat dpmin = dpfloat(1.0, 0.0);
130
131        bwd_diff(_clamp)(dpx, dpmin, dpmax, 1.0);
132        
133        outputBuffer[12] = dpx.d; // Expected: 1.0
134        outputBuffer[13] = dpmin.d; // Expected: 0.0
135        outputBuffer[14] = dpmax.d; // Expected: 0.0
136    }
137
138    // x less than min
139    {
140        dpfloat dpx = dpfloat(0.5, 0.0);
141        dpfloat dpmax = dpfloat(3.0, 0.0);
142        dpfloat dpmin = dpfloat(1.0, 0.0);
143
144        bwd_diff(_clamp)(dpx, dpmin, dpmax, 1.0);
145        
146        outputBuffer[15] = dpx.d; // Expected: 0.0
147        outputBuffer[16] = dpmin.d; // Expected: 1.0
148        outputBuffer[17] = dpmax.d; // Expected: 0.0
149    }
150
151    // x greater than max
152    {
153        dpfloat dpx = dpfloat(4.0, 0.0);
154        dpfloat dpmax = dpfloat(3.0, 0.0);
155        dpfloat dpmin = dpfloat(1.0, 0.0);
156
157        bwd_diff(_clamp)(dpx, dpmin, dpmax, 1.0);
158        
159        outputBuffer[18] = dpx.d; // Expected: 0.0
160        outputBuffer[19] = dpmin.d; // Expected: 0.0
161        outputBuffer[20] = dpmax.d; // Expected: 1.0
162    }
163
164    // float3 version with one in between, one below min and one above max.
165    {
166        dpfloat3 dpx = dpfloat3(float3(2.0, 0.5, 4.0), float3(0.0, 0.0, 0.0));
167        dpfloat3 dpmax = dpfloat3(float3(3.0, 3.0, 3.0), float3(0.0, 0.0, 0.0));
168        dpfloat3 dpmin = dpfloat3(float3(1.0, 1.0, 1.0), float3(0.0, 0.0, 0.0));
169
170        bwd_diff(_clamp3)(dpx, dpmin, dpmax, float3(0.1, 0.2, 0.3));
171        
172        outputBuffer[21] = dpx.d.x; // Expected: 0.1
173        outputBuffer[22] = dpx.d.y; // Expected: 0.0
174        outputBuffer[23] = dpx.d.z; // Expected: 0.0
175        outputBuffer[24] = dpmin.d.x; // Expected: 0.0
176        outputBuffer[25] = dpmin.d.y; // Expected: 0.2
177        outputBuffer[26] = dpmin.d.z; // Expected: 0.0
178        outputBuffer[27] = dpmax.d.x; // Expected: 0.0
179        outputBuffer[28] = dpmax.d.y; // Expected: 0.0
180        outputBuffer[29] = dpmax.d.z; // Expected: 0.3
181    }
182
183    // New tests: Forward-mode tests for derivative propagation at the edges with clamp(x, 0, 1)
184    {
185        // Lower edge: x exactly = 0
186        dpfloat dpx = dpfloat(0.0, 0.4);
187        dpfloat dpmin = dpfloat(0.0, 0.8);
188        dpfloat dpmax = dpfloat(1.0, 0.5);
189        dpfloat res = fwd_diff(_clamp)(dpx, dpmin, dpmax);
190        outputBuffer[30] = res.d; // Expected: 0.4 (propagated from x)
191    }
192
193    {
194        // Upper edge: x exactly = 1
195        dpfloat dpx = dpfloat(1.0, 0.7);
196        dpfloat dpmin = dpfloat(0.0, 0.8);
197        dpfloat dpmax = dpfloat(1.0, 0.9);
198        dpfloat res = fwd_diff(_clamp)(dpx, dpmin, dpmax);
199        outputBuffer[31] = res.d; // Expected: 0.7 (propagated from x)
200    }
201
202    // Reverse-mode tests for derivative propagation at the edges with clamp(x, 0, 1)
203    {
204        // Lower edge: x exactly = 0
205        dpfloat dpx = dpfloat(0.0, 0.0);
206        dpfloat dpmin = dpfloat(0.0, 0.0);
207        dpfloat dpmax = dpfloat(1.0, 0.0);
208        bwd_diff(_clamp)(dpx, dpmin, dpmax, 1.0);
209        outputBuffer[32] = dpx.d;   // Expected: 1.0 (propagated from x)
210        outputBuffer[33] = dpmin.d; // Expected: 0.0
211        outputBuffer[34] = dpmax.d; // Expected: 0.0
212    }
213
214    {
215        // Upper edge: x exactly = 1
216        dpfloat dpx = dpfloat(1.0, 0.0);
217        dpfloat dpmin = dpfloat(0.0, 0.0);
218        dpfloat dpmax = dpfloat(1.0, 0.0);
219        bwd_diff(_clamp)(dpx, dpmin, dpmax, 1.0);
220        outputBuffer[35] = dpx.d;   // Expected: 1.0 (propagated from x)
221        outputBuffer[36] = dpmin.d; // Expected: 0.0
222        outputBuffer[37] = dpmax.d; // Expected: 0.0
223    }
224}