wgblas 1.2.0 → 2.0.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 (70) hide show
  1. package/LICENSE +1 -1
  2. package/README.md +36 -54
  3. package/dist/wgblas.browser.js +779 -34
  4. package/index.d.mts +13 -0
  5. package/index.mjs +8 -0
  6. package/package.json +56 -3
  7. package/src/dasum/dasum.d.mts +2 -2
  8. package/src/dasum/dasum.mjs +6 -4
  9. package/src/idamax/idamax.d.mts +51 -0
  10. package/src/idamax/idamax.mjs +128 -0
  11. package/src/init.mjs +9 -1
  12. package/src/isamax/isamax.d.mts +1 -1
  13. package/src/sasum/sasum.d.mts +1 -1
  14. package/src/saxpy/saxpy.d.mts +1 -1
  15. package/src/scopy/scopy.d.mts +1 -1
  16. package/src/sdot/sdot.d.mts +1 -1
  17. package/src/sgemm/sgemm.d.mts +102 -0
  18. package/src/sgemm/sgemm.mjs +195 -0
  19. package/src/sgemmtr/sgemmtr.d.mts +104 -0
  20. package/src/sgemmtr/sgemmtr.mjs +203 -0
  21. package/src/sgemv/sgemv.d.mts +1 -38
  22. package/src/sgemv/sgemv.mjs +4 -0
  23. package/src/sger/sger.d.mts +1 -34
  24. package/src/sger/sger.mjs +2 -0
  25. package/src/shaders/block_transfer.wgsl +42 -0
  26. package/src/shaders/browser-shaders.mjs +26 -0
  27. package/src/shaders/dasum.wgsl +3 -2
  28. package/src/shaders/f64/dekker.wgsl +4 -85
  29. package/src/shaders/f64/utils/abs.wgsl +10 -0
  30. package/src/shaders/f64/utils/add.wgsl +77 -0
  31. package/src/shaders/f64/utils/equal.wgsl +7 -0
  32. package/src/shaders/f64/utils/greater.wgsl +12 -0
  33. package/src/shaders/f64/utils/multiply.wgsl +81 -0
  34. package/src/shaders/idamax.wgsl +96 -0
  35. package/src/shaders/reduction/argmaxF64.wgsl +50 -0
  36. package/src/shaders/reduction/sumF64.wgsl +2 -2
  37. package/src/shaders/sgemm_large.wgsl +117 -0
  38. package/src/shaders/sgemm_small.wgsl +112 -0
  39. package/src/shaders/sgemmtr_large.wgsl +117 -0
  40. package/src/shaders/sgemmtr_small.wgsl +110 -0
  41. package/src/shaders/symmetrize.wgsl +31 -0
  42. package/src/shaders/triangularize.wgsl +44 -0
  43. package/src/snrm2/snrm2.d.mts +1 -1
  44. package/src/srot/srot.d.mts +1 -1
  45. package/src/srotm/srotm.d.mts +1 -1
  46. package/src/sscal/sscal.d.mts +1 -1
  47. package/src/sswap/sswap.d.mts +1 -1
  48. package/src/ssymm/ssymm.d.mts +103 -0
  49. package/src/ssymm/ssymm.mjs +209 -0
  50. package/src/ssymv/ssymv.d.mts +1 -36
  51. package/src/ssymv/ssymv.mjs +2 -0
  52. package/src/ssyr/ssyr.d.mts +1 -30
  53. package/src/ssyr/ssyr.mjs +2 -0
  54. package/src/ssyr2/ssyr2.d.mts +1 -34
  55. package/src/ssyr2/ssyr2.mjs +2 -0
  56. package/src/ssyr2k/ssyr2k.d.mts +100 -0
  57. package/src/ssyr2k/ssyr2k.mjs +201 -0
  58. package/src/ssyrk/ssyrk.d.mts +90 -0
  59. package/src/ssyrk/ssyrk.mjs +176 -0
  60. package/src/strmm/strmm.d.mts +100 -0
  61. package/src/strmm/strmm.mjs +211 -0
  62. package/src/strmv/strmv.d.mts +1 -36
  63. package/src/strmv/strmv.mjs +2 -0
  64. package/src/strsm/strsm.d.mts +99 -0
  65. package/src/strsm/strsm.mjs +342 -0
  66. package/src/strsv/strsv.d.mts +1 -32
  67. package/src/strsv/strsv.mjs +2 -0
  68. package/src/util/buffer.mjs +4 -2
  69. package/src/util/compute.mjs +6 -3
  70. package/src/util/f64.mjs +3 -3
@@ -1,4 +1,5 @@
1
1
  import argmax from "./reduction/argmax.wgsl";
2
+ import argmaxF64 from "./reduction/argmaxF64.wgsl";
2
3
  import sum from "./reduction/sum.wgsl";
3
4
  import sumF64 from "./reduction/sumF64.wgsl";
4
5
  import sscal from "./sscal.wgsl";
@@ -20,13 +21,26 @@ import ssyr from "./ssyr.wgsl";
20
21
  import ssyr2 from "./ssyr2.wgsl";
21
22
  import f64add from "./f64add.wgsl";
22
23
  import dekker from "./f64/dekker.wgsl";
24
+ import ddAbs from "./f64/utils/abs.wgsl";
25
+ import ddAddUtil from "./f64/utils/add.wgsl";
26
+ import ddGreater from "./f64/utils/greater.wgsl";
27
+ import ddEqual from "./f64/utils/equal.wgsl";
23
28
  import dasum from "./dasum.wgsl";
29
+ import idamax from "./idamax.wgsl";
24
30
  import strsv_invert_block from "./strsv_invert_block.wgsl";
25
31
  import strsv_apply_inverse from "./strsv_apply_inverse.wgsl";
26
32
  import strsv_update from "./strsv_update.wgsl";
33
+ import sgemm_small from "./sgemm_small.wgsl";
34
+ import sgemm_large from "./sgemm_large.wgsl";
35
+ import sgemmtr_small from "./sgemmtr_small.wgsl";
36
+ import sgemmtr_large from "./sgemmtr_large.wgsl";
37
+ import symmetrize from "./symmetrize.wgsl";
38
+ import triangularize from "./triangularize.wgsl";
39
+ import blockTransfer from "./block_transfer.wgsl";
27
40
 
28
41
  export const shaderSources = {
29
42
  "reduction/argmax": argmax,
43
+ "reduction/argmaxF64": argmaxF64,
30
44
  "reduction/sum": sum,
31
45
  "reduction/sumF64": sumF64,
32
46
  sscal,
@@ -48,8 +62,20 @@ export const shaderSources = {
48
62
  ssyr2,
49
63
  f64add,
50
64
  "f64/dekker": dekker,
65
+ "f64/utils/abs": ddAbs,
66
+ "f64/utils/add": ddAddUtil,
67
+ "f64/utils/greater": ddGreater,
68
+ "f64/utils/equal": ddEqual,
51
69
  dasum,
70
+ idamax,
52
71
  strsv_invert_block,
53
72
  strsv_apply_inverse,
54
73
  strsv_update,
74
+ sgemm_small,
75
+ sgemm_large,
76
+ sgemmtr_small,
77
+ sgemmtr_large,
78
+ symmetrize,
79
+ triangularize,
80
+ block_transfer: blockTransfer,
55
81
  };
@@ -1,6 +1,7 @@
1
1
  // dasum: sum(|x[i]|), double-double (Dekker). Same ILP=4 shape as sasum.wgsl;
2
- // see f64/dekker.wgsl for ddAddProtected and why plain ddAdd isn't safe.
3
- // GpuVector input isn't pre-abs'd, so ddAbs() applies unconditionally below.
2
+ // see f64/utils/add.wgsl for ddAddProtected and why plain ddAdd isn't safe.
3
+ // GpuVector input isn't pre-abs'd, so ddAbs() (f64/utils/abs.wgsl) applies
4
+ // unconditionally below.
4
5
 
5
6
  @group(0) @binding(0) var<storage, read> xHi: array<f32>;
6
7
  @group(0) @binding(1) var<storage, read> xLo: array<f32>;
@@ -7,93 +7,12 @@
7
7
  //
8
8
  // No bindings, no entry point — a helper library, concatenated with a
9
9
  // consumer's own bindings/entry point by getPipeline (WGSL has no #include).
10
+ // The DD struct lives here — abs.wgsl/add.wgsl/greater.wgsl/equal.wgsl all
11
+ // use it but don't redefine it (WGSL errors on duplicate struct definitions
12
+ // once concatenated), so any consumer using those must concatenate this
13
+ // file too, first.
10
14
 
11
15
  struct DD {
12
16
  hi: f32,
13
17
  lo: f32,
14
18
  }
15
-
16
- // |a| for a double-double pair. Negation is exact (no rounding), so this is
17
- // just a sign flip on both components — hi alone determines the pair's sign.
18
- fn ddAbs(a: DD) -> DD {
19
- if (a.hi < 0.0) {
20
- return DD(-a.hi, -a.lo);
21
- }
22
- return a;
23
- }
24
-
25
- // ── A real compiler bug — read before touching anything below ──────────────
26
- //
27
- // twoSum/fastTwoSum's error term `e` should be nonzero (that's the point —
28
- // `s` is rounded). A front-end-optimizer bug zeros it anyway, on both NVIDIA
29
- // and Mesa (ANV + llvmpipe), via two different mechanisms needing two fixes:
30
- // bitcast-based subtraction (`fsub`/`negf`, fixes NVIDIA) and materializing
31
- // the sum through workgroup memory + workgroupBarrier() (fixes Mesa). Only
32
- // both together (ddAddProtected) is verified correct everywhere — the plain
33
- // twoSum/fastTwoSum/ddAdd below are reference-only, not safe to use.
34
- fn negf(x: f32) -> f32 {
35
- return bitcast<f32>(bitcast<u32>(x) ^ 0x80000000u);
36
- }
37
- fn fsub(a: f32, b: f32) -> f32 {
38
- return a + negf(b);
39
- }
40
-
41
- // Knuth/Møller's TwoSum: s = fl(a+b), e = exact rounding error, a+b == s+e.
42
- // Works for any a, b. UNPROTECTED — see header above.
43
- fn twoSum(a: f32, b: f32) -> DD {
44
- let s = a + b;
45
- let v = s - a;
46
- let e = (a - (s - v)) + (b - v);
47
- return DD(s, e);
48
- }
49
-
50
- // Dekker's Fast-Two-Sum: same contract, but only correct when |a| >= |b|.
51
- // UNPROTECTED — see header above.
52
- fn fastTwoSum(a: f32, b: f32) -> DD {
53
- let s = a + b;
54
- let e = b - (s - a);
55
- return DD(s, e);
56
- }
57
-
58
- // Double-double addition (Dekker's Add2). UNPROTECTED — see header above.
59
- fn ddAdd(a: DD, b: DD) -> DD {
60
- let s = twoSum(a.hi, b.hi);
61
- let loSum = a.lo + b.lo;
62
- return fastTwoSum(s.hi, s.lo + loSum);
63
- }
64
-
65
- // ── Protected variants — use these ──────────────────────────────────────────
66
- //
67
- // Bitcast subtraction + workgroup-barrier materialization, verified correct
68
- // on all three backends tested. Costs a real barrier: fine for O(1)-per-
69
- // thread or O(log n) reduction use, not a long per-element loop. A
70
- // workgroupBarrier() requires uniform control flow, so:
71
- // - `threadSlot` must be unique per concurrent caller (e.g. local_invocation_index).
72
- // - Every thread in the workgroup must call this the same number of times
73
- // — including ones whose result gets discarded. Compute unconditionally;
74
- // only the write-back should be conditional.
75
- var<workgroup> dekkerScratch: array<f32, 64>;
76
-
77
- fn twoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
78
- dekkerScratch[threadSlot] = a + b;
79
- workgroupBarrier();
80
- let s = dekkerScratch[threadSlot];
81
- let v = fsub(s, a);
82
- let e = fsub(a, fsub(s, v)) + fsub(b, v);
83
- return DD(s, e);
84
- }
85
-
86
- fn fastTwoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
87
- dekkerScratch[threadSlot] = a + b;
88
- workgroupBarrier();
89
- let s = dekkerScratch[threadSlot];
90
- let e = fsub(b, fsub(s, a));
91
- return DD(s, e);
92
- }
93
-
94
- // Protected double-double addition — same contract as ddAdd, but exact.
95
- fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
96
- let s = twoSumProtected(a.hi, b.hi, threadSlot);
97
- let loSum = a.lo + b.lo;
98
- return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
99
- }
@@ -0,0 +1,10 @@
1
+ // Requires f64/dekker.wgsl concatenated first for the DD struct.
2
+
3
+ // |a| for a double-double pair. Negation is exact (no rounding), so this is
4
+ // just a sign flip on both components — hi alone determines the pair's sign.
5
+ fn ddAbs(a: DD) -> DD {
6
+ if (a.hi < 0.0) {
7
+ return DD(-a.hi, -a.lo);
8
+ }
9
+ return a;
10
+ }
@@ -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
+ }
@@ -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
+ }
@@ -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>;
@@ -0,0 +1,117 @@
1
+ // sgemm_large: C = alpha * op(A) * op(B) + beta * C — large-tile half of
2
+ // the two-tier autotuned dispatch (see sgemm.mjs and sgemm_small.wgsl).
3
+ // BM=BN=64, BK=8, TM=8, TN=4 (128 threads/workgroup) — the kernel 9
4
+ // autotuning winner (temp/autotune_sweep.mjs, temp/gen_sweep_kernel.mjs,
5
+ // swept BM/BN/BK/TM/TN and warp-tiled variants), +69% over the old BM=32
6
+ // single-tier baseline at n=512, +84% at n=1024. But BM=64 loses to BM=32
7
+ // below a 6x6=36 workgroup grid (not enough workgroups to fill the GPU at
8
+ // that tile size), hence the two-tier split rather than one global config.
9
+ // Neither vectorized loads (kernel 6) nor warp-tiling (kernel 10) beat this
10
+ // at the sizes tried, including warp-tiled variants in the same sweep at
11
+ // BM=64/128.
12
+
13
+ const BM: u32 = 64u;
14
+ const BN: u32 = 64u;
15
+ const BK: u32 = 8u;
16
+ const TM: u32 = 8u;
17
+ const TN: u32 = 4u;
18
+ const THREADS_X: u32 = BN / TN;
19
+ const THREADS_Y: u32 = BM / TM;
20
+ const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 128
21
+ const STRIDE_A: u32 = NUM_THREADS / BK; // rows of As covered per load-loop step
22
+ const STRIDE_B: u32 = NUM_THREADS / BN; // rows of Bs covered per load-loop step
23
+
24
+ @group(0) @binding(0) var<storage, read> A: array<f32>;
25
+ @group(0) @binding(1) var<storage, read> B: array<f32>;
26
+ @group(0) @binding(2) var<storage, read_write> C: array<f32>;
27
+
28
+ struct Params {
29
+ m: u32,
30
+ n: u32,
31
+ k: u32,
32
+ alpha: f32,
33
+ beta: f32,
34
+ lda: u32,
35
+ ldb: u32,
36
+ ldc: u32,
37
+ transA: u32, // 0 = no-transpose, 1 = transpose
38
+ transB: u32,
39
+ }
40
+
41
+ @group(0) @binding(3) var<uniform> params: Params;
42
+
43
+ var<workgroup> As: array<f32, BM * BK>;
44
+ var<workgroup> Bs: array<f32, BK * BN>;
45
+
46
+ @compute @workgroup_size(THREADS_X, THREADS_Y)
47
+ fn main(
48
+ @builtin(workgroup_id) wid: vec3u,
49
+ @builtin(local_invocation_id) lid: vec3u,
50
+ @builtin(local_invocation_index) tid: u32,
51
+ ) {
52
+ let blockRow = wid.y * BM;
53
+ let blockCol = wid.x * BN;
54
+ let threadCol = lid.x;
55
+ let threadRow = lid.y;
56
+
57
+ // Load indices, independent of the compute thread shape — a loop since
58
+ // NUM_THREADS doesn't match the tile size 1:1 at this config.
59
+ let innerRowA = tid / BK;
60
+ let innerColA = tid % BK;
61
+ let innerRowB = tid / BN;
62
+ let innerColB = tid % BN;
63
+
64
+ var threadResults: array<f32, TM * TN>;
65
+ for (var i = 0u; i < TM * TN; i++) {
66
+ threadResults[i] = 0.0;
67
+ }
68
+ var regM: array<f32, TM>;
69
+ var regN: array<f32, TN>;
70
+
71
+ let numTiles = (params.k + BK - 1u) / BK;
72
+ for (var t = 0u; t < numTiles; t++) {
73
+ for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
74
+ let gRowA = blockRow + innerRowA + loadOffset;
75
+ let gColA = t * BK + innerColA;
76
+ let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
77
+ As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
78
+ }
79
+ for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
80
+ let gRowB = t * BK + innerRowB + loadOffset;
81
+ let gColB = blockCol + innerColB;
82
+ let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
83
+ Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
84
+ }
85
+
86
+ workgroupBarrier();
87
+
88
+ for (var dotIdx = 0u; dotIdx < BK; dotIdx++) {
89
+ for (var i = 0u; i < TM; i++) {
90
+ regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
91
+ }
92
+ for (var i = 0u; i < TN; i++) {
93
+ regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
94
+ }
95
+ for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
96
+ for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
97
+ threadResults[resIdxM * TN + resIdxN] += regM[resIdxM] * regN[resIdxN];
98
+ }
99
+ }
100
+ }
101
+
102
+ workgroupBarrier();
103
+ }
104
+
105
+ for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
106
+ let row = blockRow + threadRow * TM + resIdxM;
107
+ if (row < params.m) {
108
+ for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
109
+ let col = blockCol + threadCol * TN + resIdxN;
110
+ if (col < params.n) {
111
+ let cIdx = row * params.ldc + col;
112
+ C[cIdx] = params.alpha * threadResults[resIdxM * TN + resIdxN] + params.beta * C[cIdx];
113
+ }
114
+ }
115
+ }
116
+ }
117
+ }