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
1.5 KiB60 linesraw
1//TEST(compute):COMPARE_COMPUTE_EX(filecheck-buffer=CHECK):-slang -compute -output-using-type
2//TEST(compute,vulkan):COMPARE_COMPUTE_EX(filecheck-buffer=CHECK):-vk -slang -compute -output-using-type
3//TEST(compute):COMPARE_COMPUTE_EX(filecheck-buffer=CHECK):-cuda -compute -shaderobj -output-using-type
4
5[Differentiable]
6Optional<float> sumSquare(Optional<float> a, Optional<float> b)
7{
8    if (let x = a)
9    {
10        if (let y = b)
11        {
12            return x * x + y * y;
13        }
14        else
15        {
16            return x * x;
17        }
18    }
19    else if (let y = b)
20    {
21        return y * y;
22    }
23    return none;
24}
25
26//TEST_INPUT:ubuffer(data=[0 0 0], stride=4):out,name=outputBuffer
27RWStructuredBuffer<float> outputBuffer;
28
29[numthreads(1,1,1)]
30void computeMain()
31{
32    var dpa = diffPair<Optional<float>>(3.0f, none);
33    var dpb = diffPair<Optional<float>>(4.0f, none);
34
35    bwd_diff(sumSquare)(dpa, dpb, 1.0f);
36
37    outputBuffer[0] = -1;
38
39    // CHECK: 14.0
40    if (dpa.d.hasValue && dpb.d.hasValue)
41        outputBuffer[0] = dpa.d.value + dpb.d.value;
42
43    // CHECK: 1.0
44    dpa = diffPair<Optional<float>>(3.0f, none);
45    dpb = diffPair<Optional<float>>(4.0f, none);
46    bwd_diff(sumSquare)(dpa, dpb, none);
47    if (dpa.d.value == 0.0 && dpb.d.value == 0.0)
48    {
49        outputBuffer[1] = 1.0f;
50    }
51
52     // CHECK: 100.0
53     dpa = diffPair<Optional<float>>(none, none);
54     dpb = diffPair<Optional<float>>(4.0f, none);
55     bwd_diff(sumSquare)(dpa, dpb, 1.0);
56     if (dpa.d == none)
57     {
58         outputBuffer[2] = 100.0f;
59     }
60}