wgblas 0.1.2 → 1.0.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 (55) hide show
  1. package/README.md +3 -0
  2. package/dist/wgblas.browser.js +1078 -37
  3. package/index.d.mts +6 -0
  4. package/index.mjs +6 -0
  5. package/package.json +32 -1
  6. package/src/classes/GpuMatrix.d.mts +85 -0
  7. package/src/classes/GpuMatrix.mjs +91 -0
  8. package/src/classes/GpuVector.d.mts +11 -7
  9. package/src/classes/GpuVector.mjs +31 -5
  10. package/src/dasum/dasum.d.mts +48 -0
  11. package/src/dasum/dasum.mjs +121 -0
  12. package/src/init.mjs +3 -1
  13. package/src/isamax/isamax.mjs +66 -67
  14. package/src/random/random.d.mts +37 -0
  15. package/src/random/random.mjs +17 -0
  16. package/src/sasum/sasum.mjs +55 -51
  17. package/src/saxpy/saxpy.mjs +48 -35
  18. package/src/scopy/scopy.mjs +43 -32
  19. package/src/sdot/sdot.mjs +60 -55
  20. package/src/sgemv/sgemv.d.mts +118 -0
  21. package/src/sgemv/sgemv.mjs +141 -0
  22. package/src/shaders/browser-shaders.mjs +20 -0
  23. package/src/shaders/dasum.wgsl +98 -0
  24. package/src/shaders/f64add.wgsl +281 -0
  25. package/src/shaders/isamax.wgsl +32 -9
  26. package/src/shaders/reduction/sumF64.wgsl +49 -0
  27. package/src/shaders/sasum.wgsl +18 -4
  28. package/src/shaders/sdot.wgsl +18 -4
  29. package/src/shaders/sgemv_n.wgsl +75 -0
  30. package/src/shaders/sgemv_t.wgsl +65 -0
  31. package/src/shaders/snrm2.wgsl +22 -4
  32. package/src/shaders/ssymv.wgsl +69 -0
  33. package/src/shaders/strmv.wgsl +103 -0
  34. package/src/shaders/strsv_apply_inverse.wgsl +47 -0
  35. package/src/shaders/strsv_invert_block.wgsl +109 -0
  36. package/src/shaders/strsv_update.wgsl +75 -0
  37. package/src/snrm2/snrm2.mjs +56 -52
  38. package/src/srot/srot.mjs +57 -41
  39. package/src/srotm/srotm.mjs +54 -38
  40. package/src/sscal/sscal.mjs +43 -32
  41. package/src/sswap/sswap.mjs +49 -34
  42. package/src/ssymv/ssymv.d.mts +109 -0
  43. package/src/ssymv/ssymv.mjs +130 -0
  44. package/src/strmv/strmv.d.mts +109 -0
  45. package/src/strmv/strmv.mjs +132 -0
  46. package/src/strsv/strsv.d.mts +98 -0
  47. package/src/strsv/strsv.mjs +212 -0
  48. package/src/util/benchmark.mjs +1 -1
  49. package/src/util/bindgroup.mjs +14 -10
  50. package/src/util/buffer.mjs +7 -2
  51. package/src/util/compute.mjs +41 -15
  52. package/src/util/f64pack.mjs +152 -0
  53. package/src/util/pipeline.mjs +32 -17
  54. package/src/util/result.mjs +8 -4
  55. package/src/util/workgroup.mjs +10 -10
@@ -34,82 +34,81 @@ export async function isamax(device, n, x, incx) {
34
34
  const pipelineMain = await getPipeline(device, "isamax");
35
35
  const pipelineReduce = await getPipeline(device, "reduction/argmax");
36
36
 
37
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "isamax-x", false);
38
- const partialsValBuffer = createStorageBuffer(
39
- 2 * WGS * 4,
40
- "isamax-partials-val",
41
- ); //to hold 2*WGS partial max values of f32
42
- const partialsIdxBuffer = createStorageBuffer(
43
- 2 * WGS * 4,
44
- "isamax-partials-idx",
45
- ); //to hold 2*WGS partial max indices of u32
46
- const resultBuffer = createResultBuffer(4, "isamax-result"); // u32 index
47
- const paramsBuffer = createParamsBuffer(
48
- [
49
- { value: n, type: "u32" },
50
- { value: incx, type: "u32" },
51
- ],
52
- "isamax-params",
53
- );
37
+ let xBuffer = null;
38
+ let partialsValBuffer = null;
39
+ let partialsIdxBuffer = null;
40
+ let resultBuffer = null;
41
+ let paramsBuffer = null;
42
+ let readBuffer = null;
54
43
 
55
- const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
56
- xBuffer,
57
- partialsValBuffer,
58
- partialsIdxBuffer,
59
- paramsBuffer,
60
- ]);
61
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(
62
- pipelineMain,
63
- bgMain,
64
- 2 * WGS,
65
- ); //dispatch 2*WGS workgroups
66
-
67
- submit(enc1);
68
-
69
- const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
70
- partialsValBuffer,
71
- partialsIdxBuffer,
72
- resultBuffer,
73
- ]);
74
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(
75
- pipelineReduce,
76
- bgReduce,
77
- 1,
78
- ); // dispatch 1 workgroup to reduce the partial sums to a single result
79
- const readBuffer = stageReadback(enc2, resultBuffer);
80
-
81
- submit(enc2);
44
+ try {
45
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "isamax-x", false);
46
+ partialsValBuffer = createStorageBuffer(
47
+ 2 * WGS * 4,
48
+ "isamax-partials-val",
49
+ ); //to hold 2*WGS partial max values of f32
50
+ partialsIdxBuffer = createStorageBuffer(
51
+ 2 * WGS * 4,
52
+ "isamax-partials-idx",
53
+ ); //to hold 2*WGS partial max indices of u32
54
+ resultBuffer = createResultBuffer(4, "isamax-result"); // u32 index
55
+ paramsBuffer = createParamsBuffer(
56
+ [
57
+ { value: n, type: "u32" },
58
+ { value: incx, type: "u32" },
59
+ ],
60
+ "isamax-params",
61
+ );
82
62
 
83
- const [gpuTime1, gpuTime2, idxArr] = await Promise.all([
84
- extractTimestamp(ts1),
85
- extractTimestamp(ts2),
86
- extractResult(readBuffer, Uint32Array),
87
- ]);
63
+ const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
64
+ xBuffer,
65
+ partialsValBuffer,
66
+ partialsIdxBuffer,
67
+ paramsBuffer,
68
+ ]);
69
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
70
+ pipelineMain,
71
+ bgMain,
72
+ 2 * WGS,
73
+ ); //dispatch 2*WGS workgroups
88
74
 
89
- const index = idxArr[0];
75
+ submit(enc1);
90
76
 
91
- if (xIsGpu) {
92
- destroyBuffers(
77
+ const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
93
78
  partialsValBuffer,
94
79
  partialsIdxBuffer,
95
80
  resultBuffer,
96
- paramsBuffer,
97
- readBuffer,
98
- );
81
+ ]);
82
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
83
+ pipelineReduce,
84
+ bgReduce,
85
+ 1,
86
+ ); // dispatch 1 workgroup to reduce the partial sums to a single result
87
+ readBuffer = stageReadback(enc2, resultBuffer);
88
+
89
+ submit(enc2);
90
+
91
+ const resultPromise = extractResult(readBuffer, Uint32Array);
92
+ readBuffer = null; // ownership transferred — extractResult's own finally destroys it
93
+
94
+ const [gpuTime1, gpuTime2, idxArr] = await Promise.all([
95
+ extractTimestamp(ts1),
96
+ extractTimestamp(ts2),
97
+ resultPromise,
98
+ ]);
99
+
100
+ const index = idxArr[0];
101
+
99
102
  if (gpuTime1 !== undefined && gpuTime2 !== undefined)
100
103
  return { index, gpuTimeMs: gpuTime1 + gpuTime2 };
101
104
  return { index };
105
+ } finally {
106
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
107
+ if (partialsValBuffer) destroyBuffers(partialsValBuffer);
108
+ if (partialsIdxBuffer) destroyBuffers(partialsIdxBuffer);
109
+ if (resultBuffer) destroyBuffers(resultBuffer);
110
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
111
+ // Only reached if submit(enc2) threw before ownership was transferred above.
112
+ if (readBuffer) destroyBuffers(readBuffer);
102
113
  }
103
-
104
- destroyBuffers(
105
- xBuffer,
106
- partialsValBuffer,
107
- partialsIdxBuffer,
108
- resultBuffer,
109
- paramsBuffer,
110
- readBuffer,
111
- );
112
- if (gpuTime1 !== undefined && gpuTime2 !== undefined)
113
- return { index, gpuTimeMs: gpuTime1 + gpuTime2 };
114
- return { index };
115
114
  }
@@ -59,3 +59,40 @@ export declare function randomFloat64Array(
59
59
  low?: number,
60
60
  high?: number,
61
61
  ): Float64Array;
62
+
63
+ /**
64
+ * Returns a Float32Array of n*lda elements, row-major with leading dimension
65
+ * `lda`, representing an actual lower- or upper-triangular matrix for a
66
+ * triangular routine (`strmv`/`strsv`) example: entries in the `uplo`
67
+ * triangle are uniform in `[low, high)`, the n diagonal entries
68
+ * (`A[i*lda+i]`) are uniform in `[diagLow, diagHigh)` — kept well away from 0
69
+ * so a triangular solve doesn't divide by a near-zero pivot — and every entry
70
+ * in the other triangle is 0.
71
+ *
72
+ * @param n - matrix order (rows/cols read by the triangular routine)
73
+ * @param lda - leading dimension; throws if `lda < n`
74
+ * @param uplo - `'lower'` to fill the lower triangle, `'upper'` to fill the upper triangle (default: `'lower'`)
75
+ * @param low - lower bound for off-diagonal entries (default: -1)
76
+ * @param high - upper bound for off-diagonal entries (default: 1)
77
+ * @param diagLow - lower bound for diagonal entries (default: 5)
78
+ * @param diagHigh - upper bound for diagonal entries (default: 15)
79
+ *
80
+ * @example
81
+ * ```js
82
+ * import { randomTriangularFloat32Array } from "wgblas";
83
+ *
84
+ * const n = 4, lda = n;
85
+ * const A = randomTriangularFloat32Array(n, lda, "lower");
86
+ * ```
87
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/random/random.mjs#L13">Source code: random.mjs (L13)</a>
88
+ * @category Utilities
89
+ */
90
+ export declare function randomTriangularFloat32Array(
91
+ n: number,
92
+ lda: number,
93
+ uplo?: 'lower' | 'upper',
94
+ low?: number,
95
+ high?: number,
96
+ diagLow?: number,
97
+ diagHigh?: number,
98
+ ): Float32Array;
@@ -9,3 +9,20 @@ export function randomFloat64Array(n, low = -1, high = 1) {
9
9
  for (let i = 0; i < n; i++) x[i] = low + Math.random() * (high - low);
10
10
  return x;
11
11
  }
12
+
13
+ export function randomTriangularFloat32Array(n, lda, uplo = "lower", low = -1, high = 1, diagLow = 5, diagHigh = 15) {
14
+ if (uplo !== "lower" && uplo !== "upper")
15
+ throw new Error("uplo must be 'lower' or 'upper'.");
16
+ if (lda < n) throw new Error("lda must be >= n.");
17
+
18
+ const A = new Float32Array(n * lda);
19
+ for (let i = 0; i < n; i++) {
20
+ for (let j = 0; j < n; j++) {
21
+ if (i === j) continue;
22
+ const inTriangle = uplo === "lower" ? j < i : j > i;
23
+ if (inTriangle) A[i * lda + j] = low + Math.random() * (high - low);
24
+ }
25
+ A[i * lda + i] = diagLow + Math.random() * (diagHigh - diagLow);
26
+ }
27
+ return A;
28
+ }
@@ -34,65 +34,69 @@ export async function sasum(device, n, x, incx) {
34
34
  const pipelineMain = await getPipeline(device, "sasum");
35
35
  const pipelineReduce = await getPipeline(device, "reduction/sum");
36
36
 
37
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sasum-x", false);
38
- const partialsBuffer = createStorageBuffer(2 * WGS * 4, "sasum-partials"); // 2*WGS partial sums of f32
39
- const resultBuffer = createResultBuffer(4, "sasum-result"); // final f32 scalar
40
- const paramsBuffer = createParamsBuffer(
41
- [
42
- { value: n, type: "u32" },
43
- { value: incx, type: "u32" },
44
- ],
45
- "sasum-params",
46
- );
37
+ let xBuffer = null;
38
+ let partialsBuffer = null;
39
+ let resultBuffer = null;
40
+ let paramsBuffer = null;
41
+ let readBuffer = null;
47
42
 
48
- const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
49
- xBuffer,
50
- partialsBuffer,
51
- paramsBuffer,
52
- ]);
53
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(
54
- pipelineMain,
55
- bgMain,
56
- 2 * WGS,
57
- ); // dispatch 2*WGS workgroups
43
+ try {
44
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sasum-x", false);
45
+ partialsBuffer = createStorageBuffer(2 * WGS * 4, "sasum-partials"); // 2*WGS partial sums of f32
46
+ resultBuffer = createResultBuffer(4, "sasum-result"); // final f32 scalar
47
+ paramsBuffer = createParamsBuffer(
48
+ [
49
+ { value: n, type: "u32" },
50
+ { value: incx, type: "u32" },
51
+ ],
52
+ "sasum-params",
53
+ );
54
+
55
+ const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
56
+ xBuffer,
57
+ partialsBuffer,
58
+ paramsBuffer,
59
+ ]);
60
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
61
+ pipelineMain,
62
+ bgMain,
63
+ 2 * WGS,
64
+ ); // dispatch 2*WGS workgroups
58
65
 
59
- submit(enc1);
66
+ submit(enc1);
60
67
 
61
- const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
62
- partialsBuffer,
63
- resultBuffer,
64
- ]);
65
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(
66
- pipelineReduce,
67
- bgReduce,
68
- 1,
69
- ); // dispatch 1 workgroup to reduce the partial sums to a single result
70
- const readBuffer = stageReadback(enc2, resultBuffer);
68
+ const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
69
+ partialsBuffer,
70
+ resultBuffer,
71
+ ]);
72
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
73
+ pipelineReduce,
74
+ bgReduce,
75
+ 1,
76
+ ); // dispatch 1 workgroup to reduce the partial sums to a single result
77
+ readBuffer = stageReadback(enc2, resultBuffer);
71
78
 
72
- submit(enc2);
79
+ submit(enc2);
73
80
 
74
- const [gpuTime1, gpuTime2, asumArr] = await Promise.all([
75
- extractTimestamp(ts1),
76
- extractTimestamp(ts2),
77
- extractResult(readBuffer, Float32Array),
78
- ]);
81
+ const resultPromise = extractResult(readBuffer, Float32Array);
82
+ readBuffer = null; // ownership transferred — extractResult's own finally destroys it
79
83
 
80
- // asum is always a scalar readback — both paths return { asum }
81
- if (xIsGpu) {
82
- destroyBuffers(partialsBuffer, resultBuffer, paramsBuffer, readBuffer);
84
+ const [gpuTime1, gpuTime2, asumArr] = await Promise.all([
85
+ extractTimestamp(ts1),
86
+ extractTimestamp(ts2),
87
+ resultPromise,
88
+ ]);
89
+
90
+ // asum is always a scalar readback — both paths return { asum }
83
91
  if (gpuTime1 !== undefined && gpuTime2 !== undefined)
84
92
  return { asum: asumArr[0], gpuTimeMs: gpuTime1 + gpuTime2 };
85
93
  return { asum: asumArr[0] };
94
+ } finally {
95
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
96
+ if (partialsBuffer) destroyBuffers(partialsBuffer);
97
+ if (resultBuffer) destroyBuffers(resultBuffer);
98
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
99
+ // Only reached if submit(enc2) threw before ownership was transferred above.
100
+ if (readBuffer) destroyBuffers(readBuffer);
86
101
  }
87
-
88
- destroyBuffers(
89
- xBuffer,
90
- partialsBuffer,
91
- resultBuffer,
92
- paramsBuffer,
93
- readBuffer,
94
- );
95
- if (gpuTime1 !== undefined && gpuTime2 !== undefined)
96
- return { asum: asumArr[0], gpuTimeMs: gpuTime1 + gpuTime2 };
97
- return { asum: asumArr[0] };
98
102
  }
@@ -24,8 +24,10 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
24
24
  !Number.isInteger(incy)
25
25
  )
26
26
  throw new Error("n, incx, and incy must be integers.");
27
- if (isNaN(alpha)) throw new Error("alpha must not be NaN.");
28
- if (!isFinite(alpha)) throw new Error("alpha must be finite.");
27
+ if (typeof alpha !== "number")
28
+ throw new Error("alpha must be a number.");
29
+ if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
30
+ if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
29
31
  if (incx <= 0 || incy <= 0)
30
32
  throw new Error("incx and incy must be positive.");
31
33
  if (!xIsGpu && !(x instanceof Float32Array))
@@ -48,43 +50,54 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
48
50
 
49
51
  const pipeline = await getPipeline(device, "saxpy");
50
52
 
51
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "saxpy-x", false);
52
- const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "saxpy-y", true);
53
- const paramsBuffer = createParamsBuffer(
54
- [
55
- { value: n, type: "u32" },
56
- { value: alpha, type: "f32" },
57
- { value: incx, type: "u32" },
58
- { value: incy, type: "u32" },
59
- ],
60
- "saxpy-params",
61
- );
53
+ let xBuffer = null;
54
+ let yBuffer = null;
55
+ let paramsBuffer = null;
56
+ let readBuffer = null;
62
57
 
63
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
64
- xBuffer,
65
- yBuffer,
66
- paramsBuffer,
67
- ]);
68
- const { commandEncoder, ts } = runComputePass(
69
- pipeline,
70
- bindGroup,
71
- calcWorkgroups(n),
72
- );
73
- const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
58
+ try {
59
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "saxpy-x", false);
60
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "saxpy-y", true);
61
+ paramsBuffer = createParamsBuffer(
62
+ [
63
+ { value: n, type: "u32" },
64
+ { value: alpha, type: "f32" },
65
+ { value: incx, type: "u32" },
66
+ { value: incy, type: "u32" },
67
+ ],
68
+ "saxpy-params",
69
+ );
74
70
 
75
- submit(commandEncoder);
71
+ const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
72
+ xBuffer,
73
+ yBuffer,
74
+ paramsBuffer,
75
+ ]);
76
+ const { commandEncoder, ts } = runComputePass(
77
+ pipeline,
78
+ bindGroup,
79
+ calcWorkgroups(n),
80
+ );
81
+ readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
76
82
 
77
- const gpuTimeMs = await extractTimestamp(ts);
83
+ submit(commandEncoder);
78
84
 
79
- if (yIsGpu && xIsGpu) {
80
- destroyBuffers(paramsBuffer);
81
- if (gpuTimeMs !== undefined) return { gpuTimeMs };
82
- return {};
83
- }
85
+ const gpuTimeMs = await extractTimestamp(ts);
84
86
 
85
- const result = await extractResult(readBuffer, Float32Array);
86
- destroyBuffers(xBuffer, yBuffer, paramsBuffer, readBuffer);
87
+ if (yIsGpu && xIsGpu) {
88
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
89
+ return {};
90
+ }
87
91
 
88
- if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
89
- return { y: result };
92
+ const result = await extractResult(readBuffer, Float32Array);
93
+ readBuffer = null; // extractResult already destroyed it
94
+ if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
95
+ return { y: result };
96
+ } finally {
97
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
98
+ if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
99
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
100
+ // Only reached if extractTimestamp threw before extractResult ran.
101
+ if (readBuffer) destroyBuffers(readBuffer);
102
+ }
90
103
  }
@@ -46,42 +46,53 @@ export async function scopy(device, n, x, incx, y, incy) {
46
46
 
47
47
  const pipeline = await getPipeline(device, "scopy");
48
48
 
49
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "scopy-x", false);
50
- const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "scopy-y", true);
51
- const paramsBuffer = createParamsBuffer(
52
- [
53
- { value: n, type: "u32" },
54
- { value: incx, type: "u32" },
55
- { value: incy, type: "u32" },
56
- ],
57
- "scopy-params",
58
- );
49
+ let xBuffer = null;
50
+ let yBuffer = null;
51
+ let paramsBuffer = null;
52
+ let readBuffer = null;
59
53
 
60
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
61
- xBuffer,
62
- yBuffer,
63
- paramsBuffer,
64
- ]);
65
- const { commandEncoder, ts } = runComputePass(
66
- pipeline,
67
- bindGroup,
68
- calcWorkgroups(n),
69
- );
70
- const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
54
+ try {
55
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "scopy-x", false);
56
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "scopy-y", true);
57
+ paramsBuffer = createParamsBuffer(
58
+ [
59
+ { value: n, type: "u32" },
60
+ { value: incx, type: "u32" },
61
+ { value: incy, type: "u32" },
62
+ ],
63
+ "scopy-params",
64
+ );
71
65
 
72
- submit(commandEncoder);
66
+ const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
67
+ xBuffer,
68
+ yBuffer,
69
+ paramsBuffer,
70
+ ]);
71
+ const { commandEncoder, ts } = runComputePass(
72
+ pipeline,
73
+ bindGroup,
74
+ calcWorkgroups(n),
75
+ );
76
+ readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
73
77
 
74
- const gpuTimeMs = await extractTimestamp(ts);
78
+ submit(commandEncoder);
75
79
 
76
- if (yIsGpu && xIsGpu) {
77
- destroyBuffers(paramsBuffer);
78
- if (gpuTimeMs !== undefined) return { gpuTimeMs };
79
- return {};
80
- }
80
+ const gpuTimeMs = await extractTimestamp(ts);
81
81
 
82
- const result = await extractResult(readBuffer, Float32Array);
83
- destroyBuffers(xBuffer, yBuffer, paramsBuffer, readBuffer);
82
+ if (yIsGpu && xIsGpu) {
83
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
84
+ return {};
85
+ }
84
86
 
85
- if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
86
- return { y: result };
87
+ const result = await extractResult(readBuffer, Float32Array);
88
+ readBuffer = null; // extractResult already destroyed it
89
+ if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
90
+ return { y: result };
91
+ } finally {
92
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
93
+ if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
94
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
95
+ // Only reached if extractTimestamp threw before extractResult ran.
96
+ if (readBuffer) destroyBuffers(readBuffer);
97
+ }
87
98
  }
package/src/sdot/sdot.mjs CHANGED
@@ -50,69 +50,74 @@ export async function sdot(device, n, x, incx, y, incy) {
50
50
  const pipelineMain = await getPipeline(device, "sdot");
51
51
  const pipelineReduce = await getPipeline(device, "reduction/sum");
52
52
 
53
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sdot-x", false);
54
- const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sdot-y", false);
55
- const partialsBuffer = createStorageBuffer(2 * WGS * 4, "sdot-partials"); //to hold 2*WGS partial sums of f32
56
- const resultBuffer = createResultBuffer(4, "sdot-result"); //to hold the final float32 dot product
57
- const paramsBuffer = createParamsBuffer(
58
- [
59
- { value: n, type: "u32" },
60
- { value: incx, type: "u32" },
61
- { value: incy, type: "u32" },
62
- ],
63
- "sdot-params",
64
- );
53
+ let xBuffer = null;
54
+ let yBuffer = null;
55
+ let partialsBuffer = null;
56
+ let resultBuffer = null;
57
+ let paramsBuffer = null;
58
+ let readBuffer = null;
65
59
 
66
- const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
67
- xBuffer,
68
- yBuffer,
69
- partialsBuffer,
70
- paramsBuffer,
71
- ]);
72
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(
73
- pipelineMain,
74
- bgMain,
75
- 2 * WGS,
76
- ); //dispatch 2*WGS workgroups
60
+ try {
61
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sdot-x", false);
62
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sdot-y", false);
63
+ partialsBuffer = createStorageBuffer(2 * WGS * 4, "sdot-partials"); //to hold 2*WGS partial sums of f32
64
+ resultBuffer = createResultBuffer(4, "sdot-result"); //to hold the final float32 dot product
65
+ paramsBuffer = createParamsBuffer(
66
+ [
67
+ { value: n, type: "u32" },
68
+ { value: incx, type: "u32" },
69
+ { value: incy, type: "u32" },
70
+ ],
71
+ "sdot-params",
72
+ );
73
+
74
+ const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
75
+ xBuffer,
76
+ yBuffer,
77
+ partialsBuffer,
78
+ paramsBuffer,
79
+ ]);
80
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
81
+ pipelineMain,
82
+ bgMain,
83
+ 2 * WGS,
84
+ ); //dispatch 2*WGS workgroups
77
85
 
78
- submit(enc1);
86
+ submit(enc1);
79
87
 
80
- const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
81
- partialsBuffer,
82
- resultBuffer,
83
- ]);
84
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(
85
- pipelineReduce,
86
- bgReduce,
87
- 1,
88
- ); // dispatch 1 workgroup to reduce the partial sums to a single result
89
- const readBuffer = stageReadback(enc2, resultBuffer);
88
+ const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
89
+ partialsBuffer,
90
+ resultBuffer,
91
+ ]);
92
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
93
+ pipelineReduce,
94
+ bgReduce,
95
+ 1,
96
+ ); // dispatch 1 workgroup to reduce the partial sums to a single result
97
+ readBuffer = stageReadback(enc2, resultBuffer);
90
98
 
91
- submit(enc2);
99
+ submit(enc2);
92
100
 
93
- const [gpuTime1, gpuTime2, dotArr] = await Promise.all([
94
- extractTimestamp(ts1),
95
- extractTimestamp(ts2),
96
- extractResult(readBuffer, Float32Array),
97
- ]);
101
+ const resultPromise = extractResult(readBuffer, Float32Array);
102
+ readBuffer = null; // ownership transferred — extractResult's own finally destroys it
98
103
 
99
- // dot is always a scalar readback — both paths return { dot }
100
- if (xIsGpu && yIsGpu) {
101
- destroyBuffers(partialsBuffer, resultBuffer, paramsBuffer, readBuffer);
104
+ const [gpuTime1, gpuTime2, dotArr] = await Promise.all([
105
+ extractTimestamp(ts1),
106
+ extractTimestamp(ts2),
107
+ resultPromise,
108
+ ]);
109
+
110
+ // dot is always a scalar readback — both paths return { dot }
102
111
  if (gpuTime1 !== undefined && gpuTime2 !== undefined)
103
112
  return { dot: dotArr[0], gpuTimeMs: gpuTime1 + gpuTime2 };
104
113
  return { dot: dotArr[0] };
114
+ } finally {
115
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
116
+ if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
117
+ if (partialsBuffer) destroyBuffers(partialsBuffer);
118
+ if (resultBuffer) destroyBuffers(resultBuffer);
119
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
120
+ // Only reached if submit(enc2) threw before ownership was transferred above.
121
+ if (readBuffer) destroyBuffers(readBuffer);
105
122
  }
106
-
107
- destroyBuffers(
108
- xBuffer,
109
- yBuffer,
110
- partialsBuffer,
111
- resultBuffer,
112
- paramsBuffer,
113
- readBuffer,
114
- );
115
- if (gpuTime1 !== undefined && gpuTime2 !== undefined)
116
- return { dot: dotArr[0], gpuTimeMs: gpuTime1 + gpuTime2 };
117
- return { dot: dotArr[0] };
118
123
  }