wgblas 2.1.0 → 2.2.1

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 +1016 -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 +29 -11
  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
package/src/sger/sger.mjs CHANGED
@@ -11,18 +11,28 @@ import { extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { GpuVector } from "../classes/GpuVector.mjs";
13
13
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
14
- import { requireSameDevice } from "../util/device.mjs";
14
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
15
15
 
16
- export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout = "row-major") {
16
+ export async function sger(
17
+ device,
18
+ m,
19
+ n,
20
+ alpha,
21
+ x,
22
+ incx,
23
+ y,
24
+ incy,
25
+ A,
26
+ lda,
27
+ layout = "row-major",
28
+ ) {
17
29
  const AIsGpu = A instanceof GpuMatrix;
18
30
 
19
- if (!(device instanceof GPUDevice))
20
- throw new Error("device must be a GPUDevice.");
31
+ requireGpuDevice(device);
21
32
  requireSameDevice(device, "sger", { A, x, y });
22
33
  if (layout !== "row-major" && layout !== "column-major")
23
34
  throw new Error("layout must be 'row-major' or 'column-major'.");
24
- if (typeof alpha !== "number")
25
- throw new Error("alpha must be a number.");
35
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
26
36
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
27
37
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
28
38
  if (
@@ -79,9 +89,13 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
79
89
  "A does not have enough elements for the given m, n, and lda.",
80
90
  );
81
91
  if (x.length < (m - 1) * incx + 1)
82
- throw new Error("x does not have enough elements for the given m and incx.");
92
+ throw new Error(
93
+ "x does not have enough elements for the given m and incx.",
94
+ );
83
95
  if (y.length < (n - 1) * incy + 1)
84
- throw new Error("y does not have enough elements for the given n and incy.");
96
+ throw new Error(
97
+ "y does not have enough elements for the given n and incy.",
98
+ );
85
99
 
86
100
  const pipeline = await getPipeline(device, "sger");
87
101
 
@@ -94,14 +108,15 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
94
108
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sger-x", false);
95
109
  yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sger-y", false);
96
110
  ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sger-A", true);
97
- paramsBuffer = createParamsBuffer(device,
111
+ paramsBuffer = createParamsBuffer(
112
+ device,
98
113
  [
99
- { value: m, type: "u32" },
100
- { value: n, type: "u32" },
114
+ { value: m, type: "u32" },
115
+ { value: n, type: "u32" },
101
116
  { value: alpha, type: "f32" },
102
- { value: incx, type: "u32" },
103
- { value: incy, type: "u32" },
104
- { value: lda, type: "u32" },
117
+ { value: incx, type: "u32" },
118
+ { value: incy, type: "u32" },
119
+ { value: lda, type: "u32" },
105
120
  ],
106
121
  "sger-params",
107
122
  );
@@ -116,8 +131,15 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
116
131
  // One workgroup per row of A; clamped to device limit — the shader's
117
132
  // grid-stride loop handles remaining rows when m > dispatch count.
118
133
  const wgCount = Math.min(m, device.limits.maxComputeWorkgroupsPerDimension);
119
- const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
120
- const readBuffer = AIsGpu ? null : stageReadback(device, commandEncoder, ABuffer);
134
+ const { commandEncoder, ts } = runComputePass(
135
+ device,
136
+ pipeline,
137
+ bindGroup,
138
+ wgCount,
139
+ );
140
+ const readBuffer = AIsGpu
141
+ ? null
142
+ : stageReadback(device, commandEncoder, ABuffer);
121
143
 
122
144
  submit(device, commandEncoder);
123
145
 
@@ -0,0 +1,3 @@
1
+ // Fixture: valid shader, used only by test.pipeline.js's compile-error
2
+ // line-mapping regression test. Not referenced by any routine or by
3
+ // shaders/index.mjs's browser-bundle mapping.
@@ -0,0 +1,2 @@
1
+ // Fixture: shader with a deliberate syntax error on the next line.
2
+ this is not valid wgsl syntax;
@@ -0,0 +1,33 @@
1
+ // cscal: x := alpha * x, complex. x is one interleaved f32 array
2
+ // (re0, im0, re1, im1, ...), matching Complex32Array/GpuVector's storage
3
+ // (and cuBLAS's cuComplex / stdlib's Complex64Array) — no repacking needed
4
+ // between JS and GPU.
5
+ // (alphaRe + i*alphaIm)(re + i*im) = (alphaRe*re - alphaIm*im) + i*(alphaRe*im + alphaIm*re)
6
+
7
+ @group(0) @binding(0) var<storage, read_write> x: array<f32>;
8
+
9
+ struct Params {
10
+ n: u32,
11
+ alphaRe: f32,
12
+ alphaIm: f32,
13
+ x_inc: u32,
14
+ }
15
+
16
+ @group(0) @binding(1) var<uniform> params: Params;
17
+
18
+ const WGS: u32 = 64;
19
+
20
+ @compute @workgroup_size(64)
21
+ fn main(
22
+ @builtin(global_invocation_id) gid: vec3u,
23
+ @builtin(num_workgroups) num_wg: vec3u,
24
+ ) {
25
+ for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
26
+ let base = 2u * id * params.x_inc;
27
+ // Both new parts need both old parts, so capture them before either write.
28
+ let re = x[base];
29
+ let im = x[base + 1u];
30
+ x[base] = params.alphaRe * re - params.alphaIm * im;
31
+ x[base + 1u] = params.alphaRe * im + params.alphaIm * re;
32
+ }
33
+ }
@@ -0,0 +1,66 @@
1
+ // daxpy: y := alpha * x + y, double-double (Dekker) f64 emulation of saxpy.
2
+ // Each element costs one ddMulProtected (alpha*x[i]) then one ddAddProtected
3
+ // (+ y[i]) — the same two-protected-op shape ddot spends per term, applied
4
+ // straight to the output instead of folded into a reduction. See dscal.wgsl
5
+ // for why this is a uniform main pass plus a ragged, select-masked tail
6
+ // rather than a plain `id < params.n` grid-stride loop.
7
+
8
+ @group(0) @binding(0) var<storage, read> xHi: array<f32>;
9
+ @group(0) @binding(1) var<storage, read> 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
+ @group(0) @binding(4) var<uniform> params: Params;
13
+
14
+ struct Params {
15
+ n: u32,
16
+ alphaHi: f32,
17
+ alphaLo: f32,
18
+ x_inc: u32,
19
+ y_inc: u32,
20
+ }
21
+
22
+ const WGS: u32 = 64;
23
+
24
+ @compute @workgroup_size(64)
25
+ fn daxpy_main(
26
+ @builtin(global_invocation_id) gid: vec3u,
27
+ @builtin(local_invocation_id) lid: vec3u,
28
+ @builtin(workgroup_id) wgid: vec3u,
29
+ @builtin(num_workgroups) num_wg: vec3u,
30
+ ) {
31
+ let alpha = DD(params.alphaHi, params.alphaLo);
32
+ let stride = num_wg.x * WGS;
33
+
34
+ let n_floor = (params.n / stride) * stride;
35
+ let mainIters = n_floor / stride;
36
+ for (var iter = 0u; iter < mainIters; iter++) {
37
+ let id = gid.x + iter * stride;
38
+ let ix = id * params.x_inc;
39
+ let iy = id * params.y_inc;
40
+ let prod = ddMulProtected(alpha, DD(xHi[ix], xLo[ix]), lid.x);
41
+ let result = ddAddProtected(prod, DD(yHi[iy], yLo[iy]), lid.x);
42
+ yHi[iy] = result.hi;
43
+ yLo[iy] = result.lo;
44
+ }
45
+
46
+ // Tail is ragged (0 or 1 extra per thread) — pad to this workgroup's worst
47
+ // case so every thread still calls ddMulProtected/ddAddProtected the same
48
+ // number of times (their barriers need that), masking only the write.
49
+ let wgBaseGid = wgid.x * WGS;
50
+ var tailIters = 0u;
51
+ if (n_floor + wgBaseGid < params.n) {
52
+ tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
53
+ }
54
+ for (var iter = 0u; iter < tailIters; iter++) {
55
+ let id = n_floor + gid.x + iter * stride;
56
+ let valid = id < params.n;
57
+ let ix = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
58
+ let iy = select(0u, id * params.y_inc, valid);
59
+ let prod = ddMulProtected(alpha, DD(xHi[ix], xLo[ix]), lid.x);
60
+ let result = ddAddProtected(prod, DD(yHi[iy], yLo[iy]), lid.x);
61
+ if (valid) {
62
+ yHi[iy] = result.hi;
63
+ yLo[iy] = result.lo;
64
+ }
65
+ }
66
+ }
@@ -0,0 +1,34 @@
1
+ // dcopy: y = x, double-double (Dekker) f64 emulation of scopy. A copy is
2
+ // pure data movement — hi and lo are transferred verbatim, with no
3
+ // arithmetic 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 scopy.wgsl itself.
7
+
8
+ @group(0) @binding(0) var<storage, read> xHi: array<f32>;
9
+ @group(0) @binding(1) var<storage, read> 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
+ yHi[iy] = xHi[ix];
32
+ yLo[iy] = xLo[ix];
33
+ }
34
+ }
@@ -0,0 +1,106 @@
1
+ // ddot: sum(x[i] * y[i]), double-double (Dekker). Same ILP=4 shape as
2
+ // dasum.wgsl, which this mirrors closely — the only structural difference is
3
+ // a second input vector and a product where dasum takes an absolute value.
4
+ //
5
+ // See f64/utils/add.wgsl for ddAddProtected and why plain ddAdd isn't safe,
6
+ // and f64/utils/multiply.wgsl for ddMulProtected. The multiply itself
7
+ // (twoProdBit) needs no barrier; only its final renormalisation does, which
8
+ // is why each element costs two protected ops here against dasum's one.
9
+
10
+ @group(0) @binding(0) var<storage, read> xHi: array<f32>;
11
+ @group(0) @binding(1) var<storage, read> xLo: array<f32>;
12
+ @group(0) @binding(2) var<storage, read> yHi: array<f32>;
13
+ @group(0) @binding(3) var<storage, read> yLo: array<f32>;
14
+ @group(0) @binding(4) var<storage, read_write> partialsHi: array<f32>;
15
+ @group(0) @binding(5) var<storage, read_write> partialsLo: array<f32>;
16
+ @group(0) @binding(6) var<uniform> params: Params;
17
+
18
+ struct Params {
19
+ n: u32,
20
+ x_inc: u32,
21
+ y_inc: u32,
22
+ }
23
+
24
+ const WGS: u32 = 64;
25
+
26
+ var<workgroup> tile: array<DD, 64>;
27
+
28
+ @compute @workgroup_size(64)
29
+ fn ddot_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
+ var acc0 = DD(0.0, 0.0);
36
+ var acc1 = DD(0.0, 0.0);
37
+ var acc2 = DD(0.0, 0.0);
38
+ var acc3 = DD(0.0, 0.0);
39
+
40
+ let stride = num_wg.x * WGS;
41
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
42
+
43
+ // Same trip count for every thread, but driven by a counter, not `id`
44
+ // itself (the protected ops' barriers need a provably-uniform loop bound).
45
+ let mainIters = n4_floor / (4u * stride);
46
+ for (var iter = 0u; iter < mainIters; iter++) {
47
+ let id = gid.x + iter * 4u * stride;
48
+ let d0 = id;
49
+ let d1 = id + stride;
50
+ let d2 = id + 2u * stride;
51
+ let d3 = id + 3u * stride;
52
+
53
+ let p0 = ddMulProtected(DD(xHi[d0 * params.x_inc], xLo[d0 * params.x_inc]),
54
+ DD(yHi[d0 * params.y_inc], yLo[d0 * params.y_inc]), lid.x);
55
+ let p1 = ddMulProtected(DD(xHi[d1 * params.x_inc], xLo[d1 * params.x_inc]),
56
+ DD(yHi[d1 * params.y_inc], yLo[d1 * params.y_inc]), lid.x);
57
+ let p2 = ddMulProtected(DD(xHi[d2 * params.x_inc], xLo[d2 * params.x_inc]),
58
+ DD(yHi[d2 * params.y_inc], yLo[d2 * params.y_inc]), lid.x);
59
+ let p3 = ddMulProtected(DD(xHi[d3 * params.x_inc], xLo[d3 * params.x_inc]),
60
+ DD(yHi[d3 * params.y_inc], yLo[d3 * params.y_inc]), lid.x);
61
+
62
+ acc0 = ddAddProtected(acc0, p0, lid.x);
63
+ acc1 = ddAddProtected(acc1, p1, lid.x);
64
+ acc2 = ddAddProtected(acc2, p2, lid.x);
65
+ acc3 = ddAddProtected(acc3, p3, lid.x);
66
+ }
67
+
68
+ // Tail is ragged (0-3 extra per thread) — pad to this workgroup's worst case.
69
+ // Out-of-range lanes still run the multiply (it carries a barrier, so every
70
+ // thread must reach it) against index 0, then mask the result to zero.
71
+ let wgBaseGid = wgid.x * WGS;
72
+ var tailIters = 0u;
73
+ if (n4_floor + wgBaseGid < params.n) {
74
+ tailIters = (params.n - 1u - n4_floor - wgBaseGid) / stride + 1u;
75
+ }
76
+ for (var iter = 0u; iter < tailIters; iter++) {
77
+ let id = n4_floor + gid.x + iter * stride;
78
+ let valid = id < params.n;
79
+ let ix = select(0u, id * params.x_inc, valid);
80
+ let iy = select(0u, id * params.y_inc, valid);
81
+ let prod = ddMulProtected(DD(xHi[ix], xLo[ix]), DD(yHi[iy], yLo[iy]), lid.x);
82
+ // select() has no DD overload
83
+ let contribution = DD(select(0.0, prod.hi, valid), select(0.0, prod.lo, valid));
84
+ acc0 = ddAddProtected(acc0, contribution, lid.x);
85
+ }
86
+
87
+ let combined01 = ddAddProtected(acc0, acc1, lid.x);
88
+ let combined23 = ddAddProtected(acc2, acc3, lid.x);
89
+ tile[lid.x] = ddAddProtected(combined01, combined23, lid.x);
90
+ workgroupBarrier();
91
+
92
+ // Inactive threads combine against a throwaway partner and discard it
93
+ // (ddAddProtected must be called unconditionally by every thread).
94
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
95
+ let partner = select(lid.x, lid.x + s, lid.x < s);
96
+ let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
97
+ workgroupBarrier(); // all threads must read tile[] above before any write below
98
+ if (lid.x < s) { tile[lid.x] = combined; }
99
+ workgroupBarrier();
100
+ }
101
+
102
+ if (lid.x == 0u) {
103
+ partialsHi[wgid.x] = tile[0].hi;
104
+ partialsLo[wgid.x] = tile[0].lo;
105
+ }
106
+ }
@@ -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
+ }