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