wgblas 1.2.1 → 2.1.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 (96) hide show
  1. package/LICENSE +1 -1
  2. package/README.md +54 -72
  3. package/dist/wgblas.browser.js +1637 -858
  4. package/index.d.mts +51 -6
  5. package/index.mjs +8 -0
  6. package/package.json +56 -2
  7. package/src/classes/GpuMatrix.mjs +17 -10
  8. package/src/classes/GpuVector.mjs +28 -10
  9. package/src/dasum/dasum.d.mts +6 -6
  10. package/src/dasum/dasum.mjs +25 -21
  11. package/src/devdocs.mjs +13 -0
  12. package/src/idamax/idamax.d.mts +69 -0
  13. package/src/idamax/idamax.mjs +130 -0
  14. package/src/init.mjs +115 -49
  15. package/src/isamax/isamax.d.mts +21 -3
  16. package/src/isamax/isamax.mjs +16 -14
  17. package/src/random/random.d.mts +1 -0
  18. package/src/sasum/sasum.d.mts +3 -3
  19. package/src/sasum/sasum.mjs +15 -13
  20. package/src/saxpy/saxpy.d.mts +3 -3
  21. package/src/saxpy/saxpy.mjs +10 -8
  22. package/src/scopy/scopy.d.mts +3 -3
  23. package/src/scopy/scopy.mjs +10 -8
  24. package/src/sdot/sdot.d.mts +3 -3
  25. package/src/sdot/sdot.mjs +16 -14
  26. package/src/sgemm/sgemm.d.mts +102 -0
  27. package/src/sgemm/sgemm.mjs +208 -0
  28. package/src/sgemmtr/sgemmtr.d.mts +104 -0
  29. package/src/sgemmtr/sgemmtr.mjs +204 -0
  30. package/src/sgemv/sgemv.d.mts +1 -38
  31. package/src/sgemv/sgemv.mjs +42 -26
  32. package/src/sger/sger.d.mts +1 -34
  33. package/src/sger/sger.mjs +12 -8
  34. package/src/shaders/block_transfer.wgsl +42 -0
  35. package/src/shaders/dasum.wgsl +3 -2
  36. package/src/shaders/f64/dekker.wgsl +4 -85
  37. package/src/shaders/f64/utils/abs.wgsl +10 -0
  38. package/src/shaders/f64/utils/add.wgsl +77 -0
  39. package/src/shaders/f64/utils/equal.wgsl +7 -0
  40. package/src/shaders/f64/utils/greater.wgsl +12 -0
  41. package/src/shaders/f64/utils/multiply.wgsl +81 -0
  42. package/src/shaders/idamax.wgsl +96 -0
  43. package/src/shaders/index.mjs +164 -14
  44. package/src/shaders/reduction/argmaxF64.wgsl +50 -0
  45. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  46. package/src/shaders/reduction/sumF64.wgsl +2 -2
  47. package/src/shaders/sgemm_large.wgsl +206 -0
  48. package/src/shaders/sgemm_small.wgsl +212 -0
  49. package/src/shaders/sgemmtr_large.wgsl +120 -0
  50. package/src/shaders/sgemmtr_small.wgsl +113 -0
  51. package/src/shaders/sgemv_n.wgsl +3 -1
  52. package/src/shaders/sgemv_t.wgsl +3 -1
  53. package/src/shaders/snrm2.wgsl +72 -23
  54. package/src/shaders/ssymv.wgsl +3 -1
  55. package/src/shaders/symmetrize.wgsl +31 -0
  56. package/src/shaders/triangularize.wgsl +44 -0
  57. package/src/snrm2/snrm2.d.mts +3 -3
  58. package/src/snrm2/snrm2.mjs +33 -21
  59. package/src/srot/srot.d.mts +3 -5
  60. package/src/srot/srot.mjs +11 -9
  61. package/src/srotm/srotm.d.mts +3 -5
  62. package/src/srotm/srotm.mjs +19 -10
  63. package/src/sscal/sscal.d.mts +4 -4
  64. package/src/sscal/sscal.mjs +12 -10
  65. package/src/sswap/sswap.d.mts +3 -3
  66. package/src/sswap/sswap.mjs +11 -9
  67. package/src/ssymm/ssymm.d.mts +103 -0
  68. package/src/ssymm/ssymm.mjs +218 -0
  69. package/src/ssymv/ssymv.d.mts +1 -36
  70. package/src/ssymv/ssymv.mjs +12 -8
  71. package/src/ssyr/ssyr.d.mts +1 -30
  72. package/src/ssyr/ssyr.mjs +11 -7
  73. package/src/ssyr2/ssyr2.d.mts +1 -34
  74. package/src/ssyr2/ssyr2.mjs +12 -8
  75. package/src/ssyr2k/ssyr2k.d.mts +100 -0
  76. package/src/ssyr2k/ssyr2k.mjs +202 -0
  77. package/src/ssyrk/ssyrk.d.mts +90 -0
  78. package/src/ssyrk/ssyrk.mjs +177 -0
  79. package/src/strmm/strmm.d.mts +100 -0
  80. package/src/strmm/strmm.mjs +226 -0
  81. package/src/strmv/strmv.d.mts +1 -36
  82. package/src/strmv/strmv.mjs +12 -8
  83. package/src/strsm/strsm.d.mts +99 -0
  84. package/src/strsm/strsm.mjs +360 -0
  85. package/src/strsv/strsv.d.mts +1 -32
  86. package/src/strsv/strsv.mjs +18 -12
  87. package/src/util/benchmark.mjs +4 -6
  88. package/src/util/bindgroup.mjs +1 -3
  89. package/src/util/buffer.mjs +116 -20
  90. package/src/util/compute.mjs +12 -12
  91. package/src/util/constants.mjs +57 -0
  92. package/src/util/device.mjs +34 -0
  93. package/src/util/f64.mjs +3 -3
  94. package/src/util/pipeline.mjs +5 -6
  95. package/src/util/workgroup.mjs +55 -7
  96. package/src/shaders/browser-shaders.mjs +0 -55
@@ -0,0 +1,360 @@
1
+ import {
2
+ uploadBuffer,
3
+ createParamsBuffer,
4
+ createStorageBuffer,
5
+ stageReadback,
6
+ destroyBuffers,
7
+ vec4ViewBinding,
8
+ } from "../util/buffer.mjs";
9
+ import { createBindGroup } from "../util/bindgroup.mjs";
10
+ import { beginTimedEncoder, encodePass, submit } from "../util/compute.mjs";
11
+ import { extractResult } from "../util/result.mjs";
12
+ import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
13
+ import { getPipeline } from "../util/pipeline.mjs";
14
+ import { calcWorkgroups, requireWorkgroups, requireWorkgroupCount } from "../util/workgroup.mjs";
15
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
16
+ import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
17
+ import { BLOCK_SIZE } from "../util/constants.mjs";
18
+ import { requireSameDevice } from "../util/device.mjs";
19
+
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
+ requireSameDevice(device, "strsm", { A, B });
35
+ if (side !== "left" && side !== "right")
36
+ throw new Error("side must be 'left' or 'right'.");
37
+ if (uplo !== "lower" && uplo !== "upper")
38
+ throw new Error("uplo must be 'lower' or 'upper'.");
39
+ if (transA !== "no-transpose" && transA !== "transpose")
40
+ throw new Error("transA must be 'no-transpose' or 'transpose'.");
41
+ if (!isUnit && diag !== "non-unit")
42
+ throw new Error("diag must be 'unit' or 'non-unit'.");
43
+ if (layout !== "row-major" && layout !== "column-major")
44
+ throw new Error("layout must be 'row-major' or 'column-major'.");
45
+ if (typeof alpha !== "number")
46
+ throw new Error("alpha must be a number.");
47
+ if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
48
+ if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
49
+ if (!Number.isInteger(m) || !Number.isInteger(n) || !Number.isInteger(lda) || !Number.isInteger(ldb))
50
+ throw new Error("m, n, lda, and ldb must be integers.");
51
+ if (!AIsGpu && !(A instanceof Float32Array))
52
+ throw new Error("A must be a Float32Array or GpuMatrix.");
53
+ if (!BIsGpu && !(B instanceof Float32Array))
54
+ throw new Error("B must be a Float32Array or GpuMatrix.");
55
+ if (AIsGpu !== BIsGpu)
56
+ throw new Error("A and B must both be GpuMatrix or both be Float32Array.");
57
+ if (m < 0 || n < 0) throw new Error("m and n must be non-negative.");
58
+ if (m === 0 || n === 0) return BIsGpu ? {} : { B };
59
+
60
+ const effLayoutA = AIsGpu ? A.layout : layout;
61
+ const effLayoutB = BIsGpu ? B.layout : layout;
62
+
63
+ // A: triangular, order = m (side='left') or n (side='right').
64
+ const aOrder = side === "left" ? m : n;
65
+ if (lda < aOrder) throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
66
+ if (AIsGpu) {
67
+ if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
68
+ if (A.rows < aOrder || A.cols < aOrder) throw new Error("A is too small for the given m/n and side.");
69
+ } else if (A.length < (aOrder - 1) * lda + aOrder) {
70
+ throw new Error("A does not have enough elements for the given dimensions and lda.");
71
+ }
72
+
73
+ // B: always m x n, overwritten in place with the same ldb.
74
+ const bOuter = effLayoutB === "column-major" ? n : m;
75
+ const bInner = effLayoutB === "column-major" ? m : n;
76
+ if (ldb < bInner)
77
+ throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
78
+ if (BIsGpu) {
79
+ if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
80
+ if (B.rows < m || B.cols < n) throw new Error("B is too small for the given m and n.");
81
+ } else if (B.length < (bOuter - 1) * ldb + bInner) {
82
+ throw new Error("B does not have enough elements for the given dimensions and ldb.");
83
+ }
84
+
85
+ // A isn't symmetric: column-major = genuine transpose, so flip transA;
86
+ // transposing also swaps which triangle looks stored, so flip uplo too.
87
+ const uploEffA = effLayoutA === "column-major" ? (uplo === "lower" ? "upper" : "lower") : uplo;
88
+ const transEffA = effLayoutA === "column-major" ? (transA === "no-transpose" ? "transpose" : "no-transpose") : transA;
89
+
90
+ const otherLen = side === "left" ? n : m;
91
+ const blockIsRow = side === "left";
92
+
93
+ // Forward iff op(A) is lower (side='left') — side='right' flips this,
94
+ // since it solves via columns instead of rows.
95
+ const opIsLower = (transEffA === "no-transpose") === (uploEffA === "lower");
96
+ const forward = side === "left" ? opIsLower : !opIsLower;
97
+
98
+ const blockStarts = [];
99
+ for (let s = 0; s < aOrder; s += BLOCK_SIZE) blockStarts.push(s);
100
+ if (!forward) blockStarts.reverse();
101
+ const numBlocks = blockStarts.length;
102
+
103
+ const invertPipeline = await getPipeline(device, "strsv_invert_block");
104
+ const transferPipeline = await getPipeline(device, "block_transfer");
105
+ const scalarPipeline = await getPipeline(device, "sscal");
106
+
107
+ // Null-init here and allocate inside the try below, so a throw partway
108
+ // through the sequence still reaches finally with every handle visible
109
+ // (strsv.mjs is the reference for this pattern).
110
+ let ABuffer = null;
111
+ let BBuffer = null;
112
+ let AinvBuffer = null;
113
+
114
+ const paramsBuffers = [];
115
+ const scratchBuffers = [];
116
+ function scratch(size, label) {
117
+ const buf = createStorageBuffer(device, size, label);
118
+ scratchBuffers.push(buf);
119
+ return buf;
120
+ }
121
+ function params(entries, label) {
122
+ const buf = createParamsBuffer(device, entries, label);
123
+ paramsBuffers.push(buf);
124
+ return buf;
125
+ }
126
+
127
+ // Minimal valid length for B's own buffer (matches its own validation
128
+ // above) — NOT bOuter*ldb, which can exceed a Float32Array-path buffer's
129
+ // actual allocation (only guaranteed padded up to GpuMatrix's own size).
130
+ const bScaleLen = (bOuter - 1) * ldb + bInner;
131
+
132
+ try {
133
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strsm-A", false);
134
+ BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strsm-B", true);
135
+ AinvBuffer = createStorageBuffer(device, numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4, "strsm-Ainv");
136
+
137
+ // Pre-scale B by alpha once (reuses sscal, so no per-block alpha handling).
138
+ let preScaleBindGroup = null;
139
+ if (alpha !== 1.0) {
140
+ const scaleParams = params(
141
+ [
142
+ { value: bScaleLen, type: "u32" },
143
+ { value: alpha, type: "f32" },
144
+ { value: 1, type: "u32" },
145
+ ],
146
+ "strsm-scale-params",
147
+ );
148
+ preScaleBindGroup = createBindGroup(device, scalarPipeline.getBindGroupLayout(0), [BBuffer, scaleParams]);
149
+ }
150
+
151
+ // Every diagonal block's inverse, fully parallel, one dispatch, unchanged.
152
+ const invertParams = params(
153
+ [
154
+ { value: aOrder, type: "u32" },
155
+ { value: lda, type: "u32" },
156
+ { value: transEffA === "transpose" ? 1 : 0, type: "u32" },
157
+ { value: uploEffA === "upper" ? 1 : 0, type: "u32" },
158
+ { value: isUnit ? 1 : 0, type: "u32" },
159
+ ],
160
+ "strsm-invert-params",
161
+ );
162
+ const invertBindGroup = createBindGroup(device, invertPipeline.getBindGroupLayout(0), [ABuffer, AinvBuffer, invertParams]);
163
+
164
+ // Reusable scratch buffers, sized for the worst case, bound at offset 0.
165
+ const Bblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Bblock");
166
+ const Xblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Xblock");
167
+ const Aoff = scratch(aOrder * BLOCK_SIZE * 4, "strsm-Aoff");
168
+ const delta = scratch(aOrder * otherLen * 4, "strsm-delta");
169
+
170
+ const { commandEncoder, querySet } = beginTimedEncoder(device);
171
+
172
+ if (alpha === 0) {
173
+ // BLAS: alpha=0 means A is not referenced — skip straight to B:=0.
174
+ const zeroDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0, endOfPassWriteIndex: 1 } } : undefined;
175
+ encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(device, bScaleLen), zeroDesc);
176
+ } else {
177
+ if (preScaleBindGroup) {
178
+ encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(device, bScaleLen));
179
+ }
180
+ const invertDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
181
+ encodePass(commandEncoder, invertPipeline, invertBindGroup, { x: BLOCK_SIZE, y: numBlocks }, invertDesc);
182
+
183
+ for (let bi = 0; bi < blockStarts.length; bi++) {
184
+ const blockStart = blockStarts[bi];
185
+ const blockEnd = Math.min(blockStart + BLOCK_SIZE, aOrder);
186
+ const blockLen = blockEnd - blockStart;
187
+ const blockIndex = blockStart / BLOCK_SIZE;
188
+ const isLastPass = bi === blockStarts.length - 1;
189
+
190
+ // 1) gather B's current block into a tight scratch buffer.
191
+ const gatherBParams = params(
192
+ [
193
+ { value: blockStart, type: "u32" },
194
+ { value: blockLen, type: "u32" },
195
+ { value: 0, type: "u32" },
196
+ { value: otherLen, type: "u32" },
197
+ { value: ldb, type: "u32" },
198
+ { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
199
+ { value: blockIsRow ? 1 : 0, type: "u32" },
200
+ { value: 2, type: "u32" }, // gather
201
+ ],
202
+ "strsm-gather-B-params",
203
+ );
204
+ const gatherBBindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [Bblock, BBuffer, gatherBParams]);
205
+ encodePass(commandEncoder, transferPipeline, gatherBBindGroup, requireWorkgroups(device, "strsm", blockLen, otherLen));
206
+
207
+ // 2) apply: Xblock := op(Ainv_block) @ Bblock (side='left'), or the
208
+ // transpose-trick equivalent for side='right' (same trick strmm uses).
209
+ {
210
+ const mg = blockLen, ng = otherLen, kg = blockLen;
211
+ const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
212
+ const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
213
+ const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
214
+ const applyParams = params(
215
+ [
216
+ { value: mg, type: "u32" },
217
+ { value: ng, type: "u32" },
218
+ { value: kg, type: "u32" },
219
+ { value: 1.0, type: "f32" }, // alpha already applied to B up front
220
+ { value: 0.0, type: "f32" }, // beta — fresh output, no accumulation
221
+ { value: BLOCK_SIZE, type: "u32" }, // ldX = Ainv's own dense stride
222
+ { value: otherLen, type: "u32" }, // ldY = Bblock's own tight stride
223
+ { value: otherLen, type: "u32" }, // ldc = Xblock's own tight stride
224
+ { value: side === "right" ? 1 : 0, type: "u32" }, // transX: side='right' needs Ainv^T
225
+ { value: 0, type: "u32" }, // transY: Bblock is always read as-is
226
+ ],
227
+ "strsm-apply-params",
228
+ );
229
+ const ainvBlock = { buffer: AinvBuffer, offset: blockIndex * BLOCK_SIZE * BLOCK_SIZE * 4, size: BLOCK_SIZE * BLOCK_SIZE * 4 };
230
+ const applyBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
231
+ ainvBlock,
232
+ vec4ViewBinding(device, ainvBlock),
233
+ Bblock,
234
+ vec4ViewBinding(device, Bblock),
235
+ Xblock,
236
+ applyParams,
237
+ ]);
238
+ const wg = useLarge
239
+ ? { x: requireWorkgroupCount(device, largeWgX, "strsm", "x"), y: requireWorkgroupCount(device, largeWgY, "strsm", "y") }
240
+ : { x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "strsm", "x"), y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "strsm", "y") };
241
+ encodePass(commandEncoder, gemmPipeline, applyBindGroup, wg);
242
+ }
243
+
244
+ // 3) scatter the solved block back into B.
245
+ const rangeStart = forward ? blockEnd : 0;
246
+ const rangeEnd = forward ? aOrder : blockStart;
247
+ const hasRemaining = rangeStart < rangeEnd;
248
+ const scatterParams = params(
249
+ [
250
+ { value: blockStart, type: "u32" },
251
+ { value: blockLen, type: "u32" },
252
+ { value: 0, type: "u32" },
253
+ { value: otherLen, type: "u32" },
254
+ { value: ldb, type: "u32" },
255
+ { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
256
+ { value: blockIsRow ? 1 : 0, type: "u32" },
257
+ { value: 0, type: "u32" }, // overwrite
258
+ ],
259
+ "strsm-scatter-params",
260
+ );
261
+ const scatterBindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [Xblock, BBuffer, scatterParams]);
262
+ const scatterDesc = isLastPass && !hasRemaining && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
263
+ encodePass(commandEncoder, transferPipeline, scatterBindGroup, requireWorkgroups(device, "strsm", blockLen, otherLen), scatterDesc);
264
+
265
+ // 4) trailing update: subtract this block's contribution from B.
266
+ if (!hasRemaining) continue;
267
+ const remCount = rangeEnd - rangeStart;
268
+
269
+ const gatherAParams = params(
270
+ [
271
+ { value: rangeStart, type: "u32" },
272
+ { value: remCount, type: "u32" },
273
+ { value: blockStart, type: "u32" },
274
+ { value: blockLen, type: "u32" },
275
+ { value: lda, type: "u32" },
276
+ { value: transEffA === "transpose" ? 1 : 0, type: "u32" },
277
+ { value: blockIsRow ? 1 : 0, type: "u32" },
278
+ { value: 2, type: "u32" }, // gather
279
+ ],
280
+ "strsm-gather-A-params",
281
+ );
282
+ const gatherABindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [Aoff, ABuffer, gatherAParams]);
283
+ encodePass(commandEncoder, transferPipeline, gatherABindGroup, requireWorkgroups(device, "strsm", remCount, blockLen));
284
+
285
+ {
286
+ const mg = remCount, ng = otherLen, kg = blockLen;
287
+ const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
288
+ const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
289
+ const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
290
+ const updateParams = params(
291
+ [
292
+ { value: mg, type: "u32" },
293
+ { value: ng, type: "u32" },
294
+ { value: kg, type: "u32" },
295
+ { value: 1.0, type: "f32" },
296
+ { value: 0.0, type: "f32" },
297
+ { value: blockLen, type: "u32" }, // ldX = Aoff's own tight stride
298
+ { value: otherLen, type: "u32" }, // ldY = Xblock's own tight stride
299
+ { value: otherLen, type: "u32" }, // ldc = delta's own tight stride
300
+ { value: 0, type: "u32" }, // transX: Aoff already read in the right orientation
301
+ { value: 0, type: "u32" }, // transY: Xblock read as-is
302
+ ],
303
+ "strsm-update-params",
304
+ );
305
+ const updateBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
306
+ Aoff,
307
+ vec4ViewBinding(device, Aoff),
308
+ Xblock,
309
+ vec4ViewBinding(device, Xblock),
310
+ delta,
311
+ updateParams,
312
+ ]);
313
+ const wg = useLarge
314
+ ? { x: requireWorkgroupCount(device, largeWgX, "strsm", "x"), y: requireWorkgroupCount(device, largeWgY, "strsm", "y") }
315
+ : { x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "strsm", "x"), y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "strsm", "y") };
316
+ encodePass(commandEncoder, gemmPipeline, updateBindGroup, wg);
317
+ }
318
+
319
+ const scatterSubParams = params(
320
+ [
321
+ { value: rangeStart, type: "u32" },
322
+ { value: remCount, type: "u32" },
323
+ { value: 0, type: "u32" },
324
+ { value: otherLen, type: "u32" },
325
+ { value: ldb, type: "u32" },
326
+ { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
327
+ { value: blockIsRow ? 1 : 0, type: "u32" },
328
+ { value: 1, type: "u32" }, // subtract
329
+ ],
330
+ "strsm-scatter-sub-params",
331
+ );
332
+ const scatterSubBindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [delta, BBuffer, scatterSubParams]);
333
+ const subDesc = isLastPass && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
334
+ encodePass(commandEncoder, transferPipeline, scatterSubBindGroup, requireWorkgroups(device, "strsm", remCount, otherLen), subDesc);
335
+ }
336
+ }
337
+
338
+ const ts = resolveTimestamp(device, commandEncoder, querySet);
339
+ const readBuffer = BIsGpu ? null : stageReadback(device, commandEncoder, BBuffer);
340
+
341
+ submit(device, commandEncoder);
342
+
343
+ const gpuTimeMs = await extractTimestamp(ts);
344
+
345
+ if (BIsGpu) {
346
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
347
+ return {};
348
+ }
349
+
350
+ const result = await extractResult(readBuffer, Float32Array);
351
+ if (gpuTimeMs !== undefined) return { B: result, gpuTimeMs };
352
+ return { B: result };
353
+ } finally {
354
+ if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
355
+ if (!BIsGpu && BBuffer) destroyBuffers(BBuffer);
356
+ if (AinvBuffer) destroyBuffers(AinvBuffer);
357
+ destroyBuffers(scratchBuffers);
358
+ destroyBuffers(paramsBuffers);
359
+ }
360
+ }
@@ -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
@@ -12,9 +12,10 @@ import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
12
12
  import { getPipeline } from "../util/pipeline.mjs";
13
13
  import { GpuVector } from "../classes/GpuVector.mjs";
14
14
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
15
+ import { BLOCK_SIZE } from "../util/constants.mjs";
16
+ import { requireSameDevice } from "../util/device.mjs";
15
17
 
16
18
  // Blocked triangular solve via explicit block inversion (invert/apply/update passes) instead of barrier-per-row substitution.
17
- const BLOCK_SIZE = 64;
18
19
 
19
20
  // One shared buffer holds all blocks' params (offset blockIndex*stride) instead of one buffer per block — avoids the O(numBlocks) createBuffer/writeBuffer calls that dominated CPU time.
20
21
  function packBlockParams(numBlocks, stride, fieldsPerBlock) {
@@ -45,6 +46,7 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
45
46
 
46
47
  if (!(device instanceof GPUDevice))
47
48
  throw new Error("device must be a GPUDevice.");
49
+ requireSameDevice(device, "strsv", { A, x });
48
50
  if (uplo !== "lower" && uplo !== "upper")
49
51
  throw new Error("uplo must be 'lower' or 'upper'.");
50
52
  if (trans !== "no-transpose" && trans !== "transpose")
@@ -63,6 +65,10 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
63
65
  throw new Error("x must be a Float32Array or GpuVector.");
64
66
  if (xIsGpu && !AIsGpu)
65
67
  throw new Error("A must be a GpuMatrix when x is a GpuVector.");
68
+ if (AIsGpu && !xIsGpu)
69
+ throw new Error("x must be a GpuVector when A is a GpuMatrix.");
70
+ if (AIsGpu && xIsGpu && A._buf === x._buf)
71
+ throw new Error("A and x must not reference the same GPU buffer.");
66
72
  if (AIsGpu && lda !== A.lda)
67
73
  throw new Error("lda must match A.lda when A is a GpuMatrix.");
68
74
  if (AIsGpu && (A.rows < n || A.cols < n))
@@ -107,11 +113,11 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
107
113
  let invertParams = null;
108
114
 
109
115
  try {
110
- ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "strsv-A", false);
111
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "strsv-x", true);
116
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strsv-A", false);
117
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "strsv-x", true);
112
118
  // One BLOCK_SIZE x BLOCK_SIZE dense region per block (row-major), even
113
119
  // though only a triangular half is ever nonzero — see strsv_invert_block.wgsl.
114
- AinvBuffer = createStorageBuffer(
120
+ AinvBuffer = createStorageBuffer(device,
115
121
  numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4,
116
122
  "strsv-Ainv",
117
123
  );
@@ -133,10 +139,10 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
133
139
  });
134
140
  updateParamsBuffer = createSharedParamsBuffer(device, updateData, "strsv-update-params");
135
141
 
136
- const { commandEncoder, querySet } = beginTimedEncoder();
142
+ const { commandEncoder, querySet } = beginTimedEncoder(device);
137
143
 
138
144
  // Pre-pass: every block's inverse, fully parallel, one dispatch.
139
- invertParams = createParamsBuffer(
145
+ invertParams = createParamsBuffer(device,
140
146
  [
141
147
  { value: n, type: "u32" },
142
148
  { value: lda, type: "u32" },
@@ -146,7 +152,7 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
146
152
  ],
147
153
  "strsv-invert-params",
148
154
  );
149
- const invertBindGroup = createBindGroup(invertPipeline.getBindGroupLayout(0), [
155
+ const invertBindGroup = createBindGroup(device, invertPipeline.getBindGroupLayout(0), [
150
156
  ABuffer, AinvBuffer, invertParams,
151
157
  ]);
152
158
  const invertDesc = querySet
@@ -161,7 +167,7 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
161
167
  const isLastPass = bi === blockStarts.length - 1;
162
168
  const paramsOffset = blockIndex * stride;
163
169
 
164
- const applyBindGroup = createBindGroup(applyPipeline.getBindGroupLayout(0), [
170
+ const applyBindGroup = createBindGroup(device, applyPipeline.getBindGroupLayout(0), [
165
171
  AinvBuffer, xBuffer, { buffer: applyParamsBuffer, offset: paramsOffset, size: 16 },
166
172
  ]);
167
173
 
@@ -173,7 +179,7 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
173
179
  const remaining = forward ? n - blockEnd : blockStart;
174
180
  if (remaining === 0) continue;
175
181
 
176
- const updateBindGroup = createBindGroup(updatePipeline.getBindGroupLayout(0), [
182
+ const updateBindGroup = createBindGroup(device, updatePipeline.getBindGroupLayout(0), [
177
183
  ABuffer, xBuffer, { buffer: updateParamsBuffer, offset: paramsOffset, size: 32 },
178
184
  ]);
179
185
 
@@ -181,10 +187,10 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
181
187
  encodePass(commandEncoder, updatePipeline, updateBindGroup, wgCount);
182
188
  }
183
189
 
184
- const ts = resolveTimestamp(commandEncoder, querySet);
185
- const readBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
190
+ const ts = resolveTimestamp(device, commandEncoder, querySet);
191
+ const readBuffer = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
186
192
 
187
- submit(commandEncoder);
193
+ submit(device, commandEncoder);
188
194
 
189
195
  const gpuTimeMs = await extractTimestamp(ts);
190
196
 
@@ -1,5 +1,5 @@
1
1
  /** @module devdocs/utility-functions/benchmark */
2
- import { getDevice, isBenchmarkEnabled } from "../init.mjs";
2
+ import { isBenchmarkEnabled } from "../init.mjs";
3
3
 
4
4
  /**
5
5
  * Returns the `requestDevice` descriptor to pass to `adapter.requestDevice()`.
@@ -29,10 +29,9 @@ export function benchmarkMode(adapter, enabled) {
29
29
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUQuerySet GPUQuerySet}
30
30
  * @see {@link https://developer.chrome.com/blog/new-in-webgpu-121 Chrome 121 — timestamp queries (querySet + timestampWrites pattern, quantization caveat)}
31
31
  */
32
- export function beginTimestamp() {
33
- if (!isBenchmarkEnabled())
32
+ export function beginTimestamp(device) {
33
+ if (!isBenchmarkEnabled(device))
34
34
  return { querySet: null, passDescriptor: undefined };
35
- const device = getDevice();
36
35
  // Two slots: index 0 written when the pass begins, index 1 when it ends.
37
36
  const querySet = device.createQuerySet({ type: "timestamp", count: 2 });
38
37
  const passDescriptor = {
@@ -55,9 +54,8 @@ export function beginTimestamp() {
55
54
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/resolveQuerySet GPUCommandEncoder.resolveQuerySet()}
56
55
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/copyBufferToBuffer GPUCommandEncoder.copyBufferToBuffer()}
57
56
  */
58
- export function resolveTimestamp(commandEncoder, querySet) {
57
+ export function resolveTimestamp(device, commandEncoder, querySet) {
59
58
  if (!querySet) return null;
60
- const device = getDevice();
61
59
  // QUERY_RESOLVE and MAP_READ cannot be combined — two buffers are required.
62
60
  // resolveBuffer: GPU writes resolved nanosecond timestamps here.
63
61
  const resolveBuffer = device.createBuffer({
@@ -1,5 +1,4 @@
1
1
  /** @module devdocs/utility-functions/bindgroup */
2
- import { getDevice } from "../init.mjs";
3
2
 
4
3
  /**
5
4
  * Creates a `GPUBindGroup` by mapping each buffer to sequential binding indices
@@ -16,8 +15,7 @@ import { getDevice } from "../init.mjs";
16
15
  * @returns {GPUBindGroup}
17
16
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createBindGroup GPUDevice.createBindGroup()}
18
17
  */
19
- export function createBindGroup(layout, buffers, startBinding = 0) {
20
- const device = getDevice();
18
+ export function createBindGroup(device, layout, buffers, startBinding = 0) {
21
19
  const entries = buffers.map((buffer, i) => ({
22
20
  binding: startBinding + i,
23
21
  resource: buffer instanceof GPUBuffer ? { buffer } : buffer,