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
@@ -9,27 +9,40 @@ import { runComputePass, submit } from "../util/compute.mjs";
9
9
  import { extractResult } from "../util/result.mjs";
10
10
  import { extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
- import { calcWorkgroups } from "../util/workgroup.mjs";
12
+ import { requireWorkgroups } from "../util/workgroup.mjs";
13
13
  import { GpuVector } from "../classes/GpuVector.mjs";
14
14
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
15
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
15
16
 
16
- export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y, incy, layout = "row-major") {
17
+ export async function sgemv(
18
+ device,
19
+ trans,
20
+ m,
21
+ n,
22
+ alpha,
23
+ A,
24
+ lda,
25
+ x,
26
+ incx,
27
+ beta,
28
+ y,
29
+ incy,
30
+ layout = "row-major",
31
+ ) {
17
32
  const AIsGpu = A instanceof GpuMatrix;
18
33
  const xIsGpu = x instanceof GpuVector;
19
34
  const yIsGpu = y instanceof GpuVector;
20
35
 
21
- if (!(device instanceof GPUDevice))
22
- throw new Error("device must be a GPUDevice.");
36
+ requireGpuDevice(device);
37
+ requireSameDevice(device, "sgemv", { A, x, y });
23
38
  if (trans !== "no-transpose" && trans !== "transpose")
24
39
  throw new Error("trans must be 'no-transpose' or 'transpose'.");
25
40
  if (layout !== "row-major" && layout !== "column-major")
26
41
  throw new Error("layout must be 'row-major' or 'column-major'.");
27
- if (typeof alpha !== "number")
28
- throw new Error("alpha must be a number.");
42
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
29
43
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
30
44
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
31
- if (typeof beta !== "number")
32
- throw new Error("beta must be a number.");
45
+ if (typeof beta !== "number") throw new Error("beta must be a number.");
33
46
  if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
34
47
  if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
35
48
  if (
@@ -53,15 +66,15 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
53
66
  "x and y must be the same type (both Float32Array or both GpuVector).",
54
67
  );
55
68
  if (xIsGpu && !AIsGpu)
56
- throw new Error(
57
- "A must be a GpuMatrix when x and y are GpuVectors.",
58
- );
69
+ throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");
59
70
  if (AIsGpu && !xIsGpu)
71
+ throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");
72
+ if (xIsGpu && x._buf === y._buf)
60
73
  throw new Error(
61
- "x and y must be GpuVectors when A is a GpuMatrix.",
74
+ "x and y must not reference the same GPU buffer when both are GpuVectors.",
62
75
  );
63
- if (xIsGpu && x._buf === y._buf)
64
- throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");
76
+ if (AIsGpu && yIsGpu && A._buf === y._buf)
77
+ throw new Error("A and y must not reference the same GPU buffer.");
65
78
  if (AIsGpu && lda !== A.lda)
66
79
  throw new Error("lda must match A.lda when A is a GpuMatrix.");
67
80
  if (AIsGpu && (A.rows < m || A.cols < n))
@@ -97,40 +110,56 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
97
110
  );
98
111
 
99
112
  const shaderName = isNoTrans ? "sgemv_n" : "sgemv_t";
100
- const pipeline = await getPipeline(device, shaderName);
113
+ const pipeline = await getPipeline(device, shaderName);
101
114
 
102
- const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sgemv-A", false);
103
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sgemv-x", false);
104
- const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sgemv-y", true);
105
- const paramsBuffer = createParamsBuffer(
106
- [
107
- { value: m, type: "u32" },
108
- { value: n, type: "u32" },
109
- { value: alpha, type: "f32" },
110
- { value: beta, type: "f32" },
111
- { value: incx, type: "u32" },
112
- { value: incy, type: "u32" },
113
- { value: lda, type: "u32" },
114
- ],
115
- "sgemv-params",
116
- );
115
+ let ABuffer = null;
116
+ let xBuffer = null;
117
+ let yBuffer = null;
118
+ let paramsBuffer = null;
117
119
 
118
120
  try {
119
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
121
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemv-A", false);
122
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sgemv-x", false);
123
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sgemv-y", true);
124
+ paramsBuffer = createParamsBuffer(
125
+ device,
126
+ [
127
+ { value: m, type: "u32" },
128
+ { value: n, type: "u32" },
129
+ { value: alpha, type: "f32" },
130
+ { value: beta, type: "f32" },
131
+ { value: incx, type: "u32" },
132
+ { value: incy, type: "u32" },
133
+ { value: lda, type: "u32" },
134
+ ],
135
+ "sgemv-params",
136
+ );
137
+
138
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
120
139
  ABuffer,
121
140
  xBuffer,
122
141
  yBuffer,
123
142
  paramsBuffer,
124
143
  ]);
125
144
 
126
- // NoTrans: one workgroup per row (grid-stride handles overflow); Trans: one thread per output column.
145
+ // NoTrans: one workgroup per row — sgemv_n.wgsl is a grid-stride loop, so
146
+ // clamping here only costs parallelism. Trans: one thread per output
147
+ // column, and sgemv_t.wgsl indexes straight off global_invocation_id with
148
+ // no fallback, so an over-limit dispatch must be refused, not truncated.
127
149
  const wgCount = isNoTrans
128
150
  ? Math.min(m, device.limits.maxComputeWorkgroupsPerDimension)
129
- : calcWorkgroups(yLen);
130
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
131
- const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
151
+ : requireWorkgroups(device, "sgemv", yLen);
152
+ const { commandEncoder, ts } = runComputePass(
153
+ device,
154
+ pipeline,
155
+ bindGroup,
156
+ wgCount,
157
+ );
158
+ const readBuffer = yIsGpu
159
+ ? null
160
+ : stageReadback(device, commandEncoder, yBuffer);
132
161
 
133
- submit(commandEncoder);
162
+ submit(device, commandEncoder);
134
163
 
135
164
  const gpuTimeMs = await extractTimestamp(ts);
136
165
 
@@ -143,10 +172,9 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
143
172
  if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
144
173
  return { y: result };
145
174
  } finally {
146
- if (!AIsGpu) destroyBuffers(ABuffer);
147
- if (!xIsGpu) destroyBuffers(xBuffer);
148
- if (!yIsGpu) destroyBuffers(yBuffer);
149
- destroyBuffers(paramsBuffer);
150
-
175
+ if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
176
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
177
+ if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
178
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
151
179
  }
152
180
  }
@@ -2,7 +2,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
2
2
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
3
3
 
4
4
  /**
5
- * Performs the rank-1 update A = alpha * x * y^T + A
5
+ * Performs the rank-1 update $$A \leftarrow \alpha x y^{T} + A$$
6
6
  *
7
7
  * A is an m×n matrix stored in row-major order, updated in place. `lda` is
8
8
  * the leading dimension (number of floats between the start of consecutive
@@ -43,7 +43,7 @@ export declare function sger(
43
43
  ): Promise<{ A: Float32Array; gpuTimeMs?: number }>;
44
44
 
45
45
  /**
46
- * Performs the rank-1 update A = alpha * x * y^T + A
46
+ * Performs the rank-1 update $$A \leftarrow \alpha x y^{T} + A$$
47
47
  *
48
48
  * x, y, and A are all kept resident on the GPU. `A`'s own `layout` (set at
49
49
  * `GpuMatrix.from` time) determines the operation — there is no separate
package/src/sger/sger.mjs CHANGED
@@ -11,16 +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 { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
14
15
 
15
- 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
+ ) {
16
29
  const AIsGpu = A instanceof GpuMatrix;
17
30
 
18
- if (!(device instanceof GPUDevice))
19
- throw new Error("device must be a GPUDevice.");
31
+ requireGpuDevice(device);
32
+ requireSameDevice(device, "sger", { A, x, y });
20
33
  if (layout !== "row-major" && layout !== "column-major")
21
34
  throw new Error("layout must be 'row-major' or 'column-major'.");
22
- if (typeof alpha !== "number")
23
- throw new Error("alpha must be a number.");
35
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
24
36
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
25
37
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
26
38
  if (
@@ -77,9 +89,13 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
77
89
  "A does not have enough elements for the given m, n, and lda.",
78
90
  );
79
91
  if (x.length < (m - 1) * incx + 1)
80
- 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
+ );
81
95
  if (y.length < (n - 1) * incy + 1)
82
- 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
+ );
83
99
 
84
100
  const pipeline = await getPipeline(device, "sger");
85
101
 
@@ -89,22 +105,23 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
89
105
  let paramsBuffer = null;
90
106
 
91
107
  try {
92
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sger-x", false);
93
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sger-y", false);
94
- ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sger-A", true);
108
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sger-x", false);
109
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sger-y", false);
110
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sger-A", true);
95
111
  paramsBuffer = createParamsBuffer(
112
+ device,
96
113
  [
97
- { value: m, type: "u32" },
98
- { value: n, type: "u32" },
114
+ { value: m, type: "u32" },
115
+ { value: n, type: "u32" },
99
116
  { value: alpha, type: "f32" },
100
- { value: incx, type: "u32" },
101
- { value: incy, type: "u32" },
102
- { value: lda, type: "u32" },
117
+ { value: incx, type: "u32" },
118
+ { value: incy, type: "u32" },
119
+ { value: lda, type: "u32" },
103
120
  ],
104
121
  "sger-params",
105
122
  );
106
123
 
107
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
124
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
108
125
  xBuffer,
109
126
  yBuffer,
110
127
  ABuffer,
@@ -114,10 +131,17 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
114
131
  // One workgroup per row of A; clamped to device limit — the shader's
115
132
  // grid-stride loop handles remaining rows when m > dispatch count.
116
133
  const wgCount = Math.min(m, device.limits.maxComputeWorkgroupsPerDimension);
117
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
118
- const readBuffer = AIsGpu ? null : stageReadback(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);
119
143
 
120
- submit(commandEncoder);
144
+ submit(device, commandEncoder);
121
145
 
122
146
  const gpuTimeMs = await extractTimestamp(ts);
123
147
 
@@ -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
+ }