yum-mirror/slang

Making it easier to work with shaders

git clone https://git.yummers.dev/yum-mirror/slang

Yong HeAdd MLP training examples. (#7550)f28f67d98

master
1.1 KiB38 linesraw
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}