Skip to main content

ADAM

Constant ADAM 

Source
pub const ADAM: &str = r#"
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
@group(0) @binding(1) var<storage, read> y: array<f32>;
@group(0) @binding(2) var<storage, read_write> a: array<f32>;
@group(0) @binding(3) var<storage, read_write> b: array<f32>;
@group(0) @binding(4) var<storage, read> result: array<f32>;

@compute @workgroup_size(256) fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
    let n = bitcast<u32>(result[7]);
    let idx = gid.x;
    if (idx >= n) { return; }

    let lr = result[0];
    let beta1 = result[1];
    let beta2 = result[2];
    let eps = result[3];
    let wd = result[4];
    let bc1 = result[5];
    let bc2 = result[6];

    let p = x[idx];
    var g = y[idx];
    if (wd > 0.0) { g = g + wd * p; }

    let mi = beta1 * a[idx] + (1.0 - beta1) * g;
    let vi = beta2 * b[idx] + (1.0 - beta2) * g * g;
    a[idx] = mi;
    b[idx] = vi;

    let m_hat = mi / bc1;
    let v_hat = vi / bc2;
    x[idx] = p - lr * m_hat / (sqrt(v_hat) + eps);
}
"#;
Expand description

Adam with coupled L2 weight decay, matching optirs_core::optimizers::Adam.

bindingnamemeaning
0xparameters (read-write)
1ygradients (read)
2afirst moment m (read-write)
3bsecond moment v (read-write)
4result[lr, beta1, beta2, eps, weight_decay, bc1, bc2, n]