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.
| binding | name | meaning |
|---|---|---|
| 0 | x | parameters (read-write) |
| 1 | y | gradients (read) |
| 2 | a | first moment m (read-write) |
| 3 | b | second moment v (read-write) |
| 4 | result | [lr, beta1, beta2, eps, weight_decay, bc1, bc2, n] |