yum-mirror/slang
Making it easier to work with shaders
git clone https://git.yummers.dev/yum-mirror/slang
051ae8ace
master
1//TEST:SIMPLE(filecheck=CHECK): -target spirv 2// CHECK: OpEntryPoint 3 4module test; 5 6public enum class MaterialID : uint { invalid = 0xffffffff }; 7 8public struct Material : IDifferentiable 9{ 10 float x; 11} 12 13public struct Hit 14{ 15 MaterialID material; 16} 17 18public struct Scene 19{ 20 StructuredBuffer<Material> materials; 21 RWStructuredBuffer<Material> grads; 22 23 [Differentiable] 24 Material load(MaterialID id) { return materials[uint(id)]; } 25 26 void accumulate(MaterialID id, Material d) { grads[uint(id)].x += d.x; } 27 28 [Differentiable, BackwardDerivative(_get_material_bwd)] 29 public Material get_material(MaterialID id) { return load(id); } 30 31 public void _get_material_bwd(MaterialID id, Material d) { accumulate(id, d); } 32 33 [Differentiable] 34 public Material get_material(Hit hit) { return get_material(hit.material); } 35} 36 37[Differentiable] 38float trace(const Scene scene, Hit hit) 39{ 40 Material m = scene.get_material(hit); 41 return m.x; 42} 43 44 45[shader("compute")] 46void main( 47 uniform Scene scene, 48 uniform StructuredBuffer<uint> input, 49 uniform RWStructuredBuffer<float> output, 50 uniform RWStructuredBuffer<float> grads 51) 52{ 53 Hit hit; 54 hit.material = MaterialID(input[0]); 55 output[0] = trace(scene, hit); 56 bwd_diff(trace)(scene, hit, grads[0]); 57}