yum-mirror/slang
Making it easier to work with shaders
git clone https://git.yummers.dev/yum-mirror/slang
f28f67d98
master
1module network; 2 3import common; 4import mlp_sw; 5 6public struct MyNetwork 7{ 8 public FeedForwardLayer<4, 16> layer1; 9 public FeedForwardLayer<16, 4> layer2; 10 11 [Differentiable] 12 internal MLVec<4> encodeInput(NFloat x, NFloat y) 13 { 14 return MLVec<4>.fromArray({ 15 x, 16 y, 17 x*x, 18 y*y, 19 }); 20 } 21 22 [Differentiable] 23 internal MLVec<4> _eval(NFloat x, NFloat y) 24 { 25 let encoding = encodeInput(x, y); 26 let layer1Output = layer1.eval(encoding); 27 let leyer2Output = layer2.eval(layer1Output); 28 return leyer2Output; 29 } 30 31 [Differentiable] 32 public half4 eval(no_diff NFloat x, no_diff NFloat y) 33 { 34 let mlv = _eval(x, y); 35 let arr = mlv.toArray(); 36 return half4(arr[0], arr[1], arr[2], arr[3]); 37 } 38} 39 40[Differentiable] 41public half loss(MyNetwork* network, no_diff half x, no_diff half y) 42{ 43 let networkResult = network.eval(x, y); 44 let gt = no_diff groundtruth(x, y); 45 let diff = networkResult - gt; 46 47 return dot(diff, diff); 48} 49 50public half4 groundtruth(half x, half y) 51{ 52 return { 53 (x + y) / (1 + y * y), 54 2 * x + y, 55 0.5 * x * x + 1.2 * y, 56 x + 0.5 * y * y, 57 }; 58} 59