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
3.6 KiB128 linesraw
1//TEST(compute, vulkan):COMPARE_COMPUTE_EX:-vk -compute -shaderobj -output-using-type
2//TEST(compute):COMPARE_COMPUTE_EX:-slang -compute -shaderobj -output-using-type
3//TEST(compute):COMPARE_COMPUTE_EX:-cuda -compute -shaderobj -output-using-type
4
5//TEST_INPUT:ubuffer(data=[0 0 0 0 0], stride=4):out,name=outputBuffer
6RWStructuredBuffer<float> outputBuffer;
7
8typedef DifferentialPair<float> dpfloat;
9
10typealias IDFloat = __BuiltinFloatingPointType & IDifferentiable;
11
12namespace myintrinsiclib
13{
14    __generic<T : IDFloat>
15    __target_intrinsic(hlsl, "exp($0)")
16    __target_intrinsic(glsl, "exp($0)")
17    __target_intrinsic(cuda, "$P_exp($0)")
18    __target_intrinsic(cpp, "$P_exp($0)")
19    __target_intrinsic(spirv, "12 resultType resultId glsl450 27 _0")
20    __target_intrinsic(metal, "exp($0)")
21    __target_intrinsic(wgsl, "exp($0)")
22    [ForwardDerivative(d_myexp<T>)]
23    T myexp(T x);
24
25    __generic<T : IDFloat>
26    DifferentialPair<T> d_myexp(DifferentialPair<T> dpx)
27    {
28        return DifferentialPair<T>(
29            myexp(dpx.p),
30            T.dmul(myexp(dpx.p), dpx.d));
31    }
32
33    
34    // Sine
35    __generic<T : IDFloat>
36    __target_intrinsic(hlsl, "sin($0)")
37    __target_intrinsic(glsl, "sin($0)")
38    __target_intrinsic(metal, "sin($0)")
39    __target_intrinsic(cuda, "$P_sin($0)")
40    __target_intrinsic(cpp, "$P_sin($0)")
41    __target_intrinsic(spirv, "12 resultType resultId glsl450 13 _0")
42    __target_intrinsic(wgsl, "sin($0)")
43    [ForwardDerivative(d_mysin<T>)]
44    T mysin(T x);
45
46    __generic<T : IDFloat>
47    DifferentialPair<T> d_mysin(DifferentialPair<T> dpx)
48    {
49        return DifferentialPair<T>(
50            mysin(dpx.p),
51            T.dmul(mycos(dpx.p), dpx.d));
52    }
53
54    // Cosine
55    __generic<T : IDFloat>
56    __target_intrinsic(hlsl, "cos($0)")
57    __target_intrinsic(glsl, "cos($0)")
58    __target_intrinsic(metal, "cos($0)")
59    __target_intrinsic(cuda, "$P_cos($0)")
60    __target_intrinsic(cpp, "$P_cos($0)")
61    __target_intrinsic(spirv, "12 resultType resultId glsl450 14 _0")
62    __target_intrinsic(wgsl, "cos($0)")
63    [ForwardDerivative(d_mycos<T>)]
64    T mycos(T x);
65
66    __generic<T : IDFloat>
67    DifferentialPair<T> d_mycos(DifferentialPair<T> dpx)
68    {
69        return DifferentialPair<T>(
70            mycos(dpx.p),
71            T.dmul(-sin(dpx.p), dpx.d));
72    }
73
74    // Sine and cosine
75    __generic<T : IDFloat>
76    __target_intrinsic(hlsl, "sincos($0, $1, $2)")
77    __target_intrinsic(cuda, "$P_sincos($0, $1, $2)")
78    [ForwardDerivative(d_mysincos<T>)]
79    void mysincos(T x, out T s, out T c)
80    {
81        s = sin(x);
82        c = cos(x);
83    }
84
85    __generic<T : IDFloat>
86    void d_mysincos(DifferentialPair<T> x, out DifferentialPair<T> s, out DifferentialPair<T> c)
87    {
88        T _s;
89        T _c;
90        mysincos(x.p, _s, _c);
91
92        s = DifferentialPair<T>(_s, T.dmul(_c, x.d));
93        c = DifferentialPair<T>(_c, T.dmul(-_s, x.d));
94    }
95};
96
97[ForwardDifferentiable]
98float f(float x)
99{
100    return myintrinsiclib.myexp(x);
101}
102
103[ForwardDifferentiable]
104float g(float x)
105{
106    float s;
107    float t;
108    myintrinsiclib.mysincos(x, s, t);
109
110    return s + t;
111}
112
113[numthreads(1, 1, 1)]
114void computeMain(uint3 dispatchThreadID: SV_DispatchThreadID)
115{
116    {
117        dpfloat dpa = dpfloat(2.0, 1.0);
118
119        outputBuffer[0] = f(dpa.p);        // Expect: 7.389056
120        outputBuffer[1] = __fwd_diff(f)(dpa).d; // Expect: 7.389056
121
122        // g() needs additional handling of  IRMakeDifferentialPair(PtrType). This needs to 
123        // generate a new var, load from the individual vars and store into the pair var.
124
125        //outputBuffer[2] = g(dpa.p);        // Expect: 1.381773
126        //outputBuffer[3] = __fwd_diff(g)(dpa).d; // Expect: -0.301168
127    }
128}