optirs_gpu/shaders/
wgsl.rs1pub 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
52pub 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
94pub 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
138pub 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
194pub 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
230pub 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
258pub 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"#;