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
@@ -2,8 +2,9 @@ import { GpuVector } from "../classes/GpuVector.mjs";
2
2
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
3
3
 
4
4
  /**
5
- * Solves the triangular system op(A) * x = b for x, in place (x holds b on
6
- * input, the solution on output).
5
+ * Solves the triangular system for $x$, in place — x holds b on input, the
6
+ * solution on output:
7
+ * $$\mathrm{op}(A) x = b$$
7
8
  *
8
9
  * A is an n×n triangular matrix stored in row-major order. Only the triangle
9
10
  * specified by `uplo` is referenced; the other triangle is not accessed.
@@ -42,7 +43,8 @@ export declare function strsv(
42
43
  ): Promise<{ x: Float32Array; gpuTimeMs?: number }>;
43
44
 
44
45
  /**
45
- * Solves the triangular system op(A) * x = b for x, in place.
46
+ * Solves the triangular system for $x$, in place:
47
+ * $$\mathrm{op}(A) x = b$$
46
48
  *
47
49
  * x is kept resident on the GPU (mutated in place). A must be a GpuMatrix;
48
50
  * its own `layout` (set at `GpuMatrix.from` time) determines the operation —
@@ -13,7 +13,7 @@ import { getPipeline } from "../util/pipeline.mjs";
13
13
  import { GpuVector } from "../classes/GpuVector.mjs";
14
14
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
15
15
  import { BLOCK_SIZE } from "../util/constants.mjs";
16
- import { requireSameDevice } from "../util/device.mjs";
16
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
17
17
 
18
18
  // Blocked triangular solve via explicit block inversion (invert/apply/update passes) instead of barrier-per-row substitution.
19
19
 
@@ -39,13 +39,23 @@ function createSharedParamsBuffer(device, data, label) {
39
39
  return buffer;
40
40
  }
41
41
 
42
- export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layout = "row-major") {
42
+ export async function strsv(
43
+ device,
44
+ uplo,
45
+ trans,
46
+ diag,
47
+ n,
48
+ A,
49
+ lda,
50
+ x,
51
+ incx,
52
+ layout = "row-major",
53
+ ) {
43
54
  const xIsGpu = x instanceof GpuVector;
44
55
  const AIsGpu = A instanceof GpuMatrix;
45
56
  const isUnit = diag === "unit";
46
57
 
47
- if (!(device instanceof GPUDevice))
48
- throw new Error("device must be a GPUDevice.");
58
+ requireGpuDevice(device);
49
59
  requireSameDevice(device, "strsv", { A, x });
50
60
  if (uplo !== "lower" && uplo !== "upper")
51
61
  throw new Error("uplo must be 'lower' or 'upper'.");
@@ -77,9 +87,7 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
77
87
  if (n === 0) return xIsGpu ? {} : { x };
78
88
 
79
89
  if (!AIsGpu && A.length < (n - 1) * lda + n)
80
- throw new Error(
81
- "A does not have enough elements for the given n and lda.",
82
- );
90
+ throw new Error("A does not have enough elements for the given n and lda.");
83
91
  if (x.length < (n - 1) * incx + 1)
84
92
  throw new Error(
85
93
  "x does not have enough elements for the given n and incx.",
@@ -89,7 +97,9 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
89
97
  const effLayout = AIsGpu ? A.layout : layout;
90
98
  const isColMajor = effLayout === "column-major";
91
99
  const isLower = isColMajor ? uplo === "upper" : uplo === "lower";
92
- const isNoTrans = isColMajor ? trans === "transpose" : trans === "no-transpose";
100
+ const isNoTrans = isColMajor
101
+ ? trans === "transpose"
102
+ : trans === "no-transpose";
93
103
 
94
104
  const invertPipeline = await getPipeline(device, "strsv_invert_block");
95
105
  const applyPipeline = await getPipeline(device, "strsv_apply_inverse");
@@ -117,7 +127,8 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
117
127
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "strsv-x", true);
118
128
  // One BLOCK_SIZE x BLOCK_SIZE dense region per block (row-major), even
119
129
  // though only a triangular half is ever nonzero — see strsv_invert_block.wgsl.
120
- AinvBuffer = createStorageBuffer(device,
130
+ AinvBuffer = createStorageBuffer(
131
+ device,
121
132
  numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4,
122
133
  "strsv-Ainv",
123
134
  );
@@ -130,35 +141,60 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
130
141
  const blockEnd = Math.min(blockStart + BLOCK_SIZE, n);
131
142
  return [incx, blockIndex, blockStart, blockEnd];
132
143
  });
133
- applyParamsBuffer = createSharedParamsBuffer(device, applyData, "strsv-apply-params");
144
+ applyParamsBuffer = createSharedParamsBuffer(
145
+ device,
146
+ applyData,
147
+ "strsv-apply-params",
148
+ );
134
149
 
135
150
  const updateData = packBlockParams(numBlocks, stride, (blockIndex) => {
136
151
  const blockStart = blockIndex * BLOCK_SIZE;
137
152
  const blockEnd = Math.min(blockStart + BLOCK_SIZE, n);
138
- return [n, incx, lda, isNoTrans ? 0 : 1, isLower ? 0 : 1, blockStart, blockEnd];
153
+ return [
154
+ n,
155
+ incx,
156
+ lda,
157
+ isNoTrans ? 0 : 1,
158
+ isLower ? 0 : 1,
159
+ blockStart,
160
+ blockEnd,
161
+ ];
139
162
  });
140
- updateParamsBuffer = createSharedParamsBuffer(device, updateData, "strsv-update-params");
163
+ updateParamsBuffer = createSharedParamsBuffer(
164
+ device,
165
+ updateData,
166
+ "strsv-update-params",
167
+ );
141
168
 
142
169
  const { commandEncoder, querySet } = beginTimedEncoder(device);
143
170
 
144
171
  // Pre-pass: every block's inverse, fully parallel, one dispatch.
145
- invertParams = createParamsBuffer(device,
172
+ invertParams = createParamsBuffer(
173
+ device,
146
174
  [
147
- { value: n, type: "u32" },
148
- { value: lda, type: "u32" },
175
+ { value: n, type: "u32" },
176
+ { value: lda, type: "u32" },
149
177
  { value: isNoTrans ? 0 : 1, type: "u32" },
150
- { value: isLower ? 0 : 1, type: "u32" },
151
- { value: isUnit ? 1 : 0, type: "u32" },
178
+ { value: isLower ? 0 : 1, type: "u32" },
179
+ { value: isUnit ? 1 : 0, type: "u32" },
152
180
  ],
153
181
  "strsv-invert-params",
154
182
  );
155
- const invertBindGroup = createBindGroup(device, invertPipeline.getBindGroupLayout(0), [
156
- ABuffer, AinvBuffer, invertParams,
157
- ]);
183
+ const invertBindGroup = createBindGroup(
184
+ device,
185
+ invertPipeline.getBindGroupLayout(0),
186
+ [ABuffer, AinvBuffer, invertParams],
187
+ );
158
188
  const invertDesc = querySet
159
189
  ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
160
190
  : undefined;
161
- encodePass(commandEncoder, invertPipeline, invertBindGroup, { x: BLOCK_SIZE, y: numBlocks }, invertDesc);
191
+ encodePass(
192
+ commandEncoder,
193
+ invertPipeline,
194
+ invertBindGroup,
195
+ { x: BLOCK_SIZE, y: numBlocks },
196
+ invertDesc,
197
+ );
162
198
 
163
199
  for (let bi = 0; bi < blockStarts.length; bi++) {
164
200
  const blockStart = blockStarts[bi];
@@ -167,28 +203,43 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
167
203
  const isLastPass = bi === blockStarts.length - 1;
168
204
  const paramsOffset = blockIndex * stride;
169
205
 
170
- const applyBindGroup = createBindGroup(device, applyPipeline.getBindGroupLayout(0), [
171
- AinvBuffer, xBuffer, { buffer: applyParamsBuffer, offset: paramsOffset, size: 16 },
172
- ]);
206
+ const applyBindGroup = createBindGroup(
207
+ device,
208
+ applyPipeline.getBindGroupLayout(0),
209
+ [
210
+ AinvBuffer,
211
+ xBuffer,
212
+ { buffer: applyParamsBuffer, offset: paramsOffset, size: 16 },
213
+ ],
214
+ );
173
215
 
174
- const applyDesc = isLastPass && querySet
175
- ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
176
- : undefined;
216
+ const applyDesc =
217
+ isLastPass && querySet
218
+ ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
219
+ : undefined;
177
220
  encodePass(commandEncoder, applyPipeline, applyBindGroup, 1, applyDesc);
178
221
 
179
222
  const remaining = forward ? n - blockEnd : blockStart;
180
223
  if (remaining === 0) continue;
181
224
 
182
- const updateBindGroup = createBindGroup(device, updatePipeline.getBindGroupLayout(0), [
183
- ABuffer, xBuffer, { buffer: updateParamsBuffer, offset: paramsOffset, size: 32 },
184
- ]);
225
+ const updateBindGroup = createBindGroup(
226
+ device,
227
+ updatePipeline.getBindGroupLayout(0),
228
+ [
229
+ ABuffer,
230
+ xBuffer,
231
+ { buffer: updateParamsBuffer, offset: paramsOffset, size: 32 },
232
+ ],
233
+ );
185
234
 
186
235
  const wgCount = Math.min(remaining, maxWg);
187
236
  encodePass(commandEncoder, updatePipeline, updateBindGroup, wgCount);
188
237
  }
189
238
 
190
239
  const ts = resolveTimestamp(device, commandEncoder, querySet);
191
- const readBuffer = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
240
+ const readBuffer = xIsGpu
241
+ ? null
242
+ : stageReadback(device, commandEncoder, xBuffer);
192
243
 
193
244
  submit(device, commandEncoder);
194
245
 
@@ -71,9 +71,11 @@ export function resolveTimestamp(device, commandEncoder, querySet) {
71
71
  usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.MAP_READ,
72
72
  });
73
73
  commandEncoder.copyBufferToBuffer(
74
- resolveBuffer, 0, // src, srcOffset
75
- tsReadBuffer, 0, // dst, dstOffset
76
- 16, // full 16 bytes (both timestamps)
74
+ resolveBuffer,
75
+ 0, // src, srcOffset
76
+ tsReadBuffer,
77
+ 0, // dst, dstOffset
78
+ 16, // full 16 bytes (both timestamps)
77
79
  );
78
80
  // resolveBuffer is returned to prevent GC — the copy command is only encoded here, not yet executed.
79
81
  return { tsReadBuffer, resolveBuffer, querySet };
@@ -29,7 +29,7 @@ function requireStorageSize(device, byteSize, label) {
29
29
  if (byteSize > maxSize) {
30
30
  throw new Error(
31
31
  `Buffer "${label}" needs ${byteSize} bytes, exceeding this device's ` +
32
- `maxStorageBufferBindingSize (${maxSize} bytes). The operands are too large for this device.`,
32
+ `maxStorageBufferBindingSize (${maxSize} bytes). The operands are too large for this device.`,
33
33
  );
34
34
  }
35
35
  }
@@ -51,7 +51,12 @@ function requireStorageSize(device, byteSize, label) {
51
51
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUBuffer/unmap GPUBuffer.unmap()}
52
52
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUSupportedLimits GPUSupportedLimits} (`maxStorageBufferBindingSize`)
53
53
  */
54
- export function uploadBuffer(device, data, label = "blas-input", readback = false) {
54
+ export function uploadBuffer(
55
+ device,
56
+ data,
57
+ label = "blas-input",
58
+ readback = false,
59
+ ) {
55
60
  const byteSize = data.byteLength;
56
61
  requireStorageSize(device, byteSize, label);
57
62
 
@@ -86,7 +91,12 @@ export function uploadBuffer(device, data, label = "blas-input", readback = fals
86
91
  * @throws {Error} if `size` exceeds the device's `maxStorageBufferBindingSize`
87
92
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createBuffer GPUDevice.createBuffer()}
88
93
  */
89
- export function createStorageBuffer(device, size, label = "blas-storage", extraUsage = 0) {
94
+ export function createStorageBuffer(
95
+ device,
96
+ size,
97
+ label = "blas-storage",
98
+ extraUsage = 0,
99
+ ) {
90
100
  requireStorageSize(device, size, label);
91
101
  return device.createBuffer({
92
102
  label,
@@ -125,7 +135,6 @@ export function createResultBuffer(device, size, label = "blas-result") {
125
135
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/copyBufferToBuffer GPUCommandEncoder.copyBufferToBuffer()}
126
136
  */
127
137
  export function stageReadback(device, commandEncoder, sourceBuffer) {
128
-
129
138
  // COPY_DST: receives the copyBufferToBuffer transfer; MAP_READ: lets the CPU map and read it back.
130
139
  const readBuffer = device.createBuffer({
131
140
  label: "blas-readback",
@@ -134,9 +143,11 @@ export function stageReadback(device, commandEncoder, sourceBuffer) {
134
143
  });
135
144
 
136
145
  commandEncoder.copyBufferToBuffer(
137
- sourceBuffer, 0, // src, srcOffset
138
- readBuffer, 0, // dst, dstOffset
139
- sourceBuffer.size, // full copy, no partial reads
146
+ sourceBuffer,
147
+ 0, // src, srcOffset
148
+ readBuffer,
149
+ 0, // dst, dstOffset
150
+ sourceBuffer.size, // full copy, no partial reads
140
151
  );
141
152
 
142
153
  return readBuffer;
@@ -179,10 +190,17 @@ function vec4FallbackBuffer(device) {
179
190
  export function vec4ViewBinding(device, entry) {
180
191
  const buffer = entry instanceof GPUBuffer ? entry : entry.buffer;
181
192
  const offset = entry instanceof GPUBuffer ? 0 : (entry.offset ?? 0);
182
- const avail = (entry instanceof GPUBuffer ? entry.size : (entry.size ?? buffer.size - offset));
193
+ const avail =
194
+ entry instanceof GPUBuffer
195
+ ? entry.size
196
+ : (entry.size ?? buffer.size - offset);
183
197
  const size = Math.floor(avail / VEC4_ELEM_BYTES) * VEC4_ELEM_BYTES;
184
198
  if (size < VEC4_ELEM_BYTES) {
185
- return { buffer: vec4FallbackBuffer(device), offset: 0, size: VEC4_ELEM_BYTES };
199
+ return {
200
+ buffer: vec4FallbackBuffer(device),
201
+ offset: 0,
202
+ size: VEC4_ELEM_BYTES,
203
+ };
186
204
  }
187
205
  return { buffer, offset, size };
188
206
  }
@@ -206,12 +224,16 @@ export function vec4Usable(entry, stride, outerCount, innerCount) {
206
224
  if (stride % 4 !== 0) return false;
207
225
  const buffer = entry instanceof GPUBuffer ? entry : entry.buffer;
208
226
  const offset = entry instanceof GPUBuffer ? 0 : (entry.offset ?? 0);
209
- const avail = (entry instanceof GPUBuffer ? buffer.size : (entry.size ?? buffer.size - offset));
227
+ const avail =
228
+ entry instanceof GPUBuffer
229
+ ? buffer.size
230
+ : (entry.size ?? buffer.size - offset);
210
231
  const viewFloats = Math.floor(avail / VEC4_ELEM_BYTES) * 4;
211
232
  if (viewFloats <= 0) return false;
212
233
  // Highest flat index any masked-in component can touch; usable iff its
213
234
  // containing vec4 ends within the view.
214
- const maxFlat = (Math.max(outerCount, 1) - 1) * stride + (Math.max(innerCount, 1) - 1);
235
+ const maxFlat =
236
+ (Math.max(outerCount, 1) - 1) * stride + (Math.max(innerCount, 1) - 1);
215
237
  return Math.floor(maxFlat / 4) * 4 + 4 <= viewFloats;
216
238
  }
217
239
 
@@ -226,7 +248,6 @@ export function vec4Usable(entry, stride, outerCount, innerCount) {
226
248
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUQueue/writeBuffer GPUQueue.writeBuffer()}
227
249
  */
228
250
  export function createParamsBuffer(device, params, label = "blas-params") {
229
-
230
251
  const rawSize = params.length * 4;
231
252
  const size = Math.ceil(rawSize / 16) * 16;
232
253
 
@@ -0,0 +1,87 @@
1
+ /** @module devdocs/utility-functions/complex */
2
+ import { splitDoubleDouble, mergeDoubleDouble } from "./f64.mjs";
3
+ import { Complex64, Complex64Array } from "../classes/Complex64.mjs";
4
+
5
+ // Complex32Array <-> interleaved [re0, im0, re1, im1, ...] f32 buffer — one
6
+ // buffer, not a (hi, lo) pair like Float64Array (see f64.mjs), since re/im
7
+ // need no error compensation. Matches cuBLAS/stdlib's own complex layout.
8
+
9
+ /**
10
+ * Interleaves the first `n` elements of a Complex32Array into a flat f32
11
+ * buffer ready for `uploadBuffer`.
12
+ * @param {import("../classes/Complex32.mjs").Complex32Array} data
13
+ * @param {number} [n] - element count to interleave (default: data.length)
14
+ * @returns {Float32Array} length `2*n`, [re0, im0, re1, im1, ...]
15
+ * @public
16
+ */
17
+ export function interleaveComplex32(data, n = data.length) {
18
+ const flat = new Float32Array(n * 2);
19
+ for (let i = 0; i < n; i++) {
20
+ flat[i * 2] = data[i].re;
21
+ flat[i * 2 + 1] = data[i].im;
22
+ }
23
+ return flat;
24
+ }
25
+
26
+ // Complex64Array <-> a double-double (hi, lo) pair of interleaved f32
27
+ // buffers. re and im each need their own (hi, lo) split (see f64.mjs) —
28
+ // zipped together per channel, [reHi0,imHi0,...] / [reLo0,imLo0,...],
29
+ // rather than four separate buffers, so GpuVector/GpuMatrix's existing
30
+ // two-buffer shape covers this dtype with no restructuring.
31
+
32
+ /**
33
+ * Splits the first `n` elements of a Complex64Array into an interleaved
34
+ * double-double (hi, lo) pair of f32 buffers.
35
+ * @param {Complex64Array} data
36
+ * @param {number} [n] - element count to split (default: data.length)
37
+ * @returns {{hi: Float32Array, lo: Float32Array}} each length `2*n`,
38
+ * [reHi0, imHi0, reHi1, imHi1, ...] / [reLo0, imLo0, reLo1, imLo1, ...]
39
+ * @public
40
+ */
41
+ export function splitComplex64(data, n = data.length) {
42
+ const re = new Float64Array(n);
43
+ const im = new Float64Array(n);
44
+ for (let i = 0; i < n; i++) {
45
+ re[i] = data[i].re;
46
+ im[i] = data[i].im;
47
+ }
48
+ const { hi: reHi, lo: reLo } = splitDoubleDouble(re);
49
+ const { hi: imHi, lo: imLo } = splitDoubleDouble(im);
50
+
51
+ const hi = new Float32Array(n * 2);
52
+ const lo = new Float32Array(n * 2);
53
+ for (let i = 0; i < n; i++) {
54
+ hi[i * 2] = reHi[i];
55
+ hi[i * 2 + 1] = imHi[i];
56
+ lo[i * 2] = reLo[i];
57
+ lo[i * 2 + 1] = imLo[i];
58
+ }
59
+ return { hi, lo };
60
+ }
61
+
62
+ /**
63
+ * Reassembles a Complex64Array from an interleaved double-double (hi, lo)
64
+ * pair of f32 buffers — the inverse of splitComplex64.
65
+ * @param {Float32Array} hi - [reHi0, imHi0, reHi1, imHi1, ...]
66
+ * @param {Float32Array} lo - [reLo0, imLo0, reLo1, imLo1, ...]
67
+ * @returns {Complex64Array}
68
+ * @public
69
+ */
70
+ export function mergeComplex64(hi, lo) {
71
+ const n = hi.length / 2;
72
+ const reHi = new Float32Array(n),
73
+ reLo = new Float32Array(n);
74
+ const imHi = new Float32Array(n),
75
+ imLo = new Float32Array(n);
76
+ for (let i = 0; i < n; i++) {
77
+ reHi[i] = hi[i * 2];
78
+ imHi[i] = hi[i * 2 + 1];
79
+ reLo[i] = lo[i * 2];
80
+ imLo[i] = lo[i * 2 + 1];
81
+ }
82
+ const re = mergeDoubleDouble(reHi, reLo);
83
+ const im = mergeDoubleDouble(imHi, imLo);
84
+ const out = new Complex64Array(n);
85
+ for (let i = 0; i < n; i++) out[i] = new Complex64(re[i], im[i]);
86
+ return out;
87
+ }
@@ -1,9 +1,6 @@
1
1
  /** @module devdocs/utility-functions/compute */
2
2
  import { beginTimestamp, resolveTimestamp } from "./benchmark.mjs";
3
3
 
4
- // Anchors the pass encoder to its command encoder to prevent premature GC.
5
- const _passEncoders = new WeakMap();
6
-
7
4
  /**
8
5
  * Finalises `commandEncoder` into a command buffer and submits it to the GPU queue.
9
6
  * @param {GPUCommandEncoder} commandEncoder
@@ -40,7 +37,13 @@ export function beginTimedEncoder(device) {
40
37
  * @param {object} [passDescriptor] - passed to `beginComputePass`, e.g. for timestamp writes
41
38
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/beginComputePass GPUCommandEncoder.beginComputePass()}
42
39
  */
43
- export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, passDescriptor) {
40
+ export function encodePass(
41
+ commandEncoder,
42
+ pipeline,
43
+ bindGroup,
44
+ workgroups,
45
+ passDescriptor,
46
+ ) {
44
47
  const passEncoder = commandEncoder.beginComputePass(passDescriptor);
45
48
 
46
49
  passEncoder.setPipeline(pipeline);
@@ -50,12 +53,14 @@ export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, pass
50
53
  passEncoder.dispatchWorkgroups(workgroups);
51
54
  } else {
52
55
  // `?? 1` is load-bearing — confirmed passing undefined (from an {x,y}-only caller) crashes the process, not just no-ops.
53
- passEncoder.dispatchWorkgroups(workgroups.x, workgroups.y, workgroups.z ?? 1);
56
+ passEncoder.dispatchWorkgroups(
57
+ workgroups.x,
58
+ workgroups.y,
59
+ workgroups.z ?? 1,
60
+ );
54
61
  }
55
62
 
56
63
  passEncoder.end();
57
-
58
- _passEncoders.set(commandEncoder, passEncoder);
59
64
  }
60
65
 
61
66
  /**
@@ -69,7 +74,8 @@ export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, pass
69
74
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createCommandEncoder GPUDevice.createCommandEncoder()}
70
75
  */
71
76
  export function runComputePass(device, pipeline, bindGroup, workgroups) {
72
- const { commandEncoder, querySet, passDescriptor } = beginTimedEncoder(device);
77
+ const { commandEncoder, querySet, passDescriptor } =
78
+ beginTimedEncoder(device);
73
79
  encodePass(commandEncoder, pipeline, bindGroup, workgroups, passDescriptor);
74
80
 
75
81
  const ts = resolveTimestamp(device, commandEncoder, querySet);
@@ -2,6 +2,20 @@
2
2
  import { GpuVector } from "../classes/GpuVector.mjs";
3
3
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
4
4
 
5
+ /**
6
+ * Throws if `device` is not a `GPUDevice`.
7
+ *
8
+ * Every routine's first guard, extracted since the check and message are
9
+ * identical across all of them.
10
+ *
11
+ * @param {GPUDevice} device - the value to check
12
+ * @throws {Error} if `device` is not a `GPUDevice`
13
+ */
14
+ export function requireGpuDevice(device) {
15
+ if (!(device instanceof GPUDevice))
16
+ throw new Error("device must be a GPUDevice.");
17
+ }
18
+
5
19
  /**
6
20
  * Throws if any GPU-resident operand belongs to a device other than the one
7
21
  * the routine was called with.
@@ -22,12 +36,13 @@ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
22
36
  */
23
37
  export function requireSameDevice(device, routine, operands) {
24
38
  for (const [name, value] of Object.entries(operands)) {
25
- if (!(value instanceof GpuVector) && !(value instanceof GpuMatrix)) continue;
39
+ if (!(value instanceof GpuVector) && !(value instanceof GpuMatrix))
40
+ continue;
26
41
  if (value.device !== device) {
27
42
  throw new Error(
28
43
  `${routine}: ${name} belongs to a different GPUDevice than the one passed in. ` +
29
- "GPU buffers cannot be shared across devices — recreate the operand on this " +
30
- "device, or call the routine with the device that owns it.",
44
+ "GPU buffers cannot be shared across devices — recreate the operand on this " +
45
+ "device, or call the routine with the device that owns it.",
31
46
  );
32
47
  }
33
48
  }
@@ -23,7 +23,14 @@ export async function getPipeline(device, shaderName, entryPoint = "main") {
23
23
  const names = Array.isArray(shaderName) ? shaderName : [shaderName];
24
24
  const key = `${names.join("+")}::${entryPoint}`;
25
25
  if (!byName.has(key)) {
26
- byName.set(key, await loadShader(device, names, entryPoint));
26
+ // Cache the in-flight promise (not its resolved value) so a concurrent
27
+ // call awaits the same compile instead of starting a duplicate one;
28
+ // drop the entry on failure so a later call can retry.
29
+ const pending = loadShader(device, names, entryPoint).catch((err) => {
30
+ byName.delete(key);
31
+ throw err;
32
+ });
33
+ byName.set(key, pending);
27
34
  }
28
35
  return byName.get(key);
29
36
  }
@@ -39,7 +46,8 @@ async function loadCode(shaderName) {
39
46
  if (typeof process === "undefined" || !process.versions?.node) {
40
47
  const { shaderSources } = await import("../shaders/index.mjs");
41
48
  const src = shaderSources[shaderName];
42
- if (!src) throw new Error(`Shader "${shaderName}" not found in browser bundle.`);
49
+ if (!src)
50
+ throw new Error(`Shader "${shaderName}" not found in browser bundle.`);
43
51
  return src;
44
52
  } else {
45
53
  const { readFileSync } = await import("fs");
@@ -66,8 +74,32 @@ async function loadCode(shaderName) {
66
74
  */
67
75
  export async function loadShader(device, shaderNames, entryPoint = "main") {
68
76
  const label = shaderNames.join("+");
69
- const code = (await Promise.all(shaderNames.map(loadCode))).join("\n");
77
+ const codes = await Promise.all(shaderNames.map(loadCode));
70
78
 
79
+ // Per-file line ranges within the concatenated module, so a compile error
80
+ // reports as "<file>.wgsl:<local line>" instead of a whole-module line
81
+ // number — confusing for multi-file pipelines like dasum's.
82
+ let offset = 0;
83
+ const ranges = codes.map((c, i) => {
84
+ const lineCount = c.split("\n").length;
85
+ const range = {
86
+ name: shaderNames[i],
87
+ startLine: offset + 1,
88
+ endLine: offset + lineCount,
89
+ };
90
+ offset += lineCount;
91
+ return range;
92
+ });
93
+ const locate = (lineNum) => {
94
+ const range =
95
+ lineNum &&
96
+ ranges.find((r) => lineNum >= r.startLine && lineNum <= r.endLine);
97
+ return range
98
+ ? `${range.name}.wgsl:${lineNum - range.startLine + 1}`
99
+ : `line ${lineNum}`;
100
+ };
101
+
102
+ const code = codes.join("\n");
71
103
  const shaderModule = device.createShaderModule({ label, code });
72
104
 
73
105
  const info = await shaderModule.getCompilationInfo();
@@ -75,7 +107,7 @@ export async function loadShader(device, shaderNames, entryPoint = "main") {
75
107
  const errors = info.messages.filter((m) => m.type === "error");
76
108
  if (errors.length > 0) {
77
109
  throw new Error(
78
- `Shader "${label}" compilation failed:\n${errors.map((m) => ` line ${m.lineNum}: ${m.message}`).join("\n")}`,
110
+ `Shader "${label}" compilation failed:\n${errors.map((m) => ` ${locate(m.lineNum)}: ${m.message}`).join("\n")}`,
79
111
  );
80
112
  }
81
113
 
@@ -83,7 +115,10 @@ export async function loadShader(device, shaderNames, entryPoint = "main") {
83
115
  // this project's WebGPU backend is unstable (intermittent multi-minute hangs and wrong
84
116
  // results, confirmed by bisection) when entryPoint is set explicitly, even to the shader's
85
117
  // only/correct entry point. Auto-detecting the single entry point is the stable path.
86
- const compute = entryPoint === "main" ? { module: shaderModule } : { module: shaderModule, entryPoint };
118
+ const compute =
119
+ entryPoint === "main"
120
+ ? { module: shaderModule }
121
+ : { module: shaderModule, entryPoint };
87
122
  const pipeline = device.createComputePipeline({
88
123
  label,
89
124
  layout: "auto",
@@ -2,7 +2,10 @@
2
2
  // Fixed sizes match the shader declarations (WGS = 64 for 1D, 8×8 = 64 threads
3
3
  // for 2D) — see constants.mjs, which is where both values are defined and
4
4
  // where the WGSL cross-check hangs off.
5
- import { WGS as WORKGROUP_SIZE_1D, TILE_WG_2D as WORKGROUP_SIZE_2D } from "./constants.mjs";
5
+ import {
6
+ WGS as WORKGROUP_SIZE_1D,
7
+ TILE_WG_2D as WORKGROUP_SIZE_2D,
8
+ } from "./constants.mjs";
6
9
 
7
10
  /**
8
11
  * Calculates the number of workgroups to dispatch, clamped to the device's
@@ -53,8 +56,8 @@ export function requireWorkgroupCount(device, count, routine, dim = "x") {
53
56
  if (count > max)
54
57
  throw new Error(
55
58
  `${routine}: this problem needs ${count} workgroups in ${dim}, but the device allows ` +
56
- `${max} (maxComputeWorkgroupsPerDimension). The operands are too large for this device — ` +
57
- `split the operation into smaller blocks.`,
59
+ `${max} (maxComputeWorkgroupsPerDimension). The operands are too large for this device — ` +
60
+ `split the operation into smaller blocks.`,
58
61
  );
59
62
  return count;
60
63
  }
@@ -71,9 +74,23 @@ export function requireWorkgroupCount(device, count, routine, dim = "x") {
71
74
  */
72
75
  export function requireWorkgroups(device, routine, rows, cols) {
73
76
  if (cols === undefined)
74
- return requireWorkgroupCount(device, Math.ceil(rows / WORKGROUP_SIZE_1D), routine);
77
+ return requireWorkgroupCount(
78
+ device,
79
+ Math.ceil(rows / WORKGROUP_SIZE_1D),
80
+ routine,
81
+ );
75
82
  return {
76
- x: requireWorkgroupCount(device, Math.ceil(cols / WORKGROUP_SIZE_2D), routine, "x"),
77
- y: requireWorkgroupCount(device, Math.ceil(rows / WORKGROUP_SIZE_2D), routine, "y"),
83
+ x: requireWorkgroupCount(
84
+ device,
85
+ Math.ceil(cols / WORKGROUP_SIZE_2D),
86
+ routine,
87
+ "x",
88
+ ),
89
+ y: requireWorkgroupCount(
90
+ device,
91
+ Math.ceil(rows / WORKGROUP_SIZE_2D),
92
+ routine,
93
+ "y",
94
+ ),
78
95
  };
79
96
  }