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/strsm/strsm.mjs
CHANGED
|
@@ -4,33 +4,54 @@ 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";
|
|
10
11
|
import { extractResult } from "../util/result.mjs";
|
|
11
12
|
import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
|
|
12
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
13
|
-
import {
|
|
14
|
+
import {
|
|
15
|
+
calcWorkgroups,
|
|
16
|
+
requireWorkgroups,
|
|
17
|
+
requireWorkgroupCount,
|
|
18
|
+
} from "../util/workgroup.mjs";
|
|
14
19
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
|
|
18
|
-
|
|
19
|
-
|
|
20
|
+
import {
|
|
21
|
+
BM_SMALL,
|
|
22
|
+
BN_SMALL,
|
|
23
|
+
BM_LARGE,
|
|
24
|
+
BN_LARGE,
|
|
25
|
+
LARGE_TILE_WORKGROUP_THRESHOLD,
|
|
26
|
+
} from "../util/constants.mjs";
|
|
27
|
+
import { BLOCK_SIZE } from "../util/constants.mjs";
|
|
28
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
20
29
|
|
|
21
30
|
// strsm: B := alpha*op(A)^-1*B (side='left') or alpha*B*op(A)^-1 (side='right'),
|
|
22
31
|
// A triangular. Blocked substitution (strsv's own technique, generalized to
|
|
23
32
|
// a matrix RHS): strsv_invert_block + sgemm, unchanged; every per-block B/A
|
|
24
33
|
// access goes through block_transfer.wgsl (see that shader for why).
|
|
25
34
|
export async function strsm(
|
|
26
|
-
device,
|
|
35
|
+
device,
|
|
36
|
+
side,
|
|
37
|
+
uplo,
|
|
38
|
+
transA,
|
|
39
|
+
diag,
|
|
40
|
+
m,
|
|
41
|
+
n,
|
|
42
|
+
alpha,
|
|
43
|
+
A,
|
|
44
|
+
lda,
|
|
45
|
+
B,
|
|
46
|
+
ldb,
|
|
47
|
+
layout = "row-major",
|
|
27
48
|
) {
|
|
28
49
|
const AIsGpu = A instanceof GpuMatrix;
|
|
29
50
|
const BIsGpu = B instanceof GpuMatrix;
|
|
30
51
|
const isUnit = diag === "unit";
|
|
31
52
|
|
|
32
|
-
|
|
33
|
-
|
|
53
|
+
requireGpuDevice(device);
|
|
54
|
+
requireSameDevice(device, "strsm", { A, B });
|
|
34
55
|
if (side !== "left" && side !== "right")
|
|
35
56
|
throw new Error("side must be 'left' or 'right'.");
|
|
36
57
|
if (uplo !== "lower" && uplo !== "upper")
|
|
@@ -41,11 +62,15 @@ export async function strsm(
|
|
|
41
62
|
throw new Error("diag must be 'unit' or 'non-unit'.");
|
|
42
63
|
if (layout !== "row-major" && layout !== "column-major")
|
|
43
64
|
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
44
|
-
if (typeof alpha !== "number")
|
|
45
|
-
throw new Error("alpha must be a number.");
|
|
65
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
46
66
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
47
67
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
48
|
-
if (
|
|
68
|
+
if (
|
|
69
|
+
!Number.isInteger(m) ||
|
|
70
|
+
!Number.isInteger(n) ||
|
|
71
|
+
!Number.isInteger(lda) ||
|
|
72
|
+
!Number.isInteger(ldb)
|
|
73
|
+
)
|
|
49
74
|
throw new Error("m, n, lda, and ldb must be integers.");
|
|
50
75
|
if (!AIsGpu && !(A instanceof Float32Array))
|
|
51
76
|
throw new Error("A must be a Float32Array or GpuMatrix.");
|
|
@@ -61,30 +86,51 @@ export async function strsm(
|
|
|
61
86
|
|
|
62
87
|
// A: triangular, order = m (side='left') or n (side='right').
|
|
63
88
|
const aOrder = side === "left" ? m : n;
|
|
64
|
-
if (lda < aOrder)
|
|
89
|
+
if (lda < aOrder)
|
|
90
|
+
throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
|
|
65
91
|
if (AIsGpu) {
|
|
66
|
-
if (lda !== A.lda)
|
|
67
|
-
|
|
92
|
+
if (lda !== A.lda)
|
|
93
|
+
throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
94
|
+
if (A.rows < aOrder || A.cols < aOrder)
|
|
95
|
+
throw new Error("A is too small for the given m/n and side.");
|
|
68
96
|
} else if (A.length < (aOrder - 1) * lda + aOrder) {
|
|
69
|
-
throw new Error(
|
|
97
|
+
throw new Error(
|
|
98
|
+
"A does not have enough elements for the given dimensions and lda.",
|
|
99
|
+
);
|
|
70
100
|
}
|
|
71
101
|
|
|
72
102
|
// B: always m x n, overwritten in place with the same ldb.
|
|
73
103
|
const bOuter = effLayoutB === "column-major" ? n : m;
|
|
74
104
|
const bInner = effLayoutB === "column-major" ? m : n;
|
|
75
105
|
if (ldb < bInner)
|
|
76
|
-
throw new Error(
|
|
106
|
+
throw new Error(
|
|
107
|
+
`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
|
|
108
|
+
);
|
|
77
109
|
if (BIsGpu) {
|
|
78
|
-
if (ldb !== B.lda)
|
|
79
|
-
|
|
110
|
+
if (ldb !== B.lda)
|
|
111
|
+
throw new Error("ldb must match B.lda when B is a GpuMatrix.");
|
|
112
|
+
if (B.rows < m || B.cols < n)
|
|
113
|
+
throw new Error("B is too small for the given m and n.");
|
|
80
114
|
} else if (B.length < (bOuter - 1) * ldb + bInner) {
|
|
81
|
-
throw new Error(
|
|
115
|
+
throw new Error(
|
|
116
|
+
"B does not have enough elements for the given dimensions and ldb.",
|
|
117
|
+
);
|
|
82
118
|
}
|
|
83
119
|
|
|
84
120
|
// A isn't symmetric: column-major = genuine transpose, so flip transA;
|
|
85
121
|
// transposing also swaps which triangle looks stored, so flip uplo too.
|
|
86
|
-
const uploEffA =
|
|
87
|
-
|
|
122
|
+
const uploEffA =
|
|
123
|
+
effLayoutA === "column-major"
|
|
124
|
+
? uplo === "lower"
|
|
125
|
+
? "upper"
|
|
126
|
+
: "lower"
|
|
127
|
+
: uplo;
|
|
128
|
+
const transEffA =
|
|
129
|
+
effLayoutA === "column-major"
|
|
130
|
+
? transA === "no-transpose"
|
|
131
|
+
? "transpose"
|
|
132
|
+
: "no-transpose"
|
|
133
|
+
: transA;
|
|
88
134
|
|
|
89
135
|
const otherLen = side === "left" ? n : m;
|
|
90
136
|
const blockIsRow = side === "left";
|
|
@@ -103,19 +149,22 @@ export async function strsm(
|
|
|
103
149
|
const transferPipeline = await getPipeline(device, "block_transfer");
|
|
104
150
|
const scalarPipeline = await getPipeline(device, "sscal");
|
|
105
151
|
|
|
106
|
-
|
|
107
|
-
|
|
108
|
-
|
|
152
|
+
// Null-init here and allocate inside the try below, so a throw partway
|
|
153
|
+
// through the sequence still reaches finally with every handle visible
|
|
154
|
+
// (strsv.mjs is the reference for this pattern).
|
|
155
|
+
let ABuffer = null;
|
|
156
|
+
let BBuffer = null;
|
|
157
|
+
let AinvBuffer = null;
|
|
109
158
|
|
|
110
159
|
const paramsBuffers = [];
|
|
111
160
|
const scratchBuffers = [];
|
|
112
161
|
function scratch(size, label) {
|
|
113
|
-
const buf = createStorageBuffer(size, label);
|
|
162
|
+
const buf = createStorageBuffer(device, size, label);
|
|
114
163
|
scratchBuffers.push(buf);
|
|
115
164
|
return buf;
|
|
116
165
|
}
|
|
117
166
|
function params(entries, label) {
|
|
118
|
-
const buf = createParamsBuffer(entries, label);
|
|
167
|
+
const buf = createParamsBuffer(device, entries, label);
|
|
119
168
|
paramsBuffers.push(buf);
|
|
120
169
|
return buf;
|
|
121
170
|
}
|
|
@@ -126,9 +175,19 @@ export async function strsm(
|
|
|
126
175
|
const bScaleLen = (bOuter - 1) * ldb + bInner;
|
|
127
176
|
|
|
128
177
|
try {
|
|
178
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strsm-A", false);
|
|
179
|
+
BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strsm-B", true);
|
|
180
|
+
AinvBuffer = createStorageBuffer(
|
|
181
|
+
device,
|
|
182
|
+
numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4,
|
|
183
|
+
"strsm-Ainv",
|
|
184
|
+
);
|
|
185
|
+
|
|
129
186
|
// Pre-scale B by alpha once (reuses sscal, so no per-block alpha handling).
|
|
187
|
+
// Skipped for alpha=0 — sscal computes 0*B[i], which leaks NaN/Inf from a
|
|
188
|
+
// poisoned B; see the literal zero-write dispatched below instead.
|
|
130
189
|
let preScaleBindGroup = null;
|
|
131
|
-
if (alpha !== 1.0) {
|
|
190
|
+
if (alpha !== 1.0 && alpha !== 0) {
|
|
132
191
|
const scaleParams = params(
|
|
133
192
|
[
|
|
134
193
|
{ value: bScaleLen, type: "u32" },
|
|
@@ -137,7 +196,11 @@ export async function strsm(
|
|
|
137
196
|
],
|
|
138
197
|
"strsm-scale-params",
|
|
139
198
|
);
|
|
140
|
-
preScaleBindGroup = createBindGroup(
|
|
199
|
+
preScaleBindGroup = createBindGroup(
|
|
200
|
+
device,
|
|
201
|
+
scalarPipeline.getBindGroupLayout(0),
|
|
202
|
+
[BBuffer, scaleParams],
|
|
203
|
+
);
|
|
141
204
|
}
|
|
142
205
|
|
|
143
206
|
// Every diagonal block's inverse, fully parallel, one dispatch, unchanged.
|
|
@@ -151,7 +214,11 @@ export async function strsm(
|
|
|
151
214
|
],
|
|
152
215
|
"strsm-invert-params",
|
|
153
216
|
);
|
|
154
|
-
const invertBindGroup = createBindGroup(
|
|
217
|
+
const invertBindGroup = createBindGroup(
|
|
218
|
+
device,
|
|
219
|
+
invertPipeline.getBindGroupLayout(0),
|
|
220
|
+
[ABuffer, AinvBuffer, invertParams],
|
|
221
|
+
);
|
|
155
222
|
|
|
156
223
|
// Reusable scratch buffers, sized for the worst case, bound at offset 0.
|
|
157
224
|
const Bblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Bblock");
|
|
@@ -159,168 +226,365 @@ export async function strsm(
|
|
|
159
226
|
const Aoff = scratch(aOrder * BLOCK_SIZE * 4, "strsm-Aoff");
|
|
160
227
|
const delta = scratch(aOrder * otherLen * 4, "strsm-delta");
|
|
161
228
|
|
|
162
|
-
const { commandEncoder, querySet } = beginTimedEncoder();
|
|
229
|
+
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
163
230
|
|
|
164
231
|
if (alpha === 0) {
|
|
165
|
-
// BLAS: alpha=0 means
|
|
166
|
-
|
|
167
|
-
|
|
232
|
+
// BLAS: alpha=0 means B must become a literal zero, not 0*B (leaks
|
|
233
|
+
// NaN/Inf via sscal). Reuse sgemm with k=0 (X/Y never read, so
|
|
234
|
+
// AinvBuffer is a safe dummy for both) and beta=0.
|
|
235
|
+
const zeroLargeWgX = Math.ceil(bInner / BN_LARGE);
|
|
236
|
+
const zeroLargeWgY = Math.ceil(bOuter / BM_LARGE);
|
|
237
|
+
const zeroUseLarge =
|
|
238
|
+
zeroLargeWgX * zeroLargeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
239
|
+
const zeroPipeline = await getPipeline(
|
|
240
|
+
device,
|
|
241
|
+
zeroUseLarge ? "sgemm_large" : "sgemm_small",
|
|
242
|
+
);
|
|
243
|
+
const zeroParams = params(
|
|
244
|
+
[
|
|
245
|
+
{ value: bOuter, type: "u32" },
|
|
246
|
+
{ value: bInner, type: "u32" },
|
|
247
|
+
{ value: 0, type: "u32" }, // k forced to 0 — X/Y never read
|
|
248
|
+
{ value: 0.0, type: "f32" }, // alpha
|
|
249
|
+
{ value: 0.0, type: "f32" }, // beta
|
|
250
|
+
{ value: 1, type: "u32" }, // ldX (dummy, unread)
|
|
251
|
+
{ value: 1, type: "u32" }, // ldY (dummy, unread)
|
|
252
|
+
{ value: ldb, type: "u32" }, // ldc — B's own real stride
|
|
253
|
+
{ value: 0, type: "u32" }, // transX (dummy)
|
|
254
|
+
{ value: 0, type: "u32" }, // transY (dummy)
|
|
255
|
+
],
|
|
256
|
+
"strsm-zero-params",
|
|
257
|
+
);
|
|
258
|
+
const zeroBindGroup = createBindGroup(
|
|
259
|
+
device,
|
|
260
|
+
zeroPipeline.getBindGroupLayout(0),
|
|
261
|
+
[
|
|
262
|
+
AinvBuffer,
|
|
263
|
+
vec4ViewBinding(device, AinvBuffer),
|
|
264
|
+
AinvBuffer,
|
|
265
|
+
vec4ViewBinding(device, AinvBuffer),
|
|
266
|
+
BBuffer,
|
|
267
|
+
zeroParams,
|
|
268
|
+
],
|
|
269
|
+
);
|
|
270
|
+
const zeroWgCount = zeroUseLarge
|
|
271
|
+
? {
|
|
272
|
+
x: requireWorkgroupCount(device, zeroLargeWgX, "strsm", "x"),
|
|
273
|
+
y: requireWorkgroupCount(device, zeroLargeWgY, "strsm", "y"),
|
|
274
|
+
}
|
|
275
|
+
: {
|
|
276
|
+
x: requireWorkgroupCount(
|
|
277
|
+
device,
|
|
278
|
+
Math.ceil(bInner / BN_SMALL),
|
|
279
|
+
"strsm",
|
|
280
|
+
"x",
|
|
281
|
+
),
|
|
282
|
+
y: requireWorkgroupCount(
|
|
283
|
+
device,
|
|
284
|
+
Math.ceil(bOuter / BM_SMALL),
|
|
285
|
+
"strsm",
|
|
286
|
+
"y",
|
|
287
|
+
),
|
|
288
|
+
};
|
|
289
|
+
const zeroDesc = querySet
|
|
290
|
+
? {
|
|
291
|
+
timestampWrites: {
|
|
292
|
+
querySet,
|
|
293
|
+
beginningOfPassWriteIndex: 0,
|
|
294
|
+
endOfPassWriteIndex: 1,
|
|
295
|
+
},
|
|
296
|
+
}
|
|
297
|
+
: undefined;
|
|
298
|
+
encodePass(
|
|
299
|
+
commandEncoder,
|
|
300
|
+
zeroPipeline,
|
|
301
|
+
zeroBindGroup,
|
|
302
|
+
zeroWgCount,
|
|
303
|
+
zeroDesc,
|
|
304
|
+
);
|
|
168
305
|
} else {
|
|
169
306
|
if (preScaleBindGroup) {
|
|
170
|
-
encodePass(
|
|
307
|
+
encodePass(
|
|
308
|
+
commandEncoder,
|
|
309
|
+
scalarPipeline,
|
|
310
|
+
preScaleBindGroup,
|
|
311
|
+
calcWorkgroups(device, bScaleLen),
|
|
312
|
+
);
|
|
171
313
|
}
|
|
172
|
-
const invertDesc = querySet
|
|
173
|
-
|
|
314
|
+
const invertDesc = querySet
|
|
315
|
+
? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
|
|
316
|
+
: undefined;
|
|
317
|
+
encodePass(
|
|
318
|
+
commandEncoder,
|
|
319
|
+
invertPipeline,
|
|
320
|
+
invertBindGroup,
|
|
321
|
+
{ x: BLOCK_SIZE, y: numBlocks },
|
|
322
|
+
invertDesc,
|
|
323
|
+
);
|
|
174
324
|
|
|
175
325
|
for (let bi = 0; bi < blockStarts.length; bi++) {
|
|
176
|
-
|
|
177
|
-
|
|
178
|
-
|
|
179
|
-
|
|
180
|
-
|
|
181
|
-
|
|
182
|
-
|
|
183
|
-
|
|
184
|
-
[
|
|
185
|
-
{ value: blockStart, type: "u32" },
|
|
186
|
-
{ value: blockLen, type: "u32" },
|
|
187
|
-
{ value: 0, type: "u32" },
|
|
188
|
-
{ value: otherLen, type: "u32" },
|
|
189
|
-
{ value: ldb, type: "u32" },
|
|
190
|
-
{ value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
|
|
191
|
-
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
192
|
-
{ value: 2, type: "u32" }, // gather
|
|
193
|
-
],
|
|
194
|
-
"strsm-gather-B-params",
|
|
195
|
-
);
|
|
196
|
-
const gatherBBindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [Bblock, BBuffer, gatherBParams]);
|
|
197
|
-
encodePass(commandEncoder, transferPipeline, gatherBBindGroup, calcWorkgroups(blockLen, otherLen));
|
|
198
|
-
|
|
199
|
-
// 2) apply: Xblock := op(Ainv_block) @ Bblock (side='left'), or the
|
|
200
|
-
// transpose-trick equivalent for side='right' (same trick strmm uses).
|
|
201
|
-
{
|
|
202
|
-
const mg = blockLen, ng = otherLen, kg = blockLen;
|
|
203
|
-
const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
|
|
204
|
-
const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
205
|
-
const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
|
|
206
|
-
const applyParams = params(
|
|
326
|
+
const blockStart = blockStarts[bi];
|
|
327
|
+
const blockEnd = Math.min(blockStart + BLOCK_SIZE, aOrder);
|
|
328
|
+
const blockLen = blockEnd - blockStart;
|
|
329
|
+
const blockIndex = blockStart / BLOCK_SIZE;
|
|
330
|
+
const isLastPass = bi === blockStarts.length - 1;
|
|
331
|
+
|
|
332
|
+
// 1) gather B's current block into a tight scratch buffer.
|
|
333
|
+
const gatherBParams = params(
|
|
207
334
|
[
|
|
208
|
-
{ value:
|
|
209
|
-
{ value:
|
|
210
|
-
{ value:
|
|
211
|
-
{ value:
|
|
212
|
-
{ value:
|
|
213
|
-
{ value:
|
|
214
|
-
{ value:
|
|
215
|
-
{ value:
|
|
216
|
-
{ value: side === "right" ? 1 : 0, type: "u32" }, // transX: side='right' needs Ainv^T
|
|
217
|
-
{ value: 0, type: "u32" }, // transY: Bblock is always read as-is
|
|
335
|
+
{ value: blockStart, type: "u32" },
|
|
336
|
+
{ value: blockLen, type: "u32" },
|
|
337
|
+
{ value: 0, type: "u32" },
|
|
338
|
+
{ value: otherLen, type: "u32" },
|
|
339
|
+
{ value: ldb, type: "u32" },
|
|
340
|
+
{ value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
|
|
341
|
+
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
342
|
+
{ value: 2, type: "u32" }, // gather
|
|
218
343
|
],
|
|
219
|
-
"strsm-
|
|
344
|
+
"strsm-gather-B-params",
|
|
345
|
+
);
|
|
346
|
+
const gatherBBindGroup = createBindGroup(
|
|
347
|
+
device,
|
|
348
|
+
transferPipeline.getBindGroupLayout(0),
|
|
349
|
+
[Bblock, BBuffer, gatherBParams],
|
|
350
|
+
);
|
|
351
|
+
encodePass(
|
|
352
|
+
commandEncoder,
|
|
353
|
+
transferPipeline,
|
|
354
|
+
gatherBBindGroup,
|
|
355
|
+
requireWorkgroups(device, "strsm", blockLen, otherLen),
|
|
220
356
|
);
|
|
221
|
-
const applyBindGroup = createBindGroup(gemmPipeline.getBindGroupLayout(0), [
|
|
222
|
-
{ buffer: AinvBuffer, offset: blockIndex * BLOCK_SIZE * BLOCK_SIZE * 4, size: BLOCK_SIZE * BLOCK_SIZE * 4 },
|
|
223
|
-
Bblock,
|
|
224
|
-
Xblock,
|
|
225
|
-
applyParams,
|
|
226
|
-
]);
|
|
227
|
-
const wg = useLarge
|
|
228
|
-
? { x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension) }
|
|
229
|
-
: { x: Math.min(Math.ceil(ng / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(Math.ceil(mg / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension) };
|
|
230
|
-
encodePass(commandEncoder, gemmPipeline, applyBindGroup, wg);
|
|
231
|
-
}
|
|
232
357
|
|
|
233
|
-
|
|
234
|
-
|
|
235
|
-
|
|
236
|
-
|
|
237
|
-
|
|
238
|
-
|
|
239
|
-
|
|
240
|
-
|
|
241
|
-
|
|
242
|
-
|
|
243
|
-
|
|
244
|
-
|
|
245
|
-
|
|
246
|
-
|
|
247
|
-
|
|
248
|
-
|
|
249
|
-
|
|
250
|
-
|
|
251
|
-
|
|
252
|
-
|
|
358
|
+
// 2) apply: Xblock := op(Ainv_block) @ Bblock (side='left'), or the
|
|
359
|
+
// transpose-trick equivalent for side='right' (same trick strmm uses).
|
|
360
|
+
{
|
|
361
|
+
const mg = blockLen,
|
|
362
|
+
ng = otherLen,
|
|
363
|
+
kg = blockLen;
|
|
364
|
+
const largeWgX = Math.ceil(ng / BN_LARGE),
|
|
365
|
+
largeWgY = Math.ceil(mg / BM_LARGE);
|
|
366
|
+
const useLarge =
|
|
367
|
+
largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
368
|
+
const gemmPipeline = await getPipeline(
|
|
369
|
+
device,
|
|
370
|
+
useLarge ? "sgemm_large" : "sgemm_small",
|
|
371
|
+
);
|
|
372
|
+
const applyParams = params(
|
|
373
|
+
[
|
|
374
|
+
{ value: mg, type: "u32" },
|
|
375
|
+
{ value: ng, type: "u32" },
|
|
376
|
+
{ value: kg, type: "u32" },
|
|
377
|
+
{ value: 1.0, type: "f32" }, // alpha already applied to B up front
|
|
378
|
+
{ value: 0.0, type: "f32" }, // beta — fresh output, no accumulation
|
|
379
|
+
{ value: BLOCK_SIZE, type: "u32" }, // ldX = Ainv's own dense stride
|
|
380
|
+
{ value: otherLen, type: "u32" }, // ldY = Bblock's own tight stride
|
|
381
|
+
{ value: otherLen, type: "u32" }, // ldc = Xblock's own tight stride
|
|
382
|
+
{ value: side === "right" ? 1 : 0, type: "u32" }, // transX: side='right' needs Ainv^T
|
|
383
|
+
{ value: 0, type: "u32" }, // transY: Bblock is always read as-is
|
|
384
|
+
],
|
|
385
|
+
"strsm-apply-params",
|
|
386
|
+
);
|
|
387
|
+
const ainvBlock = {
|
|
388
|
+
buffer: AinvBuffer,
|
|
389
|
+
offset: blockIndex * BLOCK_SIZE * BLOCK_SIZE * 4,
|
|
390
|
+
size: BLOCK_SIZE * BLOCK_SIZE * 4,
|
|
391
|
+
};
|
|
392
|
+
const applyBindGroup = createBindGroup(
|
|
393
|
+
device,
|
|
394
|
+
gemmPipeline.getBindGroupLayout(0),
|
|
395
|
+
[
|
|
396
|
+
ainvBlock,
|
|
397
|
+
vec4ViewBinding(device, ainvBlock),
|
|
398
|
+
Bblock,
|
|
399
|
+
vec4ViewBinding(device, Bblock),
|
|
400
|
+
Xblock,
|
|
401
|
+
applyParams,
|
|
402
|
+
],
|
|
403
|
+
);
|
|
404
|
+
const wg = useLarge
|
|
405
|
+
? {
|
|
406
|
+
x: requireWorkgroupCount(device, largeWgX, "strsm", "x"),
|
|
407
|
+
y: requireWorkgroupCount(device, largeWgY, "strsm", "y"),
|
|
408
|
+
}
|
|
409
|
+
: {
|
|
410
|
+
x: requireWorkgroupCount(
|
|
411
|
+
device,
|
|
412
|
+
Math.ceil(ng / BN_SMALL),
|
|
413
|
+
"strsm",
|
|
414
|
+
"x",
|
|
415
|
+
),
|
|
416
|
+
y: requireWorkgroupCount(
|
|
417
|
+
device,
|
|
418
|
+
Math.ceil(mg / BM_SMALL),
|
|
419
|
+
"strsm",
|
|
420
|
+
"y",
|
|
421
|
+
),
|
|
422
|
+
};
|
|
423
|
+
encodePass(commandEncoder, gemmPipeline, applyBindGroup, wg);
|
|
424
|
+
}
|
|
425
|
+
|
|
426
|
+
// 3) scatter the solved block back into B.
|
|
427
|
+
const rangeStart = forward ? blockEnd : 0;
|
|
428
|
+
const rangeEnd = forward ? aOrder : blockStart;
|
|
429
|
+
const hasRemaining = rangeStart < rangeEnd;
|
|
430
|
+
const scatterParams = params(
|
|
431
|
+
[
|
|
432
|
+
{ value: blockStart, type: "u32" },
|
|
433
|
+
{ value: blockLen, type: "u32" },
|
|
434
|
+
{ value: 0, type: "u32" },
|
|
435
|
+
{ value: otherLen, type: "u32" },
|
|
436
|
+
{ value: ldb, type: "u32" },
|
|
437
|
+
{ value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
|
|
438
|
+
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
439
|
+
{ value: 0, type: "u32" }, // overwrite
|
|
440
|
+
],
|
|
441
|
+
"strsm-scatter-params",
|
|
442
|
+
);
|
|
443
|
+
const scatterBindGroup = createBindGroup(
|
|
444
|
+
device,
|
|
445
|
+
transferPipeline.getBindGroupLayout(0),
|
|
446
|
+
[Xblock, BBuffer, scatterParams],
|
|
447
|
+
);
|
|
448
|
+
const scatterDesc =
|
|
449
|
+
isLastPass && !hasRemaining && querySet
|
|
450
|
+
? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
|
|
451
|
+
: undefined;
|
|
452
|
+
encodePass(
|
|
453
|
+
commandEncoder,
|
|
454
|
+
transferPipeline,
|
|
455
|
+
scatterBindGroup,
|
|
456
|
+
requireWorkgroups(device, "strsm", blockLen, otherLen),
|
|
457
|
+
scatterDesc,
|
|
458
|
+
);
|
|
253
459
|
|
|
254
|
-
|
|
255
|
-
|
|
256
|
-
|
|
460
|
+
// 4) trailing update: subtract this block's contribution from B.
|
|
461
|
+
if (!hasRemaining) continue;
|
|
462
|
+
const remCount = rangeEnd - rangeStart;
|
|
257
463
|
|
|
258
|
-
|
|
259
|
-
[
|
|
260
|
-
{ value: rangeStart, type: "u32" },
|
|
261
|
-
{ value: remCount, type: "u32" },
|
|
262
|
-
{ value: blockStart, type: "u32" },
|
|
263
|
-
{ value: blockLen, type: "u32" },
|
|
264
|
-
{ value: lda, type: "u32" },
|
|
265
|
-
{ value: transEffA === "transpose" ? 1 : 0, type: "u32" },
|
|
266
|
-
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
267
|
-
{ value: 2, type: "u32" }, // gather
|
|
268
|
-
],
|
|
269
|
-
"strsm-gather-A-params",
|
|
270
|
-
);
|
|
271
|
-
const gatherABindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [Aoff, ABuffer, gatherAParams]);
|
|
272
|
-
encodePass(commandEncoder, transferPipeline, gatherABindGroup, calcWorkgroups(remCount, blockLen));
|
|
273
|
-
|
|
274
|
-
{
|
|
275
|
-
const mg = remCount, ng = otherLen, kg = blockLen;
|
|
276
|
-
const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
|
|
277
|
-
const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
278
|
-
const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
|
|
279
|
-
const updateParams = params(
|
|
464
|
+
const gatherAParams = params(
|
|
280
465
|
[
|
|
281
|
-
{ value:
|
|
282
|
-
{ value:
|
|
283
|
-
{ value:
|
|
284
|
-
{ value:
|
|
285
|
-
{ value:
|
|
286
|
-
{ value:
|
|
287
|
-
{ value:
|
|
288
|
-
{ value:
|
|
289
|
-
{ value: 0, type: "u32" }, // transX: Aoff already read in the right orientation
|
|
290
|
-
{ value: 0, type: "u32" }, // transY: Xblock read as-is
|
|
466
|
+
{ value: rangeStart, type: "u32" },
|
|
467
|
+
{ value: remCount, type: "u32" },
|
|
468
|
+
{ value: blockStart, type: "u32" },
|
|
469
|
+
{ value: blockLen, type: "u32" },
|
|
470
|
+
{ value: lda, type: "u32" },
|
|
471
|
+
{ value: transEffA === "transpose" ? 1 : 0, type: "u32" },
|
|
472
|
+
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
473
|
+
{ value: 2, type: "u32" }, // gather
|
|
291
474
|
],
|
|
292
|
-
"strsm-
|
|
475
|
+
"strsm-gather-A-params",
|
|
476
|
+
);
|
|
477
|
+
const gatherABindGroup = createBindGroup(
|
|
478
|
+
device,
|
|
479
|
+
transferPipeline.getBindGroupLayout(0),
|
|
480
|
+
[Aoff, ABuffer, gatherAParams],
|
|
481
|
+
);
|
|
482
|
+
encodePass(
|
|
483
|
+
commandEncoder,
|
|
484
|
+
transferPipeline,
|
|
485
|
+
gatherABindGroup,
|
|
486
|
+
requireWorkgroups(device, "strsm", remCount, blockLen),
|
|
293
487
|
);
|
|
294
|
-
const updateBindGroup = createBindGroup(gemmPipeline.getBindGroupLayout(0), [Aoff, Xblock, delta, updateParams]);
|
|
295
|
-
const wg = useLarge
|
|
296
|
-
? { x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension) }
|
|
297
|
-
: { x: Math.min(Math.ceil(ng / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(Math.ceil(mg / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension) };
|
|
298
|
-
encodePass(commandEncoder, gemmPipeline, updateBindGroup, wg);
|
|
299
|
-
}
|
|
300
488
|
|
|
301
|
-
|
|
302
|
-
|
|
303
|
-
|
|
304
|
-
|
|
305
|
-
|
|
306
|
-
|
|
307
|
-
|
|
308
|
-
|
|
309
|
-
|
|
310
|
-
|
|
311
|
-
|
|
312
|
-
|
|
313
|
-
|
|
314
|
-
|
|
315
|
-
|
|
316
|
-
|
|
489
|
+
{
|
|
490
|
+
const mg = remCount,
|
|
491
|
+
ng = otherLen,
|
|
492
|
+
kg = blockLen;
|
|
493
|
+
const largeWgX = Math.ceil(ng / BN_LARGE),
|
|
494
|
+
largeWgY = Math.ceil(mg / BM_LARGE);
|
|
495
|
+
const useLarge =
|
|
496
|
+
largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
497
|
+
const gemmPipeline = await getPipeline(
|
|
498
|
+
device,
|
|
499
|
+
useLarge ? "sgemm_large" : "sgemm_small",
|
|
500
|
+
);
|
|
501
|
+
const updateParams = params(
|
|
502
|
+
[
|
|
503
|
+
{ value: mg, type: "u32" },
|
|
504
|
+
{ value: ng, type: "u32" },
|
|
505
|
+
{ value: kg, type: "u32" },
|
|
506
|
+
{ value: 1.0, type: "f32" },
|
|
507
|
+
{ value: 0.0, type: "f32" },
|
|
508
|
+
{ value: blockLen, type: "u32" }, // ldX = Aoff's own tight stride
|
|
509
|
+
{ value: otherLen, type: "u32" }, // ldY = Xblock's own tight stride
|
|
510
|
+
{ value: otherLen, type: "u32" }, // ldc = delta's own tight stride
|
|
511
|
+
{ value: 0, type: "u32" }, // transX: Aoff already read in the right orientation
|
|
512
|
+
{ value: 0, type: "u32" }, // transY: Xblock read as-is
|
|
513
|
+
],
|
|
514
|
+
"strsm-update-params",
|
|
515
|
+
);
|
|
516
|
+
const updateBindGroup = createBindGroup(
|
|
517
|
+
device,
|
|
518
|
+
gemmPipeline.getBindGroupLayout(0),
|
|
519
|
+
[
|
|
520
|
+
Aoff,
|
|
521
|
+
vec4ViewBinding(device, Aoff),
|
|
522
|
+
Xblock,
|
|
523
|
+
vec4ViewBinding(device, Xblock),
|
|
524
|
+
delta,
|
|
525
|
+
updateParams,
|
|
526
|
+
],
|
|
527
|
+
);
|
|
528
|
+
const wg = useLarge
|
|
529
|
+
? {
|
|
530
|
+
x: requireWorkgroupCount(device, largeWgX, "strsm", "x"),
|
|
531
|
+
y: requireWorkgroupCount(device, largeWgY, "strsm", "y"),
|
|
532
|
+
}
|
|
533
|
+
: {
|
|
534
|
+
x: requireWorkgroupCount(
|
|
535
|
+
device,
|
|
536
|
+
Math.ceil(ng / BN_SMALL),
|
|
537
|
+
"strsm",
|
|
538
|
+
"x",
|
|
539
|
+
),
|
|
540
|
+
y: requireWorkgroupCount(
|
|
541
|
+
device,
|
|
542
|
+
Math.ceil(mg / BM_SMALL),
|
|
543
|
+
"strsm",
|
|
544
|
+
"y",
|
|
545
|
+
),
|
|
546
|
+
};
|
|
547
|
+
encodePass(commandEncoder, gemmPipeline, updateBindGroup, wg);
|
|
548
|
+
}
|
|
549
|
+
|
|
550
|
+
const scatterSubParams = params(
|
|
551
|
+
[
|
|
552
|
+
{ value: rangeStart, type: "u32" },
|
|
553
|
+
{ value: remCount, type: "u32" },
|
|
554
|
+
{ value: 0, type: "u32" },
|
|
555
|
+
{ value: otherLen, type: "u32" },
|
|
556
|
+
{ value: ldb, type: "u32" },
|
|
557
|
+
{ value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
|
|
558
|
+
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
559
|
+
{ value: 1, type: "u32" }, // subtract
|
|
560
|
+
],
|
|
561
|
+
"strsm-scatter-sub-params",
|
|
562
|
+
);
|
|
563
|
+
const scatterSubBindGroup = createBindGroup(
|
|
564
|
+
device,
|
|
565
|
+
transferPipeline.getBindGroupLayout(0),
|
|
566
|
+
[delta, BBuffer, scatterSubParams],
|
|
567
|
+
);
|
|
568
|
+
const subDesc =
|
|
569
|
+
isLastPass && querySet
|
|
570
|
+
? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
|
|
571
|
+
: undefined;
|
|
572
|
+
encodePass(
|
|
573
|
+
commandEncoder,
|
|
574
|
+
transferPipeline,
|
|
575
|
+
scatterSubBindGroup,
|
|
576
|
+
requireWorkgroups(device, "strsm", remCount, otherLen),
|
|
577
|
+
subDesc,
|
|
578
|
+
);
|
|
317
579
|
}
|
|
318
580
|
}
|
|
319
581
|
|
|
320
|
-
const ts = resolveTimestamp(commandEncoder, querySet);
|
|
321
|
-
const readBuffer = BIsGpu
|
|
582
|
+
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
583
|
+
const readBuffer = BIsGpu
|
|
584
|
+
? null
|
|
585
|
+
: stageReadback(device, commandEncoder, BBuffer);
|
|
322
586
|
|
|
323
|
-
submit(commandEncoder);
|
|
587
|
+
submit(device, commandEncoder);
|
|
324
588
|
|
|
325
589
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
326
590
|
|
|
@@ -333,9 +597,9 @@ export async function strsm(
|
|
|
333
597
|
if (gpuTimeMs !== undefined) return { B: result, gpuTimeMs };
|
|
334
598
|
return { B: result };
|
|
335
599
|
} finally {
|
|
336
|
-
if (!AIsGpu) destroyBuffers(ABuffer);
|
|
337
|
-
if (!BIsGpu) destroyBuffers(BBuffer);
|
|
338
|
-
destroyBuffers(AinvBuffer);
|
|
600
|
+
if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
|
|
601
|
+
if (!BIsGpu && BBuffer) destroyBuffers(BBuffer);
|
|
602
|
+
if (AinvBuffer) destroyBuffers(AinvBuffer);
|
|
339
603
|
destroyBuffers(scratchBuffers);
|
|
340
604
|
destroyBuffers(paramsBuffers);
|
|
341
605
|
}
|