yum-mirror/slang
Making it easier to work with shaders
git clone https://git.yummers.dev/yum-mirror/slang
f28f67d98
master
1module adam; 2 3import mlp_sw; 4import common; 5 6public struct AdamState 7{ 8 internal NFloat mean; 9 internal NFloat variance; 10 internal int iteration; 11} 12 13public struct AdamOptimizer 14{ 15 // Adam parameters 16 public static const NFloat beta1 = 0.9h; 17 public static const NFloat beta2 = 0.999h; 18 public static const NFloat epsilon = 1e-7h; 19 public static const NFloat learningRate = 0.01h; 20 21 public static void step(inout AdamState state, inout NFloat param, inout NFloat grad) 22 { 23 state.iteration++; 24 if (isinf(grad)) 25 { 26 if (grad > 0) 27 grad = 10000.0h; 28 else 29 grad = -10000.0h; 30 } 31 state.mean = beta1 * state.mean + (NFloat(1.f) - beta1) * grad; 32 state.variance = beta2 * state.variance + (NFloat(1.f) - beta2) * grad * grad; 33 NFloat meanHat = state.mean / (NFloat(1.f) - pow(beta1, NFloat(state.iteration))); 34 NFloat varianceHat = state.variance / (NFloat(1.f) - pow(beta2, NFloat(state.iteration))); 35 param -= learningRate * meanHat / (sqrt(max(NFloat(0.f), varianceHat) + epsilon)); 36 grad = NFloat(0.f); 37 } 38}