wgblas 1.2.1 → 2.1.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (96) hide show
  1. package/LICENSE +1 -1
  2. package/README.md +54 -72
  3. package/dist/wgblas.browser.js +1637 -858
  4. package/index.d.mts +51 -6
  5. package/index.mjs +8 -0
  6. package/package.json +56 -2
  7. package/src/classes/GpuMatrix.mjs +17 -10
  8. package/src/classes/GpuVector.mjs +28 -10
  9. package/src/dasum/dasum.d.mts +6 -6
  10. package/src/dasum/dasum.mjs +25 -21
  11. package/src/devdocs.mjs +13 -0
  12. package/src/idamax/idamax.d.mts +69 -0
  13. package/src/idamax/idamax.mjs +130 -0
  14. package/src/init.mjs +115 -49
  15. package/src/isamax/isamax.d.mts +21 -3
  16. package/src/isamax/isamax.mjs +16 -14
  17. package/src/random/random.d.mts +1 -0
  18. package/src/sasum/sasum.d.mts +3 -3
  19. package/src/sasum/sasum.mjs +15 -13
  20. package/src/saxpy/saxpy.d.mts +3 -3
  21. package/src/saxpy/saxpy.mjs +10 -8
  22. package/src/scopy/scopy.d.mts +3 -3
  23. package/src/scopy/scopy.mjs +10 -8
  24. package/src/sdot/sdot.d.mts +3 -3
  25. package/src/sdot/sdot.mjs +16 -14
  26. package/src/sgemm/sgemm.d.mts +102 -0
  27. package/src/sgemm/sgemm.mjs +208 -0
  28. package/src/sgemmtr/sgemmtr.d.mts +104 -0
  29. package/src/sgemmtr/sgemmtr.mjs +204 -0
  30. package/src/sgemv/sgemv.d.mts +1 -38
  31. package/src/sgemv/sgemv.mjs +42 -26
  32. package/src/sger/sger.d.mts +1 -34
  33. package/src/sger/sger.mjs +12 -8
  34. package/src/shaders/block_transfer.wgsl +42 -0
  35. package/src/shaders/dasum.wgsl +3 -2
  36. package/src/shaders/f64/dekker.wgsl +4 -85
  37. package/src/shaders/f64/utils/abs.wgsl +10 -0
  38. package/src/shaders/f64/utils/add.wgsl +77 -0
  39. package/src/shaders/f64/utils/equal.wgsl +7 -0
  40. package/src/shaders/f64/utils/greater.wgsl +12 -0
  41. package/src/shaders/f64/utils/multiply.wgsl +81 -0
  42. package/src/shaders/idamax.wgsl +96 -0
  43. package/src/shaders/index.mjs +164 -14
  44. package/src/shaders/reduction/argmaxF64.wgsl +50 -0
  45. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  46. package/src/shaders/reduction/sumF64.wgsl +2 -2
  47. package/src/shaders/sgemm_large.wgsl +206 -0
  48. package/src/shaders/sgemm_small.wgsl +212 -0
  49. package/src/shaders/sgemmtr_large.wgsl +120 -0
  50. package/src/shaders/sgemmtr_small.wgsl +113 -0
  51. package/src/shaders/sgemv_n.wgsl +3 -1
  52. package/src/shaders/sgemv_t.wgsl +3 -1
  53. package/src/shaders/snrm2.wgsl +72 -23
  54. package/src/shaders/ssymv.wgsl +3 -1
  55. package/src/shaders/symmetrize.wgsl +31 -0
  56. package/src/shaders/triangularize.wgsl +44 -0
  57. package/src/snrm2/snrm2.d.mts +3 -3
  58. package/src/snrm2/snrm2.mjs +33 -21
  59. package/src/srot/srot.d.mts +3 -5
  60. package/src/srot/srot.mjs +11 -9
  61. package/src/srotm/srotm.d.mts +3 -5
  62. package/src/srotm/srotm.mjs +19 -10
  63. package/src/sscal/sscal.d.mts +4 -4
  64. package/src/sscal/sscal.mjs +12 -10
  65. package/src/sswap/sswap.d.mts +3 -3
  66. package/src/sswap/sswap.mjs +11 -9
  67. package/src/ssymm/ssymm.d.mts +103 -0
  68. package/src/ssymm/ssymm.mjs +218 -0
  69. package/src/ssymv/ssymv.d.mts +1 -36
  70. package/src/ssymv/ssymv.mjs +12 -8
  71. package/src/ssyr/ssyr.d.mts +1 -30
  72. package/src/ssyr/ssyr.mjs +11 -7
  73. package/src/ssyr2/ssyr2.d.mts +1 -34
  74. package/src/ssyr2/ssyr2.mjs +12 -8
  75. package/src/ssyr2k/ssyr2k.d.mts +100 -0
  76. package/src/ssyr2k/ssyr2k.mjs +202 -0
  77. package/src/ssyrk/ssyrk.d.mts +90 -0
  78. package/src/ssyrk/ssyrk.mjs +177 -0
  79. package/src/strmm/strmm.d.mts +100 -0
  80. package/src/strmm/strmm.mjs +226 -0
  81. package/src/strmv/strmv.d.mts +1 -36
  82. package/src/strmv/strmv.mjs +12 -8
  83. package/src/strsm/strsm.d.mts +99 -0
  84. package/src/strsm/strsm.mjs +360 -0
  85. package/src/strsv/strsv.d.mts +1 -32
  86. package/src/strsv/strsv.mjs +18 -12
  87. package/src/util/benchmark.mjs +4 -6
  88. package/src/util/bindgroup.mjs +1 -3
  89. package/src/util/buffer.mjs +116 -20
  90. package/src/util/compute.mjs +12 -12
  91. package/src/util/constants.mjs +57 -0
  92. package/src/util/device.mjs +34 -0
  93. package/src/util/f64.mjs +3 -3
  94. package/src/util/pipeline.mjs +5 -6
  95. package/src/util/workgroup.mjs +55 -7
  96. package/src/shaders/browser-shaders.mjs +0 -55
@@ -0,0 +1,77 @@
1
+ // Requires f64/dekker.wgsl concatenated first for the DD struct.
2
+
3
+ // ── A real compiler bug — read before touching anything below ──────────────
4
+ //
5
+ // twoSum/fastTwoSum's error term `e` should be nonzero (that's the point —
6
+ // `s` is rounded). A front-end-optimizer bug zeros it anyway, on both NVIDIA
7
+ // and Mesa (ANV + llvmpipe), via two different mechanisms needing two fixes:
8
+ // bitcast-based subtraction (`fsub`/`negf`, fixes NVIDIA) and materializing
9
+ // the sum through workgroup memory + workgroupBarrier() (fixes Mesa). Only
10
+ // both together (ddAddProtected) is verified correct everywhere — the plain
11
+ // twoSum/fastTwoSum/ddAdd below are reference-only, not safe to use.
12
+ fn negf(x: f32) -> f32 {
13
+ return bitcast<f32>(bitcast<u32>(x) ^ 0x80000000u);
14
+ }
15
+ fn fsub(a: f32, b: f32) -> f32 {
16
+ return a + negf(b);
17
+ }
18
+
19
+ // Knuth/Møller's TwoSum: s = fl(a+b), e = exact rounding error, a+b == s+e.
20
+ // Works for any a, b. UNPROTECTED — see header above.
21
+ fn twoSum(a: f32, b: f32) -> DD {
22
+ let s = a + b;
23
+ let v = s - a;
24
+ let e = (a - (s - v)) + (b - v);
25
+ return DD(s, e);
26
+ }
27
+
28
+ // Dekker's Fast-Two-Sum: same contract, but only correct when |a| >= |b|.
29
+ // UNPROTECTED — see header above.
30
+ fn fastTwoSum(a: f32, b: f32) -> DD {
31
+ let s = a + b;
32
+ let e = b - (s - a);
33
+ return DD(s, e);
34
+ }
35
+
36
+ // Double-double addition (Dekker's Add2). UNPROTECTED — see header above.
37
+ fn ddAdd(a: DD, b: DD) -> DD {
38
+ let s = twoSum(a.hi, b.hi);
39
+ let loSum = a.lo + b.lo;
40
+ return fastTwoSum(s.hi, s.lo + loSum);
41
+ }
42
+
43
+ // ── Protected variants — use these ──────────────────────────────────────────
44
+ //
45
+ // Bitcast subtraction + workgroup-barrier materialization, verified correct
46
+ // on all three backends tested. Costs a real barrier: fine for O(1)-per-
47
+ // thread or O(log n) reduction use, not a long per-element loop. A
48
+ // workgroupBarrier() requires uniform control flow, so:
49
+ // - `threadSlot` must be unique per concurrent caller (e.g. local_invocation_index).
50
+ // - Every thread in the workgroup must call this the same number of times
51
+ // — including ones whose result gets discarded. Compute unconditionally;
52
+ // only the write-back should be conditional.
53
+ var<workgroup> dekkerScratch: array<f32, 64>;
54
+
55
+ fn twoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
56
+ dekkerScratch[threadSlot] = a + b;
57
+ workgroupBarrier();
58
+ let s = dekkerScratch[threadSlot];
59
+ let v = fsub(s, a);
60
+ let e = fsub(a, fsub(s, v)) + fsub(b, v);
61
+ return DD(s, e);
62
+ }
63
+
64
+ fn fastTwoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
65
+ dekkerScratch[threadSlot] = a + b;
66
+ workgroupBarrier();
67
+ let s = dekkerScratch[threadSlot];
68
+ let e = fsub(b, fsub(s, a));
69
+ return DD(s, e);
70
+ }
71
+
72
+ // Protected double-double addition — same contract as ddAdd, but exact.
73
+ fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
74
+ let s = twoSumProtected(a.hi, b.hi, threadSlot);
75
+ let loSum = a.lo + b.lo;
76
+ return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
77
+ }
@@ -0,0 +1,7 @@
1
+ // Requires f64/dekker.wgsl concatenated first for the DD struct.
2
+
3
+ // a == b for double-double pairs — exact field equality, no rounding
4
+ // involved, so (like ddGreater) this needs no protection.
5
+ fn ddEqual(a: DD, b: DD) -> bool {
6
+ return a.hi == b.hi && a.lo == b.lo;
7
+ }
@@ -0,0 +1,12 @@
1
+ // Requires f64/dekker.wgsl concatenated first for the DD struct.
2
+
3
+ // a > b for double-double pairs. hi dominates (|lo| <= ulp(hi)/2 always), so
4
+ // comparing hi alone is correct except on an exact hi tie, when lo breaks it.
5
+ // A plain comparison, not a rounding-identity subtraction — no reassociation
6
+ // risk, so unlike twoSum/fastTwoSum this needs no protection.
7
+ fn ddGreater(a: DD, b: DD) -> bool {
8
+ if (a.hi != b.hi) {
9
+ return a.hi > b.hi;
10
+ }
11
+ return a.lo > b.lo;
12
+ }
@@ -0,0 +1,81 @@
1
+ // Requires f64/dekker.wgsl concatenated first for the DD struct, and
2
+ // f64/utils/add.wgsl for fsub/negf (bitcast-based subtraction/negation) and,
3
+ // for ddMulProtected at the bottom, fastTwoSumProtected.
4
+ //
5
+ // Use twoProdBit — verified universal (0 corrupting failures across 3000+
6
+ // random trials on NVIDIA/Intel-Mesa-ANV/llvmpipe), no barrier protection
7
+ // needed. The classic approaches below (twoProd, twoProdFma) each fail on
8
+ // one backend in a way barrier materialization doesn't fix; twoProdBit
9
+ // sidesteps the bug instead by deriving the split via bitcast+bitmask
10
+ // rather than an arithmetic identity, leaving nothing for a reassociating
11
+ // compiler to fold. Intel Mesa ANV shows frequent last-bit-only diffs from
12
+ // strict ground truth (never data-corrupting) — consistent with the driver
13
+ // legitimately auto-fusing `x - y*z` into hardware FMA.
14
+ const SPLIT_CONST: f32 = 4097.0;
15
+
16
+ fn bitSplit(a: f32) -> DD {
17
+ let bits = bitcast<u32>(a);
18
+ let hiBits = bits & 0xFFFFF800u; // keep sign+exponent+top 12 mantissa bits
19
+ let hi = bitcast<f32>(hiBits);
20
+ let lo = fsub(a, hi); // exact by Sterbenz's lemma (hi, a share an exponent, are close)
21
+ return DD(hi, lo);
22
+ }
23
+
24
+ fn twoProdBit(a: f32, b: f32) -> DD {
25
+ let s = a * b;
26
+ let aSplit = bitSplit(a);
27
+ let bSplit = bitSplit(b);
28
+ let e = fsub(fsub(fsub(fsub(s, aSplit.hi * bSplit.hi), aSplit.lo * bSplit.hi), aSplit.hi * bSplit.lo), aSplit.lo * bSplit.lo);
29
+ return DD(s, negf(e));
30
+ }
31
+
32
+ // ── Unsafe historical reference — do not use ────────────────────────────
33
+ // Both broken on one backend, confirmed via isolated cross-driver testing,
34
+ // NOT fixed by barrier materialization (unlike addition's bug):
35
+ // - veltkampSplit/twoProd (Dekker's original): fails on NVIDIA — compiler
36
+ // folds `hi = c - (c - a)` to `= a` straight through fsub/negf, even
37
+ // with every intermediate barrier-materialized (11/11 fail, worse than
38
+ // unprotected's 6/11).
39
+ // - twoProdFma (Ogita/Rump/Oishi): fails on llvmpipe — its software fma()
40
+ // likely isn't genuinely fused, making `fma(a,b,-(a*b))` correctly (not
41
+ // buggily) zero. Materializing `s` doesn't change this.
42
+ fn veltkampSplit(a: f32) -> DD {
43
+ let c = SPLIT_CONST * a;
44
+ let big = fsub(c, a);
45
+ let hi = fsub(c, big);
46
+ let lo = fsub(a, hi);
47
+ return DD(hi, lo);
48
+ }
49
+
50
+ fn twoProd(a: f32, b: f32) -> DD {
51
+ let s = a * b;
52
+ let aSplit = veltkampSplit(a);
53
+ let bSplit = veltkampSplit(b);
54
+ let e = fsub(fsub(fsub(fsub(s, aSplit.hi * bSplit.hi), aSplit.lo * bSplit.hi), aSplit.hi * bSplit.lo), aSplit.lo * bSplit.lo);
55
+ return DD(s, negf(e));
56
+ }
57
+
58
+ fn twoProdFma(a: f32, b: f32) -> DD {
59
+ let s = a * b;
60
+ let e = fma(a, b, negf(s));
61
+ return DD(s, e);
62
+ }
63
+
64
+ // DD × DD product (Dekker/Bailey): twoProdBit(a.hi, b.hi) already captures
65
+ // the dominant term to full DD precision, and the cross terms are below the
66
+ // ~48-bit floor anyway, so folding them in with plain f32 loses nothing —
67
+ // only the final renormalization needs barrier protection. Split into
68
+ // ddMulRaw (unprotected) and ddMulProtected (renormalizes via
69
+ // fastTwoSumProtected) so callers with several products can batch them
70
+ // through one shared barrier. ddMulRaw's result isn't a valid DD pair on
71
+ // its own — it must be renormalized before use.
72
+ fn ddMulRaw(a: DD, b: DD) -> DD {
73
+ let p = twoProdBit(a.hi, b.hi);
74
+ let crossAndLo = p.lo + (a.hi * b.lo + a.lo * b.hi);
75
+ return DD(p.hi, crossAndLo);
76
+ }
77
+
78
+ fn ddMulProtected(a: DD, b: DD, threadSlot: u32) -> DD {
79
+ let raw = ddMulRaw(a, b);
80
+ return fastTwoSumProtected(raw.hi, raw.lo, threadSlot);
81
+ }
@@ -0,0 +1,96 @@
1
+ // idamax: returns index of element with largest absolute value (f64, double-double)
2
+ // pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/argmaxF64.wgsl.
3
+ // Concatenated after f64/dekker.wgsl (DD struct), f64/utils/abs.wgsl (ddAbs),
4
+ // f64/utils/greater.wgsl (ddGreater), and f64/utils/equal.wgsl (ddEqual).
5
+
6
+ @group(0) @binding(0) var<storage, read> xHi: array<f32>;
7
+ @group(0) @binding(1) var<storage, read> xLo: array<f32>;
8
+ @group(0) @binding(2) var<storage, read_write> partialsValHi: array<f32>;
9
+ @group(0) @binding(3) var<storage, read_write> partialsValLo: array<f32>;
10
+ @group(0) @binding(4) var<storage, read_write> partialsIdx: array<u32>;
11
+ @group(0) @binding(5) var<uniform> params: Params;
12
+
13
+ struct Params {
14
+ n: u32,
15
+ x_inc: u32,
16
+ }
17
+
18
+ const WGS: u32 = 64;
19
+
20
+ var<workgroup> tile_val: array<DD, 64>;
21
+ var<workgroup> tile_idx: array<u32, 64>;
22
+
23
+ @compute @workgroup_size(64)
24
+ fn idamax_main(
25
+ @builtin(global_invocation_id) gid: vec3u,
26
+ @builtin(local_invocation_id) lid: vec3u,
27
+ @builtin(workgroup_id) wgid: vec3u,
28
+ @builtin(num_workgroups) num_wg: vec3u,
29
+ ) {
30
+ // DD(-1.0, 0.0) is a safe sentinel: any |x[i]| >= 0 beats it,
31
+ // so workgroups with no elements lose gracefully in the epilogue.
32
+ var best_val0 = DD(-1.0, 0.0); var best_idx0: u32 = 0u;
33
+ var best_val1 = DD(-1.0, 0.0); var best_idx1: u32 = 0u;
34
+ var best_val2 = DD(-1.0, 0.0); var best_idx2: u32 = 0u;
35
+ var best_val3 = DD(-1.0, 0.0); var best_idx3: u32 = 0u;
36
+
37
+ let stride = num_wg.x * WGS;
38
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
39
+
40
+ for (var id = gid.x; id < n4_floor; id += 4u * stride) {
41
+ let i0 = id * params.x_inc;
42
+ let i1 = (id + stride) * params.x_inc;
43
+ let i2 = (id + 2u * stride) * params.x_inc;
44
+ let i3 = (id + 3u * stride) * params.x_inc;
45
+ let v0 = ddAbs(DD(xHi[i0], xLo[i0]));
46
+ let v1 = ddAbs(DD(xHi[i1], xLo[i1]));
47
+ let v2 = ddAbs(DD(xHi[i2], xLo[i2]));
48
+ let v3 = ddAbs(DD(xHi[i3], xLo[i3]));
49
+ if (ddGreater(v0, best_val0)) { best_val0 = v0; best_idx0 = id; }
50
+ if (ddGreater(v1, best_val1)) { best_val1 = v1; best_idx1 = id + stride; }
51
+ if (ddGreater(v2, best_val2)) { best_val2 = v2; best_idx2 = id + 2u * stride; }
52
+ if (ddGreater(v3, best_val3)) { best_val3 = v3; best_idx3 = id + 3u * stride; }
53
+ }
54
+ for (var id = n4_floor + gid.x; id < params.n; id += stride) {
55
+ let i = id * params.x_inc;
56
+ let v = ddAbs(DD(xHi[i], xLo[i]));
57
+ if (ddGreater(v, best_val0)) { best_val0 = v; best_idx0 = id; }
58
+ }
59
+
60
+ // merge 4 independent lanes; prefer lower index on tie (first occurrence wins)
61
+ if (ddGreater(best_val1, best_val0) ||
62
+ (ddEqual(best_val1, best_val0) && best_idx1 < best_idx0)) {
63
+ best_val0 = best_val1; best_idx0 = best_idx1;
64
+ }
65
+ if (ddGreater(best_val2, best_val0) ||
66
+ (ddEqual(best_val2, best_val0) && best_idx2 < best_idx0)) {
67
+ best_val0 = best_val2; best_idx0 = best_idx2;
68
+ }
69
+ if (ddGreater(best_val3, best_val0) ||
70
+ (ddEqual(best_val3, best_val0) && best_idx3 < best_idx0)) {
71
+ best_val0 = best_val3; best_idx0 = best_idx3;
72
+ }
73
+
74
+ tile_val[lid.x] = best_val0;
75
+ tile_idx[lid.x] = best_idx0;
76
+ workgroupBarrier();
77
+
78
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
79
+ if (lid.x < s) {
80
+ let a_val = tile_val[lid.x];
81
+ let b_val = tile_val[lid.x + s];
82
+ if (ddGreater(b_val, a_val) ||
83
+ (ddEqual(b_val, a_val) && tile_idx[lid.x + s] < tile_idx[lid.x])) {
84
+ tile_val[lid.x] = b_val;
85
+ tile_idx[lid.x] = tile_idx[lid.x + s];
86
+ }
87
+ }
88
+ workgroupBarrier();
89
+ }
90
+
91
+ if (lid.x == 0u) {
92
+ partialsValHi[wgid.x] = tile_val[0].hi;
93
+ partialsValLo[wgid.x] = tile_val[0].lo;
94
+ partialsIdx[wgid.x] = tile_idx[0];
95
+ }
96
+ }
@@ -3,25 +3,175 @@
3
3
  *
4
4
  * `shaders/*.wgsl` — one WGSL compute shader per BLAS routine (sscal, saxpy, sdot, …).
5
5
  *
6
- * `shaders/browser-shaders.mjs` — the browser's runtime shader source. In Node.js, shaders are
7
- * read directly from disk via `readFileSync`. In the browser there is no filesystem, so this file
8
- * provides all shader strings inline. Vite bundles it by importing each `.wgsl` file as a string.
6
+ * `routineShaders` below is the single source of truth: routine name → the WGSL source(s)
7
+ * its `getPipeline()` calls actually reference, verified against every `src/<routine>/<routine>.mjs`
8
+ * rather than inferred from naming convention (see its doc comment for the exceptions). Each
9
+ * shader is imported right above the line that adds it — the import *is* the mapping entry, no
10
+ * separate block to cross-reference. `shaderSources`, the flat name → source registry the
11
+ * browser bundle's runtime lookup needs, is *derived* from `routineShaders` rather than
12
+ * hand-duplicated, so the two can never drift apart. In Node.js neither is read — shaders are
13
+ * `readFileSync` from disk directly; `scripts/build-browser.mjs` inlines this module into the
14
+ * browser's IIFE bundle via esbuild instead.
9
15
  *
10
16
  * ## Cross-shader patterns
11
17
  *
12
- * **Fixed workgroup size of 64.** Every shader declares `const WGS: u32 = 64` and
13
- * `@workgroup_size(64)`. 64 is the minimum `maxComputeInvocationsPerWorkgroup` guaranteed across
14
- * all WebGPU devices, so this works everywhere without querying device limits.
18
+ * **Single bind group.** Every shader with bindings uses `@group(0)` only — the JS side always
19
+ * calls `pipeline.getBindGroupLayout(0)`, no secondary groups to track. Binding order is
20
+ * consistent too: any read-only storage buffers come before read_write ones, with the
21
+ * `uniform Params` struct always last. `@binding` indices match the position of each resource in
22
+ * the array passed to `createBindGroup`, which appends `resultBuffer` last.
15
23
  *
16
- * **Single bind group.** All bindings use `@group(0)`. This means the JS side always calls
17
- * `pipeline.getBindGroupLayout(0)` — no secondary groups to track.
24
+ * **All counts and strides are `u32`.** `n`, `x_inc`, `y_inc`, and every other index/count field
25
+ * in a `Params` struct is unsigned, avoiding implicit sign-extension in index expressions like
26
+ * `id * params.x_inc`.
18
27
  *
19
- * The `@binding` indices must match the position of each resource in the array passed to
20
- * `createBindGroup` — it assigns `binding: 0, 1, 2 …` sequentially, with `resultBuffer` appended last.
21
- *
22
- * **All counts and strides are `u32`.** `n`, `x_inc`, `y_inc`, and any other index fields in the
23
- * `Params` uniform struct are unsigned. This avoids implicit sign-extension when they appear in
24
- * index expressions like `id * params.x_inc`.
28
+ * **Entry points don't have to be named `main`.** `loadShader` (`util/pipeline.mjs`)
29
+ * auto-detects the sole `@compute` function in a module instead of requiring a fixed name, so
30
+ * `dasum_main`, `strsv_invert_block_main`, etc. work without renaming.
25
31
  *
26
32
  * @module devdocs/shaders
27
33
  */
34
+
35
+ /**
36
+ * Routine name → the WGSL source(s) its `getPipeline()` calls reference. Keys are the exact
37
+ * shader names `getPipeline(device, name)` is called with — most routines have one, some pick
38
+ * one of several conditionally (e.g. sgemv's `sgemv_n`/`sgemv_t`, by `trans`), and some have no
39
+ * dedicated shader at all:
40
+ *
41
+ * - `sgemmtr`/`ssyrk`/`ssyr2k` all dispatch through `sgemmtr_small`/`sgemmtr_large`.
42
+ * - `strsm` reuses `strsv_invert_block` and `sscal`, plus its own `block_transfer` and the
43
+ * shared `sgemm_small`/`sgemm_large`.
44
+ * - `dasum`/`idamax` concatenate several f64 utility shaders with their own — see
45
+ * `getPipeline`'s `shaderName: string[]` behaviour.
46
+ * - `random` has no entry — CPU-only, no `getPipeline()` call.
47
+ *
48
+ * Built up entry by entry so each import sits next to the mapping entry that uses it.
49
+ * @public
50
+ */
51
+ export const routineShaders = {};
52
+
53
+ import sscal from "./sscal.wgsl";
54
+ routineShaders.sscal = { sscal };
55
+
56
+ import sswap from "./sswap.wgsl";
57
+ routineShaders.sswap = { sswap };
58
+
59
+ import saxpy from "./saxpy.wgsl";
60
+ routineShaders.saxpy = { saxpy };
61
+
62
+ import scopy from "./scopy.wgsl";
63
+ routineShaders.scopy = { scopy };
64
+
65
+ import sdot from "./sdot.wgsl";
66
+ import sum from "./reduction/sum.wgsl";
67
+ routineShaders.sdot = { sdot, "reduction/sum": sum };
68
+
69
+ import sasum from "./sasum.wgsl";
70
+ routineShaders.sasum = { sasum, "reduction/sum": sum };
71
+
72
+ import snrm2 from "./snrm2.wgsl";
73
+ import scaledSum from "./reduction/scaledSum.wgsl";
74
+ routineShaders.snrm2 = { snrm2, "reduction/scaledSum": scaledSum };
75
+
76
+ import isamax from "./isamax.wgsl";
77
+ import argmax from "./reduction/argmax.wgsl";
78
+ routineShaders.isamax = { isamax, "reduction/argmax": argmax };
79
+
80
+ import dekker from "./f64/dekker.wgsl";
81
+ import ddAbs from "./f64/utils/abs.wgsl";
82
+ import ddAddUtil from "./f64/utils/add.wgsl";
83
+ import dasum from "./dasum.wgsl";
84
+ import sumF64 from "./reduction/sumF64.wgsl";
85
+ routineShaders.dasum = {
86
+ "f64/dekker": dekker,
87
+ "f64/utils/abs": ddAbs,
88
+ "f64/utils/add": ddAddUtil,
89
+ dasum,
90
+ "reduction/sumF64": sumF64,
91
+ };
92
+
93
+ import ddGreater from "./f64/utils/greater.wgsl";
94
+ import ddEqual from "./f64/utils/equal.wgsl";
95
+ import idamax from "./idamax.wgsl";
96
+ import argmaxF64 from "./reduction/argmaxF64.wgsl";
97
+ routineShaders.idamax = {
98
+ "f64/dekker": dekker,
99
+ "f64/utils/abs": ddAbs,
100
+ "f64/utils/greater": ddGreater,
101
+ "f64/utils/equal": ddEqual,
102
+ idamax,
103
+ "reduction/argmaxF64": argmaxF64,
104
+ };
105
+
106
+ import srot from "./srot.wgsl";
107
+ routineShaders.srot = { srot };
108
+
109
+ import srotm from "./srotm.wgsl";
110
+ routineShaders.srotm = { srotm };
111
+
112
+ import sgemv_n from "./sgemv_n.wgsl";
113
+ import sgemv_t from "./sgemv_t.wgsl";
114
+ routineShaders.sgemv = { sgemv_n, sgemv_t }; // one or the other, picked by trans
115
+
116
+ import ssymv from "./ssymv.wgsl";
117
+ routineShaders.ssymv = { ssymv };
118
+
119
+ import strmv from "./strmv.wgsl";
120
+ routineShaders.strmv = { strmv };
121
+
122
+ import strsv_invert_block from "./strsv_invert_block.wgsl";
123
+ import strsv_apply_inverse from "./strsv_apply_inverse.wgsl";
124
+ import strsv_update from "./strsv_update.wgsl";
125
+ routineShaders.strsv = {
126
+ strsv_invert_block,
127
+ strsv_apply_inverse,
128
+ strsv_update,
129
+ };
130
+
131
+ import sger from "./sger.wgsl";
132
+ routineShaders.sger = { sger };
133
+
134
+ import ssyr from "./ssyr.wgsl";
135
+ routineShaders.ssyr = { ssyr };
136
+
137
+ import ssyr2 from "./ssyr2.wgsl";
138
+ routineShaders.ssyr2 = { ssyr2 };
139
+
140
+ import sgemm_small from "./sgemm_small.wgsl";
141
+ import sgemm_large from "./sgemm_large.wgsl";
142
+ routineShaders.sgemm = { sgemm_small, sgemm_large }; // one or the other, picked by a tile-size threshold
143
+
144
+ import sgemmtr_small from "./sgemmtr_small.wgsl";
145
+ import sgemmtr_large from "./sgemmtr_large.wgsl";
146
+ routineShaders.sgemmtr = { sgemmtr_small, sgemmtr_large };
147
+
148
+ routineShaders.ssyrk = { sgemmtr_small, sgemmtr_large }; // no shader of its own — rides on sgemmtr's
149
+ routineShaders.ssyr2k = { sgemmtr_small, sgemmtr_large }; // no shader of its own — rides on sgemmtr's
150
+
151
+ import symmetrize from "./symmetrize.wgsl";
152
+ routineShaders.ssymm = { sgemm_small, sgemm_large, symmetrize };
153
+
154
+ import triangularize from "./triangularize.wgsl";
155
+ routineShaders.strmm = { sgemm_small, sgemm_large, triangularize };
156
+
157
+ import blockTransfer from "./block_transfer.wgsl";
158
+ routineShaders.strsm = {
159
+ strsv_invert_block,
160
+ block_transfer: blockTransfer,
161
+ sscal,
162
+ sgemm_small,
163
+ sgemm_large,
164
+ };
165
+
166
+ /**
167
+ * Flat shader-name → WGSL source-string registry — what `getPipeline()`/`loadShader()` (see
168
+ * `util/pipeline.mjs`) actually look shaders up in, in the browser. Derived from
169
+ * `routineShaders` by merging every routine's shaders together; shared shaders (e.g.
170
+ * `"reduction/sum"`, used by two different routines above) collapse harmlessly here since
171
+ * every routine's copy is the same imported string, never independently authored text.
172
+ * @public
173
+ */
174
+ export const shaderSources = Object.assign(
175
+ {},
176
+ ...Object.values(routineShaders),
177
+ );
@@ -0,0 +1,50 @@
1
+ // amax reduction (f64, double-double): collapses 2*WGS (value, index) pairs
2
+ // into one index, using ddGreater/ddEqual instead of plain f32 `>`/`==` (see
3
+ // reduction/argmax.wgsl for the f32 original this mirrors).
4
+ // dispatch: 1 workgroup of WGS threads. partialsValHi/partialsValLo/
5
+ // partialsIdx must have exactly 2*WGS entries each. Concatenated after
6
+ // f64/dekker.wgsl (DD struct), f64/utils/greater.wgsl (ddGreater), and
7
+ // f64/utils/equal.wgsl (ddEqual).
8
+
9
+ @group(0) @binding(0) var<storage, read> partialsValHi: array<f32>;
10
+ @group(0) @binding(1) var<storage, read> partialsValLo: array<f32>;
11
+ @group(0) @binding(2) var<storage, read> partialsIdx: array<u32>;
12
+ @group(0) @binding(3) var<storage, read_write> result: array<u32>;
13
+
14
+ const WGS: u32 = 64;
15
+
16
+ var<workgroup> tile_val: array<DD, 64>;
17
+ var<workgroup> tile_idx: array<u32, 64>;
18
+
19
+ @compute @workgroup_size(64)
20
+ fn reduce_f64(
21
+ @builtin(local_invocation_id) lid: vec3u,
22
+ ) {
23
+ let i = lid.x;
24
+ let a_val = DD(partialsValHi[i], partialsValLo[i]);
25
+ let b_val = DD(partialsValHi[i + WGS], partialsValLo[i + WGS]);
26
+ if (ddGreater(b_val, a_val) ||
27
+ (ddEqual(b_val, a_val) && partialsIdx[i + WGS] < partialsIdx[i])) {
28
+ tile_val[i] = b_val;
29
+ tile_idx[i] = partialsIdx[i + WGS];
30
+ } else {
31
+ tile_val[i] = a_val;
32
+ tile_idx[i] = partialsIdx[i];
33
+ }
34
+ workgroupBarrier();
35
+
36
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
37
+ if (i < s) {
38
+ let c_val = tile_val[i];
39
+ let d_val = tile_val[i + s];
40
+ if (ddGreater(d_val, c_val) ||
41
+ (ddEqual(d_val, c_val) && tile_idx[i + s] < tile_idx[i])) {
42
+ tile_val[i] = d_val;
43
+ tile_idx[i] = tile_idx[i + s];
44
+ }
45
+ }
46
+ workgroupBarrier();
47
+ }
48
+
49
+ if (i == 0u) { result[0] = tile_idx[0]; }
50
+ }
@@ -0,0 +1,65 @@
1
+ // scaledSum reduction: collapses 2*WGS (scale, ssq) partials from
2
+ // snrm2.wgsl into the final norm — sqrt(scale² · ssq) == scale · sqrt(ssq).
3
+ // Mirrors reduction/sum.wgsl's shape exactly, merging via ssqMerge (see
4
+ // snrm2.wgsl for the derivation) instead of plain `+`, and taking the final
5
+ // sqrt here rather than on the CPU — unlike sasum/sdot's plain sum, "sum of
6
+ // squares" isn't a meaningful standalone value to hand back, only
7
+ // scale·sqrt(ssq) is.
8
+ // dispatch: 1 workgroup of WGS threads.
9
+ // partialsScale/partialsSsq must have exactly 2*WGS entries each.
10
+
11
+ @group(0) @binding(0) var<storage, read> partialsScale: array<f32>;
12
+ @group(0) @binding(1) var<storage, read> partialsSsq: array<f32>;
13
+ @group(0) @binding(2) var<storage, read_write> result: array<f32>;
14
+
15
+ const WGS: u32 = 64;
16
+
17
+ // True sum-of-squares represented so far == scale² · ssq — see snrm2.wgsl.
18
+ struct ScaleSsq {
19
+ scale: f32,
20
+ ssq: f32,
21
+ }
22
+
23
+ // Associative merge of two independent (scale, ssq) partials.
24
+ fn ssqMerge(a: ScaleSsq, b: ScaleSsq) -> ScaleSsq {
25
+ if (a.scale == 0.0 && b.scale == 0.0) { return ScaleSsq(0.0, 1.0); }
26
+ if (a.scale >= b.scale) {
27
+ let r = b.scale / a.scale;
28
+ return ScaleSsq(a.scale, a.ssq + b.ssq * r * r);
29
+ }
30
+ let r = a.scale / b.scale;
31
+ return ScaleSsq(b.scale, b.ssq + a.ssq * r * r);
32
+ }
33
+
34
+ var<workgroup> tileScale: array<f32, 64>;
35
+ var<workgroup> tileSsq: array<f32, 64>;
36
+
37
+ @compute @workgroup_size(64)
38
+ fn reduce_scaled(
39
+ @builtin(local_invocation_id) lid: vec3u,
40
+ ) {
41
+ let i = lid.x;
42
+ let merged0 = ssqMerge(
43
+ ScaleSsq(partialsScale[i], partialsSsq[i]),
44
+ ScaleSsq(partialsScale[i + WGS], partialsSsq[i + WGS]),
45
+ );
46
+ tileScale[i] = merged0.scale;
47
+ tileSsq[i] = merged0.ssq;
48
+ workgroupBarrier();
49
+
50
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
51
+ if (i < s) {
52
+ let merged = ssqMerge(
53
+ ScaleSsq(tileScale[i], tileSsq[i]),
54
+ ScaleSsq(tileScale[i + s], tileSsq[i + s]),
55
+ );
56
+ tileScale[i] = merged.scale;
57
+ tileSsq[i] = merged.ssq;
58
+ }
59
+ workgroupBarrier();
60
+ }
61
+
62
+ if (i == 0u) {
63
+ result[0] = tileScale[0] * sqrt(tileSsq[0]);
64
+ }
65
+ }
@@ -2,8 +2,8 @@
2
2
  // into one, using ddAddProtected instead of plain f32 `+` (see
3
3
  // reduction/sum.wgsl for the f32 original this mirrors).
4
4
  // dispatch: 1 workgroup of WGS threads. partialsHi/partialsLo must have
5
- // exactly 2*WGS entries each. Concatenated after f64/dekker.wgsl for
6
- // DD/ddAddProtected (see it for why plain ddAdd isn't safe).
5
+ // exactly 2*WGS entries each. Concatenated after f64/dekker.wgsl (DD struct)
6
+ // and f64/utils/add.wgsl (ddAddProtected — see it for why plain ddAdd isn't safe).
7
7
 
8
8
  @group(0) @binding(0) var<storage, read> partialsHi: array<f32>;
9
9
  @group(0) @binding(1) var<storage, read> partialsLo: array<f32>;