wgblas 2.1.0 → 2.2.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.
- package/README.md +2 -0
- package/dist/wgblas.browser.js +1007 -43
- package/index.d.mts +26 -53
- package/index.mjs +11 -0
- package/package.json +132 -63
- package/src/classes/Complex32.d.mts +43 -0
- package/src/classes/Complex32.mjs +82 -0
- package/src/classes/Complex64.d.mts +44 -0
- package/src/classes/Complex64.mjs +76 -0
- package/src/classes/GpuMatrix.d.mts +41 -41
- package/src/classes/GpuMatrix.mjs +112 -10
- package/src/classes/GpuVector.d.mts +36 -40
- package/src/classes/GpuVector.mjs +39 -2
- package/src/cscal/cscal.d.mts +47 -0
- package/src/cscal/cscal.mjs +98 -0
- package/src/dasum/dasum.mjs +31 -15
- package/src/daxpy/daxpy.d.mts +56 -0
- package/src/daxpy/daxpy.mjs +150 -0
- package/src/dcopy/dcopy.d.mts +52 -0
- package/src/dcopy/dcopy.mjs +140 -0
- package/src/ddot/ddot.d.mts +62 -0
- package/src/ddot/ddot.mjs +184 -0
- package/src/dnrm2/dnrm2.d.mts +50 -0
- package/src/dnrm2/dnrm2.mjs +189 -0
- package/src/drot/drot.d.mts +67 -0
- package/src/drot/drot.mjs +170 -0
- package/src/drotm/drotm.d.mts +67 -0
- package/src/drotm/drotm.mjs +171 -0
- package/src/dscal/dscal.d.mts +52 -0
- package/src/dscal/dscal.mjs +119 -0
- package/src/dswap/dswap.d.mts +57 -0
- package/src/dswap/dswap.mjs +155 -0
- package/src/idamax/idamax.mjs +49 -19
- package/src/init.mjs +6 -3
- package/src/isamax/isamax.mjs +17 -14
- package/src/random/random.d.mts +37 -40
- package/src/random/random.mjs +39 -7
- package/src/sasum/sasum.mjs +13 -11
- package/src/saxpy/saxpy.mjs +9 -8
- package/src/scopy/scopy.mjs +8 -6
- package/src/sdot/sdot.mjs +13 -11
- package/src/sgemm/sgemm.d.mts +2 -2
- package/src/sgemm/sgemm.mjs +91 -35
- package/src/sgemmtr/sgemmtr.d.mts +3 -2
- package/src/sgemmtr/sgemmtr.mjs +92 -35
- package/src/sgemv/sgemv.d.mts +2 -2
- package/src/sgemv/sgemv.mjs +41 -25
- package/src/sger/sger.d.mts +2 -2
- package/src/sger/sger.mjs +38 -16
- package/src/shaders/__test_pipeline_a.wgsl +3 -0
- package/src/shaders/__test_pipeline_b.wgsl +2 -0
- package/src/shaders/cscal.wgsl +33 -0
- package/src/shaders/daxpy.wgsl +66 -0
- package/src/shaders/dcopy.wgsl +34 -0
- package/src/shaders/ddot.wgsl +106 -0
- package/src/shaders/dnrm2.wgsl +167 -0
- package/src/shaders/drot.wgsl +81 -0
- package/src/shaders/drotm.wgsl +99 -0
- package/src/shaders/dscal.wgsl +60 -0
- package/src/shaders/dswap.wgsl +38 -0
- package/src/shaders/f64/utils/add.wgsl +6 -0
- package/src/shaders/f64/utils/divide.wgsl +45 -0
- package/src/shaders/f64/utils/multiply.wgsl +19 -10
- package/src/shaders/f64/utils/sqrt.wgsl +45 -0
- package/src/shaders/index.mjs +69 -0
- package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
- package/src/snrm2/snrm2.mjs +20 -14
- package/src/srot/srot.mjs +10 -7
- package/src/srotm/srotm.mjs +9 -12
- package/src/sscal/sscal.mjs +7 -7
- package/src/sswap/sswap.mjs +14 -8
- package/src/ssymm/ssymm.d.mts +5 -4
- package/src/ssymm/ssymm.mjs +142 -55
- package/src/ssymv/ssymv.d.mts +2 -2
- package/src/ssymv/ssymv.mjs +42 -23
- package/src/ssyr/ssyr.d.mts +2 -2
- package/src/ssyr/ssyr.mjs +34 -15
- package/src/ssyr2/ssyr2.d.mts +2 -2
- package/src/ssyr2/ssyr2.mjs +43 -18
- package/src/ssyr2k/ssyr2k.d.mts +3 -2
- package/src/ssyr2k/ssyr2k.mjs +132 -55
- package/src/ssyrk/ssyrk.d.mts +3 -2
- package/src/ssyrk/ssyrk.mjs +84 -33
- package/src/strmm/strmm.d.mts +5 -4
- package/src/strmm/strmm.mjs +153 -54
- package/src/strmv/strmv.d.mts +2 -2
- package/src/strmv/strmv.mjs +37 -17
- package/src/strsm/strsm.d.mts +6 -4
- package/src/strsm/strsm.mjs +418 -172
- package/src/strsv/strsv.d.mts +5 -3
- package/src/strsv/strsv.mjs +82 -31
- package/src/util/benchmark.mjs +5 -3
- package/src/util/buffer.mjs +33 -12
- package/src/util/complex.mjs +87 -0
- package/src/util/compute.mjs +14 -8
- package/src/util/device.mjs +18 -3
- package/src/util/pipeline.mjs +40 -5
- package/src/util/workgroup.mjs +23 -6
- package/src/shaders/f64add.wgsl +0 -281
package/src/strmm/strmm.mjs
CHANGED
|
@@ -13,24 +13,40 @@ import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
|
|
|
13
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
14
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
15
15
|
import { requireWorkgroupCount } from "../util/workgroup.mjs";
|
|
16
|
-
import {
|
|
16
|
+
import {
|
|
17
|
+
BM_SMALL,
|
|
18
|
+
BN_SMALL,
|
|
19
|
+
BM_LARGE,
|
|
20
|
+
BN_LARGE,
|
|
21
|
+
LARGE_TILE_WORKGROUP_THRESHOLD,
|
|
22
|
+
} from "../util/constants.mjs";
|
|
17
23
|
import { TILE_WG_2D } from "../util/constants.mjs";
|
|
18
|
-
import { requireSameDevice } from "../util/device.mjs";
|
|
19
|
-
|
|
24
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
20
25
|
|
|
21
26
|
// strmm: B := alpha*op(A)*B (side='left') or alpha*B*op(A) (side='right'), A
|
|
22
27
|
// triangular. Triangularize then sgemm, one command encoder. B is both
|
|
23
28
|
// input and output, so gemm writes to a fresh buffer (no aliasing race),
|
|
24
29
|
// copied back into B (GpuMatrix) or read back directly (Float32Array).
|
|
25
30
|
export async function strmm(
|
|
26
|
-
device,
|
|
31
|
+
device,
|
|
32
|
+
side,
|
|
33
|
+
uplo,
|
|
34
|
+
transA,
|
|
35
|
+
diag,
|
|
36
|
+
m,
|
|
37
|
+
n,
|
|
38
|
+
alpha,
|
|
39
|
+
A,
|
|
40
|
+
lda,
|
|
41
|
+
B,
|
|
42
|
+
ldb,
|
|
43
|
+
layout = "row-major",
|
|
27
44
|
) {
|
|
28
45
|
const AIsGpu = A instanceof GpuMatrix;
|
|
29
46
|
const BIsGpu = B instanceof GpuMatrix;
|
|
30
47
|
const isUnit = diag === "unit";
|
|
31
48
|
|
|
32
|
-
|
|
33
|
-
throw new Error("device must be a GPUDevice.");
|
|
49
|
+
requireGpuDevice(device);
|
|
34
50
|
requireSameDevice(device, "strmm", { A, B });
|
|
35
51
|
if (side !== "left" && side !== "right")
|
|
36
52
|
throw new Error("side must be 'left' or 'right'.");
|
|
@@ -42,11 +58,15 @@ export async function strmm(
|
|
|
42
58
|
throw new Error("diag must be 'unit' or 'non-unit'.");
|
|
43
59
|
if (layout !== "row-major" && layout !== "column-major")
|
|
44
60
|
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.");
|
|
61
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
47
62
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
48
63
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
49
|
-
if (
|
|
64
|
+
if (
|
|
65
|
+
!Number.isInteger(m) ||
|
|
66
|
+
!Number.isInteger(n) ||
|
|
67
|
+
!Number.isInteger(lda) ||
|
|
68
|
+
!Number.isInteger(ldb)
|
|
69
|
+
)
|
|
50
70
|
throw new Error("m, n, lda, and ldb must be integers.");
|
|
51
71
|
if (!AIsGpu && !(A instanceof Float32Array))
|
|
52
72
|
throw new Error("A must be a Float32Array or GpuMatrix.");
|
|
@@ -62,37 +82,59 @@ export async function strmm(
|
|
|
62
82
|
|
|
63
83
|
// A: triangular, order = m (side='left') or n (side='right').
|
|
64
84
|
const aOrder = side === "left" ? m : n;
|
|
65
|
-
if (lda < aOrder)
|
|
85
|
+
if (lda < aOrder)
|
|
86
|
+
throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
|
|
66
87
|
if (AIsGpu) {
|
|
67
|
-
if (lda !== A.lda)
|
|
68
|
-
|
|
88
|
+
if (lda !== A.lda)
|
|
89
|
+
throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
90
|
+
if (A.rows < aOrder || A.cols < aOrder)
|
|
91
|
+
throw new Error("A is too small for the given m/n and side.");
|
|
69
92
|
} else if (A.length < (aOrder - 1) * lda + aOrder) {
|
|
70
|
-
throw new Error(
|
|
93
|
+
throw new Error(
|
|
94
|
+
"A does not have enough elements for the given dimensions and lda.",
|
|
95
|
+
);
|
|
71
96
|
}
|
|
72
97
|
|
|
73
98
|
// B: always m x n, overwritten in place with the same ldb.
|
|
74
99
|
const bOuter = effLayoutB === "column-major" ? n : m;
|
|
75
100
|
const bInner = effLayoutB === "column-major" ? m : n;
|
|
76
101
|
if (ldb < bInner)
|
|
77
|
-
throw new Error(
|
|
102
|
+
throw new Error(
|
|
103
|
+
`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
|
|
104
|
+
);
|
|
78
105
|
if (BIsGpu) {
|
|
79
|
-
if (ldb !== B.lda)
|
|
80
|
-
|
|
106
|
+
if (ldb !== B.lda)
|
|
107
|
+
throw new Error("ldb must match B.lda when B is a GpuMatrix.");
|
|
108
|
+
if (B.rows < m || B.cols < n)
|
|
109
|
+
throw new Error("B is too small for the given m and n.");
|
|
81
110
|
} else if (B.length < (bOuter - 1) * ldb + bInner) {
|
|
82
|
-
throw new Error(
|
|
111
|
+
throw new Error(
|
|
112
|
+
"B does not have enough elements for the given dimensions and ldb.",
|
|
113
|
+
);
|
|
83
114
|
}
|
|
84
115
|
|
|
85
116
|
// A isn't symmetric: column-major = genuine transpose, so flip transA;
|
|
86
117
|
// transposing also swaps which triangle looks stored, so flip uplo too.
|
|
87
|
-
const uploEffA =
|
|
88
|
-
|
|
118
|
+
const uploEffA =
|
|
119
|
+
effLayoutA === "column-major"
|
|
120
|
+
? uplo === "lower"
|
|
121
|
+
? "upper"
|
|
122
|
+
: "lower"
|
|
123
|
+
: uplo;
|
|
124
|
+
const transEffA =
|
|
125
|
+
effLayoutA === "column-major"
|
|
126
|
+
? transA === "no-transpose"
|
|
127
|
+
? "transpose"
|
|
128
|
+
: "no-transpose"
|
|
129
|
+
: transA;
|
|
89
130
|
|
|
90
131
|
const transB = effLayoutB === "column-major" ? "transpose" : "no-transpose";
|
|
91
132
|
const transDense = "no-transpose"; // Adense already embodies op(A)
|
|
92
133
|
|
|
93
134
|
// X*Y (X=Adense,Y=B for side='left', swapped for 'right'). Column-major
|
|
94
135
|
// output: compute (B_out)^T instead — sgemm's own trick, same as ssymm's.
|
|
95
|
-
let mg = m,
|
|
136
|
+
let mg = m,
|
|
137
|
+
ng = n;
|
|
96
138
|
const kg = aOrder;
|
|
97
139
|
let transX = side === "left" ? transDense : transB;
|
|
98
140
|
let transY = side === "left" ? transB : transDense;
|
|
@@ -108,39 +150,63 @@ export async function strmm(
|
|
|
108
150
|
const largeWgX = Math.ceil(ng / BN_LARGE);
|
|
109
151
|
const largeWgY = Math.ceil(mg / BM_LARGE);
|
|
110
152
|
const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
111
|
-
const gemmPipeline = await getPipeline(
|
|
153
|
+
const gemmPipeline = await getPipeline(
|
|
154
|
+
device,
|
|
155
|
+
useLargeTile ? "sgemm_large" : "sgemm_small",
|
|
156
|
+
);
|
|
112
157
|
const triPipeline = await getPipeline(device, "triangularize");
|
|
113
158
|
const gemmWgCount = useLargeTile
|
|
114
159
|
? {
|
|
115
|
-
|
|
116
|
-
|
|
117
|
-
|
|
160
|
+
x: requireWorkgroupCount(device, largeWgX, "strmm", "x"),
|
|
161
|
+
y: requireWorkgroupCount(device, largeWgY, "strmm", "y"),
|
|
162
|
+
}
|
|
118
163
|
: {
|
|
119
|
-
|
|
120
|
-
|
|
121
|
-
|
|
164
|
+
x: requireWorkgroupCount(
|
|
165
|
+
device,
|
|
166
|
+
Math.ceil(ng / BN_SMALL),
|
|
167
|
+
"strmm",
|
|
168
|
+
"x",
|
|
169
|
+
),
|
|
170
|
+
y: requireWorkgroupCount(
|
|
171
|
+
device,
|
|
172
|
+
Math.ceil(mg / BM_SMALL),
|
|
173
|
+
"strmm",
|
|
174
|
+
"y",
|
|
175
|
+
),
|
|
176
|
+
};
|
|
122
177
|
|
|
123
178
|
// Null-init here and allocate inside the try below, so a throw partway
|
|
124
179
|
// through the sequence still reaches finally with every handle visible
|
|
125
180
|
// (strsv.mjs is the reference for this pattern).
|
|
126
|
-
let ABuffer = null,
|
|
127
|
-
|
|
128
|
-
let
|
|
181
|
+
let ABuffer = null,
|
|
182
|
+
BBuffer = null;
|
|
183
|
+
let AdenseBuffer = null,
|
|
184
|
+
outBuffer = null;
|
|
185
|
+
let triParams = null,
|
|
186
|
+
gemmParams = null;
|
|
129
187
|
let outBufferAdopted = false; // true once B._buf is repointed at outBuffer
|
|
130
188
|
|
|
131
189
|
try {
|
|
132
190
|
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmm-A", false);
|
|
133
191
|
// readback=true (COPY_SRC): BBuffer is also the source that seeds outBuffer.
|
|
134
192
|
BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strmm-B", true);
|
|
135
|
-
AdenseBuffer = createStorageBuffer(
|
|
193
|
+
AdenseBuffer = createStorageBuffer(
|
|
194
|
+
device,
|
|
195
|
+
aOrder * ldDense * 4,
|
|
196
|
+
"strmm-Adense",
|
|
197
|
+
);
|
|
136
198
|
// COPY_DST: seeded from B's own content before gemm runs, so stride-padding
|
|
137
199
|
// gaps (never written by gemm's tight m x n loop) keep B's original bytes
|
|
138
200
|
// instead of reading back as zero. COPY_SRC: read back / adopted by B after.
|
|
139
|
-
outBuffer = createStorageBuffer(
|
|
140
|
-
|
|
201
|
+
outBuffer = createStorageBuffer(
|
|
202
|
+
device,
|
|
203
|
+
bOuter * ldb * 4,
|
|
204
|
+
"strmm-out",
|
|
205
|
+
GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST,
|
|
141
206
|
);
|
|
142
207
|
|
|
143
|
-
triParams = createParamsBuffer(
|
|
208
|
+
triParams = createParamsBuffer(
|
|
209
|
+
device,
|
|
144
210
|
[
|
|
145
211
|
{ value: aOrder, type: "u32" },
|
|
146
212
|
{ value: lda, type: "u32" },
|
|
@@ -151,7 +217,11 @@ export async function strmm(
|
|
|
151
217
|
],
|
|
152
218
|
"strmm-tri-params",
|
|
153
219
|
);
|
|
154
|
-
const triBindGroup = createBindGroup(
|
|
220
|
+
const triBindGroup = createBindGroup(
|
|
221
|
+
device,
|
|
222
|
+
triPipeline.getBindGroupLayout(0),
|
|
223
|
+
[ABuffer, AdenseBuffer, triParams],
|
|
224
|
+
);
|
|
155
225
|
|
|
156
226
|
// X/Y buffers and their own ld, matching swapXY above.
|
|
157
227
|
const XBuffer = swapXY ? BBuffer : AdenseBuffer;
|
|
@@ -159,13 +229,14 @@ export async function strmm(
|
|
|
159
229
|
const YBuffer = swapXY ? AdenseBuffer : BBuffer;
|
|
160
230
|
const ldY = swapXY ? ldDense : ldb;
|
|
161
231
|
|
|
162
|
-
gemmParams = createParamsBuffer(
|
|
232
|
+
gemmParams = createParamsBuffer(
|
|
233
|
+
device,
|
|
163
234
|
[
|
|
164
|
-
{ value: mg,
|
|
165
|
-
{ value: ng,
|
|
166
|
-
{ value: kg,
|
|
235
|
+
{ value: mg, type: "u32" },
|
|
236
|
+
{ value: ng, type: "u32" },
|
|
237
|
+
{ value: kg, type: "u32" },
|
|
167
238
|
{ value: alpha, type: "f32" },
|
|
168
|
-
{ value: 0.0,
|
|
239
|
+
{ value: 0.0, type: "f32" }, // beta — strmm has no C accumulation term
|
|
169
240
|
{ value: ldX, type: "u32" },
|
|
170
241
|
{ value: ldY, type: "u32" },
|
|
171
242
|
{ value: ldb, type: "u32" },
|
|
@@ -174,28 +245,56 @@ export async function strmm(
|
|
|
174
245
|
],
|
|
175
246
|
"strmm-gemm-params",
|
|
176
247
|
);
|
|
177
|
-
const gemmBindGroup = createBindGroup(
|
|
178
|
-
|
|
179
|
-
|
|
180
|
-
|
|
181
|
-
|
|
182
|
-
|
|
183
|
-
|
|
184
|
-
|
|
248
|
+
const gemmBindGroup = createBindGroup(
|
|
249
|
+
device,
|
|
250
|
+
gemmPipeline.getBindGroupLayout(0),
|
|
251
|
+
[
|
|
252
|
+
XBuffer,
|
|
253
|
+
vec4ViewBinding(device, XBuffer),
|
|
254
|
+
YBuffer,
|
|
255
|
+
vec4ViewBinding(device, YBuffer),
|
|
256
|
+
outBuffer,
|
|
257
|
+
gemmParams,
|
|
258
|
+
],
|
|
259
|
+
);
|
|
185
260
|
|
|
186
261
|
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
187
262
|
// Seed outBuffer with B's own bytes first, so gemm's tight m x n write
|
|
188
263
|
// leaves stride-padding gaps holding B's original content, not zero.
|
|
189
264
|
// BBuffer may be larger than outBuffer (e.g. a validation-test baseline
|
|
190
265
|
// over-provisioned for a bigger ldb it might later be substituted with).
|
|
191
|
-
commandEncoder.copyBufferToBuffer(
|
|
192
|
-
|
|
193
|
-
|
|
194
|
-
|
|
195
|
-
|
|
266
|
+
commandEncoder.copyBufferToBuffer(
|
|
267
|
+
BBuffer,
|
|
268
|
+
0,
|
|
269
|
+
outBuffer,
|
|
270
|
+
0,
|
|
271
|
+
Math.min(BBuffer.size, outBuffer.size),
|
|
272
|
+
);
|
|
273
|
+
const triDesc = querySet
|
|
274
|
+
? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
|
|
275
|
+
: undefined;
|
|
276
|
+
const gemmDesc = querySet
|
|
277
|
+
? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
|
|
278
|
+
: undefined;
|
|
279
|
+
encodePass(
|
|
280
|
+
commandEncoder,
|
|
281
|
+
triPipeline,
|
|
282
|
+
triBindGroup,
|
|
283
|
+
{ x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) },
|
|
284
|
+
triDesc,
|
|
285
|
+
);
|
|
286
|
+
encodePass(
|
|
287
|
+
commandEncoder,
|
|
288
|
+
gemmPipeline,
|
|
289
|
+
gemmBindGroup,
|
|
290
|
+
gemmWgCount,
|
|
291
|
+
gemmDesc,
|
|
292
|
+
);
|
|
196
293
|
|
|
197
294
|
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
198
|
-
const readBuffer = BIsGpu
|
|
295
|
+
const readBuffer = BIsGpu
|
|
296
|
+
? null
|
|
297
|
+
: stageReadback(device, commandEncoder, outBuffer);
|
|
199
298
|
|
|
200
299
|
submit(device, commandEncoder);
|
|
201
300
|
|
package/src/strmv/strmv.d.mts
CHANGED
|
@@ -2,7 +2,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
|
|
|
2
2
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
3
3
|
|
|
4
4
|
/**
|
|
5
|
-
* Performs the triangular matrix-vector operation y
|
|
5
|
+
* Performs the triangular matrix-vector operation $$y \leftarrow \mathrm{op}(A) x$$
|
|
6
6
|
*
|
|
7
7
|
* A is an n×n triangular matrix stored in row-major order. Only the triangle
|
|
8
8
|
* specified by `uplo` is referenced; the other triangle is not accessed.
|
|
@@ -45,7 +45,7 @@ export declare function strmv(
|
|
|
45
45
|
): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
|
|
46
46
|
|
|
47
47
|
/**
|
|
48
|
-
* Performs the triangular matrix-vector operation y
|
|
48
|
+
* Performs the triangular matrix-vector operation $$y \leftarrow \mathrm{op}(A) x$$
|
|
49
49
|
*
|
|
50
50
|
* x and y are kept resident on the GPU. A must be a GpuMatrix; its own
|
|
51
51
|
* `layout` (set at `GpuMatrix.from` time) determines the operation — there is
|
package/src/strmv/strmv.mjs
CHANGED
|
@@ -11,16 +11,28 @@ 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
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
15
15
|
|
|
16
|
-
export async function strmv(
|
|
16
|
+
export async function strmv(
|
|
17
|
+
device,
|
|
18
|
+
uplo,
|
|
19
|
+
trans,
|
|
20
|
+
diag,
|
|
21
|
+
n,
|
|
22
|
+
A,
|
|
23
|
+
lda,
|
|
24
|
+
x,
|
|
25
|
+
incx,
|
|
26
|
+
y,
|
|
27
|
+
incy,
|
|
28
|
+
layout = "row-major",
|
|
29
|
+
) {
|
|
17
30
|
const xIsGpu = x instanceof GpuVector;
|
|
18
31
|
const yIsGpu = y instanceof GpuVector;
|
|
19
32
|
const AIsGpu = A instanceof GpuMatrix;
|
|
20
33
|
const isUnit = diag === "unit";
|
|
21
34
|
|
|
22
|
-
|
|
23
|
-
throw new Error("device must be a GPUDevice.");
|
|
35
|
+
requireGpuDevice(device);
|
|
24
36
|
requireSameDevice(device, "strmv", { A, x, y });
|
|
25
37
|
if (uplo !== "lower" && uplo !== "upper")
|
|
26
38
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
@@ -68,9 +80,7 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
68
80
|
if (n === 0) return yIsGpu ? {} : { y };
|
|
69
81
|
|
|
70
82
|
if (!AIsGpu && A.length < (n - 1) * lda + n)
|
|
71
|
-
throw new Error(
|
|
72
|
-
"A does not have enough elements for the given n and lda.",
|
|
73
|
-
);
|
|
83
|
+
throw new Error("A does not have enough elements for the given n and lda.");
|
|
74
84
|
if (x.length < (n - 1) * incx + 1)
|
|
75
85
|
throw new Error(
|
|
76
86
|
"x does not have enough elements for the given n and incx.",
|
|
@@ -84,7 +94,9 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
84
94
|
const effLayout = AIsGpu ? A.layout : layout;
|
|
85
95
|
const isColMajor = effLayout === "column-major";
|
|
86
96
|
const isLower = isColMajor ? uplo === "upper" : uplo === "lower";
|
|
87
|
-
const isNoTrans = isColMajor
|
|
97
|
+
const isNoTrans = isColMajor
|
|
98
|
+
? trans === "transpose"
|
|
99
|
+
: trans === "no-transpose";
|
|
88
100
|
|
|
89
101
|
const pipeline = await getPipeline(device, "strmv");
|
|
90
102
|
|
|
@@ -97,15 +109,16 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
97
109
|
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmv-A", false);
|
|
98
110
|
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "strmv-x", false);
|
|
99
111
|
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "strmv-y", true);
|
|
100
|
-
paramsBuffer = createParamsBuffer(
|
|
112
|
+
paramsBuffer = createParamsBuffer(
|
|
113
|
+
device,
|
|
101
114
|
[
|
|
102
|
-
{ value: n,
|
|
103
|
-
{ value: incx,
|
|
104
|
-
{ value: incy,
|
|
105
|
-
{ value: lda,
|
|
115
|
+
{ value: n, type: "u32" },
|
|
116
|
+
{ value: incx, type: "u32" },
|
|
117
|
+
{ value: incy, type: "u32" },
|
|
118
|
+
{ value: lda, type: "u32" },
|
|
106
119
|
{ value: isNoTrans ? 0 : 1, type: "u32" },
|
|
107
|
-
{ value: isLower ? 0 : 1,
|
|
108
|
-
{ value: isUnit ? 1 : 0,
|
|
120
|
+
{ value: isLower ? 0 : 1, type: "u32" },
|
|
121
|
+
{ value: isUnit ? 1 : 0, type: "u32" },
|
|
109
122
|
],
|
|
110
123
|
"strmv-params",
|
|
111
124
|
);
|
|
@@ -118,8 +131,15 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
118
131
|
]);
|
|
119
132
|
|
|
120
133
|
const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
|
|
121
|
-
const { commandEncoder, ts } = runComputePass(
|
|
122
|
-
|
|
134
|
+
const { commandEncoder, ts } = runComputePass(
|
|
135
|
+
device,
|
|
136
|
+
pipeline,
|
|
137
|
+
bindGroup,
|
|
138
|
+
wgCount,
|
|
139
|
+
);
|
|
140
|
+
const readBuffer = yIsGpu
|
|
141
|
+
? null
|
|
142
|
+
: stageReadback(device, commandEncoder, yBuffer);
|
|
123
143
|
|
|
124
144
|
submit(device, commandEncoder);
|
|
125
145
|
|
package/src/strsm/strsm.d.mts
CHANGED
|
@@ -2,8 +2,9 @@ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
4
|
* Solves the triangular matrix equation
|
|
5
|
-
* op(A)
|
|
6
|
-
* X
|
|
5
|
+
* $$\mathrm{op}(A) X = \alpha B \quad (\texttt{side='left'})$$
|
|
6
|
+
* $$X \mathrm{op}(A) = \alpha B \quad (\texttt{side='right'})$$
|
|
7
|
+
* overwriting `B` with the
|
|
7
8
|
* solution `X` — `A` is triangular, only its `uplo` triangle stored; `B` is
|
|
8
9
|
* a general m×n matrix.
|
|
9
10
|
*
|
|
@@ -57,8 +58,9 @@ export declare function strsm(
|
|
|
57
58
|
|
|
58
59
|
/**
|
|
59
60
|
* Solves the triangular matrix equation
|
|
60
|
-
* op(A)
|
|
61
|
-
* X
|
|
61
|
+
* $$\mathrm{op}(A) X = \alpha B \quad (\texttt{side='left'})$$
|
|
62
|
+
* $$X \mathrm{op}(A) = \alpha B \quad (\texttt{side='right'})$$
|
|
63
|
+
* overwriting `B` in place with `X`.
|
|
62
64
|
*
|
|
63
65
|
* A and B are both kept GPU-resident. Each matrix's own `layout` (set at
|
|
64
66
|
* `GpuMatrix.from` time) determines the operation — there is no separate
|