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/strsm/strsm.mjs
CHANGED
|
@@ -11,26 +11,46 @@ import { beginTimedEncoder, encodePass, submit } from "../util/compute.mjs";
|
|
|
11
11
|
import { extractResult } from "../util/result.mjs";
|
|
12
12
|
import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
|
|
13
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
|
-
import {
|
|
14
|
+
import {
|
|
15
|
+
calcWorkgroups,
|
|
16
|
+
requireWorkgroups,
|
|
17
|
+
requireWorkgroupCount,
|
|
18
|
+
} from "../util/workgroup.mjs";
|
|
15
19
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
16
|
-
import {
|
|
20
|
+
import {
|
|
21
|
+
BM_SMALL,
|
|
22
|
+
BN_SMALL,
|
|
23
|
+
BM_LARGE,
|
|
24
|
+
BN_LARGE,
|
|
25
|
+
LARGE_TILE_WORKGROUP_THRESHOLD,
|
|
26
|
+
} from "../util/constants.mjs";
|
|
17
27
|
import { BLOCK_SIZE } from "../util/constants.mjs";
|
|
18
|
-
import { requireSameDevice } from "../util/device.mjs";
|
|
19
|
-
|
|
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
|
-
throw new Error("device must be a GPUDevice.");
|
|
53
|
+
requireGpuDevice(device);
|
|
34
54
|
requireSameDevice(device, "strsm", { A, B });
|
|
35
55
|
if (side !== "left" && side !== "right")
|
|
36
56
|
throw new Error("side must be 'left' or 'right'.");
|
|
@@ -42,11 +62,15 @@ export async function strsm(
|
|
|
42
62
|
throw new Error("diag must be 'unit' or 'non-unit'.");
|
|
43
63
|
if (layout !== "row-major" && layout !== "column-major")
|
|
44
64
|
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.");
|
|
65
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
47
66
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
48
67
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
49
|
-
if (
|
|
68
|
+
if (
|
|
69
|
+
!Number.isInteger(m) ||
|
|
70
|
+
!Number.isInteger(n) ||
|
|
71
|
+
!Number.isInteger(lda) ||
|
|
72
|
+
!Number.isInteger(ldb)
|
|
73
|
+
)
|
|
50
74
|
throw new Error("m, n, lda, and ldb must be integers.");
|
|
51
75
|
if (!AIsGpu && !(A instanceof Float32Array))
|
|
52
76
|
throw new Error("A must be a Float32Array or GpuMatrix.");
|
|
@@ -62,30 +86,51 @@ export async function strsm(
|
|
|
62
86
|
|
|
63
87
|
// A: triangular, order = m (side='left') or n (side='right').
|
|
64
88
|
const aOrder = side === "left" ? m : n;
|
|
65
|
-
if (lda < aOrder)
|
|
89
|
+
if (lda < aOrder)
|
|
90
|
+
throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
|
|
66
91
|
if (AIsGpu) {
|
|
67
|
-
if (lda !== A.lda)
|
|
68
|
-
|
|
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.");
|
|
69
96
|
} else if (A.length < (aOrder - 1) * lda + aOrder) {
|
|
70
|
-
throw new Error(
|
|
97
|
+
throw new Error(
|
|
98
|
+
"A does not have enough elements for the given dimensions and lda.",
|
|
99
|
+
);
|
|
71
100
|
}
|
|
72
101
|
|
|
73
102
|
// B: always m x n, overwritten in place with the same ldb.
|
|
74
103
|
const bOuter = effLayoutB === "column-major" ? n : m;
|
|
75
104
|
const bInner = effLayoutB === "column-major" ? m : n;
|
|
76
105
|
if (ldb < bInner)
|
|
77
|
-
throw new Error(
|
|
106
|
+
throw new Error(
|
|
107
|
+
`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
|
|
108
|
+
);
|
|
78
109
|
if (BIsGpu) {
|
|
79
|
-
if (ldb !== B.lda)
|
|
80
|
-
|
|
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.");
|
|
81
114
|
} else if (B.length < (bOuter - 1) * ldb + bInner) {
|
|
82
|
-
throw new Error(
|
|
115
|
+
throw new Error(
|
|
116
|
+
"B does not have enough elements for the given dimensions and ldb.",
|
|
117
|
+
);
|
|
83
118
|
}
|
|
84
119
|
|
|
85
120
|
// A isn't symmetric: column-major = genuine transpose, so flip transA;
|
|
86
121
|
// transposing also swaps which triangle looks stored, so flip uplo too.
|
|
87
|
-
const uploEffA =
|
|
88
|
-
|
|
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;
|
|
89
134
|
|
|
90
135
|
const otherLen = side === "left" ? n : m;
|
|
91
136
|
const blockIsRow = side === "left";
|
|
@@ -132,11 +177,17 @@ export async function strsm(
|
|
|
132
177
|
try {
|
|
133
178
|
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strsm-A", false);
|
|
134
179
|
BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strsm-B", true);
|
|
135
|
-
AinvBuffer = createStorageBuffer(
|
|
180
|
+
AinvBuffer = createStorageBuffer(
|
|
181
|
+
device,
|
|
182
|
+
numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4,
|
|
183
|
+
"strsm-Ainv",
|
|
184
|
+
);
|
|
136
185
|
|
|
137
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.
|
|
138
189
|
let preScaleBindGroup = null;
|
|
139
|
-
if (alpha !== 1.0) {
|
|
190
|
+
if (alpha !== 1.0 && alpha !== 0) {
|
|
140
191
|
const scaleParams = params(
|
|
141
192
|
[
|
|
142
193
|
{ value: bScaleLen, type: "u32" },
|
|
@@ -145,7 +196,11 @@ export async function strsm(
|
|
|
145
196
|
],
|
|
146
197
|
"strsm-scale-params",
|
|
147
198
|
);
|
|
148
|
-
preScaleBindGroup = createBindGroup(
|
|
199
|
+
preScaleBindGroup = createBindGroup(
|
|
200
|
+
device,
|
|
201
|
+
scalarPipeline.getBindGroupLayout(0),
|
|
202
|
+
[BBuffer, scaleParams],
|
|
203
|
+
);
|
|
149
204
|
}
|
|
150
205
|
|
|
151
206
|
// Every diagonal block's inverse, fully parallel, one dispatch, unchanged.
|
|
@@ -159,7 +214,11 @@ export async function strsm(
|
|
|
159
214
|
],
|
|
160
215
|
"strsm-invert-params",
|
|
161
216
|
);
|
|
162
|
-
const invertBindGroup = createBindGroup(
|
|
217
|
+
const invertBindGroup = createBindGroup(
|
|
218
|
+
device,
|
|
219
|
+
invertPipeline.getBindGroupLayout(0),
|
|
220
|
+
[ABuffer, AinvBuffer, invertParams],
|
|
221
|
+
);
|
|
163
222
|
|
|
164
223
|
// Reusable scratch buffers, sized for the worst case, bound at offset 0.
|
|
165
224
|
const Bblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Bblock");
|
|
@@ -170,173 +229,360 @@ export async function strsm(
|
|
|
170
229
|
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
171
230
|
|
|
172
231
|
if (alpha === 0) {
|
|
173
|
-
// BLAS: alpha=0 means
|
|
174
|
-
|
|
175
|
-
|
|
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
|
+
);
|
|
176
305
|
} else {
|
|
177
306
|
if (preScaleBindGroup) {
|
|
178
|
-
encodePass(
|
|
307
|
+
encodePass(
|
|
308
|
+
commandEncoder,
|
|
309
|
+
scalarPipeline,
|
|
310
|
+
preScaleBindGroup,
|
|
311
|
+
calcWorkgroups(device, bScaleLen),
|
|
312
|
+
);
|
|
179
313
|
}
|
|
180
|
-
const invertDesc = querySet
|
|
181
|
-
|
|
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
|
+
);
|
|
182
324
|
|
|
183
325
|
for (let bi = 0; bi < blockStarts.length; bi++) {
|
|
184
|
-
|
|
185
|
-
|
|
186
|
-
|
|
187
|
-
|
|
188
|
-
|
|
189
|
-
|
|
190
|
-
|
|
191
|
-
|
|
192
|
-
[
|
|
193
|
-
{ value: blockStart, type: "u32" },
|
|
194
|
-
{ value: blockLen, type: "u32" },
|
|
195
|
-
{ value: 0, type: "u32" },
|
|
196
|
-
{ value: otherLen, type: "u32" },
|
|
197
|
-
{ value: ldb, type: "u32" },
|
|
198
|
-
{ value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
|
|
199
|
-
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
200
|
-
{ value: 2, type: "u32" }, // gather
|
|
201
|
-
],
|
|
202
|
-
"strsm-gather-B-params",
|
|
203
|
-
);
|
|
204
|
-
const gatherBBindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [Bblock, BBuffer, gatherBParams]);
|
|
205
|
-
encodePass(commandEncoder, transferPipeline, gatherBBindGroup, requireWorkgroups(device, "strsm", blockLen, otherLen));
|
|
206
|
-
|
|
207
|
-
// 2) apply: Xblock := op(Ainv_block) @ Bblock (side='left'), or the
|
|
208
|
-
// transpose-trick equivalent for side='right' (same trick strmm uses).
|
|
209
|
-
{
|
|
210
|
-
const mg = blockLen, ng = otherLen, kg = blockLen;
|
|
211
|
-
const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
|
|
212
|
-
const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
213
|
-
const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
|
|
214
|
-
const applyParams = params(
|
|
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(
|
|
215
334
|
[
|
|
216
|
-
{ value:
|
|
217
|
-
{ value:
|
|
218
|
-
{ value:
|
|
219
|
-
{ value:
|
|
220
|
-
{ value:
|
|
221
|
-
{ value:
|
|
222
|
-
{ value:
|
|
223
|
-
{ value:
|
|
224
|
-
{ value: side === "right" ? 1 : 0, type: "u32" }, // transX: side='right' needs Ainv^T
|
|
225
|
-
{ value: 0, type: "u32" }, // transY: Bblock is always read as-is
|
|
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
|
|
226
343
|
],
|
|
227
|
-
"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),
|
|
228
356
|
);
|
|
229
|
-
const ainvBlock = { buffer: AinvBuffer, offset: blockIndex * BLOCK_SIZE * BLOCK_SIZE * 4, size: BLOCK_SIZE * BLOCK_SIZE * 4 };
|
|
230
|
-
const applyBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
|
|
231
|
-
ainvBlock,
|
|
232
|
-
vec4ViewBinding(device, ainvBlock),
|
|
233
|
-
Bblock,
|
|
234
|
-
vec4ViewBinding(device, Bblock),
|
|
235
|
-
Xblock,
|
|
236
|
-
applyParams,
|
|
237
|
-
]);
|
|
238
|
-
const wg = useLarge
|
|
239
|
-
? { x: requireWorkgroupCount(device, largeWgX, "strsm", "x"), y: requireWorkgroupCount(device, largeWgY, "strsm", "y") }
|
|
240
|
-
: { x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "strsm", "x"), y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "strsm", "y") };
|
|
241
|
-
encodePass(commandEncoder, gemmPipeline, applyBindGroup, wg);
|
|
242
|
-
}
|
|
243
357
|
|
|
244
|
-
|
|
245
|
-
|
|
246
|
-
|
|
247
|
-
|
|
248
|
-
|
|
249
|
-
|
|
250
|
-
|
|
251
|
-
|
|
252
|
-
|
|
253
|
-
|
|
254
|
-
|
|
255
|
-
|
|
256
|
-
|
|
257
|
-
|
|
258
|
-
|
|
259
|
-
|
|
260
|
-
|
|
261
|
-
|
|
262
|
-
|
|
263
|
-
|
|
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
|
+
);
|
|
264
459
|
|
|
265
|
-
|
|
266
|
-
|
|
267
|
-
|
|
460
|
+
// 4) trailing update: subtract this block's contribution from B.
|
|
461
|
+
if (!hasRemaining) continue;
|
|
462
|
+
const remCount = rangeEnd - rangeStart;
|
|
268
463
|
|
|
269
|
-
|
|
270
|
-
[
|
|
271
|
-
{ value: rangeStart, type: "u32" },
|
|
272
|
-
{ value: remCount, type: "u32" },
|
|
273
|
-
{ value: blockStart, type: "u32" },
|
|
274
|
-
{ value: blockLen, type: "u32" },
|
|
275
|
-
{ value: lda, type: "u32" },
|
|
276
|
-
{ value: transEffA === "transpose" ? 1 : 0, type: "u32" },
|
|
277
|
-
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
278
|
-
{ value: 2, type: "u32" }, // gather
|
|
279
|
-
],
|
|
280
|
-
"strsm-gather-A-params",
|
|
281
|
-
);
|
|
282
|
-
const gatherABindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [Aoff, ABuffer, gatherAParams]);
|
|
283
|
-
encodePass(commandEncoder, transferPipeline, gatherABindGroup, requireWorkgroups(device, "strsm", remCount, blockLen));
|
|
284
|
-
|
|
285
|
-
{
|
|
286
|
-
const mg = remCount, ng = otherLen, kg = blockLen;
|
|
287
|
-
const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
|
|
288
|
-
const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
289
|
-
const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
|
|
290
|
-
const updateParams = params(
|
|
464
|
+
const gatherAParams = params(
|
|
291
465
|
[
|
|
292
|
-
{ value:
|
|
293
|
-
{ value:
|
|
294
|
-
{ value:
|
|
295
|
-
{ value:
|
|
296
|
-
{ value:
|
|
297
|
-
{ value:
|
|
298
|
-
{ value:
|
|
299
|
-
{ value:
|
|
300
|
-
{ value: 0, type: "u32" }, // transX: Aoff already read in the right orientation
|
|
301
|
-
{ 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
|
|
302
474
|
],
|
|
303
|
-
"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),
|
|
304
487
|
);
|
|
305
|
-
const updateBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
|
|
306
|
-
Aoff,
|
|
307
|
-
vec4ViewBinding(device, Aoff),
|
|
308
|
-
Xblock,
|
|
309
|
-
vec4ViewBinding(device, Xblock),
|
|
310
|
-
delta,
|
|
311
|
-
updateParams,
|
|
312
|
-
]);
|
|
313
|
-
const wg = useLarge
|
|
314
|
-
? { x: requireWorkgroupCount(device, largeWgX, "strsm", "x"), y: requireWorkgroupCount(device, largeWgY, "strsm", "y") }
|
|
315
|
-
: { x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "strsm", "x"), y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "strsm", "y") };
|
|
316
|
-
encodePass(commandEncoder, gemmPipeline, updateBindGroup, wg);
|
|
317
|
-
}
|
|
318
488
|
|
|
319
|
-
|
|
320
|
-
|
|
321
|
-
|
|
322
|
-
|
|
323
|
-
|
|
324
|
-
|
|
325
|
-
|
|
326
|
-
|
|
327
|
-
|
|
328
|
-
|
|
329
|
-
|
|
330
|
-
|
|
331
|
-
|
|
332
|
-
|
|
333
|
-
|
|
334
|
-
|
|
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
|
+
);
|
|
335
579
|
}
|
|
336
580
|
}
|
|
337
581
|
|
|
338
582
|
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
339
|
-
const readBuffer = BIsGpu
|
|
583
|
+
const readBuffer = BIsGpu
|
|
584
|
+
? null
|
|
585
|
+
: stageReadback(device, commandEncoder, BBuffer);
|
|
340
586
|
|
|
341
587
|
submit(device, commandEncoder);
|
|
342
588
|
|