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
@@ -0,0 +1,75 @@
1
+ // strsv_update: subtracts a solved block's contribution from every
2
+ // remaining row in parallel (one workgroup per row, like strmv.wgsl) —
3
+ // this is what turns strsv's O(n) sequential stages into O(n/blockSize).
4
+ // No diag/masking needed: this region never touches the diagonal.
5
+
6
+ @group(0) @binding(0) var<storage, read> A: array<f32>;
7
+ @group(0) @binding(1) var<storage, read_write> x: array<f32>;
8
+
9
+ struct Params {
10
+ n: u32,
11
+ incx: u32,
12
+ lda: u32,
13
+ trans: u32, // 0 = no-transpose, 1 = transpose
14
+ uplo: u32, // 0 = lower, 1 = upper
15
+ blockStart: u32,
16
+ blockEnd: u32, // exclusive
17
+ }
18
+
19
+ @group(0) @binding(2) var<uniform> params: Params;
20
+
21
+ const WGS: u32 = 64u;
22
+ var<workgroup> scratch: array<f32, 64>;
23
+
24
+ @compute @workgroup_size(64)
25
+ fn strsv_update_main(
26
+ @builtin(workgroup_id) wgid: vec3u,
27
+ @builtin(local_invocation_id) lid: vec3u,
28
+ @builtin(num_workgroups) nwg: vec3u,
29
+ ) {
30
+ // forward: remaining rows are [blockEnd,n); backward: [0,blockStart).
31
+ let forward = (params.trans == 0u) == (params.uplo == 0u);
32
+
33
+ var rangeStart: u32;
34
+ var rangeEnd: u32;
35
+
36
+ if forward {
37
+ rangeStart = params.blockEnd;
38
+ rangeEnd = params.n;
39
+ } else {
40
+ rangeStart = 0u;
41
+ rangeEnd = params.blockStart;
42
+ }
43
+
44
+ if (rangeStart >= rangeEnd) { return; }
45
+ let count = rangeEnd - rangeStart;
46
+
47
+ for (var idx = wgid.x; idx < count; idx += nwg.x) {
48
+ let i = rangeStart + idx;
49
+
50
+ // No-trans reads A[i,j]; transpose reads A[j,i] — uplo only sets the range above.
51
+ var acc = 0.0f;
52
+ if params.trans == 0u {
53
+ for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
54
+ acc += A[i * params.lda + j] * x[j * params.incx];
55
+ }
56
+ } else {
57
+ for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
58
+ acc += A[j * params.lda + i] * x[j * params.incx];
59
+ }
60
+ }
61
+
62
+ // Parallel reduction: 64 → 1
63
+ scratch[lid.x] = acc;
64
+ workgroupBarrier();
65
+ for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
66
+ if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
67
+ workgroupBarrier();
68
+ }
69
+
70
+ if lid.x == 0u {
71
+ x[i * params.incx] -= scratch[0];
72
+ }
73
+ workgroupBarrier();
74
+ }
75
+ }
@@ -34,67 +34,71 @@ export async function snrm2(device, n, x, incx) {
34
34
  const pipelineMain = await getPipeline(device, "snrm2");
35
35
  const pipelineReduce = await getPipeline(device, "reduction/sum");
36
36
 
37
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "snrm2-x", false);
38
- const partialsBuffer = createStorageBuffer(2 * WGS * 4, "snrm2-partials"); // 2*WGS partial sums of f32
39
- const resultBuffer = createResultBuffer(4, "snrm2-result"); // final f32 scalar
40
- const paramsBuffer = createParamsBuffer(
41
- [
42
- { value: n, type: "u32" },
43
- { value: incx, type: "u32" },
44
- ],
45
- "snrm2-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, "snrm2-x", false);
45
+ partialsBuffer = createStorageBuffer(2 * WGS * 4, "snrm2-partials"); // 2*WGS partial sums of f32
46
+ resultBuffer = createResultBuffer(4, "snrm2-result"); // final f32 scalar
47
+ paramsBuffer = createParamsBuffer(
48
+ [
49
+ { value: n, type: "u32" },
50
+ { value: incx, type: "u32" },
51
+ ],
52
+ "snrm2-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
- ); // reduce partials to 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
+ ); // reduce partials to single result
77
+ readBuffer = stageReadback(enc2, resultBuffer);
71
78
 
72
- submit(enc2);
79
+ submit(enc2);
73
80
 
74
- const [gpuTime1, gpuTime2, sqsumArr] = 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
- // sqrt is taken on CPU after the GPU sum-of-squares reduction
81
- const nrm2 = Math.sqrt(sqsumArr[0]);
84
+ const [gpuTime1, gpuTime2, sqsumArr] = await Promise.all([
85
+ extractTimestamp(ts1),
86
+ extractTimestamp(ts2),
87
+ resultPromise,
88
+ ]);
89
+
90
+ // sqrt is taken on CPU after the GPU sum-of-squares reduction
91
+ const nrm2 = Math.sqrt(sqsumArr[0]);
82
92
 
83
- if (xIsGpu) {
84
- destroyBuffers(partialsBuffer, resultBuffer, paramsBuffer, readBuffer);
85
93
  if (gpuTime1 !== undefined && gpuTime2 !== undefined)
86
94
  return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
87
95
  return { nrm2 };
96
+ } finally {
97
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
98
+ if (partialsBuffer) destroyBuffers(partialsBuffer);
99
+ if (resultBuffer) destroyBuffers(resultBuffer);
100
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
101
+ // Only reached if submit(enc2) threw before ownership was transferred above.
102
+ if (readBuffer) destroyBuffers(readBuffer);
88
103
  }
89
-
90
- destroyBuffers(
91
- xBuffer,
92
- partialsBuffer,
93
- resultBuffer,
94
- paramsBuffer,
95
- readBuffer,
96
- );
97
- if (gpuTime1 !== undefined && gpuTime2 !== undefined)
98
- return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
99
- return { nrm2 };
100
104
  }
package/src/srot/srot.mjs CHANGED
@@ -24,9 +24,11 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
24
24
  !Number.isInteger(incy)
25
25
  )
26
26
  throw new Error("n, incx, and incy must be integers.");
27
- if (isNaN(c) || isNaN(s)) throw new Error("c and s must not be NaN.");
28
- if (!isFinite(c)) throw new Error("c must be finite.");
29
- if (!isFinite(s)) throw new Error("s must be finite.");
27
+ if (typeof c !== "number") throw new Error("c must be a number.");
28
+ if (typeof s !== "number") throw new Error("s must be a number.");
29
+ if (Number.isNaN(c) || Number.isNaN(s)) throw new Error("c and s must not be NaN.");
30
+ if (!Number.isFinite(c)) throw new Error("c must be finite.");
31
+ if (!Number.isFinite(s)) throw new Error("s must be finite.");
30
32
  if (incx <= 0 || incy <= 0)
31
33
  throw new Error("incx and incy must be positive.");
32
34
  if (!xIsGpu && !(x instanceof Float32Array))
@@ -49,47 +51,61 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
49
51
 
50
52
  const pipeline = await getPipeline(device, "srot");
51
53
 
52
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srot-x", true);
53
- const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srot-y", true);
54
- const paramsBuffer = createParamsBuffer(
55
- [
56
- { value: n, type: "u32" },
57
- { value: c, type: "f32" },
58
- { value: s, type: "f32" },
59
- { value: incx, type: "u32" },
60
- { value: incy, type: "u32" },
61
- ],
62
- "srot-params",
63
- );
54
+ let xBuffer = null;
55
+ let yBuffer = null;
56
+ let paramsBuffer = null;
57
+ let readX = null;
58
+ let readY = null;
64
59
 
65
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
66
- xBuffer,
67
- yBuffer,
68
- paramsBuffer,
69
- ]);
70
- const { commandEncoder, ts } = runComputePass(
71
- pipeline,
72
- bindGroup,
73
- calcWorkgroups(n),
74
- );
75
- const readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
76
- const readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
77
- submit(commandEncoder);
60
+ try {
61
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srot-x", true);
62
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srot-y", true);
63
+ paramsBuffer = createParamsBuffer(
64
+ [
65
+ { value: n, type: "u32" },
66
+ { value: c, type: "f32" },
67
+ { value: s, type: "f32" },
68
+ { value: incx, type: "u32" },
69
+ { value: incy, type: "u32" },
70
+ ],
71
+ "srot-params",
72
+ );
78
73
 
79
- const gpuTimeMs = await extractTimestamp(ts);
74
+ const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
75
+ xBuffer,
76
+ yBuffer,
77
+ paramsBuffer,
78
+ ]);
79
+ const { commandEncoder, ts } = runComputePass(
80
+ pipeline,
81
+ bindGroup,
82
+ calcWorkgroups(n),
83
+ );
84
+ readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
85
+ readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
86
+ submit(commandEncoder);
80
87
 
81
- if (xIsGpu && yIsGpu) {
82
- destroyBuffers(paramsBuffer);
83
- if (gpuTimeMs !== undefined) return { gpuTimeMs };
84
- return {};
85
- }
88
+ const gpuTimeMs = await extractTimestamp(ts);
86
89
 
87
- const [xResult, yResult] = await Promise.all([
88
- extractResult(readX, Float32Array),
89
- extractResult(readY, Float32Array),
90
- ]);
91
- destroyBuffers(xBuffer, yBuffer, paramsBuffer, readX, readY);
90
+ if (xIsGpu && yIsGpu) {
91
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
92
+ return {};
93
+ }
92
94
 
93
- if (gpuTimeMs !== undefined) return { x: xResult, y: yResult, gpuTimeMs };
94
- return { x: xResult, y: yResult };
95
+ const xPromise = extractResult(readX, Float32Array);
96
+ const yPromise = extractResult(readY, Float32Array);
97
+ readX = null; // ownership transferred — extractResult's own finally destroys it
98
+ readY = null;
99
+ const [xResult, yResult] = await Promise.all([xPromise, yPromise]);
100
+
101
+ if (gpuTimeMs !== undefined) return { x: xResult, y: yResult, gpuTimeMs };
102
+ return { x: xResult, y: yResult };
103
+ } finally {
104
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
105
+ if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
106
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
107
+ // Only reached if extractTimestamp threw before ownership was transferred above.
108
+ if (readX) destroyBuffers(readX);
109
+ if (readY) destroyBuffers(readY);
110
+ }
95
111
  }
@@ -48,47 +48,63 @@ export async function srotm(device, n, x, incx, y, incy, param) {
48
48
 
49
49
  const pipeline = await getPipeline(device, "srotm");
50
50
 
51
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srotm-x", true);
52
- const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srotm-y", true);
53
- const paramBuffer = uploadBuffer(param, "srotm-param", false);
54
- const paramsBuffer = createParamsBuffer(
55
- [
56
- { value: n, type: "u32" },
57
- { value: incx, type: "u32" },
58
- { value: incy, type: "u32" },
59
- ],
60
- "srotm-params",
61
- );
51
+ let xBuffer = null;
52
+ let yBuffer = null;
53
+ let paramBuffer = null;
54
+ let paramsBuffer = null;
55
+ let readX = null;
56
+ let readY = null;
62
57
 
63
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
64
- xBuffer,
65
- yBuffer,
66
- paramBuffer,
67
- paramsBuffer,
68
- ]);
69
- const { commandEncoder, ts } = runComputePass(
70
- pipeline,
71
- bindGroup,
72
- calcWorkgroups(n),
73
- );
74
- const readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
75
- const readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
76
- submit(commandEncoder);
58
+ try {
59
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srotm-x", true);
60
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srotm-y", true);
61
+ paramBuffer = uploadBuffer(param, "srotm-param", false);
62
+ paramsBuffer = createParamsBuffer(
63
+ [
64
+ { value: n, type: "u32" },
65
+ { value: incx, type: "u32" },
66
+ { value: incy, type: "u32" },
67
+ ],
68
+ "srotm-params",
69
+ );
77
70
 
78
- const gpuTimeMs = await extractTimestamp(ts);
71
+ const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
72
+ xBuffer,
73
+ yBuffer,
74
+ paramBuffer,
75
+ paramsBuffer,
76
+ ]);
77
+ const { commandEncoder, ts } = runComputePass(
78
+ pipeline,
79
+ bindGroup,
80
+ calcWorkgroups(n),
81
+ );
82
+ readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
83
+ readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
84
+ submit(commandEncoder);
79
85
 
80
- if (xIsGpu && yIsGpu) {
81
- destroyBuffers(paramBuffer, paramsBuffer);
82
- if (gpuTimeMs !== undefined) return { gpuTimeMs };
83
- return {};
84
- }
86
+ const gpuTimeMs = await extractTimestamp(ts);
85
87
 
86
- const [xResult, yResult] = await Promise.all([
87
- extractResult(readX, Float32Array),
88
- extractResult(readY, Float32Array),
89
- ]);
90
- destroyBuffers(xBuffer, yBuffer, paramBuffer, paramsBuffer, readX, readY);
88
+ if (xIsGpu && yIsGpu) {
89
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
90
+ return {};
91
+ }
91
92
 
92
- if (gpuTimeMs !== undefined) return { x: xResult, y: yResult, gpuTimeMs };
93
- return { x: xResult, y: yResult };
93
+ const xPromise = extractResult(readX, Float32Array);
94
+ const yPromise = extractResult(readY, Float32Array);
95
+ readX = null; // ownership transferred — extractResult's own finally destroys it
96
+ readY = null;
97
+ const [xResult, yResult] = await Promise.all([xPromise, yPromise]);
98
+
99
+ if (gpuTimeMs !== undefined) return { x: xResult, y: yResult, gpuTimeMs };
100
+ return { x: xResult, y: yResult };
101
+ } finally {
102
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
103
+ if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
104
+ if (paramBuffer) destroyBuffers(paramBuffer);
105
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
106
+ // Only reached if extractTimestamp threw before ownership was transferred above.
107
+ if (readX) destroyBuffers(readX);
108
+ if (readY) destroyBuffers(readY);
109
+ }
94
110
  }
@@ -19,8 +19,10 @@ export async function sscal(device, n, alpha, x, incx) {
19
19
  throw new Error("device must be a GPUDevice.");
20
20
  if (!Number.isInteger(n) || !Number.isInteger(incx))
21
21
  throw new Error("n and incx must be integers.");
22
- if (isNaN(alpha)) throw new Error("alpha must not be NaN.");
23
- if (!isFinite(alpha)) throw new Error("alpha must be finite.");
22
+ if (typeof alpha !== "number")
23
+ throw new Error("alpha must be a number.");
24
+ if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
25
+ if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
24
26
  if (incx <= 0) throw new Error("incx must be positive.");
25
27
  if (!(x instanceof Float32Array) && !(x instanceof GpuVector))
26
28
  throw new Error("x must be a Float32Array or GpuVector.");
@@ -32,40 +34,49 @@ export async function sscal(device, n, alpha, x, incx) {
32
34
 
33
35
  const pipeline = await getPipeline(device, "sscal");
34
36
 
35
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sscal-x", true);
36
- const paramsBuffer = createParamsBuffer(
37
- [
38
- { value: n, type: "u32" },
39
- { value: alpha, type: "f32" },
40
- { value: incx, type: "u32" },
41
- ],
42
- "sscal-params",
43
- );
37
+ let xBuffer = null;
38
+ let paramsBuffer = null;
39
+ let readBuffer = null;
44
40
 
45
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
46
- xBuffer,
47
- paramsBuffer,
48
- ]);
49
- const { commandEncoder, ts } = runComputePass(
50
- pipeline,
51
- bindGroup,
52
- calcWorkgroups(n),
53
- );
54
- const readBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
41
+ try {
42
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sscal-x", true);
43
+ paramsBuffer = createParamsBuffer(
44
+ [
45
+ { value: n, type: "u32" },
46
+ { value: alpha, type: "f32" },
47
+ { value: incx, type: "u32" },
48
+ ],
49
+ "sscal-params",
50
+ );
55
51
 
56
- submit(commandEncoder);
52
+ const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
53
+ xBuffer,
54
+ paramsBuffer,
55
+ ]);
56
+ const { commandEncoder, ts } = runComputePass(
57
+ pipeline,
58
+ bindGroup,
59
+ calcWorkgroups(n),
60
+ );
61
+ readBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
57
62
 
58
- const gpuTimeMs = await extractTimestamp(ts);
63
+ submit(commandEncoder);
59
64
 
60
- if (xIsGpu) {
61
- destroyBuffers(paramsBuffer);
62
- if (gpuTimeMs !== undefined) return { gpuTimeMs };
63
- return {};
64
- }
65
+ const gpuTimeMs = await extractTimestamp(ts);
65
66
 
66
- const result = await extractResult(readBuffer, Float32Array);
67
- destroyBuffers(xBuffer, paramsBuffer, readBuffer);
67
+ if (xIsGpu) {
68
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
69
+ return {};
70
+ }
68
71
 
69
- if (gpuTimeMs !== undefined) return { result, gpuTimeMs };
70
- return result;
72
+ const result = await extractResult(readBuffer, Float32Array);
73
+ readBuffer = null; // extractResult already destroyed it
74
+ if (gpuTimeMs !== undefined) return { x: result, gpuTimeMs };
75
+ return result;
76
+ } finally {
77
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
78
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
79
+ // Only reached if extractTimestamp threw before extractResult ran.
80
+ if (readBuffer) destroyBuffers(readBuffer);
81
+ }
71
82
  }
@@ -46,45 +46,60 @@ export async function sswap(device, n, x, incx, y, incy) {
46
46
 
47
47
  const pipeline = await getPipeline(device, "sswap");
48
48
 
49
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sswap-x", true);
50
- const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sswap-y", true);
51
- const paramsBuffer = createParamsBuffer(
52
- [
53
- { value: n, type: "u32" },
54
- { value: incx, type: "u32" },
55
- { value: incy, type: "u32" },
56
- ],
57
- "sswap-params",
58
- );
49
+ let xBuffer = null;
50
+ let yBuffer = null;
51
+ let paramsBuffer = null;
52
+ let xReadBuffer = null;
53
+ let yReadBuffer = null;
59
54
 
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 xReadBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
71
- const yReadBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
55
+ try {
56
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sswap-x", true);
57
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sswap-y", true);
58
+ paramsBuffer = createParamsBuffer(
59
+ [
60
+ { value: n, type: "u32" },
61
+ { value: incx, type: "u32" },
62
+ { value: incy, type: "u32" },
63
+ ],
64
+ "sswap-params",
65
+ );
72
66
 
73
- submit(commandEncoder);
67
+ const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
68
+ xBuffer,
69
+ yBuffer,
70
+ paramsBuffer,
71
+ ]);
72
+ const { commandEncoder, ts } = runComputePass(
73
+ pipeline,
74
+ bindGroup,
75
+ calcWorkgroups(n),
76
+ );
77
+ xReadBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
78
+ yReadBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
74
79
 
75
- const gpuTimeMs = await extractTimestamp(ts);
80
+ submit(commandEncoder);
76
81
 
77
- if (xIsGpu && yIsGpu) {
78
- destroyBuffers(paramsBuffer);
79
- if (gpuTimeMs !== undefined) return { gpuTimeMs };
80
- return {};
81
- }
82
+ const gpuTimeMs = await extractTimestamp(ts);
82
83
 
83
- const resultX = await extractResult(xReadBuffer, Float32Array);
84
- const resultY = await extractResult(yReadBuffer, Float32Array);
84
+ if (xIsGpu && yIsGpu) {
85
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
86
+ return {};
87
+ }
85
88
 
86
- destroyBuffers(xBuffer, xReadBuffer, yBuffer, yReadBuffer, paramsBuffer);
89
+ const resultX = await extractResult(xReadBuffer, Float32Array);
90
+ xReadBuffer = null; // extractResult already destroyed it
91
+ const resultY = await extractResult(yReadBuffer, Float32Array);
92
+ yReadBuffer = null; // extractResult already destroyed it
87
93
 
88
- if (gpuTimeMs !== undefined) return { x: resultX, y: resultY, gpuTimeMs };
89
- return { x: resultX, y: resultY };
94
+ if (gpuTimeMs !== undefined) return { x: resultX, y: resultY, gpuTimeMs };
95
+ return { x: resultX, y: resultY };
96
+ } finally {
97
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
98
+ if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
99
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
100
+ // Only reached if extractTimestamp or extractResult threw before
101
+ // clearing these — on the success path they're already null.
102
+ if (xReadBuffer) destroyBuffers(xReadBuffer);
103
+ if (yReadBuffer) destroyBuffers(yReadBuffer);
104
+ }
90
105
  }