Skip to main content

optirs_gpu/shaders/
wgsl.rs

1//! WGSL sources for the optimizer kernels (WebGPU backend).
2//!
3//! See [`super`] for the buffer-naming and scalar-packing conventions. In
4//! short: bindings are named `x`, `y`, `a`, `b`, `result`, `output` in that
5//! order, every scalar travels in a storage buffer, and integer scalars are
6//! bit-cast through an `f32` slot.
7
8/// Adam with coupled L2 weight decay, matching `optirs_core::optimizers::Adam`.
9///
10/// | binding | name | meaning |
11/// |---|---|---|
12/// | 0 | `x` | parameters (read-write) |
13/// | 1 | `y` | gradients (read) |
14/// | 2 | `a` | first moment `m` (read-write) |
15/// | 3 | `b` | second moment `v` (read-write) |
16/// | 4 | `result` | `[lr, beta1, beta2, eps, weight_decay, bc1, bc2, n]` |
17pub const ADAM: &str = r#"
18@group(0) @binding(0) var<storage, read_write> x: array<f32>;
19@group(0) @binding(1) var<storage, read> y: array<f32>;
20@group(0) @binding(2) var<storage, read_write> a: array<f32>;
21@group(0) @binding(3) var<storage, read_write> b: array<f32>;
22@group(0) @binding(4) var<storage, read> result: array<f32>;
23
24@compute @workgroup_size(256) fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
25    let n = bitcast<u32>(result[7]);
26    let idx = gid.x;
27    if (idx >= n) { return; }
28
29    let lr = result[0];
30    let beta1 = result[1];
31    let beta2 = result[2];
32    let eps = result[3];
33    let wd = result[4];
34    let bc1 = result[5];
35    let bc2 = result[6];
36
37    let p = x[idx];
38    var g = y[idx];
39    if (wd > 0.0) { g = g + wd * p; }
40
41    let mi = beta1 * a[idx] + (1.0 - beta1) * g;
42    let vi = beta2 * b[idx] + (1.0 - beta2) * g * g;
43    a[idx] = mi;
44    b[idx] = vi;
45
46    let m_hat = mi / bc1;
47    let v_hat = vi / bc2;
48    x[idx] = p - lr * m_hat / (sqrt(v_hat) + eps);
49}
50"#;
51
52/// AdamW with *decoupled* weight decay (Loshchilov & Hutter).
53///
54/// The decay never enters the moment estimates; it is applied straight to the
55/// parameter, which is exactly what separates AdamW from Adam + L2.
56///
57/// Bindings are identical to [`ADAM`].
58pub const ADAMW: &str = r#"
59@group(0) @binding(0) var<storage, read_write> x: array<f32>;
60@group(0) @binding(1) var<storage, read> y: array<f32>;
61@group(0) @binding(2) var<storage, read_write> a: array<f32>;
62@group(0) @binding(3) var<storage, read_write> b: array<f32>;
63@group(0) @binding(4) var<storage, read> result: array<f32>;
64
65@compute @workgroup_size(256) fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
66    let n = bitcast<u32>(result[7]);
67    let idx = gid.x;
68    if (idx >= n) { return; }
69
70    let lr = result[0];
71    let beta1 = result[1];
72    let beta2 = result[2];
73    let eps = result[3];
74    let wd = result[4];
75    let bc1 = result[5];
76    let bc2 = result[6];
77
78    let p = x[idx];
79    let g = y[idx];
80
81    let mi = beta1 * a[idx] + (1.0 - beta1) * g;
82    let vi = beta2 * b[idx] + (1.0 - beta2) * g * g;
83    a[idx] = mi;
84    b[idx] = vi;
85
86    let m_hat = mi / bc1;
87    let v_hat = vi / bc2;
88
89    let decayed = p - lr * wd * p;
90    x[idx] = decayed - lr * m_hat / (sqrt(v_hat) + eps);
91}
92"#;
93
94/// SGD with optional momentum, dampening, Nesterov acceleration and L2 decay.
95///
96/// | binding | name | meaning |
97/// |---|---|---|
98/// | 0 | `x` | parameters (read-write) |
99/// | 1 | `y` | gradients (read) |
100/// | 2 | `a` | momentum buffer (read-write) |
101/// | 3 | `b` | `[lr, momentum, dampening, weight_decay, nesterov, first_step, n]` |
102///
103/// `first_step` is `1.0` only on the very first update, so the momentum buffer
104/// is seeded with the raw gradient rather than a dampened one.
105pub const SGD: &str = r#"
106@group(0) @binding(0) var<storage, read_write> x: array<f32>;
107@group(0) @binding(1) var<storage, read> y: array<f32>;
108@group(0) @binding(2) var<storage, read_write> a: array<f32>;
109@group(0) @binding(3) var<storage, read> b: array<f32>;
110
111@compute @workgroup_size(256) fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
112    let n = bitcast<u32>(b[6]);
113    let idx = gid.x;
114    if (idx >= n) { return; }
115
116    let lr = b[0];
117    let momentum = b[1];
118    let dampening = b[2];
119    let wd = b[3];
120    let nesterov = b[4];
121    let first = b[5];
122
123    let p = x[idx];
124    var g = y[idx];
125    if (wd > 0.0) { g = g + wd * p; }
126
127    if (momentum > 0.0) {
128        var buf = g;
129        if (first < 0.5) { buf = momentum * a[idx] + (1.0 - dampening) * g; }
130        a[idx] = buf;
131        if (nesterov > 0.5) { g = g + momentum * buf; } else { g = buf; }
132    }
133
134    x[idx] = p - lr * g;
135}
136"#;
137
138/// RMSprop with optional centering and momentum (PyTorch semantics).
139///
140/// | binding | name | meaning |
141/// |---|---|---|
142/// | 0 | `x` | parameters (read-write) |
143/// | 1 | `y` | gradients (read) |
144/// | 2 | `a` | squared-gradient average (read-write) |
145/// | 3 | `b` | mean-gradient average, used only when centered (read-write) |
146/// | 4 | `result` | momentum buffer (read-write) |
147/// | 5 | `output` | `[lr, alpha, eps, weight_decay, momentum, centered, n]` |
148pub const RMSPROP: &str = r#"
149@group(0) @binding(0) var<storage, read_write> x: array<f32>;
150@group(0) @binding(1) var<storage, read> y: array<f32>;
151@group(0) @binding(2) var<storage, read_write> a: array<f32>;
152@group(0) @binding(3) var<storage, read_write> b: array<f32>;
153@group(0) @binding(4) var<storage, read_write> result: array<f32>;
154@group(0) @binding(5) var<storage, read> output: array<f32>;
155
156@compute @workgroup_size(256) fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
157    let n = bitcast<u32>(output[6]);
158    let idx = gid.x;
159    if (idx >= n) { return; }
160
161    let lr = output[0];
162    let alpha = output[1];
163    let eps = output[2];
164    let wd = output[3];
165    let momentum = output[4];
166    let centered = output[5];
167
168    let p = x[idx];
169    var g = y[idx];
170    if (wd > 0.0) { g = g + wd * p; }
171
172    let sq = alpha * a[idx] + (1.0 - alpha) * g * g;
173    a[idx] = sq;
174
175    var avg = sq;
176    if (centered > 0.5) {
177        let ga = alpha * b[idx] + (1.0 - alpha) * g;
178        b[idx] = ga;
179        avg = sq - ga * ga;
180    }
181
182    let denom = sqrt(max(avg, 0.0)) + eps;
183
184    if (momentum > 0.0) {
185        let buf = momentum * result[idx] + g / denom;
186        result[idx] = buf;
187        x[idx] = p - lr * buf;
188    } else {
189        x[idx] = p - lr * g / denom;
190    }
191}
192"#;
193
194/// Adagrad with learning-rate decay.
195///
196/// | binding | name | meaning |
197/// |---|---|---|
198/// | 0 | `x` | parameters (read-write) |
199/// | 1 | `y` | gradients (read) |
200/// | 2 | `a` | accumulated squared gradients (read-write) |
201/// | 3 | `b` | `[clr, eps, weight_decay, n]` |
202///
203/// The host passes the already-decayed step size `clr`; the decay denominator
204/// counts *completed* steps, so the first update uses exactly `lr`.
205pub const ADAGRAD: &str = r#"
206@group(0) @binding(0) var<storage, read_write> x: array<f32>;
207@group(0) @binding(1) var<storage, read> y: array<f32>;
208@group(0) @binding(2) var<storage, read_write> a: array<f32>;
209@group(0) @binding(3) var<storage, read> b: array<f32>;
210
211@compute @workgroup_size(256) fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
212    let n = bitcast<u32>(b[3]);
213    let idx = gid.x;
214    if (idx >= n) { return; }
215
216    let clr = b[0];
217    let eps = b[1];
218    let wd = b[2];
219
220    let p = x[idx];
221    var g = y[idx];
222    if (wd > 0.0) { g = g + wd * p; }
223
224    let s = a[idx] + g * g;
225    a[idx] = s;
226    x[idx] = p - clr * g / (sqrt(s) + eps);
227}
228"#;
229
230/// Local reduction step for a multi-GPU all-reduce-mean collective.
231///
232/// | binding | name | meaning |
233/// |---|---|---|
234/// | 0 | `x` | this device's local contribution (read-write, in place) |
235/// | 1 | `y` | `[n, num_gpus]` (both bit-cast `u32`) |
236///
237/// This divides the local buffer by the replica count. It is the *finishing*
238/// step of a sum-then-average all-reduce: the summation across physical
239/// devices itself requires a transport this crate does not have (see
240/// [`crate::multi_gpu`]), so the only replica count ever dispatched is `1`
241/// (this device's own contribution), for which the division is the
242/// mathematically exact identity — computed for real on the GPU rather than
243/// asserted on the host.
244pub const ALL_REDUCE_MEAN: &str = r#"
245@group(0) @binding(0) var<storage, read_write> x: array<f32>;
246@group(0) @binding(1) var<storage, read> y: array<f32>;
247
248@compute @workgroup_size(256) fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
249    let n = bitcast<u32>(y[0]);
250    let idx = gid.x;
251    if (idx >= n) { return; }
252
253    let num_gpus = bitcast<u32>(y[1]);
254    x[idx] = x[idx] / f32(num_gpus);
255}
256"#;
257
258/// LAMB, run as two dispatches of the *same* pipeline.
259///
260/// | binding | name | meaning |
261/// |---|---|---|
262/// | 0 | `x` | parameters (read-write) |
263/// | 1 | `y` | gradients (read) |
264/// | 2 | `a` | first moment `m` (read-write) |
265/// | 3 | `b` | second moment `v` (read-write) |
266/// | 4 | `result` | scratch: `update[0..n]` then `partials[n..n+2*groups]` |
267/// | 5 | `output` | `[lr, beta1, beta2, eps, weight_decay, bc1, bc2, n, phase, trust]` |
268///
269/// * `phase == 0` advances the moments, materialises
270///   `update = m_hat / (sqrt(v_hat) + eps) + weight_decay * p`, and emits
271///   per-workgroup partial sums of `p^2` and `update^2`.
272/// * `phase == 1` applies `p -= lr * trust * update`, with the trust ratio
273///   computed on the host from the finished reduction.
274///
275/// One pipeline serves both phases because `GpuCompiler::compile` in
276/// scirs2-core 0.6.5 registers every compiled WGSL shader under one internal
277/// name per context, so two live handles would alias.
278///
279/// The workgroup reduction sits *outside* the `idx < n` guard so that
280/// `workgroupBarrier()` is only ever reached in uniform control flow.
281pub const LAMB: &str = r#"
282@group(0) @binding(0) var<storage, read_write> x: array<f32>;
283@group(0) @binding(1) var<storage, read> y: array<f32>;
284@group(0) @binding(2) var<storage, read_write> a: array<f32>;
285@group(0) @binding(3) var<storage, read_write> b: array<f32>;
286@group(0) @binding(4) var<storage, read_write> result: array<f32>;
287@group(0) @binding(5) var<storage, read> output: array<f32>;
288
289var<workgroup> scratch_p: array<f32, 256>;
290var<workgroup> scratch_u: array<f32, 256>;
291
292@compute @workgroup_size(256) fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>, @builtin(workgroup_id) wgid: vec3<u32>) {
293    let n = bitcast<u32>(output[7]);
294    let phase = bitcast<u32>(output[8]);
295    let idx = gid.x;
296
297    let lr = output[0];
298    let beta1 = output[1];
299    let beta2 = output[2];
300    let eps = output[3];
301    let wd = output[4];
302    let bc1 = output[5];
303    let bc2 = output[6];
304    let trust = output[9];
305
306    var sum_p = 0.0;
307    var sum_u = 0.0;
308
309    if (idx < n) {
310        if (phase == 0u) {
311            let p = x[idx];
312            let g = y[idx];
313            let mi = beta1 * a[idx] + (1.0 - beta1) * g;
314            let vi = beta2 * b[idx] + (1.0 - beta2) * g * g;
315            a[idx] = mi;
316            b[idx] = vi;
317            let m_hat = mi / bc1;
318            let v_hat = vi / bc2;
319            result[idx] = m_hat / (sqrt(v_hat) + eps) + wd * p;
320        } else {
321            x[idx] = x[idx] - lr * trust * result[idx];
322        }
323        let pv = x[idx];
324        let uv = result[idx];
325        sum_p = pv * pv;
326        sum_u = uv * uv;
327    }
328
329    scratch_p[lid.x] = sum_p;
330    scratch_u[lid.x] = sum_u;
331    workgroupBarrier();
332
333    var stride = 128u;
334    loop {
335        if (lid.x < stride) {
336            scratch_p[lid.x] = scratch_p[lid.x] + scratch_p[lid.x + stride];
337            scratch_u[lid.x] = scratch_u[lid.x] + scratch_u[lid.x + stride];
338        }
339        workgroupBarrier();
340        if (stride == 1u) { break; }
341        stride = stride >> 1u;
342    }
343
344    if (lid.x == 0u) {
345        result[n + wgid.x * 2u] = scratch_p[0];
346        result[n + wgid.x * 2u + 1u] = scratch_u[0];
347    }
348}
349"#;