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,100 @@
1
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
+
3
+ /**
4
+ * Performs the triangular matrix-matrix operation
5
+ * B := alpha * op(A) * B (`side='left'`) or
6
+ * B := alpha * B * op(A) (`side='right'`) — `A` is triangular, only its
7
+ * `uplo` triangle stored; `B` is a general m×n matrix, overwritten in place.
8
+ *
9
+ * - `side='left'`: `A` is m×m — `A` premultiplies `B`
10
+ * - `side='right'`: `A` is n×n — `A` postmultiplies `B`
11
+ *
12
+ * No dedicated fused kernel — a `triangularize` pass materializes a dense
13
+ * copy of `op(A)` (zero-filling the unstored triangle, substituting the
14
+ * implicit diagonal when `diag='unit'`), then a plain `sgemm` pass
15
+ * (`sgemm_small.wgsl`/`sgemm_large.wgsl`, unmodified) does the actual
16
+ * multiply, both on one command encoder.
17
+ *
18
+ * {@includeCode ../../examples/strmm/strmm.js}
19
+ *
20
+ * **Browser (standalone HTML):**
21
+ * {@includeCode ../../examples/strmm/web/strmm.html}
22
+ *
23
+ * @param device - GPUDevice from `init()`
24
+ * @param side - `'left'` for op(A)*B, `'right'` for B*op(A)
25
+ * @param uplo - `'lower'` if only `A`'s lower triangle is stored, `'upper'` for upper
26
+ * @param transA - `'no-transpose'` for A, `'transpose'` for A^T
27
+ * @param diag - `'unit'` to treat A's diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
28
+ * @param m - rows of B
29
+ * @param n - columns of B
30
+ * @param alpha - scalar multiplier for the matrix product
31
+ * @param A - Float32Array, triangular, row-major or column-major (see `layout`)
32
+ * @param lda - leading dimension of A as stored
33
+ * @param B - Float32Array input/output matrix, overwritten with the result, row-major or column-major
34
+ * @param ldb - leading dimension of B as stored
35
+ * @param layout - storage layout shared by A/B when they're Float32Array
36
+ * (default: `'row-major'`) — column-major A is a genuine transpose (A isn't
37
+ * symmetric like ssymm's), so both `transA` and `uplo` are adjusted
38
+ * internally to compensate
39
+ * @returns updated B as a Float32Array
40
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strmm/strmm.mjs#L25">Source code: strmm.mjs (L25)</a>
41
+ * @category BLAS Level 3
42
+ */
43
+ export declare function strmm(
44
+ device: GPUDevice,
45
+ side: 'left' | 'right',
46
+ uplo: 'lower' | 'upper',
47
+ transA: 'no-transpose' | 'transpose',
48
+ diag: 'unit' | 'non-unit',
49
+ m: number,
50
+ n: number,
51
+ alpha: number,
52
+ A: Float32Array,
53
+ lda: number,
54
+ B: Float32Array,
55
+ ldb: number,
56
+ layout?: 'row-major' | 'column-major',
57
+ ): Promise<{ B: Float32Array; gpuTimeMs?: number }>;
58
+
59
+ /**
60
+ * Performs the triangular matrix-matrix operation
61
+ * B := alpha * op(A) * B (`side='left'`) or
62
+ * B := alpha * B * op(A) (`side='right'`)
63
+ *
64
+ * A and B are both kept GPU-resident; B is mutated in place. Each matrix's
65
+ * own `layout` (set at `GpuMatrix.from` time) determines the operation —
66
+ * there is no separate `layout` argument here. A and B must both be
67
+ * GpuMatrix or both be Float32Array — mixing is not supported.
68
+ *
69
+ * {@includeCode ../../examples/strmm/gpu.strmm.js}
70
+ *
71
+ * @param device - GPUDevice from `init()`
72
+ * @param side - `'left'` for op(A)*B, `'right'` for B*op(A)
73
+ * @param uplo - `'lower'` if only `A`'s lower triangle is stored, `'upper'` for upper
74
+ * @param transA - `'no-transpose'` for A, `'transpose'` for A^T
75
+ * @param diag - `'unit'` to treat A's diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
76
+ * @param m - rows of B
77
+ * @param n - columns of B
78
+ * @param alpha - scalar multiplier for the matrix product
79
+ * @param A - GpuMatrix, triangular
80
+ * @param lda - leading dimension of A (must equal A.lda)
81
+ * @param B - GpuMatrix (mutated in place)
82
+ * @param ldb - leading dimension of B (must equal B.lda)
83
+ * @returns no B — it stays GPU-resident; call `B.read()` yourself for a CPU readback (see the example)
84
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strmm/strmm.mjs#L25">Source code: strmm.mjs (L25)</a>
85
+ * @category BLAS Level 3
86
+ */
87
+ export declare function strmm(
88
+ device: GPUDevice,
89
+ side: 'left' | 'right',
90
+ uplo: 'lower' | 'upper',
91
+ transA: 'no-transpose' | 'transpose',
92
+ diag: 'unit' | 'non-unit',
93
+ m: number,
94
+ n: number,
95
+ alpha: number,
96
+ A: GpuMatrix,
97
+ lda: number,
98
+ B: GpuMatrix,
99
+ ldb: number,
100
+ ): Promise<{ gpuTimeMs?: number }>;
@@ -0,0 +1,226 @@
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 { GpuMatrix } from "../classes/GpuMatrix.mjs";
15
+ import { requireWorkgroupCount } from "../util/workgroup.mjs";
16
+ import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
17
+ import { TILE_WG_2D } from "../util/constants.mjs";
18
+ import { requireSameDevice } from "../util/device.mjs";
19
+
20
+
21
+ // strmm: B := alpha*op(A)*B (side='left') or alpha*B*op(A) (side='right'), A
22
+ // triangular. Triangularize then sgemm, one command encoder. B is both
23
+ // input and output, so gemm writes to a fresh buffer (no aliasing race),
24
+ // copied back into B (GpuMatrix) or read back directly (Float32Array).
25
+ export async function strmm(
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, "strmm", { 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 transB = effLayoutB === "column-major" ? "transpose" : "no-transpose";
91
+ const transDense = "no-transpose"; // Adense already embodies op(A)
92
+
93
+ // X*Y (X=Adense,Y=B for side='left', swapped for 'right'). Column-major
94
+ // output: compute (B_out)^T instead — sgemm's own trick, same as ssymm's.
95
+ let mg = m, ng = n;
96
+ const kg = aOrder;
97
+ let transX = side === "left" ? transDense : transB;
98
+ let transY = side === "left" ? transB : transDense;
99
+ const flip = (t) => (t === "no-transpose" ? "transpose" : "no-transpose");
100
+ let swapXY = side === "right";
101
+ if (effLayoutB === "column-major") {
102
+ [transX, transY] = [flip(transY), flip(transX)];
103
+ swapXY = !swapXY;
104
+ [mg, ng] = [ng, mg];
105
+ }
106
+
107
+ const ldDense = aOrder; // Adense is tightly packed, row-major
108
+ const largeWgX = Math.ceil(ng / BN_LARGE);
109
+ const largeWgY = Math.ceil(mg / BM_LARGE);
110
+ const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
111
+ const gemmPipeline = await getPipeline(device, useLargeTile ? "sgemm_large" : "sgemm_small");
112
+ const triPipeline = await getPipeline(device, "triangularize");
113
+ const gemmWgCount = useLargeTile
114
+ ? {
115
+ x: requireWorkgroupCount(device, largeWgX, "strmm", "x"),
116
+ y: requireWorkgroupCount(device, largeWgY, "strmm", "y"),
117
+ }
118
+ : {
119
+ x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "strmm", "x"),
120
+ y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "strmm", "y"),
121
+ };
122
+
123
+ // Null-init here and allocate inside the try below, so a throw partway
124
+ // through the sequence still reaches finally with every handle visible
125
+ // (strsv.mjs is the reference for this pattern).
126
+ let ABuffer = null, BBuffer = null;
127
+ let AdenseBuffer = null, outBuffer = null;
128
+ let triParams = null, gemmParams = null;
129
+ let outBufferAdopted = false; // true once B._buf is repointed at outBuffer
130
+
131
+ try {
132
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmm-A", false);
133
+ // readback=true (COPY_SRC): BBuffer is also the source that seeds outBuffer.
134
+ BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strmm-B", true);
135
+ AdenseBuffer = createStorageBuffer(device, aOrder * ldDense * 4, "strmm-Adense");
136
+ // COPY_DST: seeded from B's own content before gemm runs, so stride-padding
137
+ // gaps (never written by gemm's tight m x n loop) keep B's original bytes
138
+ // instead of reading back as zero. COPY_SRC: read back / adopted by B after.
139
+ outBuffer = createStorageBuffer(device,
140
+ bOuter * ldb * 4, "strmm-out", GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST,
141
+ );
142
+
143
+ triParams = createParamsBuffer(device,
144
+ [
145
+ { value: aOrder, type: "u32" },
146
+ { value: lda, type: "u32" },
147
+ { value: ldDense, type: "u32" },
148
+ { value: uploEffA === "upper" ? 1 : 0, type: "u32" },
149
+ { value: transEffA === "transpose" ? 1 : 0, type: "u32" },
150
+ { value: isUnit ? 1 : 0, type: "u32" },
151
+ ],
152
+ "strmm-tri-params",
153
+ );
154
+ const triBindGroup = createBindGroup(device, triPipeline.getBindGroupLayout(0), [ABuffer, AdenseBuffer, triParams]);
155
+
156
+ // X/Y buffers and their own ld, matching swapXY above.
157
+ const XBuffer = swapXY ? BBuffer : AdenseBuffer;
158
+ const ldX = swapXY ? ldb : ldDense;
159
+ const YBuffer = swapXY ? AdenseBuffer : BBuffer;
160
+ const ldY = swapXY ? ldDense : ldb;
161
+
162
+ gemmParams = createParamsBuffer(device,
163
+ [
164
+ { value: mg, type: "u32" },
165
+ { value: ng, type: "u32" },
166
+ { value: kg, type: "u32" },
167
+ { value: alpha, type: "f32" },
168
+ { value: 0.0, type: "f32" }, // beta — strmm has no C accumulation term
169
+ { value: ldX, type: "u32" },
170
+ { value: ldY, type: "u32" },
171
+ { value: ldb, type: "u32" },
172
+ { value: transX === "transpose" ? 1 : 0, type: "u32" },
173
+ { value: transY === "transpose" ? 1 : 0, type: "u32" },
174
+ ],
175
+ "strmm-gemm-params",
176
+ );
177
+ const gemmBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
178
+ XBuffer,
179
+ vec4ViewBinding(device, XBuffer),
180
+ YBuffer,
181
+ vec4ViewBinding(device, YBuffer),
182
+ outBuffer,
183
+ gemmParams,
184
+ ]);
185
+
186
+ const { commandEncoder, querySet } = beginTimedEncoder(device);
187
+ // Seed outBuffer with B's own bytes first, so gemm's tight m x n write
188
+ // leaves stride-padding gaps holding B's original content, not zero.
189
+ // BBuffer may be larger than outBuffer (e.g. a validation-test baseline
190
+ // over-provisioned for a bigger ldb it might later be substituted with).
191
+ commandEncoder.copyBufferToBuffer(BBuffer, 0, outBuffer, 0, Math.min(BBuffer.size, outBuffer.size));
192
+ const triDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
193
+ const gemmDesc = querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
194
+ encodePass(commandEncoder, triPipeline, triBindGroup, { x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) }, triDesc);
195
+ encodePass(commandEncoder, gemmPipeline, gemmBindGroup, gemmWgCount, gemmDesc);
196
+
197
+ const ts = resolveTimestamp(device, commandEncoder, querySet);
198
+ const readBuffer = BIsGpu ? null : stageReadback(device, commandEncoder, outBuffer);
199
+
200
+ submit(device, commandEncoder);
201
+
202
+ const gpuTimeMs = await extractTimestamp(ts);
203
+
204
+ if (BIsGpu) {
205
+ // Adopt outBuffer as B's own backing buffer instead of copying into the
206
+ // old one (B._buf has no COPY_DST usage) — cheaper and avoids needing
207
+ // an extra buffer-usage flag on every GpuMatrix for this one routine.
208
+ destroyBuffers(B._buf);
209
+ B._buf = outBuffer;
210
+ outBufferAdopted = true;
211
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
212
+ return {};
213
+ }
214
+
215
+ const result = await extractResult(readBuffer, Float32Array);
216
+ if (gpuTimeMs !== undefined) return { B: result, gpuTimeMs };
217
+ return { B: result };
218
+ } finally {
219
+ if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
220
+ if (!BIsGpu && BBuffer) destroyBuffers(BBuffer);
221
+ if (AdenseBuffer) destroyBuffers(AdenseBuffer);
222
+ if (outBuffer && !outBufferAdopted) destroyBuffers(outBuffer);
223
+ if (triParams) destroyBuffers(triParams);
224
+ if (gemmParams) destroyBuffers(gemmParams);
225
+ }
226
+ }
@@ -44,41 +44,6 @@ export declare function strmv(
44
44
  layout?: 'row-major' | 'column-major',
45
45
  ): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
46
46
 
47
- /**
48
- * Performs the triangular matrix-vector operation y = op(A) * x
49
- *
50
- * A is kept GPU-resident; x and y are CPU Float32Arrays. `A`'s own `layout`
51
- * (set at `GpuMatrix.from` time) determines the operation — there is no
52
- * separate `layout` argument here.
53
- *
54
- * @param device - GPUDevice from `init()`
55
- * @param uplo - `'lower'` to use the lower triangle, `'upper'` to use the upper triangle
56
- * @param trans - `'no-transpose'` for A, `'transpose'` for A^T
57
- * @param diag - `'unit'` to treat the diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
58
- * @param n - order of the matrix A
59
- * @param A - GpuMatrix, GPU-resident
60
- * @param lda - leading dimension of A (must equal A.lda)
61
- * @param x - Float32Array input vector
62
- * @param incx - stride for x (must be a positive integer)
63
- * @param y - Float32Array output vector
64
- * @param incy - stride for y (must be a positive integer)
65
- * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strmv/strmv.mjs#L15">Source code: strmv.mjs (L15)</a>
66
- * @category BLAS Level 2
67
- */
68
- export declare function strmv(
69
- device: GPUDevice,
70
- uplo: 'lower' | 'upper',
71
- trans: 'no-transpose' | 'transpose',
72
- diag: 'unit' | 'non-unit',
73
- n: number,
74
- A: GpuMatrix,
75
- lda: number,
76
- x: Float32Array,
77
- incx: number,
78
- y: Float32Array,
79
- incy: number,
80
- ): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
81
-
82
47
  /**
83
48
  * Performs the triangular matrix-vector operation y = op(A) * x
84
49
  *
@@ -86,7 +51,7 @@ export declare function strmv(
86
51
  * `layout` (set at `GpuMatrix.from` time) determines the operation — there is
87
52
  * no separate `layout` argument here.
88
53
  *
89
- * {@includeCode ../../examples/strmv/gpuvec.strmv.js}
54
+ * {@includeCode ../../examples/strmv/gpu.strmv.js}
90
55
  *
91
56
  * @param device - GPUDevice from `init()`
92
57
  * @param uplo - `'lower'` to use the lower triangle, `'upper'` to use the upper triangle
@@ -11,6 +11,7 @@ import { extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { GpuVector } from "../classes/GpuVector.mjs";
13
13
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
14
+ import { requireSameDevice } from "../util/device.mjs";
14
15
 
15
16
  export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, incy, layout = "row-major") {
16
17
  const xIsGpu = x instanceof GpuVector;
@@ -20,6 +21,7 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
20
21
 
21
22
  if (!(device instanceof GPUDevice))
22
23
  throw new Error("device must be a GPUDevice.");
24
+ requireSameDevice(device, "strmv", { A, x, y });
23
25
  if (uplo !== "lower" && uplo !== "upper")
24
26
  throw new Error("uplo must be 'lower' or 'upper'.");
25
27
  if (trans !== "no-transpose" && trans !== "transpose")
@@ -54,6 +56,8 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
54
56
  );
55
57
  if (xIsGpu && !AIsGpu)
56
58
  throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");
59
+ if (AIsGpu && !xIsGpu)
60
+ throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");
57
61
  if (AIsGpu && yIsGpu && A._buf === y._buf)
58
62
  throw new Error("A and y must not reference the same GPU buffer.");
59
63
  if (AIsGpu && lda !== A.lda)
@@ -90,10 +94,10 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
90
94
  let paramsBuffer = null;
91
95
 
92
96
  try {
93
- ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "strmv-A", false);
94
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "strmv-x", false);
95
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "strmv-y", true);
96
- paramsBuffer = createParamsBuffer(
97
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmv-A", false);
98
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "strmv-x", false);
99
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "strmv-y", true);
100
+ paramsBuffer = createParamsBuffer(device,
97
101
  [
98
102
  { value: n, type: "u32" },
99
103
  { value: incx, type: "u32" },
@@ -106,7 +110,7 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
106
110
  "strmv-params",
107
111
  );
108
112
 
109
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
113
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
110
114
  ABuffer,
111
115
  xBuffer,
112
116
  yBuffer,
@@ -114,10 +118,10 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
114
118
  ]);
115
119
 
116
120
  const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
117
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
118
- const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
121
+ const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
122
+ const readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
119
123
 
120
- submit(commandEncoder);
124
+ submit(device, commandEncoder);
121
125
 
122
126
  const gpuTimeMs = await extractTimestamp(ts);
123
127
 
@@ -0,0 +1,99 @@
1
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
+
3
+ /**
4
+ * Solves the triangular matrix equation
5
+ * op(A) * X = alpha * B (`side='left'`) or
6
+ * X * op(A) = alpha * B (`side='right'`), overwriting `B` with the
7
+ * solution `X` — `A` is triangular, only its `uplo` triangle stored; `B` is
8
+ * a general m×n matrix.
9
+ *
10
+ * - `side='left'`: `A` is m×m — solves `op(A)*X = alpha*B`
11
+ * - `side='right'`: `A` is n×n — solves `X*op(A) = alpha*B`
12
+ *
13
+ * Blocked substitution (strsv's own technique, generalized to a matrix
14
+ * RHS): strsv_invert_block + sgemm, both unmodified. A near-zero diagonal
15
+ * entry amplifies error, same as any triangular solve.
16
+ *
17
+ * {@includeCode ../../examples/strsm/strsm.js}
18
+ *
19
+ * **Browser (standalone HTML):**
20
+ * {@includeCode ../../examples/strsm/web/strsm.html}
21
+ *
22
+ * @param device - GPUDevice from `init()`
23
+ * @param side - `'left'` to solve op(A)*X=alpha*B, `'right'` to solve X*op(A)=alpha*B
24
+ * @param uplo - `'lower'` if only `A`'s lower triangle is stored, `'upper'` for upper
25
+ * @param transA - `'no-transpose'` for A, `'transpose'` for A^T
26
+ * @param diag - `'unit'` to treat A's diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
27
+ * @param m - rows of B
28
+ * @param n - columns of B
29
+ * @param alpha - scalar multiplier for B
30
+ * @param A - Float32Array, triangular, row-major or column-major (see `layout`)
31
+ * @param lda - leading dimension of A as stored
32
+ * @param B - Float32Array input/output matrix, overwritten with the solution, row-major or column-major
33
+ * @param ldb - leading dimension of B as stored
34
+ * @param layout - storage layout shared by A/B when they're Float32Array
35
+ * (default: `'row-major'`) — column-major A is a genuine transpose (A isn't
36
+ * symmetric like ssymm's), so both `transA` and `uplo` are adjusted
37
+ * internally to compensate
38
+ * @returns the solution X, written into B, as a Float32Array
39
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strsm/strsm.mjs#L25">Source code: strsm.mjs (L25)</a>
40
+ * @category BLAS Level 3
41
+ */
42
+ export declare function strsm(
43
+ device: GPUDevice,
44
+ side: 'left' | 'right',
45
+ uplo: 'lower' | 'upper',
46
+ transA: 'no-transpose' | 'transpose',
47
+ diag: 'unit' | 'non-unit',
48
+ m: number,
49
+ n: number,
50
+ alpha: number,
51
+ A: Float32Array,
52
+ lda: number,
53
+ B: Float32Array,
54
+ ldb: number,
55
+ layout?: 'row-major' | 'column-major',
56
+ ): Promise<{ B: Float32Array; gpuTimeMs?: number }>;
57
+
58
+ /**
59
+ * Solves the triangular matrix equation
60
+ * op(A) * X = alpha * B (`side='left'`) or
61
+ * X * op(A) = alpha * B (`side='right'`), overwriting `B` in place with `X`.
62
+ *
63
+ * A and B are both kept GPU-resident. Each matrix's own `layout` (set at
64
+ * `GpuMatrix.from` time) determines the operation — there is no separate
65
+ * `layout` argument here. A and B must both be GpuMatrix or both be
66
+ * Float32Array — mixing is not supported.
67
+ *
68
+ * {@includeCode ../../examples/strsm/gpu.strsm.js}
69
+ *
70
+ * @param device - GPUDevice from `init()`
71
+ * @param side - `'left'` to solve op(A)*X=alpha*B, `'right'` to solve X*op(A)=alpha*B
72
+ * @param uplo - `'lower'` if only `A`'s lower triangle is stored, `'upper'` for upper
73
+ * @param transA - `'no-transpose'` for A, `'transpose'` for A^T
74
+ * @param diag - `'unit'` to treat A's diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
75
+ * @param m - rows of B
76
+ * @param n - columns of B
77
+ * @param alpha - scalar multiplier for B
78
+ * @param A - GpuMatrix, triangular
79
+ * @param lda - leading dimension of A (must equal A.lda)
80
+ * @param B - GpuMatrix (overwritten in place with the solution)
81
+ * @param ldb - leading dimension of B (must equal B.lda)
82
+ * @returns no B — it stays GPU-resident; call `B.read()` yourself for a CPU readback (see the example)
83
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strsm/strsm.mjs#L25">Source code: strsm.mjs (L25)</a>
84
+ * @category BLAS Level 3
85
+ */
86
+ export declare function strsm(
87
+ device: GPUDevice,
88
+ side: 'left' | 'right',
89
+ uplo: 'lower' | 'upper',
90
+ transA: 'no-transpose' | 'transpose',
91
+ diag: 'unit' | 'non-unit',
92
+ m: number,
93
+ n: number,
94
+ alpha: number,
95
+ A: GpuMatrix,
96
+ lda: number,
97
+ B: GpuMatrix,
98
+ ldb: number,
99
+ ): Promise<{ gpuTimeMs?: number }>;