yum-mirror/slang
Making it easier to work with shaders
git clone https://git.yummers.dev/yum-mirror/slang
d10732742
master
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}