wgblas 1.2.0 → 2.0.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/LICENSE +1 -1
- package/README.md +36 -54
- package/dist/wgblas.browser.js +779 -34
- package/index.d.mts +13 -0
- package/index.mjs +8 -0
- package/package.json +56 -3
- package/src/dasum/dasum.d.mts +2 -2
- package/src/dasum/dasum.mjs +6 -4
- package/src/idamax/idamax.d.mts +51 -0
- package/src/idamax/idamax.mjs +128 -0
- package/src/init.mjs +9 -1
- package/src/isamax/isamax.d.mts +1 -1
- package/src/sasum/sasum.d.mts +1 -1
- package/src/saxpy/saxpy.d.mts +1 -1
- package/src/scopy/scopy.d.mts +1 -1
- package/src/sdot/sdot.d.mts +1 -1
- package/src/sgemm/sgemm.d.mts +102 -0
- package/src/sgemm/sgemm.mjs +195 -0
- package/src/sgemmtr/sgemmtr.d.mts +104 -0
- package/src/sgemmtr/sgemmtr.mjs +203 -0
- package/src/sgemv/sgemv.d.mts +1 -38
- package/src/sgemv/sgemv.mjs +4 -0
- package/src/sger/sger.d.mts +1 -34
- package/src/sger/sger.mjs +2 -0
- package/src/shaders/block_transfer.wgsl +42 -0
- package/src/shaders/browser-shaders.mjs +26 -0
- package/src/shaders/dasum.wgsl +3 -2
- package/src/shaders/f64/dekker.wgsl +4 -85
- package/src/shaders/f64/utils/abs.wgsl +10 -0
- package/src/shaders/f64/utils/add.wgsl +77 -0
- package/src/shaders/f64/utils/equal.wgsl +7 -0
- package/src/shaders/f64/utils/greater.wgsl +12 -0
- package/src/shaders/f64/utils/multiply.wgsl +81 -0
- package/src/shaders/idamax.wgsl +96 -0
- package/src/shaders/reduction/argmaxF64.wgsl +50 -0
- package/src/shaders/reduction/sumF64.wgsl +2 -2
- package/src/shaders/sgemm_large.wgsl +117 -0
- package/src/shaders/sgemm_small.wgsl +112 -0
- package/src/shaders/sgemmtr_large.wgsl +117 -0
- package/src/shaders/sgemmtr_small.wgsl +110 -0
- package/src/shaders/symmetrize.wgsl +31 -0
- package/src/shaders/triangularize.wgsl +44 -0
- package/src/snrm2/snrm2.d.mts +1 -1
- package/src/srot/srot.d.mts +1 -1
- package/src/srotm/srotm.d.mts +1 -1
- package/src/sscal/sscal.d.mts +1 -1
- package/src/sswap/sswap.d.mts +1 -1
- package/src/ssymm/ssymm.d.mts +103 -0
- package/src/ssymm/ssymm.mjs +209 -0
- package/src/ssymv/ssymv.d.mts +1 -36
- package/src/ssymv/ssymv.mjs +2 -0
- package/src/ssyr/ssyr.d.mts +1 -30
- package/src/ssyr/ssyr.mjs +2 -0
- package/src/ssyr2/ssyr2.d.mts +1 -34
- package/src/ssyr2/ssyr2.mjs +2 -0
- package/src/ssyr2k/ssyr2k.d.mts +100 -0
- package/src/ssyr2k/ssyr2k.mjs +201 -0
- package/src/ssyrk/ssyrk.d.mts +90 -0
- package/src/ssyrk/ssyrk.mjs +176 -0
- package/src/strmm/strmm.d.mts +100 -0
- package/src/strmm/strmm.mjs +211 -0
- package/src/strmv/strmv.d.mts +1 -36
- package/src/strmv/strmv.mjs +2 -0
- package/src/strsm/strsm.d.mts +99 -0
- package/src/strsm/strsm.mjs +342 -0
- package/src/strsv/strsv.d.mts +1 -32
- package/src/strsv/strsv.mjs +2 -0
- package/src/util/buffer.mjs +4 -2
- package/src/util/compute.mjs +6 -3
- package/src/util/f64.mjs +3 -3
|
@@ -0,0 +1,342 @@
|
|
|
1
|
+
import {
|
|
2
|
+
uploadBuffer,
|
|
3
|
+
createParamsBuffer,
|
|
4
|
+
createStorageBuffer,
|
|
5
|
+
stageReadback,
|
|
6
|
+
destroyBuffers,
|
|
7
|
+
} from "../util/buffer.mjs";
|
|
8
|
+
import { createBindGroup } from "../util/bindgroup.mjs";
|
|
9
|
+
import { beginTimedEncoder, encodePass, submit } from "../util/compute.mjs";
|
|
10
|
+
import { extractResult } from "../util/result.mjs";
|
|
11
|
+
import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
|
|
12
|
+
import { getPipeline } from "../util/pipeline.mjs";
|
|
13
|
+
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
14
|
+
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
15
|
+
|
|
16
|
+
const BLOCK_SIZE = 64; // must match strsv_invert_block.wgsl's own constant
|
|
17
|
+
const BM_SMALL = 32, BN_SMALL = 32; // sgemm_small.wgsl's block tile
|
|
18
|
+
const BM_LARGE = 64, BN_LARGE = 64; // sgemm_large.wgsl's block tile
|
|
19
|
+
const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // same threshold sgemm/ssymm/strmm use
|
|
20
|
+
|
|
21
|
+
// strsm: B := alpha*op(A)^-1*B (side='left') or alpha*B*op(A)^-1 (side='right'),
|
|
22
|
+
// A triangular. Blocked substitution (strsv's own technique, generalized to
|
|
23
|
+
// a matrix RHS): strsv_invert_block + sgemm, unchanged; every per-block B/A
|
|
24
|
+
// access goes through block_transfer.wgsl (see that shader for why).
|
|
25
|
+
export async function strsm(
|
|
26
|
+
device, side, uplo, transA, diag, m, n, alpha, A, lda, B, ldb, layout = "row-major",
|
|
27
|
+
) {
|
|
28
|
+
const AIsGpu = A instanceof GpuMatrix;
|
|
29
|
+
const BIsGpu = B instanceof GpuMatrix;
|
|
30
|
+
const isUnit = diag === "unit";
|
|
31
|
+
|
|
32
|
+
if (!(device instanceof GPUDevice))
|
|
33
|
+
throw new Error("device must be a GPUDevice.");
|
|
34
|
+
if (side !== "left" && side !== "right")
|
|
35
|
+
throw new Error("side must be 'left' or 'right'.");
|
|
36
|
+
if (uplo !== "lower" && uplo !== "upper")
|
|
37
|
+
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
38
|
+
if (transA !== "no-transpose" && transA !== "transpose")
|
|
39
|
+
throw new Error("transA must be 'no-transpose' or 'transpose'.");
|
|
40
|
+
if (!isUnit && diag !== "non-unit")
|
|
41
|
+
throw new Error("diag must be 'unit' or 'non-unit'.");
|
|
42
|
+
if (layout !== "row-major" && layout !== "column-major")
|
|
43
|
+
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.");
|
|
46
|
+
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
47
|
+
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
48
|
+
if (!Number.isInteger(m) || !Number.isInteger(n) || !Number.isInteger(lda) || !Number.isInteger(ldb))
|
|
49
|
+
throw new Error("m, n, lda, and ldb must be integers.");
|
|
50
|
+
if (!AIsGpu && !(A instanceof Float32Array))
|
|
51
|
+
throw new Error("A must be a Float32Array or GpuMatrix.");
|
|
52
|
+
if (!BIsGpu && !(B instanceof Float32Array))
|
|
53
|
+
throw new Error("B must be a Float32Array or GpuMatrix.");
|
|
54
|
+
if (AIsGpu !== BIsGpu)
|
|
55
|
+
throw new Error("A and B must both be GpuMatrix or both be Float32Array.");
|
|
56
|
+
if (m < 0 || n < 0) throw new Error("m and n must be non-negative.");
|
|
57
|
+
if (m === 0 || n === 0) return BIsGpu ? {} : { B };
|
|
58
|
+
|
|
59
|
+
const effLayoutA = AIsGpu ? A.layout : layout;
|
|
60
|
+
const effLayoutB = BIsGpu ? B.layout : layout;
|
|
61
|
+
|
|
62
|
+
// A: triangular, order = m (side='left') or n (side='right').
|
|
63
|
+
const aOrder = side === "left" ? m : n;
|
|
64
|
+
if (lda < aOrder) throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
|
|
65
|
+
if (AIsGpu) {
|
|
66
|
+
if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
67
|
+
if (A.rows < aOrder || A.cols < aOrder) throw new Error("A is too small for the given m/n and side.");
|
|
68
|
+
} else if (A.length < (aOrder - 1) * lda + aOrder) {
|
|
69
|
+
throw new Error("A does not have enough elements for the given dimensions and lda.");
|
|
70
|
+
}
|
|
71
|
+
|
|
72
|
+
// B: always m x n, overwritten in place with the same ldb.
|
|
73
|
+
const bOuter = effLayoutB === "column-major" ? n : m;
|
|
74
|
+
const bInner = effLayoutB === "column-major" ? m : n;
|
|
75
|
+
if (ldb < bInner)
|
|
76
|
+
throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
|
|
77
|
+
if (BIsGpu) {
|
|
78
|
+
if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
|
|
79
|
+
if (B.rows < m || B.cols < n) throw new Error("B is too small for the given m and n.");
|
|
80
|
+
} else if (B.length < (bOuter - 1) * ldb + bInner) {
|
|
81
|
+
throw new Error("B does not have enough elements for the given dimensions and ldb.");
|
|
82
|
+
}
|
|
83
|
+
|
|
84
|
+
// A isn't symmetric: column-major = genuine transpose, so flip transA;
|
|
85
|
+
// transposing also swaps which triangle looks stored, so flip uplo too.
|
|
86
|
+
const uploEffA = effLayoutA === "column-major" ? (uplo === "lower" ? "upper" : "lower") : uplo;
|
|
87
|
+
const transEffA = effLayoutA === "column-major" ? (transA === "no-transpose" ? "transpose" : "no-transpose") : transA;
|
|
88
|
+
|
|
89
|
+
const otherLen = side === "left" ? n : m;
|
|
90
|
+
const blockIsRow = side === "left";
|
|
91
|
+
|
|
92
|
+
// Forward iff op(A) is lower (side='left') — side='right' flips this,
|
|
93
|
+
// since it solves via columns instead of rows.
|
|
94
|
+
const opIsLower = (transEffA === "no-transpose") === (uploEffA === "lower");
|
|
95
|
+
const forward = side === "left" ? opIsLower : !opIsLower;
|
|
96
|
+
|
|
97
|
+
const blockStarts = [];
|
|
98
|
+
for (let s = 0; s < aOrder; s += BLOCK_SIZE) blockStarts.push(s);
|
|
99
|
+
if (!forward) blockStarts.reverse();
|
|
100
|
+
const numBlocks = blockStarts.length;
|
|
101
|
+
|
|
102
|
+
const invertPipeline = await getPipeline(device, "strsv_invert_block");
|
|
103
|
+
const transferPipeline = await getPipeline(device, "block_transfer");
|
|
104
|
+
const scalarPipeline = await getPipeline(device, "sscal");
|
|
105
|
+
|
|
106
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "strsm-A", false);
|
|
107
|
+
const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "strsm-B", true);
|
|
108
|
+
const AinvBuffer = createStorageBuffer(numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4, "strsm-Ainv");
|
|
109
|
+
|
|
110
|
+
const paramsBuffers = [];
|
|
111
|
+
const scratchBuffers = [];
|
|
112
|
+
function scratch(size, label) {
|
|
113
|
+
const buf = createStorageBuffer(size, label);
|
|
114
|
+
scratchBuffers.push(buf);
|
|
115
|
+
return buf;
|
|
116
|
+
}
|
|
117
|
+
function params(entries, label) {
|
|
118
|
+
const buf = createParamsBuffer(entries, label);
|
|
119
|
+
paramsBuffers.push(buf);
|
|
120
|
+
return buf;
|
|
121
|
+
}
|
|
122
|
+
|
|
123
|
+
// Minimal valid length for B's own buffer (matches its own validation
|
|
124
|
+
// above) — NOT bOuter*ldb, which can exceed a Float32Array-path buffer's
|
|
125
|
+
// actual allocation (only guaranteed padded up to GpuMatrix's own size).
|
|
126
|
+
const bScaleLen = (bOuter - 1) * ldb + bInner;
|
|
127
|
+
|
|
128
|
+
try {
|
|
129
|
+
// Pre-scale B by alpha once (reuses sscal, so no per-block alpha handling).
|
|
130
|
+
let preScaleBindGroup = null;
|
|
131
|
+
if (alpha !== 1.0) {
|
|
132
|
+
const scaleParams = params(
|
|
133
|
+
[
|
|
134
|
+
{ value: bScaleLen, type: "u32" },
|
|
135
|
+
{ value: alpha, type: "f32" },
|
|
136
|
+
{ value: 1, type: "u32" },
|
|
137
|
+
],
|
|
138
|
+
"strsm-scale-params",
|
|
139
|
+
);
|
|
140
|
+
preScaleBindGroup = createBindGroup(scalarPipeline.getBindGroupLayout(0), [BBuffer, scaleParams]);
|
|
141
|
+
}
|
|
142
|
+
|
|
143
|
+
// Every diagonal block's inverse, fully parallel, one dispatch, unchanged.
|
|
144
|
+
const invertParams = params(
|
|
145
|
+
[
|
|
146
|
+
{ value: aOrder, type: "u32" },
|
|
147
|
+
{ value: lda, type: "u32" },
|
|
148
|
+
{ value: transEffA === "transpose" ? 1 : 0, type: "u32" },
|
|
149
|
+
{ value: uploEffA === "upper" ? 1 : 0, type: "u32" },
|
|
150
|
+
{ value: isUnit ? 1 : 0, type: "u32" },
|
|
151
|
+
],
|
|
152
|
+
"strsm-invert-params",
|
|
153
|
+
);
|
|
154
|
+
const invertBindGroup = createBindGroup(invertPipeline.getBindGroupLayout(0), [ABuffer, AinvBuffer, invertParams]);
|
|
155
|
+
|
|
156
|
+
// Reusable scratch buffers, sized for the worst case, bound at offset 0.
|
|
157
|
+
const Bblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Bblock");
|
|
158
|
+
const Xblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Xblock");
|
|
159
|
+
const Aoff = scratch(aOrder * BLOCK_SIZE * 4, "strsm-Aoff");
|
|
160
|
+
const delta = scratch(aOrder * otherLen * 4, "strsm-delta");
|
|
161
|
+
|
|
162
|
+
const { commandEncoder, querySet } = beginTimedEncoder();
|
|
163
|
+
|
|
164
|
+
if (alpha === 0) {
|
|
165
|
+
// BLAS: alpha=0 means A is not referenced — skip straight to B:=0.
|
|
166
|
+
const zeroDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0, endOfPassWriteIndex: 1 } } : undefined;
|
|
167
|
+
encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(bScaleLen), zeroDesc);
|
|
168
|
+
} else {
|
|
169
|
+
if (preScaleBindGroup) {
|
|
170
|
+
encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(bScaleLen));
|
|
171
|
+
}
|
|
172
|
+
const invertDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
|
|
173
|
+
encodePass(commandEncoder, invertPipeline, invertBindGroup, { x: BLOCK_SIZE, y: numBlocks }, invertDesc);
|
|
174
|
+
|
|
175
|
+
for (let bi = 0; bi < blockStarts.length; bi++) {
|
|
176
|
+
const blockStart = blockStarts[bi];
|
|
177
|
+
const blockEnd = Math.min(blockStart + BLOCK_SIZE, aOrder);
|
|
178
|
+
const blockLen = blockEnd - blockStart;
|
|
179
|
+
const blockIndex = blockStart / BLOCK_SIZE;
|
|
180
|
+
const isLastPass = bi === blockStarts.length - 1;
|
|
181
|
+
|
|
182
|
+
// 1) gather B's current block into a tight scratch buffer.
|
|
183
|
+
const gatherBParams = params(
|
|
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(
|
|
207
|
+
[
|
|
208
|
+
{ value: mg, type: "u32" },
|
|
209
|
+
{ value: ng, type: "u32" },
|
|
210
|
+
{ value: kg, type: "u32" },
|
|
211
|
+
{ value: 1.0, type: "f32" }, // alpha already applied to B up front
|
|
212
|
+
{ value: 0.0, type: "f32" }, // beta — fresh output, no accumulation
|
|
213
|
+
{ value: BLOCK_SIZE, type: "u32" }, // ldX = Ainv's own dense stride
|
|
214
|
+
{ value: otherLen, type: "u32" }, // ldY = Bblock's own tight stride
|
|
215
|
+
{ value: otherLen, type: "u32" }, // ldc = Xblock's own tight stride
|
|
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
|
|
218
|
+
],
|
|
219
|
+
"strsm-apply-params",
|
|
220
|
+
);
|
|
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
|
+
|
|
233
|
+
// 3) scatter the solved block back into B.
|
|
234
|
+
const rangeStart = forward ? blockEnd : 0;
|
|
235
|
+
const rangeEnd = forward ? aOrder : blockStart;
|
|
236
|
+
const hasRemaining = rangeStart < rangeEnd;
|
|
237
|
+
const scatterParams = params(
|
|
238
|
+
[
|
|
239
|
+
{ value: blockStart, type: "u32" },
|
|
240
|
+
{ value: blockLen, type: "u32" },
|
|
241
|
+
{ value: 0, type: "u32" },
|
|
242
|
+
{ value: otherLen, type: "u32" },
|
|
243
|
+
{ value: ldb, type: "u32" },
|
|
244
|
+
{ value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
|
|
245
|
+
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
246
|
+
{ value: 0, type: "u32" }, // overwrite
|
|
247
|
+
],
|
|
248
|
+
"strsm-scatter-params",
|
|
249
|
+
);
|
|
250
|
+
const scatterBindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [Xblock, BBuffer, scatterParams]);
|
|
251
|
+
const scatterDesc = isLastPass && !hasRemaining && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
|
|
252
|
+
encodePass(commandEncoder, transferPipeline, scatterBindGroup, calcWorkgroups(blockLen, otherLen), scatterDesc);
|
|
253
|
+
|
|
254
|
+
// 4) trailing update: subtract this block's contribution from B.
|
|
255
|
+
if (!hasRemaining) continue;
|
|
256
|
+
const remCount = rangeEnd - rangeStart;
|
|
257
|
+
|
|
258
|
+
const gatherAParams = params(
|
|
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(
|
|
280
|
+
[
|
|
281
|
+
{ value: mg, type: "u32" },
|
|
282
|
+
{ value: ng, type: "u32" },
|
|
283
|
+
{ value: kg, type: "u32" },
|
|
284
|
+
{ value: 1.0, type: "f32" },
|
|
285
|
+
{ value: 0.0, type: "f32" },
|
|
286
|
+
{ value: blockLen, type: "u32" }, // ldX = Aoff's own tight stride
|
|
287
|
+
{ value: otherLen, type: "u32" }, // ldY = Xblock's own tight stride
|
|
288
|
+
{ value: otherLen, type: "u32" }, // ldc = delta's own tight stride
|
|
289
|
+
{ value: 0, type: "u32" }, // transX: Aoff already read in the right orientation
|
|
290
|
+
{ value: 0, type: "u32" }, // transY: Xblock read as-is
|
|
291
|
+
],
|
|
292
|
+
"strsm-update-params",
|
|
293
|
+
);
|
|
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
|
+
|
|
301
|
+
const scatterSubParams = params(
|
|
302
|
+
[
|
|
303
|
+
{ value: rangeStart, type: "u32" },
|
|
304
|
+
{ value: remCount, type: "u32" },
|
|
305
|
+
{ value: 0, type: "u32" },
|
|
306
|
+
{ value: otherLen, type: "u32" },
|
|
307
|
+
{ value: ldb, type: "u32" },
|
|
308
|
+
{ value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
|
|
309
|
+
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
310
|
+
{ value: 1, type: "u32" }, // subtract
|
|
311
|
+
],
|
|
312
|
+
"strsm-scatter-sub-params",
|
|
313
|
+
);
|
|
314
|
+
const scatterSubBindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [delta, BBuffer, scatterSubParams]);
|
|
315
|
+
const subDesc = isLastPass && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
|
|
316
|
+
encodePass(commandEncoder, transferPipeline, scatterSubBindGroup, calcWorkgroups(remCount, otherLen), subDesc);
|
|
317
|
+
}
|
|
318
|
+
}
|
|
319
|
+
|
|
320
|
+
const ts = resolveTimestamp(commandEncoder, querySet);
|
|
321
|
+
const readBuffer = BIsGpu ? null : stageReadback(commandEncoder, BBuffer);
|
|
322
|
+
|
|
323
|
+
submit(commandEncoder);
|
|
324
|
+
|
|
325
|
+
const gpuTimeMs = await extractTimestamp(ts);
|
|
326
|
+
|
|
327
|
+
if (BIsGpu) {
|
|
328
|
+
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
329
|
+
return {};
|
|
330
|
+
}
|
|
331
|
+
|
|
332
|
+
const result = await extractResult(readBuffer, Float32Array);
|
|
333
|
+
if (gpuTimeMs !== undefined) return { B: result, gpuTimeMs };
|
|
334
|
+
return { B: result };
|
|
335
|
+
} finally {
|
|
336
|
+
if (!AIsGpu) destroyBuffers(ABuffer);
|
|
337
|
+
if (!BIsGpu) destroyBuffers(BBuffer);
|
|
338
|
+
destroyBuffers(AinvBuffer);
|
|
339
|
+
destroyBuffers(scratchBuffers);
|
|
340
|
+
destroyBuffers(paramsBuffers);
|
|
341
|
+
}
|
|
342
|
+
}
|
package/src/strsv/strsv.d.mts
CHANGED
|
@@ -41,37 +41,6 @@ export declare function strsv(
|
|
|
41
41
|
layout?: 'row-major' | 'column-major',
|
|
42
42
|
): Promise<{ x: Float32Array; gpuTimeMs?: number }>;
|
|
43
43
|
|
|
44
|
-
/**
|
|
45
|
-
* Solves the triangular system op(A) * x = b for x, in place.
|
|
46
|
-
*
|
|
47
|
-
* A is kept GPU-resident; x is a CPU Float32Array. `A`'s own `layout` (set at
|
|
48
|
-
* `GpuMatrix.from` time) determines the operation — there is no separate
|
|
49
|
-
* `layout` argument here.
|
|
50
|
-
*
|
|
51
|
-
* @param device - GPUDevice from `init()`
|
|
52
|
-
* @param uplo - `'lower'` to use the lower triangle, `'upper'` to use the upper triangle
|
|
53
|
-
* @param trans - `'no-transpose'` to solve A*x=b, `'transpose'` to solve A^T*x=b
|
|
54
|
-
* @param diag - `'unit'` to treat the diagonal as all-ones (A's diagonal is not read), `'non-unit'` to read it
|
|
55
|
-
* @param n - order of the matrix A
|
|
56
|
-
* @param A - GpuMatrix, GPU-resident
|
|
57
|
-
* @param lda - leading dimension of A (must equal A.lda)
|
|
58
|
-
* @param x - Float32Array holding b on input, the solution on output
|
|
59
|
-
* @param incx - stride for x (must be a positive integer)
|
|
60
|
-
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/strsv/strsv.mjs#L15">Source code: strsv.mjs (L15)</a>
|
|
61
|
-
* @category BLAS Level 2
|
|
62
|
-
*/
|
|
63
|
-
export declare function strsv(
|
|
64
|
-
device: GPUDevice,
|
|
65
|
-
uplo: 'lower' | 'upper',
|
|
66
|
-
trans: 'no-transpose' | 'transpose',
|
|
67
|
-
diag: 'unit' | 'non-unit',
|
|
68
|
-
n: number,
|
|
69
|
-
A: GpuMatrix,
|
|
70
|
-
lda: number,
|
|
71
|
-
x: Float32Array,
|
|
72
|
-
incx: number,
|
|
73
|
-
): Promise<{ x: Float32Array; gpuTimeMs?: number }>;
|
|
74
|
-
|
|
75
44
|
/**
|
|
76
45
|
* Solves the triangular system op(A) * x = b for x, in place.
|
|
77
46
|
*
|
|
@@ -79,7 +48,7 @@ export declare function strsv(
|
|
|
79
48
|
* its own `layout` (set at `GpuMatrix.from` time) determines the operation —
|
|
80
49
|
* there is no separate `layout` argument here.
|
|
81
50
|
*
|
|
82
|
-
* {@includeCode ../../examples/strsv/
|
|
51
|
+
* {@includeCode ../../examples/strsv/gpu.strsv.js}
|
|
83
52
|
*
|
|
84
53
|
* @param device - GPUDevice from `init()`
|
|
85
54
|
* @param uplo - `'lower'` to use the lower triangle, `'upper'` to use the upper triangle
|
package/src/strsv/strsv.mjs
CHANGED
|
@@ -63,6 +63,8 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
63
63
|
throw new Error("x must be a Float32Array or GpuVector.");
|
|
64
64
|
if (xIsGpu && !AIsGpu)
|
|
65
65
|
throw new Error("A must be a GpuMatrix when x is a GpuVector.");
|
|
66
|
+
if (AIsGpu && !xIsGpu)
|
|
67
|
+
throw new Error("x must be a GpuVector when A is a GpuMatrix.");
|
|
66
68
|
if (AIsGpu && lda !== A.lda)
|
|
67
69
|
throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
68
70
|
if (AIsGpu && (A.rows < n || A.cols < n))
|
package/src/util/buffer.mjs
CHANGED
|
@@ -62,15 +62,17 @@ export function uploadBuffer(data, label = "blas-input", readback = false) {
|
|
|
62
62
|
* that are written by a shader before being read.
|
|
63
63
|
* @param {number} size - byte size
|
|
64
64
|
* @param {string} [label] - debug label visible in browser DevTools GPU inspection
|
|
65
|
+
* @param {number} [extraUsage=0] - additional `GPUBufferUsage` flags OR'd in alongside `STORAGE`
|
|
66
|
+
* (e.g. `GPUBufferUsage.COPY_DST` for a buffer that's also a `copyBufferToBuffer` destination)
|
|
65
67
|
* @returns {GPUBuffer}
|
|
66
68
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createBuffer GPUDevice.createBuffer()}
|
|
67
69
|
*/
|
|
68
|
-
export function createStorageBuffer(size, label = "blas-storage") {
|
|
70
|
+
export function createStorageBuffer(size, label = "blas-storage", extraUsage = 0) {
|
|
69
71
|
const device = getDevice();
|
|
70
72
|
return device.createBuffer({
|
|
71
73
|
label,
|
|
72
74
|
size,
|
|
73
|
-
usage: GPUBufferUsage.STORAGE,
|
|
75
|
+
usage: GPUBufferUsage.STORAGE | extraUsage,
|
|
74
76
|
});
|
|
75
77
|
}
|
|
76
78
|
|
package/src/util/compute.mjs
CHANGED
|
@@ -38,7 +38,8 @@ export function beginTimedEncoder() {
|
|
|
38
38
|
* @param {GPUCommandEncoder} commandEncoder
|
|
39
39
|
* @param {GPUComputePipeline} pipeline
|
|
40
40
|
* @param {GPUBindGroup} bindGroup
|
|
41
|
-
* @param {number | { x: number, y: number }} workgroups - workgroup count;
|
|
41
|
+
* @param {number | { x: number, y: number, z?: number }} workgroups - workgroup count;
|
|
42
|
+
* number for 1D dispatch, `{x, y}` for 2D, `{x, y, z}` for 3D (z defaults to 1)
|
|
42
43
|
* @param {object} [passDescriptor] - passed to `beginComputePass`, e.g. for timestamp writes
|
|
43
44
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/beginComputePass GPUCommandEncoder.beginComputePass()}
|
|
44
45
|
*/
|
|
@@ -51,7 +52,8 @@ export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, pass
|
|
|
51
52
|
if (typeof workgroups === "number") {
|
|
52
53
|
passEncoder.dispatchWorkgroups(workgroups);
|
|
53
54
|
} else {
|
|
54
|
-
|
|
55
|
+
// `?? 1` is load-bearing — confirmed passing undefined (from an {x,y}-only caller) crashes the process, not just no-ops.
|
|
56
|
+
passEncoder.dispatchWorkgroups(workgroups.x, workgroups.y, workgroups.z ?? 1);
|
|
55
57
|
}
|
|
56
58
|
|
|
57
59
|
passEncoder.end();
|
|
@@ -64,7 +66,8 @@ export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, pass
|
|
|
64
66
|
* dispatches workgroups, and optionally wraps the pass in GPU timestamp queries.
|
|
65
67
|
* @param {GPUComputePipeline} pipeline
|
|
66
68
|
* @param {GPUBindGroup} bindGroup
|
|
67
|
-
* @param {number | { x: number, y: number }} workgroups - workgroup count;
|
|
69
|
+
* @param {number | { x: number, y: number, z?: number }} workgroups - workgroup count;
|
|
70
|
+
* number for 1D dispatch, `{x, y}` for 2D, `{x, y, z}` for 3D (z defaults to 1)
|
|
68
71
|
* @returns {{ commandEncoder: GPUCommandEncoder, ts: any }} encoded commands and timestamp handle
|
|
69
72
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createCommandEncoder GPUDevice.createCommandEncoder()}
|
|
70
73
|
*/
|
package/src/util/f64.mjs
CHANGED
|
@@ -2,9 +2,9 @@
|
|
|
2
2
|
|
|
3
3
|
// Double-double f64 emulation — splits a double into a (hi, lo) pair of f32
|
|
4
4
|
// values with hi+lo approximating the original, hi holding the leading bits
|
|
5
|
-
// and lo the rounding error hi lost on its own. See
|
|
6
|
-
//
|
|
7
|
-
// (Dekker's algorithm).
|
|
5
|
+
// and lo the rounding error hi lost on its own. See src/shaders/f64/ (the DD
|
|
6
|
+
// struct in dekker.wgsl, operations in utils/) for the GPU-side arithmetic
|
|
7
|
+
// this pairs with (Dekker's algorithm).
|
|
8
8
|
//
|
|
9
9
|
// Not a value-preserving exact split: double-double buys roughly 2x f32's
|
|
10
10
|
// mantissa (~48 bits vs f32's 24), less than real f64's 52-bit mantissa.
|