wgblas 2.0.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 (124) hide show
  1. package/README.md +20 -18
  2. package/dist/wgblas.browser.js +2172 -1174
  3. package/index.d.mts +49 -44
  4. package/index.mjs +11 -0
  5. package/package.json +133 -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 +126 -17
  12. package/src/classes/GpuVector.d.mts +36 -40
  13. package/src/classes/GpuVector.mjs +66 -11
  14. package/src/cscal/cscal.d.mts +47 -0
  15. package/src/cscal/cscal.mjs +98 -0
  16. package/src/dasum/dasum.d.mts +4 -4
  17. package/src/dasum/dasum.mjs +38 -20
  18. package/src/daxpy/daxpy.d.mts +56 -0
  19. package/src/daxpy/daxpy.mjs +150 -0
  20. package/src/dcopy/dcopy.d.mts +52 -0
  21. package/src/dcopy/dcopy.mjs +140 -0
  22. package/src/ddot/ddot.d.mts +62 -0
  23. package/src/ddot/ddot.mjs +184 -0
  24. package/src/devdocs.mjs +13 -0
  25. package/src/dnrm2/dnrm2.d.mts +50 -0
  26. package/src/dnrm2/dnrm2.mjs +189 -0
  27. package/src/drot/drot.d.mts +67 -0
  28. package/src/drot/drot.mjs +170 -0
  29. package/src/drotm/drotm.d.mts +67 -0
  30. package/src/drotm/drotm.mjs +171 -0
  31. package/src/dscal/dscal.d.mts +52 -0
  32. package/src/dscal/dscal.mjs +119 -0
  33. package/src/dswap/dswap.d.mts +57 -0
  34. package/src/dswap/dswap.mjs +155 -0
  35. package/src/idamax/idamax.d.mts +20 -2
  36. package/src/idamax/idamax.mjs +56 -24
  37. package/src/init.mjs +117 -56
  38. package/src/isamax/isamax.d.mts +20 -2
  39. package/src/isamax/isamax.mjs +21 -16
  40. package/src/random/random.d.mts +37 -39
  41. package/src/random/random.mjs +39 -7
  42. package/src/sasum/sasum.d.mts +2 -2
  43. package/src/sasum/sasum.mjs +20 -16
  44. package/src/saxpy/saxpy.d.mts +2 -2
  45. package/src/saxpy/saxpy.mjs +14 -11
  46. package/src/scopy/scopy.d.mts +2 -2
  47. package/src/scopy/scopy.mjs +13 -9
  48. package/src/sdot/sdot.d.mts +2 -2
  49. package/src/sdot/sdot.mjs +21 -17
  50. package/src/sgemm/sgemm.d.mts +2 -2
  51. package/src/sgemm/sgemm.mjs +109 -40
  52. package/src/sgemmtr/sgemmtr.d.mts +3 -2
  53. package/src/sgemmtr/sgemmtr.mjs +98 -40
  54. package/src/sgemv/sgemv.d.mts +2 -2
  55. package/src/sgemv/sgemv.mjs +69 -41
  56. package/src/sger/sger.d.mts +2 -2
  57. package/src/sger/sger.mjs +43 -19
  58. package/src/shaders/__test_pipeline_a.wgsl +3 -0
  59. package/src/shaders/__test_pipeline_b.wgsl +2 -0
  60. package/src/shaders/cscal.wgsl +33 -0
  61. package/src/shaders/daxpy.wgsl +66 -0
  62. package/src/shaders/dcopy.wgsl +34 -0
  63. package/src/shaders/ddot.wgsl +106 -0
  64. package/src/shaders/dnrm2.wgsl +167 -0
  65. package/src/shaders/drot.wgsl +81 -0
  66. package/src/shaders/drotm.wgsl +99 -0
  67. package/src/shaders/dscal.wgsl +60 -0
  68. package/src/shaders/dswap.wgsl +38 -0
  69. package/src/shaders/f64/utils/add.wgsl +6 -0
  70. package/src/shaders/f64/utils/divide.wgsl +45 -0
  71. package/src/shaders/f64/utils/multiply.wgsl +19 -10
  72. package/src/shaders/f64/utils/sqrt.wgsl +45 -0
  73. package/src/shaders/index.mjs +233 -14
  74. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  75. package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
  76. package/src/shaders/sgemm_large.wgsl +107 -18
  77. package/src/shaders/sgemm_small.wgsl +115 -15
  78. package/src/shaders/sgemmtr_large.wgsl +4 -1
  79. package/src/shaders/sgemmtr_small.wgsl +4 -1
  80. package/src/shaders/sgemv_n.wgsl +3 -1
  81. package/src/shaders/sgemv_t.wgsl +3 -1
  82. package/src/shaders/snrm2.wgsl +72 -23
  83. package/src/shaders/ssymv.wgsl +3 -1
  84. package/src/snrm2/snrm2.d.mts +2 -2
  85. package/src/snrm2/snrm2.mjs +41 -23
  86. package/src/srot/srot.d.mts +2 -4
  87. package/src/srot/srot.mjs +16 -11
  88. package/src/srotm/srotm.d.mts +2 -4
  89. package/src/srotm/srotm.mjs +17 -11
  90. package/src/sscal/sscal.d.mts +3 -3
  91. package/src/sscal/sscal.mjs +14 -12
  92. package/src/sswap/sswap.d.mts +2 -2
  93. package/src/sswap/sswap.mjs +18 -10
  94. package/src/ssymm/ssymm.d.mts +5 -4
  95. package/src/ssymm/ssymm.mjs +150 -54
  96. package/src/ssymv/ssymv.d.mts +2 -2
  97. package/src/ssymv/ssymv.mjs +47 -26
  98. package/src/ssyr/ssyr.d.mts +2 -2
  99. package/src/ssyr/ssyr.mjs +38 -17
  100. package/src/ssyr2/ssyr2.d.mts +2 -2
  101. package/src/ssyr2/ssyr2.mjs +48 -21
  102. package/src/ssyr2k/ssyr2k.d.mts +3 -2
  103. package/src/ssyr2k/ssyr2k.mjs +140 -62
  104. package/src/ssyrk/ssyrk.d.mts +3 -2
  105. package/src/ssyrk/ssyrk.mjs +91 -39
  106. package/src/strmm/strmm.d.mts +5 -4
  107. package/src/strmm/strmm.mjs +174 -60
  108. package/src/strmv/strmv.d.mts +2 -2
  109. package/src/strmv/strmv.mjs +42 -20
  110. package/src/strsm/strsm.d.mts +6 -4
  111. package/src/strsm/strsm.mjs +438 -174
  112. package/src/strsv/strsv.d.mts +5 -3
  113. package/src/strsv/strsv.mjs +89 -34
  114. package/src/util/benchmark.mjs +9 -9
  115. package/src/util/bindgroup.mjs +1 -3
  116. package/src/util/buffer.mjs +139 -24
  117. package/src/util/complex.mjs +87 -0
  118. package/src/util/compute.mjs +19 -16
  119. package/src/util/constants.mjs +57 -0
  120. package/src/util/device.mjs +49 -0
  121. package/src/util/pipeline.mjs +44 -10
  122. package/src/util/workgroup.mjs +72 -7
  123. package/src/shaders/browser-shaders.mjs +0 -81
  124. package/src/shaders/f64add.wgsl +0 -281
@@ -0,0 +1,167 @@
1
+ // dnrm2: result = sqrt(sum(x[i] * x[i])), double-double (Dekker) f64
2
+ // emulation of snrm2 — same scaled accumulation (Blue's algorithm), just
3
+ // with `scale`/`ssq` as DD pairs (via ddDivProtected/ddMulProtected/
4
+ // ddAddProtected/ddSqrtProtected) instead of plain f32. Squaring still
5
+ // saturates an f32 hi component above ~1.8e19 regardless of DD precision
6
+ // (DD widens the mantissa, not the exponent range), so the scaling is
7
+ // still needed here for the same reason it was in snrm2.
8
+ //
9
+ // snrm2.wgsl's ssqAccum/ssqMerge each branch on which operand is bigger —
10
+ // can't carry over directly, since a protected op's workgroupBarrier()
11
+ // needs every thread to reach the same call site, and here different
12
+ // threads could take different branches. Both formulas are computed
13
+ // unconditionally below; only the final combine (`ddSelect`) differs per
14
+ // thread — same fix shape as drot's/drotm's own per-dispatch flags, just
15
+ // applied to a per-element branch instead.
16
+ //
17
+ // pass 1 dispatches 2*WGS workgroups; pass 2 (reduction/scaledSumF64.wgsl)
18
+ // duplicates ssqAccumProtected/ssqMergeProtected rather than sharing them,
19
+ // same as the plain-f32 pair already does.
20
+
21
+ @group(0) @binding(0) var<storage, read> xHi: array<f32>;
22
+ @group(0) @binding(1) var<storage, read> xLo: array<f32>;
23
+ @group(0) @binding(2) var<storage, read_write> partialsScaleHi: array<f32>;
24
+ @group(0) @binding(3) var<storage, read_write> partialsScaleLo: array<f32>;
25
+ @group(0) @binding(4) var<storage, read_write> partialsSsqHi: array<f32>;
26
+ @group(0) @binding(5) var<storage, read_write> partialsSsqLo: array<f32>;
27
+ @group(0) @binding(6) var<uniform> params: Params;
28
+
29
+ struct Params {
30
+ n: u32,
31
+ x_inc: u32,
32
+ }
33
+
34
+ const WGS: u32 = 64;
35
+
36
+ struct ScaleSsq {
37
+ scale: DD,
38
+ ssq: DD,
39
+ }
40
+
41
+ fn ddSelect(a: DD, b: DD, cond: bool) -> DD {
42
+ return DD(select(a.hi, b.hi, cond), select(a.lo, b.lo, cond));
43
+ }
44
+
45
+ // Folds one more |value| (DD) into a running (scale, ssq) pair — branch-free,
46
+ // see file header. `bigger`/`smaller` name the two operands by magnitude
47
+ // (not by which one was "acc" vs "new"), and biggerIsZero==true only when
48
+ // both scale and absxi are still exactly zero (the very first zero
49
+ // elements, before any nonzero value has been seen) — substituting a safe
50
+ // denominator there avoids a 0/0 without needing a separate branch/return;
51
+ // the arithmetic already reduces to a correct no-op in that case.
52
+ fn ssqAccumProtected(acc: ScaleSsq, absxi: DD, threadSlot: u32) -> ScaleSsq {
53
+ let isBigger = ddGreater(absxi, acc.scale);
54
+ let bigger = ddSelect(acc.scale, absxi, isBigger);
55
+ let smaller = ddSelect(absxi, acc.scale, isBigger);
56
+ let biggerIsZero = bigger.hi == 0.0;
57
+ let safeBigger = ddSelect(bigger, DD(1.0, 0.0), biggerIsZero);
58
+ let r = ddDivProtected(smaller, safeBigger, threadSlot);
59
+ let rsq = ddMulProtected(r, r, threadSlot);
60
+ let ssqTimesRsq = ddMulProtected(acc.ssq, rsq, threadSlot);
61
+ let sumIfBigger = ddAddProtected(DD(1.0, 0.0), ssqTimesRsq, threadSlot);
62
+ let sumIfNotBigger = ddAddProtected(acc.ssq, rsq, threadSlot);
63
+ let newSsq = ddSelect(sumIfNotBigger, sumIfBigger, isBigger);
64
+ return ScaleSsq(bigger, newSsq);
65
+ }
66
+
67
+ // Associative merge of two independent (scale, ssq) partials — same
68
+ // branch-free shape, for combining ILP lanes and the tree reduction.
69
+ fn ssqMergeProtected(a: ScaleSsq, b: ScaleSsq, threadSlot: u32) -> ScaleSsq {
70
+ let isBigger = !ddGreater(b.scale, a.scale); // a.scale >= b.scale
71
+ let bigger = ddSelect(b.scale, a.scale, isBigger);
72
+ let smaller = ddSelect(a.scale, b.scale, isBigger);
73
+ let biggerSsq = ddSelect(b.ssq, a.ssq, isBigger);
74
+ let smallerSsq = ddSelect(a.ssq, b.ssq, isBigger);
75
+ let biggerIsZero = bigger.hi == 0.0;
76
+ let safeBigger = ddSelect(bigger, DD(1.0, 0.0), biggerIsZero);
77
+ let r = ddDivProtected(smaller, safeBigger, threadSlot);
78
+ let rsq = ddMulProtected(r, r, threadSlot);
79
+ let smallerSsqTimesRsq = ddMulProtected(smallerSsq, rsq, threadSlot);
80
+ let newSsq = ddAddProtected(biggerSsq, smallerSsqTimesRsq, threadSlot);
81
+ return ScaleSsq(bigger, newSsq);
82
+ }
83
+
84
+ var<workgroup> tileScaleHi: array<f32, 64>;
85
+ var<workgroup> tileScaleLo: array<f32, 64>;
86
+ var<workgroup> tileSsqHi: array<f32, 64>;
87
+ var<workgroup> tileSsqLo: array<f32, 64>;
88
+
89
+ @compute @workgroup_size(64)
90
+ fn dnrm2_main(
91
+ @builtin(global_invocation_id) gid: vec3u,
92
+ @builtin(local_invocation_id) lid: vec3u,
93
+ @builtin(workgroup_id) wgid: vec3u,
94
+ @builtin(num_workgroups) num_wg: vec3u,
95
+ ) {
96
+ var acc0 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
97
+ var acc1 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
98
+ var acc2 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
99
+ var acc3 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
100
+
101
+ let stride = num_wg.x * WGS;
102
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
103
+
104
+ // Same trip count for every thread, driven by a counter (protected ops'
105
+ // barriers need a provably-uniform loop bound) — see dasum.wgsl.
106
+ let mainIters = n4_floor / (4u * stride);
107
+ for (var iter = 0u; iter < mainIters; iter++) {
108
+ let id = gid.x + iter * 4u * stride;
109
+ let i0 = id * params.x_inc;
110
+ let i1 = (id + stride) * params.x_inc;
111
+ let i2 = (id + 2u * stride) * params.x_inc;
112
+ let i3 = (id + 3u * stride) * params.x_inc;
113
+ acc0 = ssqAccumProtected(acc0, ddAbs(DD(xHi[i0], xLo[i0])), lid.x);
114
+ acc1 = ssqAccumProtected(acc1, ddAbs(DD(xHi[i1], xLo[i1])), lid.x);
115
+ acc2 = ssqAccumProtected(acc2, ddAbs(DD(xHi[i2], xLo[i2])), lid.x);
116
+ acc3 = ssqAccumProtected(acc3, ddAbs(DD(xHi[i3], xLo[i3])), lid.x);
117
+ }
118
+
119
+ // Tail is ragged (0-3 extra per thread) — pad to this workgroup's worst
120
+ // case, masking an invalid element to exactly 0 (contributes nothing).
121
+ let wgBaseGid = wgid.x * WGS;
122
+ var tailIters = 0u;
123
+ if (n4_floor + wgBaseGid < params.n) {
124
+ tailIters = (params.n - 1u - n4_floor - wgBaseGid) / stride + 1u;
125
+ }
126
+ for (var iter = 0u; iter < tailIters; iter++) {
127
+ let id = n4_floor + gid.x + iter * stride;
128
+ let valid = id < params.n;
129
+ let i = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
130
+ let loaded = ddAbs(DD(xHi[i], xLo[i]));
131
+ let contribution = DD(select(0.0, loaded.hi, valid), select(0.0, loaded.lo, valid));
132
+ acc0 = ssqAccumProtected(acc0, contribution, lid.x);
133
+ }
134
+
135
+ let combined01 = ssqMergeProtected(acc0, acc1, lid.x);
136
+ let combined23 = ssqMergeProtected(acc2, acc3, lid.x);
137
+ let combined = ssqMergeProtected(combined01, combined23, lid.x);
138
+ tileScaleHi[lid.x] = combined.scale.hi;
139
+ tileScaleLo[lid.x] = combined.scale.lo;
140
+ tileSsqHi[lid.x] = combined.ssq.hi;
141
+ tileSsqLo[lid.x] = combined.ssq.lo;
142
+ workgroupBarrier();
143
+
144
+ // Inactive threads merge against a throwaway partner and discard it
145
+ // (ssqMergeProtected must be called unconditionally by every thread).
146
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
147
+ let partner = select(lid.x, lid.x + s, lid.x < s);
148
+ let a = ScaleSsq(DD(tileScaleHi[lid.x], tileScaleLo[lid.x]), DD(tileSsqHi[lid.x], tileSsqLo[lid.x]));
149
+ let b = ScaleSsq(DD(tileScaleHi[partner], tileScaleLo[partner]), DD(tileSsqHi[partner], tileSsqLo[partner]));
150
+ let merged = ssqMergeProtected(a, b, lid.x);
151
+ workgroupBarrier(); // all threads must read tile[] above before any write below
152
+ if (lid.x < s) {
153
+ tileScaleHi[lid.x] = merged.scale.hi;
154
+ tileScaleLo[lid.x] = merged.scale.lo;
155
+ tileSsqHi[lid.x] = merged.ssq.hi;
156
+ tileSsqLo[lid.x] = merged.ssq.lo;
157
+ }
158
+ workgroupBarrier();
159
+ }
160
+
161
+ if (lid.x == 0u) {
162
+ partialsScaleHi[wgid.x] = tileScaleHi[0];
163
+ partialsScaleLo[wgid.x] = tileScaleLo[0];
164
+ partialsSsqHi[wgid.x] = tileSsqHi[0];
165
+ partialsSsqLo[wgid.x] = tileSsqLo[0];
166
+ }
167
+ }
@@ -0,0 +1,81 @@
1
+ // drot: x = c*x + s*y, y = -s*x + c*y — double-double (Dekker) f64 emulation
2
+ // of srot. c, s, x, and y are each split into an f32 (hi, lo) pair. Each
3
+ // element costs four ddMulProtected (c*x, s*y, -s*x, c*y) then two
4
+ // ddAddProtected (the two sums) — negS is computed once outside the loop
5
+ // via bitcast negation (exact, no rounding, so no barrier needed there)
6
+ // rather than adding a DD-subtract helper. See dscal.wgsl for why this is a
7
+ // uniform main pass plus a ragged, select-masked tail rather than a plain
8
+ // `id < params.n` grid-stride loop.
9
+
10
+ @group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
11
+ @group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
12
+ @group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
13
+ @group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
14
+ @group(0) @binding(4) var<uniform> params: Params;
15
+
16
+ struct Params {
17
+ n: u32,
18
+ cHi: f32,
19
+ cLo: f32,
20
+ sHi: f32,
21
+ sLo: f32,
22
+ x_inc: u32,
23
+ y_inc: u32,
24
+ }
25
+
26
+ const WGS: u32 = 64;
27
+
28
+ @compute @workgroup_size(64)
29
+ fn drot_main(
30
+ @builtin(global_invocation_id) gid: vec3u,
31
+ @builtin(local_invocation_id) lid: vec3u,
32
+ @builtin(workgroup_id) wgid: vec3u,
33
+ @builtin(num_workgroups) num_wg: vec3u,
34
+ ) {
35
+ let c = DD(params.cHi, params.cLo);
36
+ let s = DD(params.sHi, params.sLo);
37
+ let negS = DD(negf(params.sHi), negf(params.sLo));
38
+ let stride = num_wg.x * WGS;
39
+
40
+ let n_floor = (params.n / stride) * stride;
41
+ let mainIters = n_floor / stride;
42
+ for (var iter = 0u; iter < mainIters; iter++) {
43
+ let id = gid.x + iter * stride;
44
+ let ix = id * params.x_inc;
45
+ let iy = id * params.y_inc;
46
+ let xi = DD(xHi[ix], xLo[ix]);
47
+ let yi = DD(yHi[iy], yLo[iy]);
48
+ let xNew = ddAddProtected(ddMulProtected(c, xi, lid.x), ddMulProtected(s, yi, lid.x), lid.x);
49
+ let yNew = ddAddProtected(ddMulProtected(negS, xi, lid.x), ddMulProtected(c, yi, lid.x), lid.x);
50
+ xHi[ix] = xNew.hi;
51
+ xLo[ix] = xNew.lo;
52
+ yHi[iy] = yNew.hi;
53
+ yLo[iy] = yNew.lo;
54
+ }
55
+
56
+ // Tail is ragged (0 or 1 extra per thread) — pad to this workgroup's worst
57
+ // case so every thread in the workgroup still calls ddMulProtected/
58
+ // ddAddProtected the same number of times (their barriers need that),
59
+ // masking only the write.
60
+ let wgBaseGid = wgid.x * WGS;
61
+ var tailIters = 0u;
62
+ if (n_floor + wgBaseGid < params.n) {
63
+ tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
64
+ }
65
+ for (var iter = 0u; iter < tailIters; iter++) {
66
+ let id = n_floor + gid.x + iter * stride;
67
+ let valid = id < params.n;
68
+ let ix = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
69
+ let iy = select(0u, id * params.y_inc, valid);
70
+ let xi = DD(xHi[ix], xLo[ix]);
71
+ let yi = DD(yHi[iy], yLo[iy]);
72
+ let xNew = ddAddProtected(ddMulProtected(c, xi, lid.x), ddMulProtected(s, yi, lid.x), lid.x);
73
+ let yNew = ddAddProtected(ddMulProtected(negS, xi, lid.x), ddMulProtected(c, yi, lid.x), lid.x);
74
+ if (valid) {
75
+ xHi[ix] = xNew.hi;
76
+ xLo[ix] = xNew.lo;
77
+ yHi[iy] = yNew.hi;
78
+ yLo[iy] = yNew.lo;
79
+ }
80
+ }
81
+ }
@@ -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
+ }