wgblas 2.0.0 → 2.2.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +20 -18
- package/dist/wgblas.browser.js +2172 -1174
- package/index.d.mts +49 -44
- package/index.mjs +11 -0
- package/package.json +133 -63
- package/src/classes/Complex32.d.mts +43 -0
- package/src/classes/Complex32.mjs +82 -0
- package/src/classes/Complex64.d.mts +44 -0
- package/src/classes/Complex64.mjs +76 -0
- package/src/classes/GpuMatrix.d.mts +41 -41
- package/src/classes/GpuMatrix.mjs +126 -17
- package/src/classes/GpuVector.d.mts +36 -40
- package/src/classes/GpuVector.mjs +66 -11
- package/src/cscal/cscal.d.mts +47 -0
- package/src/cscal/cscal.mjs +98 -0
- package/src/dasum/dasum.d.mts +4 -4
- package/src/dasum/dasum.mjs +38 -20
- package/src/daxpy/daxpy.d.mts +56 -0
- package/src/daxpy/daxpy.mjs +150 -0
- package/src/dcopy/dcopy.d.mts +52 -0
- package/src/dcopy/dcopy.mjs +140 -0
- package/src/ddot/ddot.d.mts +62 -0
- package/src/ddot/ddot.mjs +184 -0
- package/src/devdocs.mjs +13 -0
- package/src/dnrm2/dnrm2.d.mts +50 -0
- package/src/dnrm2/dnrm2.mjs +189 -0
- package/src/drot/drot.d.mts +67 -0
- package/src/drot/drot.mjs +170 -0
- package/src/drotm/drotm.d.mts +67 -0
- package/src/drotm/drotm.mjs +171 -0
- package/src/dscal/dscal.d.mts +52 -0
- package/src/dscal/dscal.mjs +119 -0
- package/src/dswap/dswap.d.mts +57 -0
- package/src/dswap/dswap.mjs +155 -0
- package/src/idamax/idamax.d.mts +20 -2
- package/src/idamax/idamax.mjs +56 -24
- package/src/init.mjs +117 -56
- package/src/isamax/isamax.d.mts +20 -2
- package/src/isamax/isamax.mjs +21 -16
- package/src/random/random.d.mts +37 -39
- package/src/random/random.mjs +39 -7
- package/src/sasum/sasum.d.mts +2 -2
- package/src/sasum/sasum.mjs +20 -16
- package/src/saxpy/saxpy.d.mts +2 -2
- package/src/saxpy/saxpy.mjs +14 -11
- package/src/scopy/scopy.d.mts +2 -2
- package/src/scopy/scopy.mjs +13 -9
- package/src/sdot/sdot.d.mts +2 -2
- package/src/sdot/sdot.mjs +21 -17
- package/src/sgemm/sgemm.d.mts +2 -2
- package/src/sgemm/sgemm.mjs +109 -40
- package/src/sgemmtr/sgemmtr.d.mts +3 -2
- package/src/sgemmtr/sgemmtr.mjs +98 -40
- package/src/sgemv/sgemv.d.mts +2 -2
- package/src/sgemv/sgemv.mjs +69 -41
- package/src/sger/sger.d.mts +2 -2
- package/src/sger/sger.mjs +43 -19
- package/src/shaders/__test_pipeline_a.wgsl +3 -0
- package/src/shaders/__test_pipeline_b.wgsl +2 -0
- package/src/shaders/cscal.wgsl +33 -0
- package/src/shaders/daxpy.wgsl +66 -0
- package/src/shaders/dcopy.wgsl +34 -0
- package/src/shaders/ddot.wgsl +106 -0
- package/src/shaders/dnrm2.wgsl +167 -0
- package/src/shaders/drot.wgsl +81 -0
- package/src/shaders/drotm.wgsl +99 -0
- package/src/shaders/dscal.wgsl +60 -0
- package/src/shaders/dswap.wgsl +38 -0
- package/src/shaders/f64/utils/add.wgsl +6 -0
- package/src/shaders/f64/utils/divide.wgsl +45 -0
- package/src/shaders/f64/utils/multiply.wgsl +19 -10
- package/src/shaders/f64/utils/sqrt.wgsl +45 -0
- package/src/shaders/index.mjs +233 -14
- package/src/shaders/reduction/scaledSum.wgsl +65 -0
- package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
- package/src/shaders/sgemm_large.wgsl +107 -18
- package/src/shaders/sgemm_small.wgsl +115 -15
- package/src/shaders/sgemmtr_large.wgsl +4 -1
- package/src/shaders/sgemmtr_small.wgsl +4 -1
- package/src/shaders/sgemv_n.wgsl +3 -1
- package/src/shaders/sgemv_t.wgsl +3 -1
- package/src/shaders/snrm2.wgsl +72 -23
- package/src/shaders/ssymv.wgsl +3 -1
- package/src/snrm2/snrm2.d.mts +2 -2
- package/src/snrm2/snrm2.mjs +41 -23
- package/src/srot/srot.d.mts +2 -4
- package/src/srot/srot.mjs +16 -11
- package/src/srotm/srotm.d.mts +2 -4
- package/src/srotm/srotm.mjs +17 -11
- package/src/sscal/sscal.d.mts +3 -3
- package/src/sscal/sscal.mjs +14 -12
- package/src/sswap/sswap.d.mts +2 -2
- package/src/sswap/sswap.mjs +18 -10
- package/src/ssymm/ssymm.d.mts +5 -4
- package/src/ssymm/ssymm.mjs +150 -54
- package/src/ssymv/ssymv.d.mts +2 -2
- package/src/ssymv/ssymv.mjs +47 -26
- package/src/ssyr/ssyr.d.mts +2 -2
- package/src/ssyr/ssyr.mjs +38 -17
- package/src/ssyr2/ssyr2.d.mts +2 -2
- package/src/ssyr2/ssyr2.mjs +48 -21
- package/src/ssyr2k/ssyr2k.d.mts +3 -2
- package/src/ssyr2k/ssyr2k.mjs +140 -62
- package/src/ssyrk/ssyrk.d.mts +3 -2
- package/src/ssyrk/ssyrk.mjs +91 -39
- package/src/strmm/strmm.d.mts +5 -4
- package/src/strmm/strmm.mjs +174 -60
- package/src/strmv/strmv.d.mts +2 -2
- package/src/strmv/strmv.mjs +42 -20
- package/src/strsm/strsm.d.mts +6 -4
- package/src/strsm/strsm.mjs +438 -174
- package/src/strsv/strsv.d.mts +5 -3
- package/src/strsv/strsv.mjs +89 -34
- package/src/util/benchmark.mjs +9 -9
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +139 -24
- package/src/util/complex.mjs +87 -0
- package/src/util/compute.mjs +19 -16
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +49 -0
- package/src/util/pipeline.mjs +44 -10
- package/src/util/workgroup.mjs +72 -7
- package/src/shaders/browser-shaders.mjs +0 -81
- package/src/shaders/f64add.wgsl +0 -281
package/src/ssyr2k/ssyr2k.mjs
CHANGED
|
@@ -10,39 +10,58 @@ import { extractResult } from "../util/result.mjs";
|
|
|
10
10
|
import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
13
|
-
|
|
14
|
-
|
|
15
|
-
|
|
16
|
-
|
|
13
|
+
import { requireWorkgroupCount } from "../util/workgroup.mjs";
|
|
14
|
+
import {
|
|
15
|
+
BM_SMALL,
|
|
16
|
+
BN_SMALL,
|
|
17
|
+
BM_LARGE,
|
|
18
|
+
BN_LARGE,
|
|
19
|
+
LARGE_TILE_WORKGROUP_THRESHOLD,
|
|
20
|
+
} from "../util/constants.mjs";
|
|
21
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
17
22
|
|
|
18
23
|
// ssyr2k: C := uplo(alpha*op(A)*op(B)^T + alpha*op(B)*op(A)^T + beta*C). No
|
|
19
24
|
// dedicated shader — two sgemmtr passes on one encoder, second with beta=1.
|
|
20
25
|
export async function ssyr2k(
|
|
21
|
-
device,
|
|
26
|
+
device,
|
|
27
|
+
uplo,
|
|
28
|
+
trans,
|
|
29
|
+
n,
|
|
30
|
+
k,
|
|
31
|
+
alpha,
|
|
32
|
+
A,
|
|
33
|
+
lda,
|
|
34
|
+
B,
|
|
35
|
+
ldb,
|
|
36
|
+
beta,
|
|
37
|
+
C,
|
|
38
|
+
ldc,
|
|
39
|
+
layout = "row-major",
|
|
22
40
|
) {
|
|
23
41
|
const AIsGpu = A instanceof GpuMatrix;
|
|
24
42
|
const BIsGpu = B instanceof GpuMatrix;
|
|
25
43
|
const CIsGpu = C instanceof GpuMatrix;
|
|
26
44
|
|
|
27
|
-
|
|
28
|
-
|
|
45
|
+
requireGpuDevice(device);
|
|
46
|
+
requireSameDevice(device, "ssyr2k", { A, B, C });
|
|
29
47
|
if (uplo !== "lower" && uplo !== "upper")
|
|
30
48
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
31
49
|
if (trans !== "no-transpose" && trans !== "transpose")
|
|
32
50
|
throw new Error("trans must be 'no-transpose' or 'transpose'.");
|
|
33
51
|
if (layout !== "row-major" && layout !== "column-major")
|
|
34
52
|
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
35
|
-
if (typeof alpha !== "number")
|
|
36
|
-
throw new Error("alpha must be a number.");
|
|
53
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
37
54
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
38
55
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
39
|
-
if (typeof beta !== "number")
|
|
40
|
-
throw new Error("beta must be a number.");
|
|
56
|
+
if (typeof beta !== "number") throw new Error("beta must be a number.");
|
|
41
57
|
if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
|
|
42
58
|
if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
|
|
43
59
|
if (
|
|
44
|
-
!Number.isInteger(n) ||
|
|
45
|
-
!Number.isInteger(
|
|
60
|
+
!Number.isInteger(n) ||
|
|
61
|
+
!Number.isInteger(k) ||
|
|
62
|
+
!Number.isInteger(lda) ||
|
|
63
|
+
!Number.isInteger(ldb) ||
|
|
64
|
+
!Number.isInteger(ldc)
|
|
46
65
|
)
|
|
47
66
|
throw new Error("n, k, lda, ldb, and ldc must be integers.");
|
|
48
67
|
if (!AIsGpu && !(A instanceof Float32Array))
|
|
@@ -56,6 +75,8 @@ export async function ssyr2k(
|
|
|
56
75
|
if (CIsGpu && (!AIsGpu || !BIsGpu))
|
|
57
76
|
throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");
|
|
58
77
|
if (n < 0 || k < 0) throw new Error("n and k must be non-negative.");
|
|
78
|
+
if (lda <= 0 || ldb <= 0 || ldc <= 0)
|
|
79
|
+
throw new Error("lda, ldb, and ldc must be positive.");
|
|
59
80
|
if (n === 0) return CIsGpu ? {} : { C };
|
|
60
81
|
|
|
61
82
|
const effLayoutA = AIsGpu ? A.layout : layout;
|
|
@@ -68,14 +89,19 @@ export async function ssyr2k(
|
|
|
68
89
|
const aOuter = trans === "no-transpose" ? aRows : aCols;
|
|
69
90
|
const aInner = trans === "no-transpose" ? aCols : aRows;
|
|
70
91
|
if (lda < aInner)
|
|
71
|
-
throw new Error(
|
|
92
|
+
throw new Error(
|
|
93
|
+
`lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`,
|
|
94
|
+
);
|
|
72
95
|
if (AIsGpu) {
|
|
73
|
-
if (lda !== A.lda)
|
|
96
|
+
if (lda !== A.lda)
|
|
97
|
+
throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
74
98
|
const [aLogRows, aLogCols] = trans === "no-transpose" ? [n, k] : [k, n];
|
|
75
99
|
if (A.rows < aLogRows || A.cols < aLogCols)
|
|
76
100
|
throw new Error("A is too small for the given n, k, and trans.");
|
|
77
101
|
} else if (A.length < (aOuter - 1) * lda + aInner) {
|
|
78
|
-
throw new Error(
|
|
102
|
+
throw new Error(
|
|
103
|
+
"A does not have enough elements for the given dimensions and lda.",
|
|
104
|
+
);
|
|
79
105
|
}
|
|
80
106
|
|
|
81
107
|
// B: same shape rule as A — netlib's syr2k shares one trans across both operands.
|
|
@@ -84,23 +110,32 @@ export async function ssyr2k(
|
|
|
84
110
|
const bOuter = trans === "no-transpose" ? bRows : bCols;
|
|
85
111
|
const bInner = trans === "no-transpose" ? bCols : bRows;
|
|
86
112
|
if (ldb < bInner)
|
|
87
|
-
throw new Error(
|
|
113
|
+
throw new Error(
|
|
114
|
+
`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
|
|
115
|
+
);
|
|
88
116
|
if (BIsGpu) {
|
|
89
|
-
if (ldb !== B.lda)
|
|
117
|
+
if (ldb !== B.lda)
|
|
118
|
+
throw new Error("ldb must match B.lda when B is a GpuMatrix.");
|
|
90
119
|
const [bLogRows, bLogCols] = trans === "no-transpose" ? [n, k] : [k, n];
|
|
91
120
|
if (B.rows < bLogRows || B.cols < bLogCols)
|
|
92
121
|
throw new Error("B is too small for the given n, k, and trans.");
|
|
93
122
|
} else if (B.length < (bOuter - 1) * ldb + bInner) {
|
|
94
|
-
throw new Error(
|
|
123
|
+
throw new Error(
|
|
124
|
+
"B does not have enough elements for the given dimensions and ldb.",
|
|
125
|
+
);
|
|
95
126
|
}
|
|
96
127
|
|
|
97
128
|
// C: always n x n symmetric — layout only affects storage order, not size.
|
|
98
129
|
if (ldc < n) throw new Error("ldc must be >= n.");
|
|
99
130
|
if (CIsGpu) {
|
|
100
|
-
if (ldc !== C.lda)
|
|
101
|
-
|
|
131
|
+
if (ldc !== C.lda)
|
|
132
|
+
throw new Error("ldc must match C.lda when C is a GpuMatrix.");
|
|
133
|
+
if (C.rows < n || C.cols < n)
|
|
134
|
+
throw new Error("C is too small for the given n.");
|
|
102
135
|
} else if (C.length < (n - 1) * ldc + n) {
|
|
103
|
-
throw new Error(
|
|
136
|
+
throw new Error(
|
|
137
|
+
"C does not have enough elements for the given dimensions and ldc.",
|
|
138
|
+
);
|
|
104
139
|
}
|
|
105
140
|
|
|
106
141
|
// Column-major A/B reinterpreted row-major is A^T/B^T — flip trans per operand.
|
|
@@ -111,75 +146,118 @@ export async function ssyr2k(
|
|
|
111
146
|
if (effLayoutB === "column-major")
|
|
112
147
|
effTransB = effTransB === "no-transpose" ? "transpose" : "no-transpose";
|
|
113
148
|
|
|
114
|
-
const uploEff =
|
|
115
|
-
|
|
116
|
-
|
|
149
|
+
const uploEff =
|
|
150
|
+
effLayoutC === "column-major"
|
|
151
|
+
? uplo === "lower"
|
|
152
|
+
? "upper"
|
|
153
|
+
: "lower"
|
|
154
|
+
: uplo;
|
|
117
155
|
const flip = (t) => (t === "no-transpose" ? "transpose" : "no-transpose");
|
|
118
156
|
|
|
119
157
|
// One pass: op(X) as-is, op(Y) transposed. transOwnY is explicit, not inferred.
|
|
120
158
|
function passShape(transOwnX, X, ldX, transOwnY, Y, ldY) {
|
|
121
159
|
const transX = transOwnX;
|
|
122
160
|
const transY = flip(transOwnY);
|
|
123
|
-
if (effLayoutC !== "column-major")
|
|
124
|
-
|
|
161
|
+
if (effLayoutC !== "column-major")
|
|
162
|
+
return { transX, X, ldX, transY, Y, ldY };
|
|
163
|
+
return {
|
|
164
|
+
transX: flip(transY),
|
|
165
|
+
X: Y,
|
|
166
|
+
ldX: ldY,
|
|
167
|
+
transY: flip(transX),
|
|
168
|
+
Y: X,
|
|
169
|
+
ldY: ldX,
|
|
170
|
+
};
|
|
125
171
|
}
|
|
126
172
|
|
|
127
173
|
// Shape-based auto-select — see sgemmtr_small.wgsl/sgemmtr_large.wgsl. m=n=n here (square C).
|
|
128
174
|
const largeWgX = Math.ceil(n / BN_LARGE);
|
|
129
175
|
const largeWgY = Math.ceil(n / BM_LARGE);
|
|
130
176
|
const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
131
|
-
const pipeline = await getPipeline(
|
|
177
|
+
const pipeline = await getPipeline(
|
|
178
|
+
device,
|
|
179
|
+
useLargeTile ? "sgemmtr_large" : "sgemmtr_small",
|
|
180
|
+
);
|
|
132
181
|
const wgCount = useLargeTile
|
|
133
182
|
? {
|
|
134
|
-
|
|
135
|
-
|
|
136
|
-
|
|
183
|
+
x: requireWorkgroupCount(device, largeWgX, "ssyr2k", "x"),
|
|
184
|
+
y: requireWorkgroupCount(device, largeWgY, "ssyr2k", "y"),
|
|
185
|
+
}
|
|
137
186
|
: {
|
|
138
|
-
|
|
139
|
-
|
|
140
|
-
|
|
187
|
+
x: requireWorkgroupCount(
|
|
188
|
+
device,
|
|
189
|
+
Math.ceil(n / BN_SMALL),
|
|
190
|
+
"ssyr2k",
|
|
191
|
+
"x",
|
|
192
|
+
),
|
|
193
|
+
y: requireWorkgroupCount(
|
|
194
|
+
device,
|
|
195
|
+
Math.ceil(n / BM_SMALL),
|
|
196
|
+
"ssyr2k",
|
|
197
|
+
"y",
|
|
198
|
+
),
|
|
199
|
+
};
|
|
141
200
|
|
|
142
|
-
const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyr2k-A", false);
|
|
143
|
-
const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "ssyr2k-B", false);
|
|
144
|
-
const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "ssyr2k-C", true);
|
|
145
|
-
let paramsBuffer1 = null,
|
|
201
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyr2k-A", false);
|
|
202
|
+
const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "ssyr2k-B", false);
|
|
203
|
+
const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssyr2k-C", true);
|
|
204
|
+
let paramsBuffer1 = null,
|
|
205
|
+
paramsBuffer2 = null;
|
|
146
206
|
|
|
147
207
|
try {
|
|
148
208
|
const pass1 = passShape(effTransA, ABuffer, lda, effTransB, BBuffer, ldb);
|
|
149
209
|
const pass2 = passShape(effTransB, BBuffer, ldb, effTransA, ABuffer, lda);
|
|
150
210
|
|
|
151
|
-
const makeParams = (p, betaVal) =>
|
|
152
|
-
|
|
153
|
-
|
|
154
|
-
|
|
155
|
-
|
|
156
|
-
|
|
157
|
-
|
|
158
|
-
|
|
159
|
-
|
|
160
|
-
|
|
161
|
-
|
|
162
|
-
|
|
163
|
-
|
|
164
|
-
|
|
165
|
-
|
|
166
|
-
|
|
211
|
+
const makeParams = (p, betaVal) =>
|
|
212
|
+
createParamsBuffer(
|
|
213
|
+
device,
|
|
214
|
+
[
|
|
215
|
+
{ value: n, type: "u32" },
|
|
216
|
+
{ value: n, type: "u32" },
|
|
217
|
+
{ value: k, type: "u32" },
|
|
218
|
+
{ value: alpha, type: "f32" },
|
|
219
|
+
{ value: betaVal, type: "f32" },
|
|
220
|
+
{ value: p.ldX, type: "u32" },
|
|
221
|
+
{ value: p.ldY, type: "u32" },
|
|
222
|
+
{ value: ldc, type: "u32" },
|
|
223
|
+
{ value: p.transX === "transpose" ? 1 : 0, type: "u32" },
|
|
224
|
+
{ value: p.transY === "transpose" ? 1 : 0, type: "u32" },
|
|
225
|
+
{ value: uploEff === "upper" ? 1 : 0, type: "u32" },
|
|
226
|
+
],
|
|
227
|
+
"ssyr2k-params",
|
|
228
|
+
);
|
|
167
229
|
paramsBuffer1 = makeParams(pass1, beta);
|
|
168
230
|
paramsBuffer2 = makeParams(pass2, 1.0);
|
|
169
231
|
|
|
170
|
-
const bindGroup1 = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
171
|
-
|
|
232
|
+
const bindGroup1 = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
233
|
+
pass1.X,
|
|
234
|
+
pass1.Y,
|
|
235
|
+
CBuffer,
|
|
236
|
+
paramsBuffer1,
|
|
237
|
+
]);
|
|
238
|
+
const bindGroup2 = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
239
|
+
pass2.X,
|
|
240
|
+
pass2.Y,
|
|
241
|
+
CBuffer,
|
|
242
|
+
paramsBuffer2,
|
|
243
|
+
]);
|
|
172
244
|
|
|
173
|
-
const { commandEncoder, querySet } = beginTimedEncoder();
|
|
174
|
-
const desc1 = querySet
|
|
175
|
-
|
|
245
|
+
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
246
|
+
const desc1 = querySet
|
|
247
|
+
? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
|
|
248
|
+
: undefined;
|
|
249
|
+
const desc2 = querySet
|
|
250
|
+
? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
|
|
251
|
+
: undefined;
|
|
176
252
|
encodePass(commandEncoder, pipeline, bindGroup1, wgCount, desc1);
|
|
177
253
|
encodePass(commandEncoder, pipeline, bindGroup2, wgCount, desc2);
|
|
178
254
|
|
|
179
|
-
const ts = resolveTimestamp(commandEncoder, querySet);
|
|
180
|
-
const readBuffer = CIsGpu
|
|
255
|
+
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
256
|
+
const readBuffer = CIsGpu
|
|
257
|
+
? null
|
|
258
|
+
: stageReadback(device, commandEncoder, CBuffer);
|
|
181
259
|
|
|
182
|
-
submit(commandEncoder);
|
|
260
|
+
submit(device, commandEncoder);
|
|
183
261
|
|
|
184
262
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
185
263
|
|
package/src/ssyrk/ssyrk.d.mts
CHANGED
|
@@ -1,7 +1,8 @@
|
|
|
1
1
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
|
-
* Performs the symmetric rank-k update C
|
|
4
|
+
* Performs the symmetric rank-k update $$C \leftarrow \mathrm{uplo}(\alpha \mathrm{op}(A) \mathrm{op}(A)^{T} + \beta C)$$
|
|
5
|
+
*
|
|
5
6
|
* only the triangle of C named by `uplo` is read or written (`'lower'`: `col <= row`,
|
|
6
7
|
* `'upper'`: `col >= row`). C is always n×n.
|
|
7
8
|
*
|
|
@@ -51,7 +52,7 @@ export declare function ssyrk(
|
|
|
51
52
|
): Promise<{ C: Float32Array; gpuTimeMs?: number }>;
|
|
52
53
|
|
|
53
54
|
/**
|
|
54
|
-
* Performs the symmetric rank-k update C
|
|
55
|
+
* Performs the symmetric rank-k update $$C \leftarrow \mathrm{uplo}(\alpha \mathrm{op}(A) \mathrm{op}(A)^{T} + \beta C)$$
|
|
55
56
|
*
|
|
56
57
|
* A and C are both kept GPU-resident. Each matrix's own `layout` (set at
|
|
57
58
|
* `GpuMatrix.from` time) determines the operation — there is no separate
|
package/src/ssyrk/ssyrk.mjs
CHANGED
|
@@ -11,36 +11,55 @@ import { extractResult } from "../util/result.mjs";
|
|
|
11
11
|
import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
|
|
12
12
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
13
13
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
14
|
-
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
|
|
14
|
+
import { requireWorkgroupCount } from "../util/workgroup.mjs";
|
|
15
|
+
import {
|
|
16
|
+
BM_SMALL,
|
|
17
|
+
BN_SMALL,
|
|
18
|
+
BM_LARGE,
|
|
19
|
+
BN_LARGE,
|
|
20
|
+
LARGE_TILE_WORKGROUP_THRESHOLD,
|
|
21
|
+
} from "../util/constants.mjs";
|
|
22
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
18
23
|
|
|
19
24
|
// ssyrk: C := uplo(alpha*op(A)*op(A)^T + beta*C). No dedicated shader —
|
|
20
25
|
// sgemmtr's kernel with A duplicated into a separate B buffer (B := A).
|
|
21
26
|
export async function ssyrk(
|
|
22
|
-
device,
|
|
27
|
+
device,
|
|
28
|
+
uplo,
|
|
29
|
+
trans,
|
|
30
|
+
n,
|
|
31
|
+
k,
|
|
32
|
+
alpha,
|
|
33
|
+
A,
|
|
34
|
+
lda,
|
|
35
|
+
beta,
|
|
36
|
+
C,
|
|
37
|
+
ldc,
|
|
38
|
+
layout = "row-major",
|
|
23
39
|
) {
|
|
24
40
|
const AIsGpu = A instanceof GpuMatrix;
|
|
25
41
|
const CIsGpu = C instanceof GpuMatrix;
|
|
26
42
|
|
|
27
|
-
|
|
28
|
-
|
|
43
|
+
requireGpuDevice(device);
|
|
44
|
+
requireSameDevice(device, "ssyrk", { A, C });
|
|
29
45
|
if (uplo !== "lower" && uplo !== "upper")
|
|
30
46
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
31
47
|
if (trans !== "no-transpose" && trans !== "transpose")
|
|
32
48
|
throw new Error("trans must be 'no-transpose' or 'transpose'.");
|
|
33
49
|
if (layout !== "row-major" && layout !== "column-major")
|
|
34
50
|
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
35
|
-
if (typeof alpha !== "number")
|
|
36
|
-
throw new Error("alpha must be a number.");
|
|
51
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
37
52
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
38
53
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
39
|
-
if (typeof beta !== "number")
|
|
40
|
-
throw new Error("beta must be a number.");
|
|
54
|
+
if (typeof beta !== "number") throw new Error("beta must be a number.");
|
|
41
55
|
if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
|
|
42
56
|
if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
|
|
43
|
-
if (
|
|
57
|
+
if (
|
|
58
|
+
!Number.isInteger(n) ||
|
|
59
|
+
!Number.isInteger(k) ||
|
|
60
|
+
!Number.isInteger(lda) ||
|
|
61
|
+
!Number.isInteger(ldc)
|
|
62
|
+
)
|
|
44
63
|
throw new Error("n, k, lda, and ldc must be integers.");
|
|
45
64
|
if (!AIsGpu && !(A instanceof Float32Array))
|
|
46
65
|
throw new Error("A must be a Float32Array or GpuMatrix.");
|
|
@@ -51,6 +70,7 @@ export async function ssyrk(
|
|
|
51
70
|
if (CIsGpu && !AIsGpu)
|
|
52
71
|
throw new Error("A must be a GpuMatrix when C is a GpuMatrix.");
|
|
53
72
|
if (n < 0 || k < 0) throw new Error("n and k must be non-negative.");
|
|
73
|
+
if (lda <= 0 || ldc <= 0) throw new Error("lda and ldc must be positive.");
|
|
54
74
|
if (n === 0) return CIsGpu ? {} : { C };
|
|
55
75
|
|
|
56
76
|
// GpuMatrix's own .layout wins over the shared `layout` argument.
|
|
@@ -63,23 +83,32 @@ export async function ssyrk(
|
|
|
63
83
|
const aOuter = trans === "no-transpose" ? aRows : aCols;
|
|
64
84
|
const aInner = trans === "no-transpose" ? aCols : aRows;
|
|
65
85
|
if (lda < aInner)
|
|
66
|
-
throw new Error(
|
|
86
|
+
throw new Error(
|
|
87
|
+
`lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`,
|
|
88
|
+
);
|
|
67
89
|
if (AIsGpu) {
|
|
68
|
-
if (lda !== A.lda)
|
|
90
|
+
if (lda !== A.lda)
|
|
91
|
+
throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
69
92
|
const [aLogRows, aLogCols] = trans === "no-transpose" ? [n, k] : [k, n];
|
|
70
93
|
if (A.rows < aLogRows || A.cols < aLogCols)
|
|
71
94
|
throw new Error("A is too small for the given n, k, and trans.");
|
|
72
95
|
} else if (A.length < (aOuter - 1) * lda + aInner) {
|
|
73
|
-
throw new Error(
|
|
96
|
+
throw new Error(
|
|
97
|
+
"A does not have enough elements for the given dimensions and lda.",
|
|
98
|
+
);
|
|
74
99
|
}
|
|
75
100
|
|
|
76
101
|
// C: always n x n symmetric — layout only affects storage order, not size.
|
|
77
102
|
if (ldc < n) throw new Error("ldc must be >= n.");
|
|
78
103
|
if (CIsGpu) {
|
|
79
|
-
if (ldc !== C.lda)
|
|
80
|
-
|
|
104
|
+
if (ldc !== C.lda)
|
|
105
|
+
throw new Error("ldc must match C.lda when C is a GpuMatrix.");
|
|
106
|
+
if (C.rows < n || C.cols < n)
|
|
107
|
+
throw new Error("C is too small for the given n.");
|
|
81
108
|
} else if (C.length < (n - 1) * ldc + n) {
|
|
82
|
-
throw new Error(
|
|
109
|
+
throw new Error(
|
|
110
|
+
"C does not have enough elements for the given dimensions and ldc.",
|
|
111
|
+
);
|
|
83
112
|
}
|
|
84
113
|
|
|
85
114
|
// Column-major A reinterpreted row-major is A^T — flip trans, same trick sgemm/sgemmtr use.
|
|
@@ -104,22 +133,31 @@ export async function ssyrk(
|
|
|
104
133
|
const largeWgY = Math.ceil(n / BM_LARGE);
|
|
105
134
|
const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
106
135
|
|
|
107
|
-
const pipeline = await getPipeline(
|
|
136
|
+
const pipeline = await getPipeline(
|
|
137
|
+
device,
|
|
138
|
+
useLargeTile ? "sgemmtr_large" : "sgemmtr_small",
|
|
139
|
+
);
|
|
108
140
|
|
|
109
|
-
const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyrk-A", false);
|
|
110
|
-
const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "ssyrk-C", true);
|
|
141
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyrk-A", false);
|
|
142
|
+
const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssyrk-C", true);
|
|
111
143
|
// B := A, but as a genuinely separate buffer — re-uploaded for Float32Array,
|
|
112
144
|
// GPU-copied (see below) for GpuMatrix, since the caller owns ABuffer.
|
|
113
145
|
const BBuffer = AIsGpu
|
|
114
|
-
? createStorageBuffer(
|
|
115
|
-
|
|
146
|
+
? createStorageBuffer(
|
|
147
|
+
device,
|
|
148
|
+
ABuffer.size,
|
|
149
|
+
"ssyrk-B",
|
|
150
|
+
GPUBufferUsage.COPY_DST,
|
|
151
|
+
)
|
|
152
|
+
: uploadBuffer(device, A, "ssyrk-B", false);
|
|
116
153
|
const paramsBuffer = createParamsBuffer(
|
|
154
|
+
device,
|
|
117
155
|
[
|
|
118
|
-
{ value: n,
|
|
119
|
-
{ value: n,
|
|
120
|
-
{ value: k,
|
|
156
|
+
{ value: n, type: "u32" }, // gemmtr's m
|
|
157
|
+
{ value: n, type: "u32" }, // gemmtr's n
|
|
158
|
+
{ value: k, type: "u32" },
|
|
121
159
|
{ value: alpha, type: "f32" },
|
|
122
|
-
{ value: beta,
|
|
160
|
+
{ value: beta, type: "f32" },
|
|
123
161
|
{ value: lda, type: "u32" }, // gemmtr's lda
|
|
124
162
|
{ value: lda, type: "u32" }, // gemmtr's ldb — B := A, same lda
|
|
125
163
|
{ value: ldc, type: "u32" },
|
|
@@ -131,7 +169,7 @@ export async function ssyrk(
|
|
|
131
169
|
);
|
|
132
170
|
|
|
133
171
|
try {
|
|
134
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
172
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
135
173
|
ABuffer,
|
|
136
174
|
BBuffer,
|
|
137
175
|
CBuffer,
|
|
@@ -140,22 +178,36 @@ export async function ssyrk(
|
|
|
140
178
|
|
|
141
179
|
const wgCount = useLargeTile
|
|
142
180
|
? {
|
|
143
|
-
|
|
144
|
-
|
|
145
|
-
|
|
181
|
+
x: requireWorkgroupCount(device, largeWgX, "ssyrk", "x"),
|
|
182
|
+
y: requireWorkgroupCount(device, largeWgY, "ssyrk", "y"),
|
|
183
|
+
}
|
|
146
184
|
: {
|
|
147
|
-
|
|
148
|
-
|
|
149
|
-
|
|
185
|
+
x: requireWorkgroupCount(
|
|
186
|
+
device,
|
|
187
|
+
Math.ceil(n / BN_SMALL),
|
|
188
|
+
"ssyrk",
|
|
189
|
+
"x",
|
|
190
|
+
),
|
|
191
|
+
y: requireWorkgroupCount(
|
|
192
|
+
device,
|
|
193
|
+
Math.ceil(n / BM_SMALL),
|
|
194
|
+
"ssyrk",
|
|
195
|
+
"y",
|
|
196
|
+
),
|
|
197
|
+
};
|
|
150
198
|
// Manual encoder (not runComputePass) so the A->B duplicate copy lands
|
|
151
199
|
// on the same command encoder, strictly before the compute pass reads B.
|
|
152
|
-
const { commandEncoder, querySet, passDescriptor } =
|
|
153
|
-
|
|
200
|
+
const { commandEncoder, querySet, passDescriptor } =
|
|
201
|
+
beginTimedEncoder(device);
|
|
202
|
+
if (AIsGpu)
|
|
203
|
+
commandEncoder.copyBufferToBuffer(ABuffer, 0, BBuffer, 0, ABuffer.size);
|
|
154
204
|
encodePass(commandEncoder, pipeline, bindGroup, wgCount, passDescriptor);
|
|
155
|
-
const ts = resolveTimestamp(commandEncoder, querySet);
|
|
156
|
-
const readBuffer = CIsGpu
|
|
205
|
+
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
206
|
+
const readBuffer = CIsGpu
|
|
207
|
+
? null
|
|
208
|
+
: stageReadback(device, commandEncoder, CBuffer);
|
|
157
209
|
|
|
158
|
-
submit(commandEncoder);
|
|
210
|
+
submit(device, commandEncoder);
|
|
159
211
|
|
|
160
212
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
161
213
|
|
package/src/strmm/strmm.d.mts
CHANGED
|
@@ -2,8 +2,9 @@ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
4
|
* Performs the triangular matrix-matrix operation
|
|
5
|
-
* B
|
|
6
|
-
* B
|
|
5
|
+
* $$B \leftarrow \alpha \mathrm{op}(A) B \quad (\texttt{side='left'})$$
|
|
6
|
+
* $$B \leftarrow \alpha B \mathrm{op}(A) \quad (\texttt{side='right'})$$
|
|
7
|
+
* `A` is triangular, only its
|
|
7
8
|
* `uplo` triangle stored; `B` is a general m×n matrix, overwritten in place.
|
|
8
9
|
*
|
|
9
10
|
* - `side='left'`: `A` is m×m — `A` premultiplies `B`
|
|
@@ -58,8 +59,8 @@ export declare function strmm(
|
|
|
58
59
|
|
|
59
60
|
/**
|
|
60
61
|
* Performs the triangular matrix-matrix operation
|
|
61
|
-
* B
|
|
62
|
-
* B
|
|
62
|
+
* $$B \leftarrow \alpha \mathrm{op}(A) B \quad (\texttt{side='left'})$$
|
|
63
|
+
* $$B \leftarrow \alpha B \mathrm{op}(A) \quad (\texttt{side='right'})$$
|
|
63
64
|
*
|
|
64
65
|
* A and B are both kept GPU-resident; B is mutated in place. Each matrix's
|
|
65
66
|
* own `layout` (set at `GpuMatrix.from` time) determines the operation —
|