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(filecheck-buffer=CHECK):-slang -compute -shaderobj -output-using-type 2//TEST(compute):COMPARE_COMPUTE_EX(filecheck-buffer=CHECK):-cuda -compute -shaderobj -output-using-type 3//TEST(compute, vulkan):COMPARE_COMPUTE_EX(filecheck-buffer=CHECK):-vk -compute -shaderobj -output-using-type 4 5enum MyEnum { A, B, C }; 6 7[BackwardDerivative(mDiff)] 8float m<let M : MyEnum>(float x) 9{ 10 switch (M) 11 { 12 case MyEnum.A: 13 return x * x; 14 case MyEnum.B: 15 return x; 16 case MyEnum.C: 17 return 3 * x; 18 default: 19 return 0; 20 } 21} 22 23void mDiff<let M : MyEnum>(inout DifferentialPair<float> x, float dResult) 24{ 25 switch (M) 26 { 27 case MyEnum.A: 28 updateDiff(x, 2 * dResult * x.p); 29 break; 30 case MyEnum.B: 31 updateDiff(x, dResult); 32 break; 33 case MyEnum.C: 34 updateDiff(x, 3 * dResult); 35 break; 36 default: 37 updateDiff(x, 0); 38 break; 39 } 40} 41 42[Differentiable] 43float test(float x) 44{ 45 return m<MyEnum.A>(x); 46} 47 48//TEST_INPUT:ubuffer(data=[0 0 0 0], stride=4):out,name=outputBuffer 49RWStructuredBuffer<float> outputBuffer; 50 51[numthreads(1, 1, 1)] 52void computeMain(uint3 dispatchThreadID: SV_DispatchThreadID) 53{ 54 var a = diffPair(3.0); 55 __bwd_diff(test)(a, 1.0); 56 outputBuffer[dispatchThreadID.x] = a.d; 57 // CHECK: 6.0 58}