wgblas 1.2.0 → 2.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 (70) hide show
  1. package/LICENSE +1 -1
  2. package/README.md +36 -54
  3. package/dist/wgblas.browser.js +779 -34
  4. package/index.d.mts +13 -0
  5. package/index.mjs +8 -0
  6. package/package.json +56 -3
  7. package/src/dasum/dasum.d.mts +2 -2
  8. package/src/dasum/dasum.mjs +6 -4
  9. package/src/idamax/idamax.d.mts +51 -0
  10. package/src/idamax/idamax.mjs +128 -0
  11. package/src/init.mjs +9 -1
  12. package/src/isamax/isamax.d.mts +1 -1
  13. package/src/sasum/sasum.d.mts +1 -1
  14. package/src/saxpy/saxpy.d.mts +1 -1
  15. package/src/scopy/scopy.d.mts +1 -1
  16. package/src/sdot/sdot.d.mts +1 -1
  17. package/src/sgemm/sgemm.d.mts +102 -0
  18. package/src/sgemm/sgemm.mjs +195 -0
  19. package/src/sgemmtr/sgemmtr.d.mts +104 -0
  20. package/src/sgemmtr/sgemmtr.mjs +203 -0
  21. package/src/sgemv/sgemv.d.mts +1 -38
  22. package/src/sgemv/sgemv.mjs +4 -0
  23. package/src/sger/sger.d.mts +1 -34
  24. package/src/sger/sger.mjs +2 -0
  25. package/src/shaders/block_transfer.wgsl +42 -0
  26. package/src/shaders/browser-shaders.mjs +26 -0
  27. package/src/shaders/dasum.wgsl +3 -2
  28. package/src/shaders/f64/dekker.wgsl +4 -85
  29. package/src/shaders/f64/utils/abs.wgsl +10 -0
  30. package/src/shaders/f64/utils/add.wgsl +77 -0
  31. package/src/shaders/f64/utils/equal.wgsl +7 -0
  32. package/src/shaders/f64/utils/greater.wgsl +12 -0
  33. package/src/shaders/f64/utils/multiply.wgsl +81 -0
  34. package/src/shaders/idamax.wgsl +96 -0
  35. package/src/shaders/reduction/argmaxF64.wgsl +50 -0
  36. package/src/shaders/reduction/sumF64.wgsl +2 -2
  37. package/src/shaders/sgemm_large.wgsl +117 -0
  38. package/src/shaders/sgemm_small.wgsl +112 -0
  39. package/src/shaders/sgemmtr_large.wgsl +117 -0
  40. package/src/shaders/sgemmtr_small.wgsl +110 -0
  41. package/src/shaders/symmetrize.wgsl +31 -0
  42. package/src/shaders/triangularize.wgsl +44 -0
  43. package/src/snrm2/snrm2.d.mts +1 -1
  44. package/src/srot/srot.d.mts +1 -1
  45. package/src/srotm/srotm.d.mts +1 -1
  46. package/src/sscal/sscal.d.mts +1 -1
  47. package/src/sswap/sswap.d.mts +1 -1
  48. package/src/ssymm/ssymm.d.mts +103 -0
  49. package/src/ssymm/ssymm.mjs +209 -0
  50. package/src/ssymv/ssymv.d.mts +1 -36
  51. package/src/ssymv/ssymv.mjs +2 -0
  52. package/src/ssyr/ssyr.d.mts +1 -30
  53. package/src/ssyr/ssyr.mjs +2 -0
  54. package/src/ssyr2/ssyr2.d.mts +1 -34
  55. package/src/ssyr2/ssyr2.mjs +2 -0
  56. package/src/ssyr2k/ssyr2k.d.mts +100 -0
  57. package/src/ssyr2k/ssyr2k.mjs +201 -0
  58. package/src/ssyrk/ssyrk.d.mts +90 -0
  59. package/src/ssyrk/ssyrk.mjs +176 -0
  60. package/src/strmm/strmm.d.mts +100 -0
  61. package/src/strmm/strmm.mjs +211 -0
  62. package/src/strmv/strmv.d.mts +1 -36
  63. package/src/strmv/strmv.mjs +2 -0
  64. package/src/strsm/strsm.d.mts +99 -0
  65. package/src/strsm/strsm.mjs +342 -0
  66. package/src/strsv/strsv.d.mts +1 -32
  67. package/src/strsv/strsv.mjs +2 -0
  68. package/src/util/buffer.mjs +4 -2
  69. package/src/util/compute.mjs +6 -3
  70. package/src/util/f64.mjs +3 -3
@@ -0,0 +1,342 @@
1
+ import {
2
+ uploadBuffer,
3
+ createParamsBuffer,
4
+ createStorageBuffer,
5
+ stageReadback,
6
+ destroyBuffers,
7
+ } from "../util/buffer.mjs";
8
+ import { createBindGroup } from "../util/bindgroup.mjs";
9
+ import { beginTimedEncoder, encodePass, submit } from "../util/compute.mjs";
10
+ import { extractResult } from "../util/result.mjs";
11
+ import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
12
+ import { getPipeline } from "../util/pipeline.mjs";
13
+ import { calcWorkgroups } from "../util/workgroup.mjs";
14
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
15
+
16
+ const BLOCK_SIZE = 64; // must match strsv_invert_block.wgsl's own constant
17
+ const BM_SMALL = 32, BN_SMALL = 32; // sgemm_small.wgsl's block tile
18
+ const BM_LARGE = 64, BN_LARGE = 64; // sgemm_large.wgsl's block tile
19
+ const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // same threshold sgemm/ssymm/strmm use
20
+
21
+ // strsm: B := alpha*op(A)^-1*B (side='left') or alpha*B*op(A)^-1 (side='right'),
22
+ // A triangular. Blocked substitution (strsv's own technique, generalized to
23
+ // a matrix RHS): strsv_invert_block + sgemm, unchanged; every per-block B/A
24
+ // access goes through block_transfer.wgsl (see that shader for why).
25
+ export async function strsm(
26
+ device, side, uplo, transA, diag, m, n, alpha, A, lda, B, ldb, layout = "row-major",
27
+ ) {
28
+ const AIsGpu = A instanceof GpuMatrix;
29
+ const BIsGpu = B instanceof GpuMatrix;
30
+ const isUnit = diag === "unit";
31
+
32
+ if (!(device instanceof GPUDevice))
33
+ throw new Error("device must be a GPUDevice.");
34
+ if (side !== "left" && side !== "right")
35
+ throw new Error("side must be 'left' or 'right'.");
36
+ if (uplo !== "lower" && uplo !== "upper")
37
+ throw new Error("uplo must be 'lower' or 'upper'.");
38
+ if (transA !== "no-transpose" && transA !== "transpose")
39
+ throw new Error("transA must be 'no-transpose' or 'transpose'.");
40
+ if (!isUnit && diag !== "non-unit")
41
+ throw new Error("diag must be 'unit' or 'non-unit'.");
42
+ if (layout !== "row-major" && layout !== "column-major")
43
+ throw new Error("layout must be 'row-major' or 'column-major'.");
44
+ if (typeof alpha !== "number")
45
+ throw new Error("alpha must be a number.");
46
+ if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
47
+ if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
48
+ if (!Number.isInteger(m) || !Number.isInteger(n) || !Number.isInteger(lda) || !Number.isInteger(ldb))
49
+ throw new Error("m, n, lda, and ldb must be integers.");
50
+ if (!AIsGpu && !(A instanceof Float32Array))
51
+ throw new Error("A must be a Float32Array or GpuMatrix.");
52
+ if (!BIsGpu && !(B instanceof Float32Array))
53
+ throw new Error("B must be a Float32Array or GpuMatrix.");
54
+ if (AIsGpu !== BIsGpu)
55
+ throw new Error("A and B must both be GpuMatrix or both be Float32Array.");
56
+ if (m < 0 || n < 0) throw new Error("m and n must be non-negative.");
57
+ if (m === 0 || n === 0) return BIsGpu ? {} : { B };
58
+
59
+ const effLayoutA = AIsGpu ? A.layout : layout;
60
+ const effLayoutB = BIsGpu ? B.layout : layout;
61
+
62
+ // A: triangular, order = m (side='left') or n (side='right').
63
+ const aOrder = side === "left" ? m : n;
64
+ if (lda < aOrder) throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
65
+ if (AIsGpu) {
66
+ if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
67
+ if (A.rows < aOrder || A.cols < aOrder) throw new Error("A is too small for the given m/n and side.");
68
+ } else if (A.length < (aOrder - 1) * lda + aOrder) {
69
+ throw new Error("A does not have enough elements for the given dimensions and lda.");
70
+ }
71
+
72
+ // B: always m x n, overwritten in place with the same ldb.
73
+ const bOuter = effLayoutB === "column-major" ? n : m;
74
+ const bInner = effLayoutB === "column-major" ? m : n;
75
+ if (ldb < bInner)
76
+ throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
77
+ if (BIsGpu) {
78
+ if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
79
+ if (B.rows < m || B.cols < n) throw new Error("B is too small for the given m and n.");
80
+ } else if (B.length < (bOuter - 1) * ldb + bInner) {
81
+ throw new Error("B does not have enough elements for the given dimensions and ldb.");
82
+ }
83
+
84
+ // A isn't symmetric: column-major = genuine transpose, so flip transA;
85
+ // transposing also swaps which triangle looks stored, so flip uplo too.
86
+ const uploEffA = effLayoutA === "column-major" ? (uplo === "lower" ? "upper" : "lower") : uplo;
87
+ const transEffA = effLayoutA === "column-major" ? (transA === "no-transpose" ? "transpose" : "no-transpose") : transA;
88
+
89
+ const otherLen = side === "left" ? n : m;
90
+ const blockIsRow = side === "left";
91
+
92
+ // Forward iff op(A) is lower (side='left') — side='right' flips this,
93
+ // since it solves via columns instead of rows.
94
+ const opIsLower = (transEffA === "no-transpose") === (uploEffA === "lower");
95
+ const forward = side === "left" ? opIsLower : !opIsLower;
96
+
97
+ const blockStarts = [];
98
+ for (let s = 0; s < aOrder; s += BLOCK_SIZE) blockStarts.push(s);
99
+ if (!forward) blockStarts.reverse();
100
+ const numBlocks = blockStarts.length;
101
+
102
+ const invertPipeline = await getPipeline(device, "strsv_invert_block");
103
+ const transferPipeline = await getPipeline(device, "block_transfer");
104
+ const scalarPipeline = await getPipeline(device, "sscal");
105
+
106
+ const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "strsm-A", false);
107
+ const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "strsm-B", true);
108
+ const AinvBuffer = createStorageBuffer(numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4, "strsm-Ainv");
109
+
110
+ const paramsBuffers = [];
111
+ const scratchBuffers = [];
112
+ function scratch(size, label) {
113
+ const buf = createStorageBuffer(size, label);
114
+ scratchBuffers.push(buf);
115
+ return buf;
116
+ }
117
+ function params(entries, label) {
118
+ const buf = createParamsBuffer(entries, label);
119
+ paramsBuffers.push(buf);
120
+ return buf;
121
+ }
122
+
123
+ // Minimal valid length for B's own buffer (matches its own validation
124
+ // above) — NOT bOuter*ldb, which can exceed a Float32Array-path buffer's
125
+ // actual allocation (only guaranteed padded up to GpuMatrix's own size).
126
+ const bScaleLen = (bOuter - 1) * ldb + bInner;
127
+
128
+ try {
129
+ // Pre-scale B by alpha once (reuses sscal, so no per-block alpha handling).
130
+ let preScaleBindGroup = null;
131
+ if (alpha !== 1.0) {
132
+ const scaleParams = params(
133
+ [
134
+ { value: bScaleLen, type: "u32" },
135
+ { value: alpha, type: "f32" },
136
+ { value: 1, type: "u32" },
137
+ ],
138
+ "strsm-scale-params",
139
+ );
140
+ preScaleBindGroup = createBindGroup(scalarPipeline.getBindGroupLayout(0), [BBuffer, scaleParams]);
141
+ }
142
+
143
+ // Every diagonal block's inverse, fully parallel, one dispatch, unchanged.
144
+ const invertParams = params(
145
+ [
146
+ { value: aOrder, type: "u32" },
147
+ { value: lda, type: "u32" },
148
+ { value: transEffA === "transpose" ? 1 : 0, type: "u32" },
149
+ { value: uploEffA === "upper" ? 1 : 0, type: "u32" },
150
+ { value: isUnit ? 1 : 0, type: "u32" },
151
+ ],
152
+ "strsm-invert-params",
153
+ );
154
+ const invertBindGroup = createBindGroup(invertPipeline.getBindGroupLayout(0), [ABuffer, AinvBuffer, invertParams]);
155
+
156
+ // Reusable scratch buffers, sized for the worst case, bound at offset 0.
157
+ const Bblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Bblock");
158
+ const Xblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Xblock");
159
+ const Aoff = scratch(aOrder * BLOCK_SIZE * 4, "strsm-Aoff");
160
+ const delta = scratch(aOrder * otherLen * 4, "strsm-delta");
161
+
162
+ const { commandEncoder, querySet } = beginTimedEncoder();
163
+
164
+ if (alpha === 0) {
165
+ // BLAS: alpha=0 means A is not referenced — skip straight to B:=0.
166
+ const zeroDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0, endOfPassWriteIndex: 1 } } : undefined;
167
+ encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(bScaleLen), zeroDesc);
168
+ } else {
169
+ if (preScaleBindGroup) {
170
+ encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(bScaleLen));
171
+ }
172
+ const invertDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
173
+ encodePass(commandEncoder, invertPipeline, invertBindGroup, { x: BLOCK_SIZE, y: numBlocks }, invertDesc);
174
+
175
+ for (let bi = 0; bi < blockStarts.length; bi++) {
176
+ const blockStart = blockStarts[bi];
177
+ const blockEnd = Math.min(blockStart + BLOCK_SIZE, aOrder);
178
+ const blockLen = blockEnd - blockStart;
179
+ const blockIndex = blockStart / BLOCK_SIZE;
180
+ const isLastPass = bi === blockStarts.length - 1;
181
+
182
+ // 1) gather B's current block into a tight scratch buffer.
183
+ const gatherBParams = params(
184
+ [
185
+ { value: blockStart, type: "u32" },
186
+ { value: blockLen, type: "u32" },
187
+ { value: 0, type: "u32" },
188
+ { value: otherLen, type: "u32" },
189
+ { value: ldb, type: "u32" },
190
+ { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
191
+ { value: blockIsRow ? 1 : 0, type: "u32" },
192
+ { value: 2, type: "u32" }, // gather
193
+ ],
194
+ "strsm-gather-B-params",
195
+ );
196
+ const gatherBBindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [Bblock, BBuffer, gatherBParams]);
197
+ encodePass(commandEncoder, transferPipeline, gatherBBindGroup, calcWorkgroups(blockLen, otherLen));
198
+
199
+ // 2) apply: Xblock := op(Ainv_block) @ Bblock (side='left'), or the
200
+ // transpose-trick equivalent for side='right' (same trick strmm uses).
201
+ {
202
+ const mg = blockLen, ng = otherLen, kg = blockLen;
203
+ const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
204
+ const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
205
+ const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
206
+ const applyParams = params(
207
+ [
208
+ { value: mg, type: "u32" },
209
+ { value: ng, type: "u32" },
210
+ { value: kg, type: "u32" },
211
+ { value: 1.0, type: "f32" }, // alpha already applied to B up front
212
+ { value: 0.0, type: "f32" }, // beta — fresh output, no accumulation
213
+ { value: BLOCK_SIZE, type: "u32" }, // ldX = Ainv's own dense stride
214
+ { value: otherLen, type: "u32" }, // ldY = Bblock's own tight stride
215
+ { value: otherLen, type: "u32" }, // ldc = Xblock's own tight stride
216
+ { value: side === "right" ? 1 : 0, type: "u32" }, // transX: side='right' needs Ainv^T
217
+ { value: 0, type: "u32" }, // transY: Bblock is always read as-is
218
+ ],
219
+ "strsm-apply-params",
220
+ );
221
+ const applyBindGroup = createBindGroup(gemmPipeline.getBindGroupLayout(0), [
222
+ { buffer: AinvBuffer, offset: blockIndex * BLOCK_SIZE * BLOCK_SIZE * 4, size: BLOCK_SIZE * BLOCK_SIZE * 4 },
223
+ Bblock,
224
+ Xblock,
225
+ applyParams,
226
+ ]);
227
+ const wg = useLarge
228
+ ? { x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension) }
229
+ : { x: Math.min(Math.ceil(ng / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(Math.ceil(mg / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension) };
230
+ encodePass(commandEncoder, gemmPipeline, applyBindGroup, wg);
231
+ }
232
+
233
+ // 3) scatter the solved block back into B.
234
+ const rangeStart = forward ? blockEnd : 0;
235
+ const rangeEnd = forward ? aOrder : blockStart;
236
+ const hasRemaining = rangeStart < rangeEnd;
237
+ const scatterParams = params(
238
+ [
239
+ { value: blockStart, type: "u32" },
240
+ { value: blockLen, type: "u32" },
241
+ { value: 0, type: "u32" },
242
+ { value: otherLen, type: "u32" },
243
+ { value: ldb, type: "u32" },
244
+ { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
245
+ { value: blockIsRow ? 1 : 0, type: "u32" },
246
+ { value: 0, type: "u32" }, // overwrite
247
+ ],
248
+ "strsm-scatter-params",
249
+ );
250
+ const scatterBindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [Xblock, BBuffer, scatterParams]);
251
+ const scatterDesc = isLastPass && !hasRemaining && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
252
+ encodePass(commandEncoder, transferPipeline, scatterBindGroup, calcWorkgroups(blockLen, otherLen), scatterDesc);
253
+
254
+ // 4) trailing update: subtract this block's contribution from B.
255
+ if (!hasRemaining) continue;
256
+ const remCount = rangeEnd - rangeStart;
257
+
258
+ const gatherAParams = params(
259
+ [
260
+ { value: rangeStart, type: "u32" },
261
+ { value: remCount, type: "u32" },
262
+ { value: blockStart, type: "u32" },
263
+ { value: blockLen, type: "u32" },
264
+ { value: lda, type: "u32" },
265
+ { value: transEffA === "transpose" ? 1 : 0, type: "u32" },
266
+ { value: blockIsRow ? 1 : 0, type: "u32" },
267
+ { value: 2, type: "u32" }, // gather
268
+ ],
269
+ "strsm-gather-A-params",
270
+ );
271
+ const gatherABindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [Aoff, ABuffer, gatherAParams]);
272
+ encodePass(commandEncoder, transferPipeline, gatherABindGroup, calcWorkgroups(remCount, blockLen));
273
+
274
+ {
275
+ const mg = remCount, ng = otherLen, kg = blockLen;
276
+ const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
277
+ const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
278
+ const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
279
+ const updateParams = params(
280
+ [
281
+ { value: mg, type: "u32" },
282
+ { value: ng, type: "u32" },
283
+ { value: kg, type: "u32" },
284
+ { value: 1.0, type: "f32" },
285
+ { value: 0.0, type: "f32" },
286
+ { value: blockLen, type: "u32" }, // ldX = Aoff's own tight stride
287
+ { value: otherLen, type: "u32" }, // ldY = Xblock's own tight stride
288
+ { value: otherLen, type: "u32" }, // ldc = delta's own tight stride
289
+ { value: 0, type: "u32" }, // transX: Aoff already read in the right orientation
290
+ { value: 0, type: "u32" }, // transY: Xblock read as-is
291
+ ],
292
+ "strsm-update-params",
293
+ );
294
+ const updateBindGroup = createBindGroup(gemmPipeline.getBindGroupLayout(0), [Aoff, Xblock, delta, updateParams]);
295
+ const wg = useLarge
296
+ ? { x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension) }
297
+ : { x: Math.min(Math.ceil(ng / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(Math.ceil(mg / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension) };
298
+ encodePass(commandEncoder, gemmPipeline, updateBindGroup, wg);
299
+ }
300
+
301
+ const scatterSubParams = params(
302
+ [
303
+ { value: rangeStart, type: "u32" },
304
+ { value: remCount, type: "u32" },
305
+ { value: 0, type: "u32" },
306
+ { value: otherLen, type: "u32" },
307
+ { value: ldb, type: "u32" },
308
+ { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
309
+ { value: blockIsRow ? 1 : 0, type: "u32" },
310
+ { value: 1, type: "u32" }, // subtract
311
+ ],
312
+ "strsm-scatter-sub-params",
313
+ );
314
+ const scatterSubBindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [delta, BBuffer, scatterSubParams]);
315
+ const subDesc = isLastPass && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
316
+ encodePass(commandEncoder, transferPipeline, scatterSubBindGroup, calcWorkgroups(remCount, otherLen), subDesc);
317
+ }
318
+ }
319
+
320
+ const ts = resolveTimestamp(commandEncoder, querySet);
321
+ const readBuffer = BIsGpu ? null : stageReadback(commandEncoder, BBuffer);
322
+
323
+ submit(commandEncoder);
324
+
325
+ const gpuTimeMs = await extractTimestamp(ts);
326
+
327
+ if (BIsGpu) {
328
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
329
+ return {};
330
+ }
331
+
332
+ const result = await extractResult(readBuffer, Float32Array);
333
+ if (gpuTimeMs !== undefined) return { B: result, gpuTimeMs };
334
+ return { B: result };
335
+ } finally {
336
+ if (!AIsGpu) destroyBuffers(ABuffer);
337
+ if (!BIsGpu) destroyBuffers(BBuffer);
338
+ destroyBuffers(AinvBuffer);
339
+ destroyBuffers(scratchBuffers);
340
+ destroyBuffers(paramsBuffers);
341
+ }
342
+ }
@@ -41,37 +41,6 @@ export declare function strsv(
41
41
  layout?: 'row-major' | 'column-major',
42
42
  ): Promise<{ x: Float32Array; gpuTimeMs?: number }>;
43
43
 
44
- /**
45
- * Solves the triangular system op(A) * x = b for x, in place.
46
- *
47
- * A is kept GPU-resident; x is a CPU Float32Array. `A`'s own `layout` (set at
48
- * `GpuMatrix.from` time) determines the operation — there is no separate
49
- * `layout` argument here.
50
- *
51
- * @param device - GPUDevice from `init()`
52
- * @param uplo - `'lower'` to use the lower triangle, `'upper'` to use the upper triangle
53
- * @param trans - `'no-transpose'` to solve A*x=b, `'transpose'` to solve A^T*x=b
54
- * @param diag - `'unit'` to treat the diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
55
- * @param n - order of the matrix A
56
- * @param A - GpuMatrix, GPU-resident
57
- * @param lda - leading dimension of A (must equal A.lda)
58
- * @param x - Float32Array holding b on input, the solution on output
59
- * @param incx - stride for x (must be a positive integer)
60
- * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strsv/strsv.mjs#L15">Source code: strsv.mjs (L15)</a>
61
- * @category BLAS Level 2
62
- */
63
- export declare function strsv(
64
- device: GPUDevice,
65
- uplo: 'lower' | 'upper',
66
- trans: 'no-transpose' | 'transpose',
67
- diag: 'unit' | 'non-unit',
68
- n: number,
69
- A: GpuMatrix,
70
- lda: number,
71
- x: Float32Array,
72
- incx: number,
73
- ): Promise<{ x: Float32Array; gpuTimeMs?: number }>;
74
-
75
44
  /**
76
45
  * Solves the triangular system op(A) * x = b for x, in place.
77
46
  *
@@ -79,7 +48,7 @@ export declare function strsv(
79
48
  * its own `layout` (set at `GpuMatrix.from` time) determines the operation —
80
49
  * there is no separate `layout` argument here.
81
50
  *
82
- * {@includeCode ../../examples/strsv/gpuvec.strsv.js}
51
+ * {@includeCode ../../examples/strsv/gpu.strsv.js}
83
52
  *
84
53
  * @param device - GPUDevice from `init()`
85
54
  * @param uplo - `'lower'` to use the lower triangle, `'upper'` to use the upper triangle
@@ -63,6 +63,8 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
63
63
  throw new Error("x must be a Float32Array or GpuVector.");
64
64
  if (xIsGpu && !AIsGpu)
65
65
  throw new Error("A must be a GpuMatrix when x is a GpuVector.");
66
+ if (AIsGpu && !xIsGpu)
67
+ throw new Error("x must be a GpuVector when A is a GpuMatrix.");
66
68
  if (AIsGpu && lda !== A.lda)
67
69
  throw new Error("lda must match A.lda when A is a GpuMatrix.");
68
70
  if (AIsGpu && (A.rows < n || A.cols < n))
@@ -62,15 +62,17 @@ export function uploadBuffer(data, label = "blas-input", readback = false) {
62
62
  * that are written by a shader before being read.
63
63
  * @param {number} size - byte size
64
64
  * @param {string} [label] - debug label visible in browser DevTools GPU inspection
65
+ * @param {number} [extraUsage=0] - additional `GPUBufferUsage` flags OR'd in alongside `STORAGE`
66
+ * (e.g. `GPUBufferUsage.COPY_DST` for a buffer that's also a `copyBufferToBuffer` destination)
65
67
  * @returns {GPUBuffer}
66
68
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createBuffer GPUDevice.createBuffer()}
67
69
  */
68
- export function createStorageBuffer(size, label = "blas-storage") {
70
+ export function createStorageBuffer(size, label = "blas-storage", extraUsage = 0) {
69
71
  const device = getDevice();
70
72
  return device.createBuffer({
71
73
  label,
72
74
  size,
73
- usage: GPUBufferUsage.STORAGE,
75
+ usage: GPUBufferUsage.STORAGE | extraUsage,
74
76
  });
75
77
  }
76
78
 
@@ -38,7 +38,8 @@ export function beginTimedEncoder() {
38
38
  * @param {GPUCommandEncoder} commandEncoder
39
39
  * @param {GPUComputePipeline} pipeline
40
40
  * @param {GPUBindGroup} bindGroup
41
- * @param {number | { x: number, y: number }} workgroups - workgroup count; number for 1D dispatch, `{x, y}` for 2D
41
+ * @param {number | { x: number, y: number, z?: number }} workgroups - workgroup count;
42
+ * number for 1D dispatch, `{x, y}` for 2D, `{x, y, z}` for 3D (z defaults to 1)
42
43
  * @param {object} [passDescriptor] - passed to `beginComputePass`, e.g. for timestamp writes
43
44
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/beginComputePass GPUCommandEncoder.beginComputePass()}
44
45
  */
@@ -51,7 +52,8 @@ export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, pass
51
52
  if (typeof workgroups === "number") {
52
53
  passEncoder.dispatchWorkgroups(workgroups);
53
54
  } else {
54
- passEncoder.dispatchWorkgroups(workgroups.x, workgroups.y);
55
+ // `?? 1` is load-bearing — confirmed passing undefined (from an {x,y}-only caller) crashes the process, not just no-ops.
56
+ passEncoder.dispatchWorkgroups(workgroups.x, workgroups.y, workgroups.z ?? 1);
55
57
  }
56
58
 
57
59
  passEncoder.end();
@@ -64,7 +66,8 @@ export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, pass
64
66
  * dispatches workgroups, and optionally wraps the pass in GPU timestamp queries.
65
67
  * @param {GPUComputePipeline} pipeline
66
68
  * @param {GPUBindGroup} bindGroup
67
- * @param {number | { x: number, y: number }} workgroups - workgroup count; number for 1D dispatch, `{x, y}` for 2D
69
+ * @param {number | { x: number, y: number, z?: number }} workgroups - workgroup count;
70
+ * number for 1D dispatch, `{x, y}` for 2D, `{x, y, z}` for 3D (z defaults to 1)
68
71
  * @returns {{ commandEncoder: GPUCommandEncoder, ts: any }} encoded commands and timestamp handle
69
72
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createCommandEncoder GPUDevice.createCommandEncoder()}
70
73
  */
package/src/util/f64.mjs CHANGED
@@ -2,9 +2,9 @@
2
2
 
3
3
  // Double-double f64 emulation — splits a double into a (hi, lo) pair of f32
4
4
  // values with hi+lo approximating the original, hi holding the leading bits
5
- // and lo the rounding error hi lost on its own. See
6
- // src/shaders/f64/dekker.wgsl for the GPU-side arithmetic this pairs with
7
- // (Dekker's algorithm).
5
+ // and lo the rounding error hi lost on its own. See src/shaders/f64/ (the DD
6
+ // struct in dekker.wgsl, operations in utils/) for the GPU-side arithmetic
7
+ // this pairs with (Dekker's algorithm).
8
8
  //
9
9
  // Not a value-preserving exact split: double-double buys roughly 2x f32's
10
10
  // mantissa (~48 bits vs f32's 24), less than real f64's 52-bit mantissa.