wgblas 2.1.0 → 2.2.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 (99) hide show
  1. package/README.md +2 -0
  2. package/dist/wgblas.browser.js +1007 -43
  3. package/index.d.mts +26 -53
  4. package/index.mjs +11 -0
  5. package/package.json +132 -63
  6. package/src/classes/Complex32.d.mts +43 -0
  7. package/src/classes/Complex32.mjs +82 -0
  8. package/src/classes/Complex64.d.mts +44 -0
  9. package/src/classes/Complex64.mjs +76 -0
  10. package/src/classes/GpuMatrix.d.mts +41 -41
  11. package/src/classes/GpuMatrix.mjs +112 -10
  12. package/src/classes/GpuVector.d.mts +36 -40
  13. package/src/classes/GpuVector.mjs +39 -2
  14. package/src/cscal/cscal.d.mts +47 -0
  15. package/src/cscal/cscal.mjs +98 -0
  16. package/src/dasum/dasum.mjs +31 -15
  17. package/src/daxpy/daxpy.d.mts +56 -0
  18. package/src/daxpy/daxpy.mjs +150 -0
  19. package/src/dcopy/dcopy.d.mts +52 -0
  20. package/src/dcopy/dcopy.mjs +140 -0
  21. package/src/ddot/ddot.d.mts +62 -0
  22. package/src/ddot/ddot.mjs +184 -0
  23. package/src/dnrm2/dnrm2.d.mts +50 -0
  24. package/src/dnrm2/dnrm2.mjs +189 -0
  25. package/src/drot/drot.d.mts +67 -0
  26. package/src/drot/drot.mjs +170 -0
  27. package/src/drotm/drotm.d.mts +67 -0
  28. package/src/drotm/drotm.mjs +171 -0
  29. package/src/dscal/dscal.d.mts +52 -0
  30. package/src/dscal/dscal.mjs +119 -0
  31. package/src/dswap/dswap.d.mts +57 -0
  32. package/src/dswap/dswap.mjs +155 -0
  33. package/src/idamax/idamax.mjs +49 -19
  34. package/src/init.mjs +6 -3
  35. package/src/isamax/isamax.mjs +17 -14
  36. package/src/random/random.d.mts +37 -40
  37. package/src/random/random.mjs +39 -7
  38. package/src/sasum/sasum.mjs +13 -11
  39. package/src/saxpy/saxpy.mjs +9 -8
  40. package/src/scopy/scopy.mjs +8 -6
  41. package/src/sdot/sdot.mjs +13 -11
  42. package/src/sgemm/sgemm.d.mts +2 -2
  43. package/src/sgemm/sgemm.mjs +91 -35
  44. package/src/sgemmtr/sgemmtr.d.mts +3 -2
  45. package/src/sgemmtr/sgemmtr.mjs +92 -35
  46. package/src/sgemv/sgemv.d.mts +2 -2
  47. package/src/sgemv/sgemv.mjs +41 -25
  48. package/src/sger/sger.d.mts +2 -2
  49. package/src/sger/sger.mjs +38 -16
  50. package/src/shaders/__test_pipeline_a.wgsl +3 -0
  51. package/src/shaders/__test_pipeline_b.wgsl +2 -0
  52. package/src/shaders/cscal.wgsl +33 -0
  53. package/src/shaders/daxpy.wgsl +66 -0
  54. package/src/shaders/dcopy.wgsl +34 -0
  55. package/src/shaders/ddot.wgsl +106 -0
  56. package/src/shaders/dnrm2.wgsl +167 -0
  57. package/src/shaders/drot.wgsl +81 -0
  58. package/src/shaders/drotm.wgsl +99 -0
  59. package/src/shaders/dscal.wgsl +60 -0
  60. package/src/shaders/dswap.wgsl +38 -0
  61. package/src/shaders/f64/utils/add.wgsl +6 -0
  62. package/src/shaders/f64/utils/divide.wgsl +45 -0
  63. package/src/shaders/f64/utils/multiply.wgsl +19 -10
  64. package/src/shaders/f64/utils/sqrt.wgsl +45 -0
  65. package/src/shaders/index.mjs +69 -0
  66. package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
  67. package/src/snrm2/snrm2.mjs +20 -14
  68. package/src/srot/srot.mjs +10 -7
  69. package/src/srotm/srotm.mjs +9 -12
  70. package/src/sscal/sscal.mjs +7 -7
  71. package/src/sswap/sswap.mjs +14 -8
  72. package/src/ssymm/ssymm.d.mts +5 -4
  73. package/src/ssymm/ssymm.mjs +142 -55
  74. package/src/ssymv/ssymv.d.mts +2 -2
  75. package/src/ssymv/ssymv.mjs +42 -23
  76. package/src/ssyr/ssyr.d.mts +2 -2
  77. package/src/ssyr/ssyr.mjs +34 -15
  78. package/src/ssyr2/ssyr2.d.mts +2 -2
  79. package/src/ssyr2/ssyr2.mjs +43 -18
  80. package/src/ssyr2k/ssyr2k.d.mts +3 -2
  81. package/src/ssyr2k/ssyr2k.mjs +132 -55
  82. package/src/ssyrk/ssyrk.d.mts +3 -2
  83. package/src/ssyrk/ssyrk.mjs +84 -33
  84. package/src/strmm/strmm.d.mts +5 -4
  85. package/src/strmm/strmm.mjs +153 -54
  86. package/src/strmv/strmv.d.mts +2 -2
  87. package/src/strmv/strmv.mjs +37 -17
  88. package/src/strsm/strsm.d.mts +6 -4
  89. package/src/strsm/strsm.mjs +418 -172
  90. package/src/strsv/strsv.d.mts +5 -3
  91. package/src/strsv/strsv.mjs +82 -31
  92. package/src/util/benchmark.mjs +5 -3
  93. package/src/util/buffer.mjs +33 -12
  94. package/src/util/complex.mjs +87 -0
  95. package/src/util/compute.mjs +14 -8
  96. package/src/util/device.mjs +18 -3
  97. package/src/util/pipeline.mjs +40 -5
  98. package/src/util/workgroup.mjs +23 -6
  99. package/src/shaders/f64add.wgsl +0 -281
@@ -0,0 +1,99 @@
1
+ // drotm: applies a modified Givens rotation H to vectors x and y — double-
2
+ // double (Dekker) f64 emulation of srotm. paramHi/paramLo[0] = flag: -1
3
+ // (full H), 0 (unit diagonal), 1 (unit off-diagonal). param = [ flag, h11,
4
+ // h21, h12, h22 ], each entry an f32 (hi, lo) pair.
5
+ // flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
6
+ //
7
+ // h11/h12/h21/h22 are resolved once, outside the loop, from the (uniform
8
+ // across every thread) flag — same shape as srot's c/s, so no barrier is
9
+ // needed for that selection itself. Each element then costs four
10
+ // ddMulProtected + two ddAddProtected, same as drot.
11
+
12
+ @group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
13
+ @group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
14
+ @group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
15
+ @group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
16
+ @group(0) @binding(4) var<storage, read> paramHi: array<f32>;
17
+ @group(0) @binding(5) var<storage, read> paramLo: array<f32>;
18
+ @group(0) @binding(6) var<uniform> params: Params;
19
+
20
+ struct Params {
21
+ n: u32,
22
+ x_inc: u32,
23
+ y_inc: u32,
24
+ }
25
+
26
+ const WGS: u32 = 64;
27
+ const ONE: DD = DD(1.0, 0.0);
28
+ const NEG_ONE: DD = DD(-1.0, 0.0);
29
+
30
+ @compute @workgroup_size(64)
31
+ fn drotm_main(
32
+ @builtin(global_invocation_id) gid: vec3u,
33
+ @builtin(local_invocation_id) lid: vec3u,
34
+ @builtin(workgroup_id) wgid: vec3u,
35
+ @builtin(num_workgroups) num_wg: vec3u,
36
+ ) {
37
+ let flag = paramHi[0]; // exact small integer (-1, 0, or 1) — lo is always 0
38
+
39
+ var h11: DD; var h12: DD;
40
+ var h21: DD; var h22: DD;
41
+
42
+ if (flag == -1.0) {
43
+ // full 2x2 matrix
44
+ h11 = DD(paramHi[1], paramLo[1]); h21 = DD(paramHi[2], paramLo[2]);
45
+ h12 = DD(paramHi[3], paramLo[3]); h22 = DD(paramHi[4], paramLo[4]);
46
+ } else if (flag == 0.0) {
47
+ // diagonal fixed at 1
48
+ h11 = ONE; h21 = DD(paramHi[2], paramLo[2]);
49
+ h12 = DD(paramHi[3], paramLo[3]); h22 = ONE;
50
+ } else {
51
+ // flag == 1.0: off-diagonal fixed at +1 / -1
52
+ h11 = DD(paramHi[1], paramLo[1]); h21 = NEG_ONE;
53
+ h12 = ONE; h22 = DD(paramHi[4], paramLo[4]);
54
+ }
55
+
56
+ let stride = num_wg.x * WGS;
57
+
58
+ let n_floor = (params.n / stride) * stride;
59
+ let mainIters = n_floor / stride;
60
+ for (var iter = 0u; iter < mainIters; iter++) {
61
+ let id = gid.x + iter * stride;
62
+ let ix = id * params.x_inc;
63
+ let iy = id * params.y_inc;
64
+ let xi = DD(xHi[ix], xLo[ix]);
65
+ let yi = DD(yHi[iy], yLo[iy]);
66
+ let xNew = ddAddProtected(ddMulProtected(h11, xi, lid.x), ddMulProtected(h12, yi, lid.x), lid.x);
67
+ let yNew = ddAddProtected(ddMulProtected(h21, xi, lid.x), ddMulProtected(h22, yi, lid.x), lid.x);
68
+ xHi[ix] = xNew.hi;
69
+ xLo[ix] = xNew.lo;
70
+ yHi[iy] = yNew.hi;
71
+ yLo[iy] = yNew.lo;
72
+ }
73
+
74
+ // Tail is ragged (0 or 1 extra per thread) — pad to this workgroup's worst
75
+ // case so every thread in the workgroup still calls ddMulProtected/
76
+ // ddAddProtected the same number of times (their barriers need that),
77
+ // masking only the write.
78
+ let wgBaseGid = wgid.x * WGS;
79
+ var tailIters = 0u;
80
+ if (n_floor + wgBaseGid < params.n) {
81
+ tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
82
+ }
83
+ for (var iter = 0u; iter < tailIters; iter++) {
84
+ let id = n_floor + gid.x + iter * stride;
85
+ let valid = id < params.n;
86
+ let ix = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
87
+ let iy = select(0u, id * params.y_inc, valid);
88
+ let xi = DD(xHi[ix], xLo[ix]);
89
+ let yi = DD(yHi[iy], yLo[iy]);
90
+ let xNew = ddAddProtected(ddMulProtected(h11, xi, lid.x), ddMulProtected(h12, yi, lid.x), lid.x);
91
+ let yNew = ddAddProtected(ddMulProtected(h21, xi, lid.x), ddMulProtected(h22, yi, lid.x), lid.x);
92
+ if (valid) {
93
+ xHi[ix] = xNew.hi;
94
+ xLo[ix] = xNew.lo;
95
+ yHi[iy] = yNew.hi;
96
+ yLo[iy] = yNew.lo;
97
+ }
98
+ }
99
+ }
@@ -0,0 +1,60 @@
1
+ // dscal: x := alpha * x, double-double (Dekker) f64 emulation of sscal.
2
+ // alpha and x are each an f32 (hi, lo) pair. See f64/utils/multiply.wgsl for
3
+ // ddMulProtected and why plain ddMulRaw isn't safe without a renormalizing
4
+ // barrier — that barrier needs a provably uniform loop trip count across
5
+ // every thread in the workgroup, so (like dasum.wgsl's reduction loop) this
6
+ // splits into a uniform main pass plus a ragged, select-masked tail rather
7
+ // than a plain `id < params.n` grid-stride loop.
8
+
9
+ @group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
10
+ @group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
11
+ @group(0) @binding(2) var<uniform> params: Params;
12
+
13
+ struct Params {
14
+ n: u32,
15
+ alphaHi: f32,
16
+ alphaLo: f32,
17
+ x_inc: u32,
18
+ }
19
+
20
+ const WGS: u32 = 64;
21
+
22
+ @compute @workgroup_size(64)
23
+ fn dscal_main(
24
+ @builtin(global_invocation_id) gid: vec3u,
25
+ @builtin(local_invocation_id) lid: vec3u,
26
+ @builtin(workgroup_id) wgid: vec3u,
27
+ @builtin(num_workgroups) num_wg: vec3u,
28
+ ) {
29
+ let alpha = DD(params.alphaHi, params.alphaLo);
30
+ let stride = num_wg.x * WGS;
31
+
32
+ let n_floor = (params.n / stride) * stride;
33
+ let mainIters = n_floor / stride;
34
+ for (var iter = 0u; iter < mainIters; iter++) {
35
+ let id = gid.x + iter * stride;
36
+ let i = id * params.x_inc;
37
+ let result = ddMulProtected(alpha, DD(xHi[i], xLo[i]), lid.x);
38
+ xHi[i] = result.hi;
39
+ xLo[i] = result.lo;
40
+ }
41
+
42
+ // Tail is ragged (0 or 1 extra per thread) — pad to this workgroup's worst
43
+ // case so every thread in the workgroup still calls ddMulProtected the
44
+ // same number of times (its barrier needs that), masking only the write.
45
+ let wgBaseGid = wgid.x * WGS;
46
+ var tailIters = 0u;
47
+ if (n_floor + wgBaseGid < params.n) {
48
+ tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
49
+ }
50
+ for (var iter = 0u; iter < tailIters; iter++) {
51
+ let id = n_floor + gid.x + iter * stride;
52
+ let valid = id < params.n;
53
+ let i = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
54
+ let result = ddMulProtected(alpha, DD(xHi[i], xLo[i]), lid.x);
55
+ if (valid) {
56
+ xHi[i] = result.hi;
57
+ xLo[i] = result.lo;
58
+ }
59
+ }
60
+ }
@@ -0,0 +1,38 @@
1
+ // dswap: x <-> y, double-double (Dekker) f64 emulation of sswap. A swap is
2
+ // pure data movement — hi and lo are exchanged verbatim, with no arithmetic
3
+ // at all — so (unlike dscal/daxpy/ddot) this needs no
4
+ // ddMulProtected/ddAddProtected renormalizing barrier, and so no
5
+ // ragged-tail-with-barrier split; a plain grid-stride loop is safe, same
6
+ // shape as sswap.wgsl itself.
7
+
8
+ @group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
9
+ @group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
10
+ @group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
11
+ @group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
12
+
13
+ struct Params {
14
+ n: u32,
15
+ x_inc: u32,
16
+ y_inc: u32,
17
+ }
18
+
19
+ @group(0) @binding(4) var<uniform> params: Params;
20
+
21
+ const WGS: u32 = 64;
22
+
23
+ @compute @workgroup_size(64)
24
+ fn main(
25
+ @builtin(global_invocation_id) gid: vec3u,
26
+ @builtin(num_workgroups) num_wg: vec3u,
27
+ ) {
28
+ for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
29
+ let ix = id * params.x_inc;
30
+ let iy = id * params.y_inc;
31
+ let tempHi = xHi[ix];
32
+ let tempLo = xLo[ix];
33
+ xHi[ix] = yHi[iy];
34
+ xLo[ix] = yLo[iy];
35
+ yHi[iy] = tempHi;
36
+ yLo[iy] = tempLo;
37
+ }
38
+ }
@@ -75,3 +75,9 @@ fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
75
75
  let loSum = a.lo + b.lo;
76
76
  return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
77
77
  }
78
+
79
+ // Double-double subtraction — a - b, via exact negation (a sign-bit flip,
80
+ // no rounding) then ddAddProtected. Same protection contract.
81
+ fn ddSubProtected(a: DD, b: DD, threadSlot: u32) -> DD {
82
+ return ddAddProtected(a, DD(negf(b.hi), negf(b.lo)), threadSlot);
83
+ }
@@ -0,0 +1,45 @@
1
+ // Requires f64/dekker.wgsl concatenated first for the DD struct, and
2
+ // f64/utils/add.wgsl (ddSubProtected/ddAddProtected/negf) and
3
+ // f64/utils/multiply.wgsl (ddMulProtected).
4
+ //
5
+ // Verified empirically on real hardware (NVIDIA GTX 1650 + Intel Mesa
6
+ // integrated GPU, ~8000 random trials each plus explicit edge cases): no
7
+ // compiler-reassociation-style corruption of the kind that broke twoSum/
8
+ // twoProd (see add.wgsl's/multiply.wgsl's own headers) — every failure
9
+ // found was a genuine algorithm gap, not a driver miscompile, and both are
10
+ // fixed below (the b.hi==0.0 guard). One real, expected-shape difference
11
+ // from every other protected op here: the low-power backend's observed
12
+ // forward-error factor for this op specifically runs noticeably higher
13
+ // (~7-14000x eps, vs ~3-9x on high-performance) than ddSqrtProtected's
14
+ // (~3-4x on both) — division inherently amplifies input imprecision more
15
+ // than a sum/product does, so a real routine built on this needs its own
16
+ // backend-calibrated threshold, same as every other f64 arithmetic routine
17
+ // in this codebase (see e.g. tests/drot/src/test.drot.js's THRESHOLDS).
18
+ //
19
+ // One Newton-style long-division refinement (Bailey/QD-style): q1 = a.hi /
20
+ // b.hi is a plain f32 quotient, accurate to ~24 bits. Computing the residual
21
+ // a - q1*b in DD arithmetic (not f32) recovers the bits q1 lost, and a
22
+ // second plain division of that residual resolves them into a correction
23
+ // term — combining q1 + q2 gives roughly double a lone f32 divide's
24
+ // precision, matching this scheme's ~48-bit double-double target (already
25
+ // short of real f64's 52 bits, so a second refinement step would chase
26
+ // precision this representation has no room for).
27
+ // b.hi == 0.0 makes q1 = a.hi/0.0 already the IEEE-754-correct answer
28
+ // (±Infinity, or NaN for 0/0) via plain float division, but the refinement
29
+ // below would corrupt it: p1 = q1*b multiplies that Infinity by a zero
30
+ // divisor, and Infinity*0 is NaN by definition, poisoning everything after.
31
+ // Substituting a safe non-zero denominator via select() — rather than
32
+ // branching/returning early — keeps every thread calling ddMulProtected/
33
+ // ddSubProtected/ddAddProtected unconditionally, which their internal
34
+ // workgroupBarrier() requires; only the final result is selected between
35
+ // the refined value and q1's own already-correct answer.
36
+ fn ddDivProtected(a: DD, b: DD, threadSlot: u32) -> DD {
37
+ let bIsZero = b.hi == 0.0;
38
+ let q1 = a.hi / b.hi;
39
+ let safeB = DD(select(b.hi, 1.0, bIsZero), select(b.lo, 0.0, bIsZero));
40
+ let p1 = ddMulProtected(DD(q1, 0.0), safeB, threadSlot);
41
+ let r1 = ddSubProtected(a, p1, threadSlot);
42
+ let q2 = r1.hi / safeB.hi;
43
+ let refined = ddAddProtected(DD(q1, 0.0), DD(q2, 0.0), threadSlot);
44
+ return DD(select(refined.hi, q1, bIsZero), select(refined.lo, 0.0, bIsZero));
45
+ }
@@ -15,7 +15,11 @@ const SPLIT_CONST: f32 = 4097.0;
15
15
 
16
16
  fn bitSplit(a: f32) -> DD {
17
17
  let bits = bitcast<u32>(a);
18
- let hiBits = bits & 0xFFFFF800u; // keep sign+exponent+top 12 mantissa bits
18
+ // Top 11 mantissa bits, so hi carries 12 significant bits with the implicit
19
+ // leading 1 — the halves are multiplied pairwise and f32 holds 24, so a
20
+ // wider split rounds those products and the "exact" error term goes wrong.
21
+ // Matches SPLIT_CONST = 2^12+1 used by the Veltkamp path below.
22
+ let hiBits = bits & 0xFFFFF000u;
19
23
  let hi = bitcast<f32>(hiBits);
20
24
  let lo = fsub(a, hi); // exact by Sterbenz's lemma (hi, a share an exponent, are close)
21
25
  return DD(hi, lo);
@@ -63,19 +67,24 @@ fn twoProdFma(a: f32, b: f32) -> DD {
63
67
 
64
68
  // DD × DD product (Dekker/Bailey): twoProdBit(a.hi, b.hi) already captures
65
69
  // 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 {
70
+ // ~48-bit floor anyway, so folding them in with plain f32 loses nothing.
71
+ //
72
+ // Another real compiler bug, distinct from add.wgsl's twoSum one — confirmed
73
+ // on Intel Mesa ANV: when p.lo feeds straight into `crossAndLo` unobserved,
74
+ // the compiler folds it away entirely. Materializing p.lo itself through
75
+ // workgroup memory + workgroupBarrier() (like twoSumProtected does for its
76
+ // sum) is what fixes it, so ddMulRaw now takes threadSlot and always pays
77
+ // that barrier — no longer a plain unprotected batchable helper.
78
+ fn ddMulRaw(a: DD, b: DD, threadSlot: u32) -> DD {
73
79
  let p = twoProdBit(a.hi, b.hi);
74
- let crossAndLo = p.lo + (a.hi * b.lo + a.lo * b.hi);
80
+ dekkerScratch[threadSlot] = p.lo;
81
+ workgroupBarrier();
82
+ let pLo = dekkerScratch[threadSlot];
83
+ let crossAndLo = pLo + (a.hi * b.lo + a.lo * b.hi);
75
84
  return DD(p.hi, crossAndLo);
76
85
  }
77
86
 
78
87
  fn ddMulProtected(a: DD, b: DD, threadSlot: u32) -> DD {
79
- let raw = ddMulRaw(a, b);
88
+ let raw = ddMulRaw(a, b, threadSlot);
80
89
  return fastTwoSumProtected(raw.hi, raw.lo, threadSlot);
81
90
  }
@@ -0,0 +1,45 @@
1
+ // Requires f64/dekker.wgsl concatenated first for the DD struct, and
2
+ // f64/utils/add.wgsl (ddSubProtected/ddAddProtected) and
3
+ // f64/utils/multiply.wgsl (twoProdBit — squaring a plain f32 needs no
4
+ // barrier, per multiply.wgsl's own note that twoProdBit is universally safe
5
+ // unprotected).
6
+ //
7
+ // Verified empirically on real hardware (NVIDIA GTX 1650 + Intel Mesa
8
+ // integrated GPU, ~8000 random trials each plus explicit edge cases): no
9
+ // compiler-reassociation-style corruption of the kind that broke twoSum/
10
+ // twoProd (see add.wgsl's/multiply.wgsl's own headers) — every failure
11
+ // found was a genuine algorithm gap (the a.hi==0.0 case below), not a
12
+ // driver miscompile, and is fixed. Observed forward-error factor against a
13
+ // true f64 reference stayed ~3-4x eps on both backends across every random
14
+ // trial — noticeably tighter than ddDivProtected's own low-power spread
15
+ // (see divide.wgsl's header), since sqrt has no denominator to be unlucky
16
+ // about.
17
+ //
18
+ // One Newton refinement step (the classic extended-precision sqrt trick):
19
+ // x0 = sqrt(a.hi) is a plain f32 approximation; the residual a - x0^2,
20
+ // computed in DD arithmetic, recovers what x0 lost, and linearizing sqrt
21
+ // around x0 (dividing that residual by 2*x0) gives a correction term
22
+ // roughly doubling the precision — same ~48-bit target as ddDivProtected,
23
+ // so one step is enough.
24
+ //
25
+ // Undefined for a.hi < 0.0, same as plain sqrt() — callers must guard
26
+ // themselves; this never checks.
27
+ //
28
+ // a.hi == 0.0 (a genuinely zero input, not an underflowed one — zero is
29
+ // exactly representable in f32, unlike this scheme's real range limits;
30
+ // see splitDoubleDouble's own doc comment) makes x0 = sqrt(0) = 0, and the
31
+ // correction step would divide by 2*x0 = 0. Substituting a safe non-zero
32
+ // denominator via select() — rather than branching/returning early — keeps
33
+ // every thread calling ddSubProtected/ddAddProtected unconditionally, which
34
+ // their internal workgroupBarrier() requires; only the final result is
35
+ // selected between the computed value and the exact DD(0,0) answer.
36
+ fn ddSqrtProtected(a: DD, threadSlot: u32) -> DD {
37
+ let isZero = a.hi == 0.0;
38
+ let x0 = sqrt(a.hi);
39
+ let x0sq = twoProdBit(x0, x0);
40
+ let r = ddSubProtected(a, x0sq, threadSlot);
41
+ let safeDenom = select(2.0 * x0, 1.0, isZero);
42
+ let correction = r.hi / safeDenom;
43
+ let result = ddAddProtected(DD(x0, 0.0), DD(correction, 0.0), threadSlot);
44
+ return DD(select(result.hi, 0.0, isZero), select(result.lo, 0.0, isZero));
45
+ }
@@ -53,15 +53,24 @@ export const routineShaders = {};
53
53
  import sscal from "./sscal.wgsl";
54
54
  routineShaders.sscal = { sscal };
55
55
 
56
+ import cscal from "./cscal.wgsl";
57
+ routineShaders.cscal = { cscal };
58
+
56
59
  import sswap from "./sswap.wgsl";
57
60
  routineShaders.sswap = { sswap };
58
61
 
62
+ import dswap from "./dswap.wgsl"; // f64 sibling of sswap — pure data movement, no arithmetic, no reduction/barrier shader needed
63
+ routineShaders.dswap = { dswap };
64
+
59
65
  import saxpy from "./saxpy.wgsl";
60
66
  routineShaders.saxpy = { saxpy };
61
67
 
62
68
  import scopy from "./scopy.wgsl";
63
69
  routineShaders.scopy = { scopy };
64
70
 
71
+ import dcopy from "./dcopy.wgsl"; // f64 sibling of scopy — pure data movement, no arithmetic, no reduction/barrier shader needed
72
+ routineShaders.dcopy = { dcopy };
73
+
65
74
  import sdot from "./sdot.wgsl";
66
75
  import sum from "./reduction/sum.wgsl";
67
76
  routineShaders.sdot = { sdot, "reduction/sum": sum };
@@ -90,6 +99,34 @@ routineShaders.dasum = {
90
99
  "reduction/sumF64": sumF64,
91
100
  };
92
101
 
102
+ import ddMulUtil from "./f64/utils/multiply.wgsl";
103
+ import ddot from "./ddot.wgsl";
104
+ // multiply.wgsl needs dekker's DD struct and add.wgsl's fsub/negf and
105
+ // fastTwoSumProtected, so those two precede it here.
106
+ routineShaders.ddot = {
107
+ "f64/dekker": dekker,
108
+ "f64/utils/add": ddAddUtil,
109
+ "f64/utils/multiply": ddMulUtil,
110
+ ddot,
111
+ "reduction/sumF64": sumF64,
112
+ };
113
+
114
+ import dscal from "./dscal.wgsl"; // f64 sibling of sscal — no reduction shader needed, unlike dasum/ddot
115
+ routineShaders.dscal = {
116
+ "f64/dekker": dekker,
117
+ "f64/utils/add": ddAddUtil,
118
+ "f64/utils/multiply": ddMulUtil,
119
+ dscal,
120
+ };
121
+
122
+ import daxpy from "./daxpy.wgsl"; // f64 sibling of saxpy — one ddMulProtected + one ddAddProtected per element, no reduction shader needed
123
+ routineShaders.daxpy = {
124
+ "f64/dekker": dekker,
125
+ "f64/utils/add": ddAddUtil,
126
+ "f64/utils/multiply": ddMulUtil,
127
+ daxpy,
128
+ };
129
+
93
130
  import ddGreater from "./f64/utils/greater.wgsl";
94
131
  import ddEqual from "./f64/utils/equal.wgsl";
95
132
  import idamax from "./idamax.wgsl";
@@ -106,9 +143,41 @@ routineShaders.idamax = {
106
143
  import srot from "./srot.wgsl";
107
144
  routineShaders.srot = { srot };
108
145
 
146
+ import drot from "./drot.wgsl"; // f64 sibling of srot — four ddMulProtected + two ddAddProtected per element, no reduction shader needed
147
+ routineShaders.drot = {
148
+ "f64/dekker": dekker,
149
+ "f64/utils/add": ddAddUtil,
150
+ "f64/utils/multiply": ddMulUtil,
151
+ drot,
152
+ };
153
+
109
154
  import srotm from "./srotm.wgsl";
110
155
  routineShaders.srotm = { srotm };
111
156
 
157
+ import drotm from "./drotm.wgsl"; // f64 sibling of srotm — four ddMulProtected + two ddAddProtected per element, no reduction shader needed
158
+ routineShaders.drotm = {
159
+ "f64/dekker": dekker,
160
+ "f64/utils/add": ddAddUtil,
161
+ "f64/utils/multiply": ddMulUtil,
162
+ drotm,
163
+ };
164
+
165
+ import ddDivUtil from "./f64/utils/divide.wgsl";
166
+ import ddSqrtUtil from "./f64/utils/sqrt.wgsl";
167
+ import dnrm2 from "./dnrm2.wgsl"; // f64 sibling of snrm2 — scaled accumulation (Blue's algorithm) ported to double-double, branch-free (select()) since ddDivProtected/ddMulProtected/ddAddProtected's barriers need every thread to take the same path
168
+ import scaledSumF64 from "./reduction/scaledSumF64.wgsl";
169
+ routineShaders.dnrm2 = {
170
+ "f64/dekker": dekker,
171
+ "f64/utils/abs": ddAbs,
172
+ "f64/utils/greater": ddGreater,
173
+ "f64/utils/add": ddAddUtil,
174
+ "f64/utils/multiply": ddMulUtil,
175
+ "f64/utils/divide": ddDivUtil,
176
+ "f64/utils/sqrt": ddSqrtUtil,
177
+ dnrm2,
178
+ "reduction/scaledSumF64": scaledSumF64,
179
+ };
180
+
112
181
  import sgemv_n from "./sgemv_n.wgsl";
113
182
  import sgemv_t from "./sgemv_t.wgsl";
114
183
  routineShaders.sgemv = { sgemv_n, sgemv_t }; // one or the other, picked by trans
@@ -0,0 +1,93 @@
1
+ // scaledSum reduction (f64, double-double): collapses 2*WGS (scale, ssq) DD
2
+ // partials from dnrm2.wgsl into the final norm — sqrt(scale² · ssq) ==
3
+ // scale · sqrt(ssq), via ddMulProtected/ddSqrtProtected. Mirrors
4
+ // reduction/scaledSum.wgsl's shape exactly; ssqMergeProtected is duplicated
5
+ // from dnrm2.wgsl rather than shared via f64/utils/ — see that file's own
6
+ // header for why (same convention the f32 pair already uses).
7
+ // dispatch: 1 workgroup of WGS threads.
8
+ // partialsScale*/partialsSsq* must have exactly 2*WGS entries each.
9
+
10
+ @group(0) @binding(0) var<storage, read> partialsScaleHi: array<f32>;
11
+ @group(0) @binding(1) var<storage, read> partialsScaleLo: array<f32>;
12
+ @group(0) @binding(2) var<storage, read> partialsSsqHi: array<f32>;
13
+ @group(0) @binding(3) var<storage, read> partialsSsqLo: array<f32>;
14
+ @group(0) @binding(4) var<storage, read_write> resultHi: array<f32, 1>;
15
+ @group(0) @binding(5) var<storage, read_write> resultLo: array<f32, 1>;
16
+
17
+ const WGS: u32 = 64;
18
+
19
+ struct ScaleSsq {
20
+ scale: DD,
21
+ ssq: DD,
22
+ }
23
+
24
+ fn ddSelect(a: DD, b: DD, cond: bool) -> DD {
25
+ return DD(select(a.hi, b.hi, cond), select(a.lo, b.lo, cond));
26
+ }
27
+
28
+ // Associative merge of two independent (scale, ssq) partials — see
29
+ // dnrm2.wgsl for the derivation and why this is branch-free.
30
+ fn ssqMergeProtected(a: ScaleSsq, b: ScaleSsq, threadSlot: u32) -> ScaleSsq {
31
+ let isBigger = !ddGreater(b.scale, a.scale); // a.scale >= b.scale
32
+ let bigger = ddSelect(b.scale, a.scale, isBigger);
33
+ let smaller = ddSelect(a.scale, b.scale, isBigger);
34
+ let biggerSsq = ddSelect(b.ssq, a.ssq, isBigger);
35
+ let smallerSsq = ddSelect(a.ssq, b.ssq, isBigger);
36
+ let biggerIsZero = bigger.hi == 0.0;
37
+ let safeBigger = ddSelect(bigger, DD(1.0, 0.0), biggerIsZero);
38
+ let r = ddDivProtected(smaller, safeBigger, threadSlot);
39
+ let rsq = ddMulProtected(r, r, threadSlot);
40
+ let smallerSsqTimesRsq = ddMulProtected(smallerSsq, rsq, threadSlot);
41
+ let newSsq = ddAddProtected(biggerSsq, smallerSsqTimesRsq, threadSlot);
42
+ return ScaleSsq(bigger, newSsq);
43
+ }
44
+
45
+ var<workgroup> tileScaleHi: array<f32, 64>;
46
+ var<workgroup> tileScaleLo: array<f32, 64>;
47
+ var<workgroup> tileSsqHi: array<f32, 64>;
48
+ var<workgroup> tileSsqLo: array<f32, 64>;
49
+
50
+ @compute @workgroup_size(64)
51
+ fn reduce_scaled_f64(
52
+ @builtin(local_invocation_id) lid: vec3u,
53
+ ) {
54
+ let i = lid.x;
55
+ let a = ScaleSsq(DD(partialsScaleHi[i], partialsScaleLo[i]), DD(partialsSsqHi[i], partialsSsqLo[i]));
56
+ let b = ScaleSsq(DD(partialsScaleHi[i + WGS], partialsScaleLo[i + WGS]), DD(partialsSsqHi[i + WGS], partialsSsqLo[i + WGS]));
57
+ let merged0 = ssqMergeProtected(a, b, i);
58
+ tileScaleHi[i] = merged0.scale.hi;
59
+ tileScaleLo[i] = merged0.scale.lo;
60
+ tileSsqHi[i] = merged0.ssq.hi;
61
+ tileSsqLo[i] = merged0.ssq.lo;
62
+ workgroupBarrier();
63
+
64
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
65
+ let partner = select(i, i + s, i < s);
66
+ let ai = ScaleSsq(DD(tileScaleHi[i], tileScaleLo[i]), DD(tileSsqHi[i], tileSsqLo[i]));
67
+ let bi = ScaleSsq(DD(tileScaleHi[partner], tileScaleLo[partner]), DD(tileSsqHi[partner], tileSsqLo[partner]));
68
+ let merged = ssqMergeProtected(ai, bi, i);
69
+ workgroupBarrier();
70
+ if (i < s) {
71
+ tileScaleHi[i] = merged.scale.hi;
72
+ tileScaleLo[i] = merged.scale.lo;
73
+ tileSsqHi[i] = merged.ssq.hi;
74
+ tileSsqLo[i] = merged.ssq.lo;
75
+ }
76
+ workgroupBarrier();
77
+ }
78
+
79
+ // ddSqrtProtected/ddMulProtected's own workgroupBarrier()s need every
80
+ // thread to call them — every thread redundantly computes the same final
81
+ // scale·sqrt(ssq) from tile[0] (still visible to all after the reduction
82
+ // above), and only the write-back is conditional. Guarding the calls
83
+ // themselves behind `if (i == 0u)` (as the plain-f32 original safely
84
+ // does with its unprotected `sqrt()`) would leave 63 threads never
85
+ // reaching a barrier the one remaining thread still needs.
86
+ let scale = DD(tileScaleHi[0], tileScaleLo[0]);
87
+ let ssq = DD(tileSsqHi[0], tileSsqLo[0]);
88
+ let result = ddMulProtected(scale, ddSqrtProtected(ssq, i), i);
89
+ if (i == 0u) {
90
+ resultHi[0] = result.hi;
91
+ resultLo[0] = result.lo;
92
+ }
93
+ }
@@ -13,14 +13,12 @@ import { extractResult } from "../util/result.mjs";
13
13
  import { getPipeline } from "../util/pipeline.mjs";
14
14
  import { GpuVector } from "../classes/GpuVector.mjs";
15
15
  import { WGS } from "../util/constants.mjs";
16
- import { requireSameDevice } from "../util/device.mjs";
17
-
16
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
18
17
 
19
18
  export async function snrm2(device, n, x, incx) {
20
19
  const xIsGpu = x instanceof GpuVector;
21
20
 
22
- if (!(device instanceof GPUDevice))
23
- throw new Error("device must be a GPUDevice.");
21
+ requireGpuDevice(device);
24
22
  requireSameDevice(device, "snrm2", { x });
25
23
  if (!Number.isInteger(n) || !Number.isInteger(incx))
26
24
  throw new Error("n and incx must be integers.");
@@ -46,13 +44,19 @@ export async function snrm2(device, n, x, incx) {
46
44
  try {
47
45
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "snrm2-x", false);
48
46
  // 2*WGS partial (scale, ssq) pairs — see snrm2.wgsl for what they represent.
49
- partialsScaleBuffer = createStorageBuffer(device,
47
+ partialsScaleBuffer = createStorageBuffer(
48
+ device,
50
49
  2 * WGS * 4,
51
50
  "snrm2-partials-scale",
52
51
  );
53
- partialsSsqBuffer = createStorageBuffer(device, 2 * WGS * 4, "snrm2-partials-ssq");
52
+ partialsSsqBuffer = createStorageBuffer(
53
+ device,
54
+ 2 * WGS * 4,
55
+ "snrm2-partials-ssq",
56
+ );
54
57
  resultBuffer = createResultBuffer(device, 4, "snrm2-result"); // final f32 scalar
55
- paramsBuffer = createParamsBuffer(device,
58
+ paramsBuffer = createParamsBuffer(
59
+ device,
56
60
  [
57
61
  { value: n, type: "u32" },
58
62
  { value: incx, type: "u32" },
@@ -66,7 +70,8 @@ export async function snrm2(device, n, x, incx) {
66
70
  partialsSsqBuffer,
67
71
  paramsBuffer,
68
72
  ]);
69
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
73
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
74
+ device,
70
75
  pipelineMain,
71
76
  bgMain,
72
77
  2 * WGS,
@@ -74,12 +79,13 @@ export async function snrm2(device, n, x, incx) {
74
79
 
75
80
  submit(device, enc1);
76
81
 
77
- const bgReduce = createBindGroup(device, pipelineReduce.getBindGroupLayout(0), [
78
- partialsScaleBuffer,
79
- partialsSsqBuffer,
80
- resultBuffer,
81
- ]);
82
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
82
+ const bgReduce = createBindGroup(
83
+ device,
84
+ pipelineReduce.getBindGroupLayout(0),
85
+ [partialsScaleBuffer, partialsSsqBuffer, resultBuffer],
86
+ );
87
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
88
+ device,
83
89
  pipelineReduce,
84
90
  bgReduce,
85
91
  1,