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.1 KiB143 linesraw
1// Test calling differentiable function through dynamic dispatch.
2
3//TEST(compute):COMPARE_COMPUTE_EX:-slang -compute -shaderobj -output-using-type
4//TEST(compute, vulkan):COMPARE_COMPUTE_EX:-vk -compute -shaderobj -output-using-type
5//TEST(compute):COMPARE_COMPUTE_EX:-cuda -compute -shaderobj -output-using-type
6
7//TEST_INPUT:ubuffer(data=[0 0 0 0 0], stride=4):out,name=outputBuffer
8RWStructuredBuffer<float> outputBuffer;
9
10//TEST_INPUT: set g_materials = new StructuredBuffer<MaterialDataBlob>[new MaterialDataBlob{new MaterialHeader{[0, 0, 0, 0]}, new MaterialPayload{[1.0, 1.2, 0.3, 0.5]}}];
11RWStructuredBuffer<MaterialDataBlob> g_materials;
12
13public struct ShadingInput
14{
15    public float scale;
16}
17
18struct MaterialHeader
19{
20    uint4 header;
21};
22struct MaterialPayload
23{
24    float4 data;
25};
26struct MaterialDataBlob
27{
28    MaterialHeader header;   // 16B
29    MaterialPayload payload; // 16B
30};
31
32interface IMaterial : IDifferentiable
33{
34    associatedtype MaterialInstance : IMaterialInstance;
35
36    [Differentiable]
37    MaterialInstance setupMaterialInstance( ShadingInput input );
38}
39
40interface IMaterialInstance : IDifferentiable
41{
42    [Differentiable]
43    float eval( float x );
44}
45
46
47[BackwardDerivative(getMaterial_bwd)]
48IMaterial getMaterial(int id)
49{
50    return createDynamicObject<IMaterial, MaterialDataBlob>(id, g_materials[id]);
51}
52
53void getMaterial_bwd(int id, IDifferentiable d)
54{
55    // Something random
56    outputBuffer[id] = 2.f;
57}
58
59struct Material1: IMaterial
60{
61    typedef MaterialInstance1 MaterialInstance;
62
63    MaterialHeader header;
64    float a;
65    float b;
66    float c;
67
68    [Differentiable]
69    MaterialInstance1 setupMaterialInstance( ShadingInput input )
70    {
71        MaterialInstance1 instance;
72        instance.a = a * input.scale;
73        instance.b = b * input.scale;
74        instance.c = c * input.scale;
75        return instance;
76    }
77
78}
79struct MaterialInstance1: IMaterialInstance
80{
81    float a;
82    float b;
83    float c;
84
85    [Differentiable]
86    float eval( float x )
87    {
88        return a * x * x + b * x + c;
89    }
90}
91
92struct Material2: IMaterial
93{
94    typedef MaterialInstance2 MaterialInstance;
95
96    MaterialHeader header;
97    float a;
98    float b;
99
100    [Differentiable]
101    MaterialInstance2 setupMaterialInstance( ShadingInput input )
102    {
103        MaterialInstance2 instance;
104        instance.a = a * input.scale * input.scale;
105        instance.b = b * input.scale * input.scale;
106        return instance;
107    }
108
109}
110public struct MaterialInstance2: IMaterialInstance
111{
112    float a;
113    float b;
114
115    [Differentiable]
116    public float eval( float x )
117    {
118        return a * x + b;
119    }
120}
121
122[Differentiable]
123public float shade(int material, ShadingInput input, float x)
124{
125    IMaterial m = getMaterial(material);
126    IMaterialInstance mi = m.setupMaterialInstance(input);
127    return mi.eval(x);
128}
129
130//TEST_INPUT: type_conformance Material1:IMaterial = 0
131//TEST_INPUT: type_conformance Material2:IMaterial = 1
132
133[shader("compute")]
134void computeMain(uint3 dispatchThreadID : SV_DispatchThreadID)
135{
136    outputBuffer[0] = shade(0, {0.5}, 0.6);
137
138    // TODO: VERIFY
139    DifferentialPair<float> dpx = diffPair(3.0);
140    bwd_diff(shade)(0, {0.5}, dpx, 1.0);
141
142    outputBuffer[3] = dpx.d;
143}