wgblas 1.2.1 → 2.1.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 +54 -72
- package/dist/wgblas.browser.js +1637 -858
- package/index.d.mts +51 -6
- package/index.mjs +8 -0
- package/package.json +56 -2
- package/src/classes/GpuMatrix.mjs +17 -10
- package/src/classes/GpuVector.mjs +28 -10
- package/src/dasum/dasum.d.mts +6 -6
- package/src/dasum/dasum.mjs +25 -21
- package/src/devdocs.mjs +13 -0
- package/src/idamax/idamax.d.mts +69 -0
- package/src/idamax/idamax.mjs +130 -0
- package/src/init.mjs +115 -49
- package/src/isamax/isamax.d.mts +21 -3
- package/src/isamax/isamax.mjs +16 -14
- package/src/random/random.d.mts +1 -0
- package/src/sasum/sasum.d.mts +3 -3
- package/src/sasum/sasum.mjs +15 -13
- package/src/saxpy/saxpy.d.mts +3 -3
- package/src/saxpy/saxpy.mjs +10 -8
- package/src/scopy/scopy.d.mts +3 -3
- package/src/scopy/scopy.mjs +10 -8
- package/src/sdot/sdot.d.mts +3 -3
- package/src/sdot/sdot.mjs +16 -14
- package/src/sgemm/sgemm.d.mts +102 -0
- package/src/sgemm/sgemm.mjs +208 -0
- package/src/sgemmtr/sgemmtr.d.mts +104 -0
- package/src/sgemmtr/sgemmtr.mjs +204 -0
- package/src/sgemv/sgemv.d.mts +1 -38
- package/src/sgemv/sgemv.mjs +42 -26
- package/src/sger/sger.d.mts +1 -34
- package/src/sger/sger.mjs +12 -8
- package/src/shaders/block_transfer.wgsl +42 -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/index.mjs +164 -14
- package/src/shaders/reduction/argmaxF64.wgsl +50 -0
- package/src/shaders/reduction/scaledSum.wgsl +65 -0
- package/src/shaders/reduction/sumF64.wgsl +2 -2
- package/src/shaders/sgemm_large.wgsl +206 -0
- package/src/shaders/sgemm_small.wgsl +212 -0
- package/src/shaders/sgemmtr_large.wgsl +120 -0
- package/src/shaders/sgemmtr_small.wgsl +113 -0
- 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/shaders/symmetrize.wgsl +31 -0
- package/src/shaders/triangularize.wgsl +44 -0
- package/src/snrm2/snrm2.d.mts +3 -3
- package/src/snrm2/snrm2.mjs +33 -21
- package/src/srot/srot.d.mts +3 -5
- package/src/srot/srot.mjs +11 -9
- package/src/srotm/srotm.d.mts +3 -5
- package/src/srotm/srotm.mjs +19 -10
- package/src/sscal/sscal.d.mts +4 -4
- package/src/sscal/sscal.mjs +12 -10
- package/src/sswap/sswap.d.mts +3 -3
- package/src/sswap/sswap.mjs +11 -9
- package/src/ssymm/ssymm.d.mts +103 -0
- package/src/ssymm/ssymm.mjs +218 -0
- package/src/ssymv/ssymv.d.mts +1 -36
- package/src/ssymv/ssymv.mjs +12 -8
- package/src/ssyr/ssyr.d.mts +1 -30
- package/src/ssyr/ssyr.mjs +11 -7
- package/src/ssyr2/ssyr2.d.mts +1 -34
- package/src/ssyr2/ssyr2.mjs +12 -8
- package/src/ssyr2k/ssyr2k.d.mts +100 -0
- package/src/ssyr2k/ssyr2k.mjs +202 -0
- package/src/ssyrk/ssyrk.d.mts +90 -0
- package/src/ssyrk/ssyrk.mjs +177 -0
- package/src/strmm/strmm.d.mts +100 -0
- package/src/strmm/strmm.mjs +226 -0
- package/src/strmv/strmv.d.mts +1 -36
- package/src/strmv/strmv.mjs +12 -8
- package/src/strsm/strsm.d.mts +99 -0
- package/src/strsm/strsm.mjs +360 -0
- package/src/strsv/strsv.d.mts +1 -32
- package/src/strsv/strsv.mjs +18 -12
- package/src/util/benchmark.mjs +4 -6
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +116 -20
- package/src/util/compute.mjs +12 -12
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +34 -0
- package/src/util/f64.mjs +3 -3
- package/src/util/pipeline.mjs +5 -6
- package/src/util/workgroup.mjs +55 -7
- package/src/shaders/browser-shaders.mjs +0 -55
|
@@ -0,0 +1,360 @@
|
|
|
1
|
+
import {
|
|
2
|
+
uploadBuffer,
|
|
3
|
+
createParamsBuffer,
|
|
4
|
+
createStorageBuffer,
|
|
5
|
+
stageReadback,
|
|
6
|
+
destroyBuffers,
|
|
7
|
+
vec4ViewBinding,
|
|
8
|
+
} from "../util/buffer.mjs";
|
|
9
|
+
import { createBindGroup } from "../util/bindgroup.mjs";
|
|
10
|
+
import { beginTimedEncoder, encodePass, submit } from "../util/compute.mjs";
|
|
11
|
+
import { extractResult } from "../util/result.mjs";
|
|
12
|
+
import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
|
|
13
|
+
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
|
+
import { calcWorkgroups, requireWorkgroups, requireWorkgroupCount } from "../util/workgroup.mjs";
|
|
15
|
+
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
16
|
+
import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
|
|
17
|
+
import { BLOCK_SIZE } from "../util/constants.mjs";
|
|
18
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
19
|
+
|
|
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
|
+
requireSameDevice(device, "strsm", { A, B });
|
|
35
|
+
if (side !== "left" && side !== "right")
|
|
36
|
+
throw new Error("side must be 'left' or 'right'.");
|
|
37
|
+
if (uplo !== "lower" && uplo !== "upper")
|
|
38
|
+
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
39
|
+
if (transA !== "no-transpose" && transA !== "transpose")
|
|
40
|
+
throw new Error("transA must be 'no-transpose' or 'transpose'.");
|
|
41
|
+
if (!isUnit && diag !== "non-unit")
|
|
42
|
+
throw new Error("diag must be 'unit' or 'non-unit'.");
|
|
43
|
+
if (layout !== "row-major" && layout !== "column-major")
|
|
44
|
+
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.");
|
|
47
|
+
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
48
|
+
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
49
|
+
if (!Number.isInteger(m) || !Number.isInteger(n) || !Number.isInteger(lda) || !Number.isInteger(ldb))
|
|
50
|
+
throw new Error("m, n, lda, and ldb must be integers.");
|
|
51
|
+
if (!AIsGpu && !(A instanceof Float32Array))
|
|
52
|
+
throw new Error("A must be a Float32Array or GpuMatrix.");
|
|
53
|
+
if (!BIsGpu && !(B instanceof Float32Array))
|
|
54
|
+
throw new Error("B must be a Float32Array or GpuMatrix.");
|
|
55
|
+
if (AIsGpu !== BIsGpu)
|
|
56
|
+
throw new Error("A and B must both be GpuMatrix or both be Float32Array.");
|
|
57
|
+
if (m < 0 || n < 0) throw new Error("m and n must be non-negative.");
|
|
58
|
+
if (m === 0 || n === 0) return BIsGpu ? {} : { B };
|
|
59
|
+
|
|
60
|
+
const effLayoutA = AIsGpu ? A.layout : layout;
|
|
61
|
+
const effLayoutB = BIsGpu ? B.layout : layout;
|
|
62
|
+
|
|
63
|
+
// A: triangular, order = m (side='left') or n (side='right').
|
|
64
|
+
const aOrder = side === "left" ? m : n;
|
|
65
|
+
if (lda < aOrder) throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
|
|
66
|
+
if (AIsGpu) {
|
|
67
|
+
if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
68
|
+
if (A.rows < aOrder || A.cols < aOrder) throw new Error("A is too small for the given m/n and side.");
|
|
69
|
+
} else if (A.length < (aOrder - 1) * lda + aOrder) {
|
|
70
|
+
throw new Error("A does not have enough elements for the given dimensions and lda.");
|
|
71
|
+
}
|
|
72
|
+
|
|
73
|
+
// B: always m x n, overwritten in place with the same ldb.
|
|
74
|
+
const bOuter = effLayoutB === "column-major" ? n : m;
|
|
75
|
+
const bInner = effLayoutB === "column-major" ? m : n;
|
|
76
|
+
if (ldb < bInner)
|
|
77
|
+
throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
|
|
78
|
+
if (BIsGpu) {
|
|
79
|
+
if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
|
|
80
|
+
if (B.rows < m || B.cols < n) throw new Error("B is too small for the given m and n.");
|
|
81
|
+
} else if (B.length < (bOuter - 1) * ldb + bInner) {
|
|
82
|
+
throw new Error("B does not have enough elements for the given dimensions and ldb.");
|
|
83
|
+
}
|
|
84
|
+
|
|
85
|
+
// A isn't symmetric: column-major = genuine transpose, so flip transA;
|
|
86
|
+
// transposing also swaps which triangle looks stored, so flip uplo too.
|
|
87
|
+
const uploEffA = effLayoutA === "column-major" ? (uplo === "lower" ? "upper" : "lower") : uplo;
|
|
88
|
+
const transEffA = effLayoutA === "column-major" ? (transA === "no-transpose" ? "transpose" : "no-transpose") : transA;
|
|
89
|
+
|
|
90
|
+
const otherLen = side === "left" ? n : m;
|
|
91
|
+
const blockIsRow = side === "left";
|
|
92
|
+
|
|
93
|
+
// Forward iff op(A) is lower (side='left') — side='right' flips this,
|
|
94
|
+
// since it solves via columns instead of rows.
|
|
95
|
+
const opIsLower = (transEffA === "no-transpose") === (uploEffA === "lower");
|
|
96
|
+
const forward = side === "left" ? opIsLower : !opIsLower;
|
|
97
|
+
|
|
98
|
+
const blockStarts = [];
|
|
99
|
+
for (let s = 0; s < aOrder; s += BLOCK_SIZE) blockStarts.push(s);
|
|
100
|
+
if (!forward) blockStarts.reverse();
|
|
101
|
+
const numBlocks = blockStarts.length;
|
|
102
|
+
|
|
103
|
+
const invertPipeline = await getPipeline(device, "strsv_invert_block");
|
|
104
|
+
const transferPipeline = await getPipeline(device, "block_transfer");
|
|
105
|
+
const scalarPipeline = await getPipeline(device, "sscal");
|
|
106
|
+
|
|
107
|
+
// Null-init here and allocate inside the try below, so a throw partway
|
|
108
|
+
// through the sequence still reaches finally with every handle visible
|
|
109
|
+
// (strsv.mjs is the reference for this pattern).
|
|
110
|
+
let ABuffer = null;
|
|
111
|
+
let BBuffer = null;
|
|
112
|
+
let AinvBuffer = null;
|
|
113
|
+
|
|
114
|
+
const paramsBuffers = [];
|
|
115
|
+
const scratchBuffers = [];
|
|
116
|
+
function scratch(size, label) {
|
|
117
|
+
const buf = createStorageBuffer(device, size, label);
|
|
118
|
+
scratchBuffers.push(buf);
|
|
119
|
+
return buf;
|
|
120
|
+
}
|
|
121
|
+
function params(entries, label) {
|
|
122
|
+
const buf = createParamsBuffer(device, entries, label);
|
|
123
|
+
paramsBuffers.push(buf);
|
|
124
|
+
return buf;
|
|
125
|
+
}
|
|
126
|
+
|
|
127
|
+
// Minimal valid length for B's own buffer (matches its own validation
|
|
128
|
+
// above) — NOT bOuter*ldb, which can exceed a Float32Array-path buffer's
|
|
129
|
+
// actual allocation (only guaranteed padded up to GpuMatrix's own size).
|
|
130
|
+
const bScaleLen = (bOuter - 1) * ldb + bInner;
|
|
131
|
+
|
|
132
|
+
try {
|
|
133
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strsm-A", false);
|
|
134
|
+
BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strsm-B", true);
|
|
135
|
+
AinvBuffer = createStorageBuffer(device, numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4, "strsm-Ainv");
|
|
136
|
+
|
|
137
|
+
// Pre-scale B by alpha once (reuses sscal, so no per-block alpha handling).
|
|
138
|
+
let preScaleBindGroup = null;
|
|
139
|
+
if (alpha !== 1.0) {
|
|
140
|
+
const scaleParams = params(
|
|
141
|
+
[
|
|
142
|
+
{ value: bScaleLen, type: "u32" },
|
|
143
|
+
{ value: alpha, type: "f32" },
|
|
144
|
+
{ value: 1, type: "u32" },
|
|
145
|
+
],
|
|
146
|
+
"strsm-scale-params",
|
|
147
|
+
);
|
|
148
|
+
preScaleBindGroup = createBindGroup(device, scalarPipeline.getBindGroupLayout(0), [BBuffer, scaleParams]);
|
|
149
|
+
}
|
|
150
|
+
|
|
151
|
+
// Every diagonal block's inverse, fully parallel, one dispatch, unchanged.
|
|
152
|
+
const invertParams = params(
|
|
153
|
+
[
|
|
154
|
+
{ value: aOrder, type: "u32" },
|
|
155
|
+
{ value: lda, type: "u32" },
|
|
156
|
+
{ value: transEffA === "transpose" ? 1 : 0, type: "u32" },
|
|
157
|
+
{ value: uploEffA === "upper" ? 1 : 0, type: "u32" },
|
|
158
|
+
{ value: isUnit ? 1 : 0, type: "u32" },
|
|
159
|
+
],
|
|
160
|
+
"strsm-invert-params",
|
|
161
|
+
);
|
|
162
|
+
const invertBindGroup = createBindGroup(device, invertPipeline.getBindGroupLayout(0), [ABuffer, AinvBuffer, invertParams]);
|
|
163
|
+
|
|
164
|
+
// Reusable scratch buffers, sized for the worst case, bound at offset 0.
|
|
165
|
+
const Bblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Bblock");
|
|
166
|
+
const Xblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Xblock");
|
|
167
|
+
const Aoff = scratch(aOrder * BLOCK_SIZE * 4, "strsm-Aoff");
|
|
168
|
+
const delta = scratch(aOrder * otherLen * 4, "strsm-delta");
|
|
169
|
+
|
|
170
|
+
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
171
|
+
|
|
172
|
+
if (alpha === 0) {
|
|
173
|
+
// BLAS: alpha=0 means A is not referenced — skip straight to B:=0.
|
|
174
|
+
const zeroDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0, endOfPassWriteIndex: 1 } } : undefined;
|
|
175
|
+
encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(device, bScaleLen), zeroDesc);
|
|
176
|
+
} else {
|
|
177
|
+
if (preScaleBindGroup) {
|
|
178
|
+
encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(device, bScaleLen));
|
|
179
|
+
}
|
|
180
|
+
const invertDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
|
|
181
|
+
encodePass(commandEncoder, invertPipeline, invertBindGroup, { x: BLOCK_SIZE, y: numBlocks }, invertDesc);
|
|
182
|
+
|
|
183
|
+
for (let bi = 0; bi < blockStarts.length; bi++) {
|
|
184
|
+
const blockStart = blockStarts[bi];
|
|
185
|
+
const blockEnd = Math.min(blockStart + BLOCK_SIZE, aOrder);
|
|
186
|
+
const blockLen = blockEnd - blockStart;
|
|
187
|
+
const blockIndex = blockStart / BLOCK_SIZE;
|
|
188
|
+
const isLastPass = bi === blockStarts.length - 1;
|
|
189
|
+
|
|
190
|
+
// 1) gather B's current block into a tight scratch buffer.
|
|
191
|
+
const gatherBParams = params(
|
|
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(
|
|
215
|
+
[
|
|
216
|
+
{ value: mg, type: "u32" },
|
|
217
|
+
{ value: ng, type: "u32" },
|
|
218
|
+
{ value: kg, type: "u32" },
|
|
219
|
+
{ value: 1.0, type: "f32" }, // alpha already applied to B up front
|
|
220
|
+
{ value: 0.0, type: "f32" }, // beta — fresh output, no accumulation
|
|
221
|
+
{ value: BLOCK_SIZE, type: "u32" }, // ldX = Ainv's own dense stride
|
|
222
|
+
{ value: otherLen, type: "u32" }, // ldY = Bblock's own tight stride
|
|
223
|
+
{ value: otherLen, type: "u32" }, // ldc = Xblock's own tight stride
|
|
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
|
|
226
|
+
],
|
|
227
|
+
"strsm-apply-params",
|
|
228
|
+
);
|
|
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
|
+
|
|
244
|
+
// 3) scatter the solved block back into B.
|
|
245
|
+
const rangeStart = forward ? blockEnd : 0;
|
|
246
|
+
const rangeEnd = forward ? aOrder : blockStart;
|
|
247
|
+
const hasRemaining = rangeStart < rangeEnd;
|
|
248
|
+
const scatterParams = params(
|
|
249
|
+
[
|
|
250
|
+
{ value: blockStart, type: "u32" },
|
|
251
|
+
{ value: blockLen, type: "u32" },
|
|
252
|
+
{ value: 0, type: "u32" },
|
|
253
|
+
{ value: otherLen, type: "u32" },
|
|
254
|
+
{ value: ldb, type: "u32" },
|
|
255
|
+
{ value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
|
|
256
|
+
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
257
|
+
{ value: 0, type: "u32" }, // overwrite
|
|
258
|
+
],
|
|
259
|
+
"strsm-scatter-params",
|
|
260
|
+
);
|
|
261
|
+
const scatterBindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [Xblock, BBuffer, scatterParams]);
|
|
262
|
+
const scatterDesc = isLastPass && !hasRemaining && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
|
|
263
|
+
encodePass(commandEncoder, transferPipeline, scatterBindGroup, requireWorkgroups(device, "strsm", blockLen, otherLen), scatterDesc);
|
|
264
|
+
|
|
265
|
+
// 4) trailing update: subtract this block's contribution from B.
|
|
266
|
+
if (!hasRemaining) continue;
|
|
267
|
+
const remCount = rangeEnd - rangeStart;
|
|
268
|
+
|
|
269
|
+
const gatherAParams = params(
|
|
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(
|
|
291
|
+
[
|
|
292
|
+
{ value: mg, type: "u32" },
|
|
293
|
+
{ value: ng, type: "u32" },
|
|
294
|
+
{ value: kg, type: "u32" },
|
|
295
|
+
{ value: 1.0, type: "f32" },
|
|
296
|
+
{ value: 0.0, type: "f32" },
|
|
297
|
+
{ value: blockLen, type: "u32" }, // ldX = Aoff's own tight stride
|
|
298
|
+
{ value: otherLen, type: "u32" }, // ldY = Xblock's own tight stride
|
|
299
|
+
{ value: otherLen, type: "u32" }, // ldc = delta's own tight stride
|
|
300
|
+
{ value: 0, type: "u32" }, // transX: Aoff already read in the right orientation
|
|
301
|
+
{ value: 0, type: "u32" }, // transY: Xblock read as-is
|
|
302
|
+
],
|
|
303
|
+
"strsm-update-params",
|
|
304
|
+
);
|
|
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
|
+
|
|
319
|
+
const scatterSubParams = params(
|
|
320
|
+
[
|
|
321
|
+
{ value: rangeStart, type: "u32" },
|
|
322
|
+
{ value: remCount, type: "u32" },
|
|
323
|
+
{ value: 0, type: "u32" },
|
|
324
|
+
{ value: otherLen, type: "u32" },
|
|
325
|
+
{ value: ldb, type: "u32" },
|
|
326
|
+
{ value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
|
|
327
|
+
{ value: blockIsRow ? 1 : 0, type: "u32" },
|
|
328
|
+
{ value: 1, type: "u32" }, // subtract
|
|
329
|
+
],
|
|
330
|
+
"strsm-scatter-sub-params",
|
|
331
|
+
);
|
|
332
|
+
const scatterSubBindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [delta, BBuffer, scatterSubParams]);
|
|
333
|
+
const subDesc = isLastPass && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
|
|
334
|
+
encodePass(commandEncoder, transferPipeline, scatterSubBindGroup, requireWorkgroups(device, "strsm", remCount, otherLen), subDesc);
|
|
335
|
+
}
|
|
336
|
+
}
|
|
337
|
+
|
|
338
|
+
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
339
|
+
const readBuffer = BIsGpu ? null : stageReadback(device, commandEncoder, BBuffer);
|
|
340
|
+
|
|
341
|
+
submit(device, commandEncoder);
|
|
342
|
+
|
|
343
|
+
const gpuTimeMs = await extractTimestamp(ts);
|
|
344
|
+
|
|
345
|
+
if (BIsGpu) {
|
|
346
|
+
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
347
|
+
return {};
|
|
348
|
+
}
|
|
349
|
+
|
|
350
|
+
const result = await extractResult(readBuffer, Float32Array);
|
|
351
|
+
if (gpuTimeMs !== undefined) return { B: result, gpuTimeMs };
|
|
352
|
+
return { B: result };
|
|
353
|
+
} finally {
|
|
354
|
+
if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
|
|
355
|
+
if (!BIsGpu && BBuffer) destroyBuffers(BBuffer);
|
|
356
|
+
if (AinvBuffer) destroyBuffers(AinvBuffer);
|
|
357
|
+
destroyBuffers(scratchBuffers);
|
|
358
|
+
destroyBuffers(paramsBuffers);
|
|
359
|
+
}
|
|
360
|
+
}
|
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
|
@@ -12,9 +12,10 @@ import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
|
|
|
12
12
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
14
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
15
|
+
import { BLOCK_SIZE } from "../util/constants.mjs";
|
|
16
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
15
17
|
|
|
16
18
|
// Blocked triangular solve via explicit block inversion (invert/apply/update passes) instead of barrier-per-row substitution.
|
|
17
|
-
const BLOCK_SIZE = 64;
|
|
18
19
|
|
|
19
20
|
// One shared buffer holds all blocks' params (offset blockIndex*stride) instead of one buffer per block — avoids the O(numBlocks) createBuffer/writeBuffer calls that dominated CPU time.
|
|
20
21
|
function packBlockParams(numBlocks, stride, fieldsPerBlock) {
|
|
@@ -45,6 +46,7 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
45
46
|
|
|
46
47
|
if (!(device instanceof GPUDevice))
|
|
47
48
|
throw new Error("device must be a GPUDevice.");
|
|
49
|
+
requireSameDevice(device, "strsv", { A, x });
|
|
48
50
|
if (uplo !== "lower" && uplo !== "upper")
|
|
49
51
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
50
52
|
if (trans !== "no-transpose" && trans !== "transpose")
|
|
@@ -63,6 +65,10 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
63
65
|
throw new Error("x must be a Float32Array or GpuVector.");
|
|
64
66
|
if (xIsGpu && !AIsGpu)
|
|
65
67
|
throw new Error("A must be a GpuMatrix when x is a GpuVector.");
|
|
68
|
+
if (AIsGpu && !xIsGpu)
|
|
69
|
+
throw new Error("x must be a GpuVector when A is a GpuMatrix.");
|
|
70
|
+
if (AIsGpu && xIsGpu && A._buf === x._buf)
|
|
71
|
+
throw new Error("A and x must not reference the same GPU buffer.");
|
|
66
72
|
if (AIsGpu && lda !== A.lda)
|
|
67
73
|
throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
68
74
|
if (AIsGpu && (A.rows < n || A.cols < n))
|
|
@@ -107,11 +113,11 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
107
113
|
let invertParams = null;
|
|
108
114
|
|
|
109
115
|
try {
|
|
110
|
-
ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "strsv-A", false);
|
|
111
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "strsv-x", true);
|
|
116
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strsv-A", false);
|
|
117
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "strsv-x", true);
|
|
112
118
|
// One BLOCK_SIZE x BLOCK_SIZE dense region per block (row-major), even
|
|
113
119
|
// though only a triangular half is ever nonzero — see strsv_invert_block.wgsl.
|
|
114
|
-
AinvBuffer = createStorageBuffer(
|
|
120
|
+
AinvBuffer = createStorageBuffer(device,
|
|
115
121
|
numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4,
|
|
116
122
|
"strsv-Ainv",
|
|
117
123
|
);
|
|
@@ -133,10 +139,10 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
133
139
|
});
|
|
134
140
|
updateParamsBuffer = createSharedParamsBuffer(device, updateData, "strsv-update-params");
|
|
135
141
|
|
|
136
|
-
const { commandEncoder, querySet } = beginTimedEncoder();
|
|
142
|
+
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
137
143
|
|
|
138
144
|
// Pre-pass: every block's inverse, fully parallel, one dispatch.
|
|
139
|
-
invertParams = createParamsBuffer(
|
|
145
|
+
invertParams = createParamsBuffer(device,
|
|
140
146
|
[
|
|
141
147
|
{ value: n, type: "u32" },
|
|
142
148
|
{ value: lda, type: "u32" },
|
|
@@ -146,7 +152,7 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
146
152
|
],
|
|
147
153
|
"strsv-invert-params",
|
|
148
154
|
);
|
|
149
|
-
const invertBindGroup = createBindGroup(invertPipeline.getBindGroupLayout(0), [
|
|
155
|
+
const invertBindGroup = createBindGroup(device, invertPipeline.getBindGroupLayout(0), [
|
|
150
156
|
ABuffer, AinvBuffer, invertParams,
|
|
151
157
|
]);
|
|
152
158
|
const invertDesc = querySet
|
|
@@ -161,7 +167,7 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
161
167
|
const isLastPass = bi === blockStarts.length - 1;
|
|
162
168
|
const paramsOffset = blockIndex * stride;
|
|
163
169
|
|
|
164
|
-
const applyBindGroup = createBindGroup(applyPipeline.getBindGroupLayout(0), [
|
|
170
|
+
const applyBindGroup = createBindGroup(device, applyPipeline.getBindGroupLayout(0), [
|
|
165
171
|
AinvBuffer, xBuffer, { buffer: applyParamsBuffer, offset: paramsOffset, size: 16 },
|
|
166
172
|
]);
|
|
167
173
|
|
|
@@ -173,7 +179,7 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
173
179
|
const remaining = forward ? n - blockEnd : blockStart;
|
|
174
180
|
if (remaining === 0) continue;
|
|
175
181
|
|
|
176
|
-
const updateBindGroup = createBindGroup(updatePipeline.getBindGroupLayout(0), [
|
|
182
|
+
const updateBindGroup = createBindGroup(device, updatePipeline.getBindGroupLayout(0), [
|
|
177
183
|
ABuffer, xBuffer, { buffer: updateParamsBuffer, offset: paramsOffset, size: 32 },
|
|
178
184
|
]);
|
|
179
185
|
|
|
@@ -181,10 +187,10 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
181
187
|
encodePass(commandEncoder, updatePipeline, updateBindGroup, wgCount);
|
|
182
188
|
}
|
|
183
189
|
|
|
184
|
-
const ts = resolveTimestamp(commandEncoder, querySet);
|
|
185
|
-
const readBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
190
|
+
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
191
|
+
const readBuffer = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
|
|
186
192
|
|
|
187
|
-
submit(commandEncoder);
|
|
193
|
+
submit(device, commandEncoder);
|
|
188
194
|
|
|
189
195
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
190
196
|
|
package/src/util/benchmark.mjs
CHANGED
|
@@ -1,5 +1,5 @@
|
|
|
1
1
|
/** @module devdocs/utility-functions/benchmark */
|
|
2
|
-
import {
|
|
2
|
+
import { isBenchmarkEnabled } from "../init.mjs";
|
|
3
3
|
|
|
4
4
|
/**
|
|
5
5
|
* Returns the `requestDevice` descriptor to pass to `adapter.requestDevice()`.
|
|
@@ -29,10 +29,9 @@ export function benchmarkMode(adapter, enabled) {
|
|
|
29
29
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUQuerySet GPUQuerySet}
|
|
30
30
|
* @see {@link https://developer.chrome.com/blog/new-in-webgpu-121 Chrome 121 — timestamp queries (querySet + timestampWrites pattern, quantization caveat)}
|
|
31
31
|
*/
|
|
32
|
-
export function beginTimestamp() {
|
|
33
|
-
if (!isBenchmarkEnabled())
|
|
32
|
+
export function beginTimestamp(device) {
|
|
33
|
+
if (!isBenchmarkEnabled(device))
|
|
34
34
|
return { querySet: null, passDescriptor: undefined };
|
|
35
|
-
const device = getDevice();
|
|
36
35
|
// Two slots: index 0 written when the pass begins, index 1 when it ends.
|
|
37
36
|
const querySet = device.createQuerySet({ type: "timestamp", count: 2 });
|
|
38
37
|
const passDescriptor = {
|
|
@@ -55,9 +54,8 @@ export function beginTimestamp() {
|
|
|
55
54
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/resolveQuerySet GPUCommandEncoder.resolveQuerySet()}
|
|
56
55
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/copyBufferToBuffer GPUCommandEncoder.copyBufferToBuffer()}
|
|
57
56
|
*/
|
|
58
|
-
export function resolveTimestamp(commandEncoder, querySet) {
|
|
57
|
+
export function resolveTimestamp(device, commandEncoder, querySet) {
|
|
59
58
|
if (!querySet) return null;
|
|
60
|
-
const device = getDevice();
|
|
61
59
|
// QUERY_RESOLVE and MAP_READ cannot be combined — two buffers are required.
|
|
62
60
|
// resolveBuffer: GPU writes resolved nanosecond timestamps here.
|
|
63
61
|
const resolveBuffer = device.createBuffer({
|
package/src/util/bindgroup.mjs
CHANGED
|
@@ -1,5 +1,4 @@
|
|
|
1
1
|
/** @module devdocs/utility-functions/bindgroup */
|
|
2
|
-
import { getDevice } from "../init.mjs";
|
|
3
2
|
|
|
4
3
|
/**
|
|
5
4
|
* Creates a `GPUBindGroup` by mapping each buffer to sequential binding indices
|
|
@@ -16,8 +15,7 @@ import { getDevice } from "../init.mjs";
|
|
|
16
15
|
* @returns {GPUBindGroup}
|
|
17
16
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createBindGroup GPUDevice.createBindGroup()}
|
|
18
17
|
*/
|
|
19
|
-
export function createBindGroup(layout, buffers, startBinding = 0) {
|
|
20
|
-
const device = getDevice();
|
|
18
|
+
export function createBindGroup(device, layout, buffers, startBinding = 0) {
|
|
21
19
|
const entries = buffers.map((buffer, i) => ({
|
|
22
20
|
binding: startBinding + i,
|
|
23
21
|
resource: buffer instanceof GPUBuffer ? { buffer } : buffer,
|