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/sdot/sdot.mjs
CHANGED
|
@@ -12,8 +12,9 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
12
12
|
import { extractResult } from "../util/result.mjs";
|
|
13
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
14
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
15
|
+
import { WGS } from "../util/constants.mjs";
|
|
16
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
15
17
|
|
|
16
|
-
const WGS = 64; //workgroup size
|
|
17
18
|
|
|
18
19
|
export async function sdot(device, n, x, incx, y, incy) {
|
|
19
20
|
const xIsGpu = x instanceof GpuVector;
|
|
@@ -21,6 +22,7 @@ export async function sdot(device, n, x, incx, y, incy) {
|
|
|
21
22
|
|
|
22
23
|
if (!(device instanceof GPUDevice))
|
|
23
24
|
throw new Error("device must be a GPUDevice.");
|
|
25
|
+
requireSameDevice(device, "sdot", { x, y });
|
|
24
26
|
if (
|
|
25
27
|
!Number.isInteger(n) ||
|
|
26
28
|
!Number.isInteger(incx) ||
|
|
@@ -58,11 +60,11 @@ export async function sdot(device, n, x, incx, y, incy) {
|
|
|
58
60
|
let readBuffer = null;
|
|
59
61
|
|
|
60
62
|
try {
|
|
61
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sdot-x", false);
|
|
62
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sdot-y", false);
|
|
63
|
-
partialsBuffer = createStorageBuffer(2 * WGS * 4, "sdot-partials"); //to hold 2*WGS partial sums of f32
|
|
64
|
-
resultBuffer = createResultBuffer(4, "sdot-result"); //to hold the final float32 dot product
|
|
65
|
-
paramsBuffer = createParamsBuffer(
|
|
63
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sdot-x", false);
|
|
64
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sdot-y", false);
|
|
65
|
+
partialsBuffer = createStorageBuffer(device, 2 * WGS * 4, "sdot-partials"); //to hold 2*WGS partial sums of f32
|
|
66
|
+
resultBuffer = createResultBuffer(device, 4, "sdot-result"); //to hold the final float32 dot product
|
|
67
|
+
paramsBuffer = createParamsBuffer(device,
|
|
66
68
|
[
|
|
67
69
|
{ value: n, type: "u32" },
|
|
68
70
|
{ value: incx, type: "u32" },
|
|
@@ -71,32 +73,32 @@ export async function sdot(device, n, x, incx, y, incy) {
|
|
|
71
73
|
"sdot-params",
|
|
72
74
|
);
|
|
73
75
|
|
|
74
|
-
const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
|
|
76
|
+
const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
|
|
75
77
|
xBuffer,
|
|
76
78
|
yBuffer,
|
|
77
79
|
partialsBuffer,
|
|
78
80
|
paramsBuffer,
|
|
79
81
|
]);
|
|
80
|
-
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
82
|
+
const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
|
|
81
83
|
pipelineMain,
|
|
82
84
|
bgMain,
|
|
83
85
|
2 * WGS,
|
|
84
86
|
); //dispatch 2*WGS workgroups
|
|
85
87
|
|
|
86
|
-
submit(enc1);
|
|
88
|
+
submit(device, enc1);
|
|
87
89
|
|
|
88
|
-
const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
|
|
90
|
+
const bgReduce = createBindGroup(device, pipelineReduce.getBindGroupLayout(0), [
|
|
89
91
|
partialsBuffer,
|
|
90
92
|
resultBuffer,
|
|
91
93
|
]);
|
|
92
|
-
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
94
|
+
const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
|
|
93
95
|
pipelineReduce,
|
|
94
96
|
bgReduce,
|
|
95
97
|
1,
|
|
96
98
|
); // dispatch 1 workgroup to reduce the partial sums to a single result
|
|
97
|
-
readBuffer = stageReadback(enc2, resultBuffer);
|
|
99
|
+
readBuffer = stageReadback(device, enc2, resultBuffer);
|
|
98
100
|
|
|
99
|
-
submit(enc2);
|
|
101
|
+
submit(device, enc2);
|
|
100
102
|
|
|
101
103
|
const resultPromise = extractResult(readBuffer, Float32Array);
|
|
102
104
|
readBuffer = null; // ownership transferred — extractResult's own finally destroys it
|
|
@@ -117,7 +119,7 @@ export async function sdot(device, n, x, incx, y, incy) {
|
|
|
117
119
|
if (partialsBuffer) destroyBuffers(partialsBuffer);
|
|
118
120
|
if (resultBuffer) destroyBuffers(resultBuffer);
|
|
119
121
|
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
120
|
-
// Only reached if submit(enc2) threw before ownership was transferred above.
|
|
122
|
+
// Only reached if submit(device, enc2) threw before ownership was transferred above.
|
|
121
123
|
if (readBuffer) destroyBuffers(readBuffer);
|
|
122
124
|
}
|
|
123
125
|
}
|
package/src/sgemm/sgemm.mjs
CHANGED
|
@@ -3,6 +3,8 @@ import {
|
|
|
3
3
|
createParamsBuffer,
|
|
4
4
|
stageReadback,
|
|
5
5
|
destroyBuffers,
|
|
6
|
+
vec4ViewBinding,
|
|
7
|
+
vec4Usable,
|
|
6
8
|
} from "../util/buffer.mjs";
|
|
7
9
|
import { createBindGroup } from "../util/bindgroup.mjs";
|
|
8
10
|
import { runComputePass, submit } from "../util/compute.mjs";
|
|
@@ -10,10 +12,10 @@ import { extractResult } from "../util/result.mjs";
|
|
|
10
12
|
import { extractTimestamp } from "../util/benchmark.mjs";
|
|
11
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
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 { requireSameDevice } from "../util/device.mjs";
|
|
13
18
|
|
|
14
|
-
const BM_SMALL = 32, BN_SMALL = 32; // sgemm_small.wgsl's block tile
|
|
15
|
-
const BM_LARGE = 64, BN_LARGE = 64; // sgemm_large.wgsl's block tile
|
|
16
|
-
const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // large tile needs >= a 6x6 grid of its own tiles to beat the small tile
|
|
17
19
|
|
|
18
20
|
export async function sgemm(
|
|
19
21
|
device, transA, transB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
|
|
@@ -24,6 +26,7 @@ export async function sgemm(
|
|
|
24
26
|
|
|
25
27
|
if (!(device instanceof GPUDevice))
|
|
26
28
|
throw new Error("device must be a GPUDevice.");
|
|
29
|
+
requireSameDevice(device, "sgemm", { A, B, C });
|
|
27
30
|
if (transA !== "no-transpose" && transA !== "transpose")
|
|
28
31
|
throw new Error("transA must be 'no-transpose' or 'transpose'.");
|
|
29
32
|
if (transB !== "no-transpose" && transB !== "transpose")
|
|
@@ -135,10 +138,16 @@ export async function sgemm(
|
|
|
135
138
|
|
|
136
139
|
const pipeline = await getPipeline(device, useLargeTile ? "sgemm_large" : "sgemm_small");
|
|
137
140
|
|
|
138
|
-
const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sgemm-A", false);
|
|
139
|
-
const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "sgemm-B", false);
|
|
140
|
-
const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "sgemm-C", true);
|
|
141
|
-
|
|
141
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemm-A", false);
|
|
142
|
+
const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "sgemm-B", false);
|
|
143
|
+
const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "sgemm-C", true);
|
|
144
|
+
// Vectorized-load enablement — kernel-side view after the column-major swap.
|
|
145
|
+
// op(A) is m×k (no-trans) or k×m (trans); op(B) is k×n or n×k.
|
|
146
|
+
const aNot = transA === "no-transpose";
|
|
147
|
+
const bNot = transB === "no-transpose";
|
|
148
|
+
const useVecA = aNot && vec4Usable(ABuffer, lda, m, k); // transposed-A vec stores bank-conflict smem — measured slower than scalar
|
|
149
|
+
const useVecB = vec4Usable(BBuffer, ldb, bNot ? k : n, bNot ? n : k);
|
|
150
|
+
const paramsBuffer = createParamsBuffer(device,
|
|
142
151
|
[
|
|
143
152
|
{ value: m, type: "u32" },
|
|
144
153
|
{ value: n, type: "u32" },
|
|
@@ -150,31 +159,35 @@ export async function sgemm(
|
|
|
150
159
|
{ value: ldc, type: "u32" },
|
|
151
160
|
{ value: transA === "transpose" ? 1 : 0, type: "u32" },
|
|
152
161
|
{ value: transB === "transpose" ? 1 : 0, type: "u32" },
|
|
162
|
+
{ value: useVecA ? 1 : 0, type: "u32" },
|
|
163
|
+
{ value: useVecB ? 1 : 0, type: "u32" },
|
|
153
164
|
],
|
|
154
165
|
"sgemm-params",
|
|
155
166
|
);
|
|
156
167
|
|
|
157
168
|
try {
|
|
158
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
169
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
159
170
|
ABuffer,
|
|
171
|
+
vec4ViewBinding(device, ABuffer),
|
|
160
172
|
BBuffer,
|
|
173
|
+
vec4ViewBinding(device, BBuffer),
|
|
161
174
|
CBuffer,
|
|
162
175
|
paramsBuffer,
|
|
163
176
|
]);
|
|
164
177
|
|
|
165
178
|
const wgCount = useLargeTile
|
|
166
179
|
? {
|
|
167
|
-
x:
|
|
168
|
-
y:
|
|
180
|
+
x: requireWorkgroupCount(device, largeWgX, "sgemm", "x"),
|
|
181
|
+
y: requireWorkgroupCount(device, largeWgY, "sgemm", "y"),
|
|
169
182
|
}
|
|
170
183
|
: {
|
|
171
|
-
x:
|
|
172
|
-
y:
|
|
184
|
+
x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "sgemm", "x"),
|
|
185
|
+
y: requireWorkgroupCount(device, Math.ceil(m / BM_SMALL), "sgemm", "y"),
|
|
173
186
|
};
|
|
174
|
-
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
175
|
-
const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
|
|
187
|
+
const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
|
|
188
|
+
const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
|
|
176
189
|
|
|
177
|
-
submit(commandEncoder);
|
|
190
|
+
submit(device, commandEncoder);
|
|
178
191
|
|
|
179
192
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
180
193
|
|
package/src/sgemmtr/sgemmtr.mjs
CHANGED
|
@@ -10,10 +10,10 @@ import { extractResult } from "../util/result.mjs";
|
|
|
10
10
|
import { 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 uses — see sgemm.mjs
|
|
17
17
|
|
|
18
18
|
export async function sgemmtr(
|
|
19
19
|
device, uplo, transA, transB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
|
|
@@ -24,6 +24,7 @@ export async function sgemmtr(
|
|
|
24
24
|
|
|
25
25
|
if (!(device instanceof GPUDevice))
|
|
26
26
|
throw new Error("device must be a GPUDevice.");
|
|
27
|
+
requireSameDevice(device, "sgemmtr", { A, B, C });
|
|
27
28
|
if (uplo !== "lower" && uplo !== "upper")
|
|
28
29
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
29
30
|
if (transA !== "no-transpose" && transA !== "transpose")
|
|
@@ -142,10 +143,10 @@ export async function sgemmtr(
|
|
|
142
143
|
|
|
143
144
|
const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
|
|
144
145
|
|
|
145
|
-
const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sgemmtr-A", false);
|
|
146
|
-
const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "sgemmtr-B", false);
|
|
147
|
-
const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "sgemmtr-C", true);
|
|
148
|
-
const paramsBuffer = createParamsBuffer(
|
|
146
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemmtr-A", false);
|
|
147
|
+
const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "sgemmtr-B", false);
|
|
148
|
+
const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "sgemmtr-C", true);
|
|
149
|
+
const paramsBuffer = createParamsBuffer(device,
|
|
149
150
|
[
|
|
150
151
|
{ value: m, type: "u32" },
|
|
151
152
|
{ value: n, type: "u32" },
|
|
@@ -163,7 +164,7 @@ export async function sgemmtr(
|
|
|
163
164
|
);
|
|
164
165
|
|
|
165
166
|
try {
|
|
166
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
167
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
167
168
|
ABuffer,
|
|
168
169
|
BBuffer,
|
|
169
170
|
CBuffer,
|
|
@@ -172,17 +173,17 @@ export async function sgemmtr(
|
|
|
172
173
|
|
|
173
174
|
const wgCount = useLargeTile
|
|
174
175
|
? {
|
|
175
|
-
x:
|
|
176
|
-
y:
|
|
176
|
+
x: requireWorkgroupCount(device, largeWgX, "sgemmtr", "x"),
|
|
177
|
+
y: requireWorkgroupCount(device, largeWgY, "sgemmtr", "y"),
|
|
177
178
|
}
|
|
178
179
|
: {
|
|
179
|
-
x:
|
|
180
|
-
y:
|
|
180
|
+
x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "sgemmtr", "x"),
|
|
181
|
+
y: requireWorkgroupCount(device, Math.ceil(m / BM_SMALL), "sgemmtr", "y"),
|
|
181
182
|
};
|
|
182
|
-
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
183
|
-
const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
|
|
183
|
+
const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
|
|
184
|
+
const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
|
|
184
185
|
|
|
185
|
-
submit(commandEncoder);
|
|
186
|
+
submit(device, commandEncoder);
|
|
186
187
|
|
|
187
188
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
188
189
|
|
package/src/sgemv/sgemv.mjs
CHANGED
|
@@ -9,9 +9,10 @@ import { runComputePass, submit } from "../util/compute.mjs";
|
|
|
9
9
|
import { extractResult } from "../util/result.mjs";
|
|
10
10
|
import { extractTimestamp } from "../util/benchmark.mjs";
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
|
-
import {
|
|
12
|
+
import { requireWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
14
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
15
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
15
16
|
|
|
16
17
|
export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y, incy, layout = "row-major") {
|
|
17
18
|
const AIsGpu = A instanceof GpuMatrix;
|
|
@@ -20,6 +21,7 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
|
|
|
20
21
|
|
|
21
22
|
if (!(device instanceof GPUDevice))
|
|
22
23
|
throw new Error("device must be a GPUDevice.");
|
|
24
|
+
requireSameDevice(device, "sgemv", { A, x, y });
|
|
23
25
|
if (trans !== "no-transpose" && trans !== "transpose")
|
|
24
26
|
throw new Error("trans must be 'no-transpose' or 'transpose'.");
|
|
25
27
|
if (layout !== "row-major" && layout !== "column-major")
|
|
@@ -62,6 +64,8 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
|
|
|
62
64
|
);
|
|
63
65
|
if (xIsGpu && x._buf === y._buf)
|
|
64
66
|
throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");
|
|
67
|
+
if (AIsGpu && yIsGpu && A._buf === y._buf)
|
|
68
|
+
throw new Error("A and y must not reference the same GPU buffer.");
|
|
65
69
|
if (AIsGpu && lda !== A.lda)
|
|
66
70
|
throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
67
71
|
if (AIsGpu && (A.rows < m || A.cols < n))
|
|
@@ -99,38 +103,46 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
|
|
|
99
103
|
const shaderName = isNoTrans ? "sgemv_n" : "sgemv_t";
|
|
100
104
|
const pipeline = await getPipeline(device, shaderName);
|
|
101
105
|
|
|
102
|
-
|
|
103
|
-
|
|
104
|
-
|
|
105
|
-
|
|
106
|
-
[
|
|
107
|
-
{ value: m, type: "u32" },
|
|
108
|
-
{ value: n, type: "u32" },
|
|
109
|
-
{ value: alpha, type: "f32" },
|
|
110
|
-
{ value: beta, type: "f32" },
|
|
111
|
-
{ value: incx, type: "u32" },
|
|
112
|
-
{ value: incy, type: "u32" },
|
|
113
|
-
{ value: lda, type: "u32" },
|
|
114
|
-
],
|
|
115
|
-
"sgemv-params",
|
|
116
|
-
);
|
|
106
|
+
let ABuffer = null;
|
|
107
|
+
let xBuffer = null;
|
|
108
|
+
let yBuffer = null;
|
|
109
|
+
let paramsBuffer = null;
|
|
117
110
|
|
|
118
111
|
try {
|
|
119
|
-
|
|
112
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemv-A", false);
|
|
113
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sgemv-x", false);
|
|
114
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sgemv-y", true);
|
|
115
|
+
paramsBuffer = createParamsBuffer(device,
|
|
116
|
+
[
|
|
117
|
+
{ value: m, type: "u32" },
|
|
118
|
+
{ value: n, type: "u32" },
|
|
119
|
+
{ value: alpha, type: "f32" },
|
|
120
|
+
{ value: beta, type: "f32" },
|
|
121
|
+
{ value: incx, type: "u32" },
|
|
122
|
+
{ value: incy, type: "u32" },
|
|
123
|
+
{ value: lda, type: "u32" },
|
|
124
|
+
],
|
|
125
|
+
"sgemv-params",
|
|
126
|
+
);
|
|
127
|
+
|
|
128
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
120
129
|
ABuffer,
|
|
121
130
|
xBuffer,
|
|
122
131
|
yBuffer,
|
|
123
132
|
paramsBuffer,
|
|
124
133
|
]);
|
|
125
134
|
|
|
126
|
-
// NoTrans: one workgroup per row
|
|
135
|
+
// NoTrans: one workgroup per row — sgemv_n.wgsl is a grid-stride loop, so
|
|
136
|
+
// clamping here only costs parallelism. Trans: one thread per output
|
|
137
|
+
// column, and sgemv_t.wgsl indexes straight off global_invocation_id with
|
|
138
|
+
// no fallback, so an over-limit dispatch must be refused, not truncated.
|
|
127
139
|
const wgCount = isNoTrans
|
|
128
140
|
? Math.min(m, device.limits.maxComputeWorkgroupsPerDimension)
|
|
129
|
-
:
|
|
130
|
-
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
131
|
-
const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
141
|
+
: requireWorkgroups(device, "sgemv", yLen);
|
|
142
|
+
const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
|
|
143
|
+
const readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
132
144
|
|
|
133
|
-
submit(commandEncoder);
|
|
145
|
+
submit(device, commandEncoder);
|
|
134
146
|
|
|
135
147
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
136
148
|
|
|
@@ -143,10 +155,10 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
|
|
|
143
155
|
if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
|
|
144
156
|
return { y: result };
|
|
145
157
|
} finally {
|
|
146
|
-
if (!AIsGpu) destroyBuffers(ABuffer);
|
|
147
|
-
if (!xIsGpu) destroyBuffers(xBuffer);
|
|
148
|
-
if (!yIsGpu) destroyBuffers(yBuffer);
|
|
149
|
-
destroyBuffers(paramsBuffer);
|
|
158
|
+
if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
|
|
159
|
+
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
160
|
+
if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
|
|
161
|
+
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
150
162
|
|
|
151
163
|
}
|
|
152
164
|
}
|
package/src/sger/sger.mjs
CHANGED
|
@@ -11,12 +11,14 @@ 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 sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout = "row-major") {
|
|
16
17
|
const AIsGpu = A instanceof GpuMatrix;
|
|
17
18
|
|
|
18
19
|
if (!(device instanceof GPUDevice))
|
|
19
20
|
throw new Error("device must be a GPUDevice.");
|
|
21
|
+
requireSameDevice(device, "sger", { A, x, y });
|
|
20
22
|
if (layout !== "row-major" && layout !== "column-major")
|
|
21
23
|
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
22
24
|
if (typeof alpha !== "number")
|
|
@@ -89,10 +91,10 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
|
|
|
89
91
|
let paramsBuffer = null;
|
|
90
92
|
|
|
91
93
|
try {
|
|
92
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sger-x", false);
|
|
93
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sger-y", false);
|
|
94
|
-
ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sger-A", true);
|
|
95
|
-
paramsBuffer = createParamsBuffer(
|
|
94
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sger-x", false);
|
|
95
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sger-y", false);
|
|
96
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sger-A", true);
|
|
97
|
+
paramsBuffer = createParamsBuffer(device,
|
|
96
98
|
[
|
|
97
99
|
{ value: m, type: "u32" },
|
|
98
100
|
{ value: n, type: "u32" },
|
|
@@ -104,7 +106,7 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
|
|
|
104
106
|
"sger-params",
|
|
105
107
|
);
|
|
106
108
|
|
|
107
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
109
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
108
110
|
xBuffer,
|
|
109
111
|
yBuffer,
|
|
110
112
|
ABuffer,
|
|
@@ -114,10 +116,10 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
|
|
|
114
116
|
// One workgroup per row of A; clamped to device limit — the shader's
|
|
115
117
|
// grid-stride loop handles remaining rows when m > dispatch count.
|
|
116
118
|
const wgCount = Math.min(m, device.limits.maxComputeWorkgroupsPerDimension);
|
|
117
|
-
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
118
|
-
const readBuffer = AIsGpu ? null : stageReadback(commandEncoder, ABuffer);
|
|
119
|
+
const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
|
|
120
|
+
const readBuffer = AIsGpu ? null : stageReadback(device, commandEncoder, ABuffer);
|
|
119
121
|
|
|
120
|
-
submit(commandEncoder);
|
|
122
|
+
submit(device, commandEncoder);
|
|
121
123
|
|
|
122
124
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
123
125
|
|
package/src/shaders/index.mjs
CHANGED
|
@@ -3,25 +3,175 @@
|
|
|
3
3
|
*
|
|
4
4
|
* `shaders/*.wgsl` — one WGSL compute shader per BLAS routine (sscal, saxpy, sdot, …).
|
|
5
5
|
*
|
|
6
|
-
* `
|
|
7
|
-
*
|
|
8
|
-
*
|
|
6
|
+
* `routineShaders` below is the single source of truth: routine name → the WGSL source(s)
|
|
7
|
+
* its `getPipeline()` calls actually reference, verified against every `src/<routine>/<routine>.mjs`
|
|
8
|
+
* rather than inferred from naming convention (see its doc comment for the exceptions). Each
|
|
9
|
+
* shader is imported right above the line that adds it — the import *is* the mapping entry, no
|
|
10
|
+
* separate block to cross-reference. `shaderSources`, the flat name → source registry the
|
|
11
|
+
* browser bundle's runtime lookup needs, is *derived* from `routineShaders` rather than
|
|
12
|
+
* hand-duplicated, so the two can never drift apart. In Node.js neither is read — shaders are
|
|
13
|
+
* `readFileSync` from disk directly; `scripts/build-browser.mjs` inlines this module into the
|
|
14
|
+
* browser's IIFE bundle via esbuild instead.
|
|
9
15
|
*
|
|
10
16
|
* ## Cross-shader patterns
|
|
11
17
|
*
|
|
12
|
-
* **
|
|
13
|
-
*
|
|
14
|
-
*
|
|
18
|
+
* **Single bind group.** Every shader with bindings uses `@group(0)` only — the JS side always
|
|
19
|
+
* calls `pipeline.getBindGroupLayout(0)`, no secondary groups to track. Binding order is
|
|
20
|
+
* consistent too: any read-only storage buffers come before read_write ones, with the
|
|
21
|
+
* `uniform Params` struct always last. `@binding` indices match the position of each resource in
|
|
22
|
+
* the array passed to `createBindGroup`, which appends `resultBuffer` last.
|
|
15
23
|
*
|
|
16
|
-
* **
|
|
17
|
-
* `
|
|
24
|
+
* **All counts and strides are `u32`.** `n`, `x_inc`, `y_inc`, and every other index/count field
|
|
25
|
+
* in a `Params` struct is unsigned, avoiding implicit sign-extension in index expressions like
|
|
26
|
+
* `id * params.x_inc`.
|
|
18
27
|
*
|
|
19
|
-
*
|
|
20
|
-
*
|
|
21
|
-
*
|
|
22
|
-
* **All counts and strides are `u32`.** `n`, `x_inc`, `y_inc`, and any other index fields in the
|
|
23
|
-
* `Params` uniform struct are unsigned. This avoids implicit sign-extension when they appear in
|
|
24
|
-
* index expressions like `id * params.x_inc`.
|
|
28
|
+
* **Entry points don't have to be named `main`.** `loadShader` (`util/pipeline.mjs`)
|
|
29
|
+
* auto-detects the sole `@compute` function in a module instead of requiring a fixed name, so
|
|
30
|
+
* `dasum_main`, `strsv_invert_block_main`, etc. work without renaming.
|
|
25
31
|
*
|
|
26
32
|
* @module devdocs/shaders
|
|
27
33
|
*/
|
|
34
|
+
|
|
35
|
+
/**
|
|
36
|
+
* Routine name → the WGSL source(s) its `getPipeline()` calls reference. Keys are the exact
|
|
37
|
+
* shader names `getPipeline(device, name)` is called with — most routines have one, some pick
|
|
38
|
+
* one of several conditionally (e.g. sgemv's `sgemv_n`/`sgemv_t`, by `trans`), and some have no
|
|
39
|
+
* dedicated shader at all:
|
|
40
|
+
*
|
|
41
|
+
* - `sgemmtr`/`ssyrk`/`ssyr2k` all dispatch through `sgemmtr_small`/`sgemmtr_large`.
|
|
42
|
+
* - `strsm` reuses `strsv_invert_block` and `sscal`, plus its own `block_transfer` and the
|
|
43
|
+
* shared `sgemm_small`/`sgemm_large`.
|
|
44
|
+
* - `dasum`/`idamax` concatenate several f64 utility shaders with their own — see
|
|
45
|
+
* `getPipeline`'s `shaderName: string[]` behaviour.
|
|
46
|
+
* - `random` has no entry — CPU-only, no `getPipeline()` call.
|
|
47
|
+
*
|
|
48
|
+
* Built up entry by entry so each import sits next to the mapping entry that uses it.
|
|
49
|
+
* @public
|
|
50
|
+
*/
|
|
51
|
+
export const routineShaders = {};
|
|
52
|
+
|
|
53
|
+
import sscal from "./sscal.wgsl";
|
|
54
|
+
routineShaders.sscal = { sscal };
|
|
55
|
+
|
|
56
|
+
import sswap from "./sswap.wgsl";
|
|
57
|
+
routineShaders.sswap = { sswap };
|
|
58
|
+
|
|
59
|
+
import saxpy from "./saxpy.wgsl";
|
|
60
|
+
routineShaders.saxpy = { saxpy };
|
|
61
|
+
|
|
62
|
+
import scopy from "./scopy.wgsl";
|
|
63
|
+
routineShaders.scopy = { scopy };
|
|
64
|
+
|
|
65
|
+
import sdot from "./sdot.wgsl";
|
|
66
|
+
import sum from "./reduction/sum.wgsl";
|
|
67
|
+
routineShaders.sdot = { sdot, "reduction/sum": sum };
|
|
68
|
+
|
|
69
|
+
import sasum from "./sasum.wgsl";
|
|
70
|
+
routineShaders.sasum = { sasum, "reduction/sum": sum };
|
|
71
|
+
|
|
72
|
+
import snrm2 from "./snrm2.wgsl";
|
|
73
|
+
import scaledSum from "./reduction/scaledSum.wgsl";
|
|
74
|
+
routineShaders.snrm2 = { snrm2, "reduction/scaledSum": scaledSum };
|
|
75
|
+
|
|
76
|
+
import isamax from "./isamax.wgsl";
|
|
77
|
+
import argmax from "./reduction/argmax.wgsl";
|
|
78
|
+
routineShaders.isamax = { isamax, "reduction/argmax": argmax };
|
|
79
|
+
|
|
80
|
+
import dekker from "./f64/dekker.wgsl";
|
|
81
|
+
import ddAbs from "./f64/utils/abs.wgsl";
|
|
82
|
+
import ddAddUtil from "./f64/utils/add.wgsl";
|
|
83
|
+
import dasum from "./dasum.wgsl";
|
|
84
|
+
import sumF64 from "./reduction/sumF64.wgsl";
|
|
85
|
+
routineShaders.dasum = {
|
|
86
|
+
"f64/dekker": dekker,
|
|
87
|
+
"f64/utils/abs": ddAbs,
|
|
88
|
+
"f64/utils/add": ddAddUtil,
|
|
89
|
+
dasum,
|
|
90
|
+
"reduction/sumF64": sumF64,
|
|
91
|
+
};
|
|
92
|
+
|
|
93
|
+
import ddGreater from "./f64/utils/greater.wgsl";
|
|
94
|
+
import ddEqual from "./f64/utils/equal.wgsl";
|
|
95
|
+
import idamax from "./idamax.wgsl";
|
|
96
|
+
import argmaxF64 from "./reduction/argmaxF64.wgsl";
|
|
97
|
+
routineShaders.idamax = {
|
|
98
|
+
"f64/dekker": dekker,
|
|
99
|
+
"f64/utils/abs": ddAbs,
|
|
100
|
+
"f64/utils/greater": ddGreater,
|
|
101
|
+
"f64/utils/equal": ddEqual,
|
|
102
|
+
idamax,
|
|
103
|
+
"reduction/argmaxF64": argmaxF64,
|
|
104
|
+
};
|
|
105
|
+
|
|
106
|
+
import srot from "./srot.wgsl";
|
|
107
|
+
routineShaders.srot = { srot };
|
|
108
|
+
|
|
109
|
+
import srotm from "./srotm.wgsl";
|
|
110
|
+
routineShaders.srotm = { srotm };
|
|
111
|
+
|
|
112
|
+
import sgemv_n from "./sgemv_n.wgsl";
|
|
113
|
+
import sgemv_t from "./sgemv_t.wgsl";
|
|
114
|
+
routineShaders.sgemv = { sgemv_n, sgemv_t }; // one or the other, picked by trans
|
|
115
|
+
|
|
116
|
+
import ssymv from "./ssymv.wgsl";
|
|
117
|
+
routineShaders.ssymv = { ssymv };
|
|
118
|
+
|
|
119
|
+
import strmv from "./strmv.wgsl";
|
|
120
|
+
routineShaders.strmv = { strmv };
|
|
121
|
+
|
|
122
|
+
import strsv_invert_block from "./strsv_invert_block.wgsl";
|
|
123
|
+
import strsv_apply_inverse from "./strsv_apply_inverse.wgsl";
|
|
124
|
+
import strsv_update from "./strsv_update.wgsl";
|
|
125
|
+
routineShaders.strsv = {
|
|
126
|
+
strsv_invert_block,
|
|
127
|
+
strsv_apply_inverse,
|
|
128
|
+
strsv_update,
|
|
129
|
+
};
|
|
130
|
+
|
|
131
|
+
import sger from "./sger.wgsl";
|
|
132
|
+
routineShaders.sger = { sger };
|
|
133
|
+
|
|
134
|
+
import ssyr from "./ssyr.wgsl";
|
|
135
|
+
routineShaders.ssyr = { ssyr };
|
|
136
|
+
|
|
137
|
+
import ssyr2 from "./ssyr2.wgsl";
|
|
138
|
+
routineShaders.ssyr2 = { ssyr2 };
|
|
139
|
+
|
|
140
|
+
import sgemm_small from "./sgemm_small.wgsl";
|
|
141
|
+
import sgemm_large from "./sgemm_large.wgsl";
|
|
142
|
+
routineShaders.sgemm = { sgemm_small, sgemm_large }; // one or the other, picked by a tile-size threshold
|
|
143
|
+
|
|
144
|
+
import sgemmtr_small from "./sgemmtr_small.wgsl";
|
|
145
|
+
import sgemmtr_large from "./sgemmtr_large.wgsl";
|
|
146
|
+
routineShaders.sgemmtr = { sgemmtr_small, sgemmtr_large };
|
|
147
|
+
|
|
148
|
+
routineShaders.ssyrk = { sgemmtr_small, sgemmtr_large }; // no shader of its own — rides on sgemmtr's
|
|
149
|
+
routineShaders.ssyr2k = { sgemmtr_small, sgemmtr_large }; // no shader of its own — rides on sgemmtr's
|
|
150
|
+
|
|
151
|
+
import symmetrize from "./symmetrize.wgsl";
|
|
152
|
+
routineShaders.ssymm = { sgemm_small, sgemm_large, symmetrize };
|
|
153
|
+
|
|
154
|
+
import triangularize from "./triangularize.wgsl";
|
|
155
|
+
routineShaders.strmm = { sgemm_small, sgemm_large, triangularize };
|
|
156
|
+
|
|
157
|
+
import blockTransfer from "./block_transfer.wgsl";
|
|
158
|
+
routineShaders.strsm = {
|
|
159
|
+
strsv_invert_block,
|
|
160
|
+
block_transfer: blockTransfer,
|
|
161
|
+
sscal,
|
|
162
|
+
sgemm_small,
|
|
163
|
+
sgemm_large,
|
|
164
|
+
};
|
|
165
|
+
|
|
166
|
+
/**
|
|
167
|
+
* Flat shader-name → WGSL source-string registry — what `getPipeline()`/`loadShader()` (see
|
|
168
|
+
* `util/pipeline.mjs`) actually look shaders up in, in the browser. Derived from
|
|
169
|
+
* `routineShaders` by merging every routine's shaders together; shared shaders (e.g.
|
|
170
|
+
* `"reduction/sum"`, used by two different routines above) collapse harmlessly here since
|
|
171
|
+
* every routine's copy is the same imported string, never independently authored text.
|
|
172
|
+
* @public
|
|
173
|
+
*/
|
|
174
|
+
export const shaderSources = Object.assign(
|
|
175
|
+
{},
|
|
176
|
+
...Object.values(routineShaders),
|
|
177
|
+
);
|