wgblas 2.0.0 → 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/README.md +18 -18
- package/dist/wgblas.browser.js +1273 -1239
- package/index.d.mts +38 -6
- package/package.json +2 -1
- package/src/classes/GpuMatrix.mjs +17 -10
- package/src/classes/GpuVector.mjs +28 -10
- package/src/dasum/dasum.d.mts +4 -4
- package/src/dasum/dasum.mjs +19 -17
- package/src/devdocs.mjs +13 -0
- package/src/idamax/idamax.d.mts +20 -2
- package/src/idamax/idamax.mjs +18 -16
- package/src/init.mjs +114 -56
- package/src/isamax/isamax.d.mts +20 -2
- package/src/isamax/isamax.mjs +16 -14
- package/src/random/random.d.mts +1 -0
- package/src/sasum/sasum.d.mts +2 -2
- package/src/sasum/sasum.mjs +15 -13
- package/src/saxpy/saxpy.d.mts +2 -2
- package/src/saxpy/saxpy.mjs +10 -8
- package/src/scopy/scopy.d.mts +2 -2
- package/src/scopy/scopy.mjs +10 -8
- package/src/sdot/sdot.d.mts +2 -2
- package/src/sdot/sdot.mjs +16 -14
- package/src/sgemm/sgemm.mjs +28 -15
- package/src/sgemmtr/sgemmtr.mjs +16 -15
- package/src/sgemv/sgemv.mjs +38 -26
- package/src/sger/sger.mjs +10 -8
- package/src/shaders/index.mjs +164 -14
- package/src/shaders/reduction/scaledSum.wgsl +65 -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 +33 -21
- package/src/srot/srot.d.mts +2 -4
- package/src/srot/srot.mjs +11 -9
- package/src/srotm/srotm.d.mts +2 -4
- package/src/srotm/srotm.mjs +19 -10
- package/src/sscal/sscal.d.mts +3 -3
- package/src/sscal/sscal.mjs +12 -10
- package/src/sswap/sswap.d.mts +2 -2
- package/src/sswap/sswap.mjs +11 -9
- package/src/ssymm/ssymm.mjs +31 -22
- package/src/ssymv/ssymv.mjs +10 -8
- package/src/ssyr/ssyr.mjs +9 -7
- package/src/ssyr2/ssyr2.mjs +10 -8
- package/src/ssyr2k/ssyr2k.mjs +18 -17
- package/src/ssyrk/ssyrk.mjs +18 -17
- package/src/strmm/strmm.mjs +47 -32
- package/src/strmv/strmv.mjs +10 -8
- package/src/strsm/strsm.mjs +54 -36
- package/src/strsv/strsv.mjs +16 -12
- package/src/util/benchmark.mjs +4 -6
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +113 -19
- package/src/util/compute.mjs +6 -9
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +34 -0
- package/src/util/pipeline.mjs +5 -6
- package/src/util/workgroup.mjs +55 -7
- package/src/shaders/browser-shaders.mjs +0 -81
package/src/ssyr/ssyr.mjs
CHANGED
|
@@ -11,6 +11,7 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
13
13
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
14
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "row-major") {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
@@ -18,6 +19,7 @@ export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "ro
|
|
|
18
19
|
|
|
19
20
|
if (!(device instanceof GPUDevice))
|
|
20
21
|
throw new Error("device must be a GPUDevice.");
|
|
22
|
+
requireSameDevice(device, "ssyr", { A, x });
|
|
21
23
|
if (uplo !== "lower" && uplo !== "upper")
|
|
22
24
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
23
25
|
if (layout !== "row-major" && layout !== "column-major")
|
|
@@ -63,9 +65,9 @@ export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "ro
|
|
|
63
65
|
let paramsBuffer = null;
|
|
64
66
|
|
|
65
67
|
try {
|
|
66
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "ssyr-x", false);
|
|
67
|
-
ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyr-A", true);
|
|
68
|
-
paramsBuffer = createParamsBuffer(
|
|
68
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "ssyr-x", false);
|
|
69
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyr-A", true);
|
|
70
|
+
paramsBuffer = createParamsBuffer(device,
|
|
69
71
|
[
|
|
70
72
|
{ value: n, type: "u32" },
|
|
71
73
|
{ value: alpha, type: "f32" },
|
|
@@ -76,7 +78,7 @@ export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "ro
|
|
|
76
78
|
"ssyr-params",
|
|
77
79
|
);
|
|
78
80
|
|
|
79
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
81
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
80
82
|
xBuffer,
|
|
81
83
|
ABuffer,
|
|
82
84
|
paramsBuffer,
|
|
@@ -85,10 +87,10 @@ export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "ro
|
|
|
85
87
|
// One workgroup per row of A; clamped to device limit — the shader's
|
|
86
88
|
// grid-stride loop handles remaining rows when n > dispatch count.
|
|
87
89
|
const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
|
|
88
|
-
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
89
|
-
const readBuffer = AIsGpu ? null : stageReadback(commandEncoder, ABuffer);
|
|
90
|
+
const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
|
|
91
|
+
const readBuffer = AIsGpu ? null : stageReadback(device, commandEncoder, ABuffer);
|
|
90
92
|
|
|
91
|
-
submit(commandEncoder);
|
|
93
|
+
submit(device, commandEncoder);
|
|
92
94
|
|
|
93
95
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
94
96
|
|
package/src/ssyr2/ssyr2.mjs
CHANGED
|
@@ -11,6 +11,7 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
13
13
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
14
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, layout = "row-major") {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
@@ -19,6 +20,7 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
|
|
|
19
20
|
|
|
20
21
|
if (!(device instanceof GPUDevice))
|
|
21
22
|
throw new Error("device must be a GPUDevice.");
|
|
23
|
+
requireSameDevice(device, "ssyr2", { A, x, y });
|
|
22
24
|
if (uplo !== "lower" && uplo !== "upper")
|
|
23
25
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
24
26
|
if (layout !== "row-major" && layout !== "column-major")
|
|
@@ -83,10 +85,10 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
|
|
|
83
85
|
let paramsBuffer = null;
|
|
84
86
|
|
|
85
87
|
try {
|
|
86
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "ssyr2-x", false);
|
|
87
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "ssyr2-y", false);
|
|
88
|
-
ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyr2-A", true);
|
|
89
|
-
paramsBuffer = createParamsBuffer(
|
|
88
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "ssyr2-x", false);
|
|
89
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "ssyr2-y", false);
|
|
90
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyr2-A", true);
|
|
91
|
+
paramsBuffer = createParamsBuffer(device,
|
|
90
92
|
[
|
|
91
93
|
{ value: n, type: "u32" },
|
|
92
94
|
{ value: alpha, type: "f32" },
|
|
@@ -98,7 +100,7 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
|
|
|
98
100
|
"ssyr2-params",
|
|
99
101
|
);
|
|
100
102
|
|
|
101
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
103
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
102
104
|
xBuffer,
|
|
103
105
|
yBuffer,
|
|
104
106
|
ABuffer,
|
|
@@ -108,10 +110,10 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
|
|
|
108
110
|
// One workgroup per row of A; clamped to device limit — the shader's
|
|
109
111
|
// grid-stride loop handles remaining rows when n > dispatch count.
|
|
110
112
|
const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
|
|
111
|
-
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
112
|
-
const readBuffer = AIsGpu ? null : stageReadback(commandEncoder, ABuffer);
|
|
113
|
+
const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
|
|
114
|
+
const readBuffer = AIsGpu ? null : stageReadback(device, commandEncoder, ABuffer);
|
|
113
115
|
|
|
114
|
-
submit(commandEncoder);
|
|
116
|
+
submit(device, commandEncoder);
|
|
115
117
|
|
|
116
118
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
117
119
|
|
package/src/ssyr2k/ssyr2k.mjs
CHANGED
|
@@ -10,10 +10,10 @@ 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
|
+
import { requireWorkgroupCount } from "../util/workgroup.mjs";
|
|
14
|
+
import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
|
|
15
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
13
16
|
|
|
14
|
-
const BM_SMALL = 32, BN_SMALL = 32; // sgemmtr_small.wgsl's block tile
|
|
15
|
-
const BM_LARGE = 64, BN_LARGE = 64; // sgemmtr_large.wgsl's block tile
|
|
16
|
-
const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // same threshold sgemm/sgemmtr/ssyrk use
|
|
17
17
|
|
|
18
18
|
// ssyr2k: C := uplo(alpha*op(A)*op(B)^T + alpha*op(B)*op(A)^T + beta*C). No
|
|
19
19
|
// dedicated shader — two sgemmtr passes on one encoder, second with beta=1.
|
|
@@ -26,6 +26,7 @@ export async function ssyr2k(
|
|
|
26
26
|
|
|
27
27
|
if (!(device instanceof GPUDevice))
|
|
28
28
|
throw new Error("device must be a GPUDevice.");
|
|
29
|
+
requireSameDevice(device, "ssyr2k", { A, B, C });
|
|
29
30
|
if (uplo !== "lower" && uplo !== "upper")
|
|
30
31
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
31
32
|
if (trans !== "no-transpose" && trans !== "transpose")
|
|
@@ -131,24 +132,24 @@ export async function ssyr2k(
|
|
|
131
132
|
const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
|
|
132
133
|
const wgCount = useLargeTile
|
|
133
134
|
? {
|
|
134
|
-
x:
|
|
135
|
-
y:
|
|
135
|
+
x: requireWorkgroupCount(device, largeWgX, "ssyr2k", "x"),
|
|
136
|
+
y: requireWorkgroupCount(device, largeWgY, "ssyr2k", "y"),
|
|
136
137
|
}
|
|
137
138
|
: {
|
|
138
|
-
x:
|
|
139
|
-
y:
|
|
139
|
+
x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "ssyr2k", "x"),
|
|
140
|
+
y: requireWorkgroupCount(device, Math.ceil(n / BM_SMALL), "ssyr2k", "y"),
|
|
140
141
|
};
|
|
141
142
|
|
|
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);
|
|
143
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyr2k-A", false);
|
|
144
|
+
const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "ssyr2k-B", false);
|
|
145
|
+
const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssyr2k-C", true);
|
|
145
146
|
let paramsBuffer1 = null, paramsBuffer2 = null;
|
|
146
147
|
|
|
147
148
|
try {
|
|
148
149
|
const pass1 = passShape(effTransA, ABuffer, lda, effTransB, BBuffer, ldb);
|
|
149
150
|
const pass2 = passShape(effTransB, BBuffer, ldb, effTransA, ABuffer, lda);
|
|
150
151
|
|
|
151
|
-
const makeParams = (p, betaVal) => createParamsBuffer(
|
|
152
|
+
const makeParams = (p, betaVal) => createParamsBuffer(device,
|
|
152
153
|
[
|
|
153
154
|
{ value: n, type: "u32" },
|
|
154
155
|
{ value: n, type: "u32" },
|
|
@@ -167,19 +168,19 @@ export async function ssyr2k(
|
|
|
167
168
|
paramsBuffer1 = makeParams(pass1, beta);
|
|
168
169
|
paramsBuffer2 = makeParams(pass2, 1.0);
|
|
169
170
|
|
|
170
|
-
const bindGroup1 = createBindGroup(pipeline.getBindGroupLayout(0), [pass1.X, pass1.Y, CBuffer, paramsBuffer1]);
|
|
171
|
-
const bindGroup2 = createBindGroup(pipeline.getBindGroupLayout(0), [pass2.X, pass2.Y, CBuffer, paramsBuffer2]);
|
|
171
|
+
const bindGroup1 = createBindGroup(device, pipeline.getBindGroupLayout(0), [pass1.X, pass1.Y, CBuffer, paramsBuffer1]);
|
|
172
|
+
const bindGroup2 = createBindGroup(device, pipeline.getBindGroupLayout(0), [pass2.X, pass2.Y, CBuffer, paramsBuffer2]);
|
|
172
173
|
|
|
173
|
-
const { commandEncoder, querySet } = beginTimedEncoder();
|
|
174
|
+
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
174
175
|
const desc1 = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
|
|
175
176
|
const desc2 = querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
|
|
176
177
|
encodePass(commandEncoder, pipeline, bindGroup1, wgCount, desc1);
|
|
177
178
|
encodePass(commandEncoder, pipeline, bindGroup2, wgCount, desc2);
|
|
178
179
|
|
|
179
|
-
const ts = resolveTimestamp(commandEncoder, querySet);
|
|
180
|
-
const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
|
|
180
|
+
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
181
|
+
const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
|
|
181
182
|
|
|
182
|
-
submit(commandEncoder);
|
|
183
|
+
submit(device, commandEncoder);
|
|
183
184
|
|
|
184
185
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
185
186
|
|
package/src/ssyrk/ssyrk.mjs
CHANGED
|
@@ -11,10 +11,10 @@ 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
|
+
import { requireWorkgroupCount } from "../util/workgroup.mjs";
|
|
15
|
+
import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
|
|
16
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
17
|
|
|
15
|
-
const BM_SMALL = 32, BN_SMALL = 32; // sgemmtr_small.wgsl's block tile
|
|
16
|
-
const BM_LARGE = 64, BN_LARGE = 64; // sgemmtr_large.wgsl's block tile
|
|
17
|
-
const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // same threshold sgemm/sgemmtr use
|
|
18
18
|
|
|
19
19
|
// ssyrk: C := uplo(alpha*op(A)*op(A)^T + beta*C). No dedicated shader —
|
|
20
20
|
// sgemmtr's kernel with A duplicated into a separate B buffer (B := A).
|
|
@@ -26,6 +26,7 @@ export async function ssyrk(
|
|
|
26
26
|
|
|
27
27
|
if (!(device instanceof GPUDevice))
|
|
28
28
|
throw new Error("device must be a GPUDevice.");
|
|
29
|
+
requireSameDevice(device, "ssyrk", { A, C });
|
|
29
30
|
if (uplo !== "lower" && uplo !== "upper")
|
|
30
31
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
31
32
|
if (trans !== "no-transpose" && trans !== "transpose")
|
|
@@ -106,14 +107,14 @@ export async function ssyrk(
|
|
|
106
107
|
|
|
107
108
|
const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
|
|
108
109
|
|
|
109
|
-
const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyrk-A", false);
|
|
110
|
-
const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "ssyrk-C", true);
|
|
110
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyrk-A", false);
|
|
111
|
+
const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssyrk-C", true);
|
|
111
112
|
// B := A, but as a genuinely separate buffer — re-uploaded for Float32Array,
|
|
112
113
|
// GPU-copied (see below) for GpuMatrix, since the caller owns ABuffer.
|
|
113
114
|
const BBuffer = AIsGpu
|
|
114
|
-
? createStorageBuffer(ABuffer.size, "ssyrk-B", GPUBufferUsage.COPY_DST)
|
|
115
|
-
: uploadBuffer(A, "ssyrk-B", false);
|
|
116
|
-
const paramsBuffer = createParamsBuffer(
|
|
115
|
+
? createStorageBuffer(device, ABuffer.size, "ssyrk-B", GPUBufferUsage.COPY_DST)
|
|
116
|
+
: uploadBuffer(device, A, "ssyrk-B", false);
|
|
117
|
+
const paramsBuffer = createParamsBuffer(device,
|
|
117
118
|
[
|
|
118
119
|
{ value: n, type: "u32" }, // gemmtr's m
|
|
119
120
|
{ value: n, type: "u32" }, // gemmtr's n
|
|
@@ -131,7 +132,7 @@ export async function ssyrk(
|
|
|
131
132
|
);
|
|
132
133
|
|
|
133
134
|
try {
|
|
134
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
135
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
135
136
|
ABuffer,
|
|
136
137
|
BBuffer,
|
|
137
138
|
CBuffer,
|
|
@@ -140,22 +141,22 @@ export async function ssyrk(
|
|
|
140
141
|
|
|
141
142
|
const wgCount = useLargeTile
|
|
142
143
|
? {
|
|
143
|
-
x:
|
|
144
|
-
y:
|
|
144
|
+
x: requireWorkgroupCount(device, largeWgX, "ssyrk", "x"),
|
|
145
|
+
y: requireWorkgroupCount(device, largeWgY, "ssyrk", "y"),
|
|
145
146
|
}
|
|
146
147
|
: {
|
|
147
|
-
x:
|
|
148
|
-
y:
|
|
148
|
+
x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "ssyrk", "x"),
|
|
149
|
+
y: requireWorkgroupCount(device, Math.ceil(n / BM_SMALL), "ssyrk", "y"),
|
|
149
150
|
};
|
|
150
151
|
// Manual encoder (not runComputePass) so the A->B duplicate copy lands
|
|
151
152
|
// on the same command encoder, strictly before the compute pass reads B.
|
|
152
|
-
const { commandEncoder, querySet, passDescriptor } = beginTimedEncoder();
|
|
153
|
+
const { commandEncoder, querySet, passDescriptor } = beginTimedEncoder(device);
|
|
153
154
|
if (AIsGpu) commandEncoder.copyBufferToBuffer(ABuffer, 0, BBuffer, 0, ABuffer.size);
|
|
154
155
|
encodePass(commandEncoder, pipeline, bindGroup, wgCount, passDescriptor);
|
|
155
|
-
const ts = resolveTimestamp(commandEncoder, querySet);
|
|
156
|
-
const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
|
|
156
|
+
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
157
|
+
const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
|
|
157
158
|
|
|
158
|
-
submit(commandEncoder);
|
|
159
|
+
submit(device, commandEncoder);
|
|
159
160
|
|
|
160
161
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
161
162
|
|
package/src/strmm/strmm.mjs
CHANGED
|
@@ -4,6 +4,7 @@ import {
|
|
|
4
4
|
createStorageBuffer,
|
|
5
5
|
stageReadback,
|
|
6
6
|
destroyBuffers,
|
|
7
|
+
vec4ViewBinding,
|
|
7
8
|
} from "../util/buffer.mjs";
|
|
8
9
|
import { createBindGroup } from "../util/bindgroup.mjs";
|
|
9
10
|
import { beginTimedEncoder, encodePass, submit } from "../util/compute.mjs";
|
|
@@ -11,11 +12,11 @@ import { extractResult } from "../util/result.mjs";
|
|
|
11
12
|
import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
|
|
12
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
13
14
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
15
|
+
import { requireWorkgroupCount } from "../util/workgroup.mjs";
|
|
16
|
+
import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
|
|
17
|
+
import { TILE_WG_2D } from "../util/constants.mjs";
|
|
18
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
19
|
|
|
15
|
-
const BM_SMALL = 32, BN_SMALL = 32; // sgemm_small.wgsl's block tile
|
|
16
|
-
const BM_LARGE = 64, BN_LARGE = 64; // sgemm_large.wgsl's block tile
|
|
17
|
-
const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // same threshold sgemm/sgemmtr/ssyrk/ssymm use
|
|
18
|
-
const TRI_WG = 8; // triangularize.wgsl's @workgroup_size(8, 8)
|
|
19
20
|
|
|
20
21
|
// strmm: B := alpha*op(A)*B (side='left') or alpha*B*op(A) (side='right'), A
|
|
21
22
|
// triangular. Triangularize then sgemm, one command encoder. B is both
|
|
@@ -30,6 +31,7 @@ export async function strmm(
|
|
|
30
31
|
|
|
31
32
|
if (!(device instanceof GPUDevice))
|
|
32
33
|
throw new Error("device must be a GPUDevice.");
|
|
34
|
+
requireSameDevice(device, "strmm", { A, B });
|
|
33
35
|
if (side !== "left" && side !== "right")
|
|
34
36
|
throw new Error("side must be 'left' or 'right'.");
|
|
35
37
|
if (uplo !== "lower" && uplo !== "upper")
|
|
@@ -110,29 +112,35 @@ export async function strmm(
|
|
|
110
112
|
const triPipeline = await getPipeline(device, "triangularize");
|
|
111
113
|
const gemmWgCount = useLargeTile
|
|
112
114
|
? {
|
|
113
|
-
x:
|
|
114
|
-
y:
|
|
115
|
+
x: requireWorkgroupCount(device, largeWgX, "strmm", "x"),
|
|
116
|
+
y: requireWorkgroupCount(device, largeWgY, "strmm", "y"),
|
|
115
117
|
}
|
|
116
118
|
: {
|
|
117
|
-
x:
|
|
118
|
-
y:
|
|
119
|
+
x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "strmm", "x"),
|
|
120
|
+
y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "strmm", "y"),
|
|
119
121
|
};
|
|
120
122
|
|
|
121
|
-
|
|
122
|
-
//
|
|
123
|
-
|
|
124
|
-
|
|
125
|
-
|
|
126
|
-
// gaps (never written by gemm's tight m x n loop) keep B's original bytes
|
|
127
|
-
// instead of reading back as zero. COPY_SRC: read back / adopted by B after.
|
|
128
|
-
const outBuffer = createStorageBuffer(
|
|
129
|
-
bOuter * ldb * 4, "strmm-out", GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST,
|
|
130
|
-
);
|
|
123
|
+
// Null-init here and allocate inside the try below, so a throw partway
|
|
124
|
+
// through the sequence still reaches finally with every handle visible
|
|
125
|
+
// (strsv.mjs is the reference for this pattern).
|
|
126
|
+
let ABuffer = null, BBuffer = null;
|
|
127
|
+
let AdenseBuffer = null, outBuffer = null;
|
|
131
128
|
let triParams = null, gemmParams = null;
|
|
132
129
|
let outBufferAdopted = false; // true once B._buf is repointed at outBuffer
|
|
133
130
|
|
|
134
131
|
try {
|
|
135
|
-
|
|
132
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmm-A", false);
|
|
133
|
+
// readback=true (COPY_SRC): BBuffer is also the source that seeds outBuffer.
|
|
134
|
+
BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strmm-B", true);
|
|
135
|
+
AdenseBuffer = createStorageBuffer(device, aOrder * ldDense * 4, "strmm-Adense");
|
|
136
|
+
// COPY_DST: seeded from B's own content before gemm runs, so stride-padding
|
|
137
|
+
// gaps (never written by gemm's tight m x n loop) keep B's original bytes
|
|
138
|
+
// instead of reading back as zero. COPY_SRC: read back / adopted by B after.
|
|
139
|
+
outBuffer = createStorageBuffer(device,
|
|
140
|
+
bOuter * ldb * 4, "strmm-out", GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST,
|
|
141
|
+
);
|
|
142
|
+
|
|
143
|
+
triParams = createParamsBuffer(device,
|
|
136
144
|
[
|
|
137
145
|
{ value: aOrder, type: "u32" },
|
|
138
146
|
{ value: lda, type: "u32" },
|
|
@@ -143,7 +151,7 @@ export async function strmm(
|
|
|
143
151
|
],
|
|
144
152
|
"strmm-tri-params",
|
|
145
153
|
);
|
|
146
|
-
const triBindGroup = createBindGroup(triPipeline.getBindGroupLayout(0), [ABuffer, AdenseBuffer, triParams]);
|
|
154
|
+
const triBindGroup = createBindGroup(device, triPipeline.getBindGroupLayout(0), [ABuffer, AdenseBuffer, triParams]);
|
|
147
155
|
|
|
148
156
|
// X/Y buffers and their own ld, matching swapXY above.
|
|
149
157
|
const XBuffer = swapXY ? BBuffer : AdenseBuffer;
|
|
@@ -151,7 +159,7 @@ export async function strmm(
|
|
|
151
159
|
const YBuffer = swapXY ? AdenseBuffer : BBuffer;
|
|
152
160
|
const ldY = swapXY ? ldDense : ldb;
|
|
153
161
|
|
|
154
|
-
gemmParams = createParamsBuffer(
|
|
162
|
+
gemmParams = createParamsBuffer(device,
|
|
155
163
|
[
|
|
156
164
|
{ value: mg, type: "u32" },
|
|
157
165
|
{ value: ng, type: "u32" },
|
|
@@ -166,9 +174,16 @@ export async function strmm(
|
|
|
166
174
|
],
|
|
167
175
|
"strmm-gemm-params",
|
|
168
176
|
);
|
|
169
|
-
const gemmBindGroup = createBindGroup(gemmPipeline.getBindGroupLayout(0), [
|
|
170
|
-
|
|
171
|
-
|
|
177
|
+
const gemmBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
|
|
178
|
+
XBuffer,
|
|
179
|
+
vec4ViewBinding(device, XBuffer),
|
|
180
|
+
YBuffer,
|
|
181
|
+
vec4ViewBinding(device, YBuffer),
|
|
182
|
+
outBuffer,
|
|
183
|
+
gemmParams,
|
|
184
|
+
]);
|
|
185
|
+
|
|
186
|
+
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
172
187
|
// Seed outBuffer with B's own bytes first, so gemm's tight m x n write
|
|
173
188
|
// leaves stride-padding gaps holding B's original content, not zero.
|
|
174
189
|
// BBuffer may be larger than outBuffer (e.g. a validation-test baseline
|
|
@@ -176,13 +191,13 @@ export async function strmm(
|
|
|
176
191
|
commandEncoder.copyBufferToBuffer(BBuffer, 0, outBuffer, 0, Math.min(BBuffer.size, outBuffer.size));
|
|
177
192
|
const triDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
|
|
178
193
|
const gemmDesc = querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
|
|
179
|
-
encodePass(commandEncoder, triPipeline, triBindGroup, { x: Math.ceil(aOrder /
|
|
194
|
+
encodePass(commandEncoder, triPipeline, triBindGroup, { x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) }, triDesc);
|
|
180
195
|
encodePass(commandEncoder, gemmPipeline, gemmBindGroup, gemmWgCount, gemmDesc);
|
|
181
196
|
|
|
182
|
-
const ts = resolveTimestamp(commandEncoder, querySet);
|
|
183
|
-
const readBuffer = BIsGpu ? null : stageReadback(commandEncoder, outBuffer);
|
|
197
|
+
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
198
|
+
const readBuffer = BIsGpu ? null : stageReadback(device, commandEncoder, outBuffer);
|
|
184
199
|
|
|
185
|
-
submit(commandEncoder);
|
|
200
|
+
submit(device, commandEncoder);
|
|
186
201
|
|
|
187
202
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
188
203
|
|
|
@@ -201,10 +216,10 @@ export async function strmm(
|
|
|
201
216
|
if (gpuTimeMs !== undefined) return { B: result, gpuTimeMs };
|
|
202
217
|
return { B: result };
|
|
203
218
|
} finally {
|
|
204
|
-
if (!AIsGpu) destroyBuffers(ABuffer);
|
|
205
|
-
if (!BIsGpu) destroyBuffers(BBuffer);
|
|
206
|
-
destroyBuffers(AdenseBuffer);
|
|
207
|
-
if (!outBufferAdopted) destroyBuffers(outBuffer);
|
|
219
|
+
if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
|
|
220
|
+
if (!BIsGpu && BBuffer) destroyBuffers(BBuffer);
|
|
221
|
+
if (AdenseBuffer) destroyBuffers(AdenseBuffer);
|
|
222
|
+
if (outBuffer && !outBufferAdopted) destroyBuffers(outBuffer);
|
|
208
223
|
if (triParams) destroyBuffers(triParams);
|
|
209
224
|
if (gemmParams) destroyBuffers(gemmParams);
|
|
210
225
|
}
|
package/src/strmv/strmv.mjs
CHANGED
|
@@ -11,6 +11,7 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
13
13
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
14
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, incy, layout = "row-major") {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
@@ -20,6 +21,7 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
20
21
|
|
|
21
22
|
if (!(device instanceof GPUDevice))
|
|
22
23
|
throw new Error("device must be a GPUDevice.");
|
|
24
|
+
requireSameDevice(device, "strmv", { A, x, y });
|
|
23
25
|
if (uplo !== "lower" && uplo !== "upper")
|
|
24
26
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
25
27
|
if (trans !== "no-transpose" && trans !== "transpose")
|
|
@@ -92,10 +94,10 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
92
94
|
let paramsBuffer = null;
|
|
93
95
|
|
|
94
96
|
try {
|
|
95
|
-
ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "strmv-A", false);
|
|
96
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "strmv-x", false);
|
|
97
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "strmv-y", true);
|
|
98
|
-
paramsBuffer = createParamsBuffer(
|
|
97
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmv-A", false);
|
|
98
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "strmv-x", false);
|
|
99
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "strmv-y", true);
|
|
100
|
+
paramsBuffer = createParamsBuffer(device,
|
|
99
101
|
[
|
|
100
102
|
{ value: n, type: "u32" },
|
|
101
103
|
{ value: incx, type: "u32" },
|
|
@@ -108,7 +110,7 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
108
110
|
"strmv-params",
|
|
109
111
|
);
|
|
110
112
|
|
|
111
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
113
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
112
114
|
ABuffer,
|
|
113
115
|
xBuffer,
|
|
114
116
|
yBuffer,
|
|
@@ -116,10 +118,10 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
|
|
|
116
118
|
]);
|
|
117
119
|
|
|
118
120
|
const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
|
|
119
|
-
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
120
|
-
const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
121
|
+
const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
|
|
122
|
+
const readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
121
123
|
|
|
122
|
-
submit(commandEncoder);
|
|
124
|
+
submit(device, commandEncoder);
|
|
123
125
|
|
|
124
126
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
125
127
|
|