yum-mirror/slang
Making it easier to work with shaders
git clone https://git.yummers.dev/yum-mirror/slang
f28f67d98
master
1module kernels; 2 3import common; 4import mlp; 5import network; 6import adam; 7 8[numthreads(256, 1, 1)] 9[require(spvGroupNonUniformBallot, spvGroupNonUniformArithmetic, spvCooperativeVectorNV)] 10void learnGradient( 11 uint32_t tid : SV_DispatchThreadID, 12 uniform MyNetwork* network, 13 uniform Atomic<uint32_t>* lossBuffer, 14 uniform float2* inputs, 15 uniform uint32_t count) 16{ 17 if (tid >= count) 18 return; 19 20 var input = (half2)inputs[tid]; 21 bwd_diff(loss)(network, input.x, input.y, 1.0h); 22 let thisLoss = (float)loss(network, input.x, input.y); 23 let maxLoss = WaveActiveMax(thisLoss); 24 if (WaveIsFirstLane()) 25 { 26 lossBuffer.max(bit_cast<uint32_t>(maxLoss)); 27 } 28} 29 30[numthreads(256, 1, 1)] 31void adjustParameters(uint32_t tid : SV_DispatchThreadID, uniform AdamState* states, uniform NFloat* params, uniform NFloat* gradients, uniform uint32_t count) 32{ 33 if (tid >= count) 34 return; 35 if (isnan(gradients[tid])) 36 { 37 gradients[tid] = 0.0h; 38 return; 39 } 40 AdamOptimizer::step(states[tid], params[tid], gradients[tid]); 41}