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/snrm2/snrm2.mjs
CHANGED
|
@@ -12,14 +12,16 @@ 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 snrm2(device, n, x, incx) {
|
|
19
20
|
const xIsGpu = x instanceof GpuVector;
|
|
20
21
|
|
|
21
22
|
if (!(device instanceof GPUDevice))
|
|
22
23
|
throw new Error("device must be a GPUDevice.");
|
|
24
|
+
requireSameDevice(device, "snrm2", { x });
|
|
23
25
|
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
24
26
|
throw new Error("n and incx must be integers.");
|
|
25
27
|
if (incx <= 0) throw new Error("incx must be positive.");
|
|
@@ -32,19 +34,25 @@ export async function snrm2(device, n, x, incx) {
|
|
|
32
34
|
);
|
|
33
35
|
|
|
34
36
|
const pipelineMain = await getPipeline(device, "snrm2");
|
|
35
|
-
const pipelineReduce = await getPipeline(device, "reduction/
|
|
37
|
+
const pipelineReduce = await getPipeline(device, "reduction/scaledSum");
|
|
36
38
|
|
|
37
39
|
let xBuffer = null;
|
|
38
|
-
let
|
|
40
|
+
let partialsScaleBuffer = null;
|
|
41
|
+
let partialsSsqBuffer = null;
|
|
39
42
|
let resultBuffer = null;
|
|
40
43
|
let paramsBuffer = null;
|
|
41
44
|
let readBuffer = null;
|
|
42
45
|
|
|
43
46
|
try {
|
|
44
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "snrm2-x", false);
|
|
45
|
-
|
|
46
|
-
|
|
47
|
-
|
|
47
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "snrm2-x", false);
|
|
48
|
+
// 2*WGS partial (scale, ssq) pairs — see snrm2.wgsl for what they represent.
|
|
49
|
+
partialsScaleBuffer = createStorageBuffer(device,
|
|
50
|
+
2 * WGS * 4,
|
|
51
|
+
"snrm2-partials-scale",
|
|
52
|
+
);
|
|
53
|
+
partialsSsqBuffer = createStorageBuffer(device, 2 * WGS * 4, "snrm2-partials-ssq");
|
|
54
|
+
resultBuffer = createResultBuffer(device, 4, "snrm2-result"); // final f32 scalar
|
|
55
|
+
paramsBuffer = createParamsBuffer(device,
|
|
48
56
|
[
|
|
49
57
|
{ value: n, type: "u32" },
|
|
50
58
|
{ value: incx, type: "u32" },
|
|
@@ -52,53 +60,57 @@ export async function snrm2(device, n, x, incx) {
|
|
|
52
60
|
"snrm2-params",
|
|
53
61
|
);
|
|
54
62
|
|
|
55
|
-
const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
|
|
63
|
+
const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
|
|
56
64
|
xBuffer,
|
|
57
|
-
|
|
65
|
+
partialsScaleBuffer,
|
|
66
|
+
partialsSsqBuffer,
|
|
58
67
|
paramsBuffer,
|
|
59
68
|
]);
|
|
60
|
-
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
69
|
+
const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
|
|
61
70
|
pipelineMain,
|
|
62
71
|
bgMain,
|
|
63
72
|
2 * WGS,
|
|
64
73
|
); // dispatch 2*WGS workgroups
|
|
65
74
|
|
|
66
|
-
submit(enc1);
|
|
75
|
+
submit(device, enc1);
|
|
67
76
|
|
|
68
|
-
const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
|
|
69
|
-
|
|
77
|
+
const bgReduce = createBindGroup(device, pipelineReduce.getBindGroupLayout(0), [
|
|
78
|
+
partialsScaleBuffer,
|
|
79
|
+
partialsSsqBuffer,
|
|
70
80
|
resultBuffer,
|
|
71
81
|
]);
|
|
72
|
-
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
82
|
+
const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
|
|
73
83
|
pipelineReduce,
|
|
74
84
|
bgReduce,
|
|
75
85
|
1,
|
|
76
86
|
); // reduce partials to single result
|
|
77
|
-
readBuffer = stageReadback(enc2, resultBuffer);
|
|
87
|
+
readBuffer = stageReadback(device, enc2, resultBuffer);
|
|
78
88
|
|
|
79
|
-
submit(enc2);
|
|
89
|
+
submit(device, enc2);
|
|
80
90
|
|
|
81
91
|
const resultPromise = extractResult(readBuffer, Float32Array);
|
|
82
92
|
readBuffer = null; // ownership transferred — extractResult's own finally destroys it
|
|
83
93
|
|
|
84
|
-
const [gpuTime1, gpuTime2,
|
|
94
|
+
const [gpuTime1, gpuTime2, resultArr] = await Promise.all([
|
|
85
95
|
extractTimestamp(ts1),
|
|
86
96
|
extractTimestamp(ts2),
|
|
87
97
|
resultPromise,
|
|
88
98
|
]);
|
|
89
99
|
|
|
90
|
-
//
|
|
91
|
-
|
|
100
|
+
// reduction/scaledSum.wgsl already computes scale·sqrt(ssq) on the GPU —
|
|
101
|
+
// unlike the old naive-sum version, there's no separate sqrt step here.
|
|
102
|
+
const nrm2 = resultArr[0];
|
|
92
103
|
|
|
93
104
|
if (gpuTime1 !== undefined && gpuTime2 !== undefined)
|
|
94
105
|
return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
|
|
95
106
|
return { nrm2 };
|
|
96
107
|
} finally {
|
|
97
108
|
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
98
|
-
if (
|
|
109
|
+
if (partialsScaleBuffer) destroyBuffers(partialsScaleBuffer);
|
|
110
|
+
if (partialsSsqBuffer) destroyBuffers(partialsSsqBuffer);
|
|
99
111
|
if (resultBuffer) destroyBuffers(resultBuffer);
|
|
100
112
|
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
101
|
-
// Only reached if submit(enc2) threw before ownership was transferred above.
|
|
113
|
+
// Only reached if submit(device, enc2) threw before ownership was transferred above.
|
|
102
114
|
if (readBuffer) destroyBuffers(readBuffer);
|
|
103
115
|
}
|
|
104
116
|
}
|
package/src/srot/srot.d.mts
CHANGED
|
@@ -2,8 +2,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
|
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
4
|
* Applies a Givens plane rotation to vectors x and y:
|
|
5
|
-
*
|
|
6
|
-
* y = -s*x + c*y
|
|
5
|
+
* $$\begin{aligned} x &\leftarrow cx + sy \\\\ y &\leftarrow -sx + cy \end{aligned}$$
|
|
7
6
|
*
|
|
8
7
|
* {@includeCode ../../examples/srot/srot.js}
|
|
9
8
|
*
|
|
@@ -34,8 +33,7 @@ export declare function srot(
|
|
|
34
33
|
|
|
35
34
|
/**
|
|
36
35
|
* Applies a Givens plane rotation to vectors x and y:
|
|
37
|
-
*
|
|
38
|
-
* y = -s*x + c*y
|
|
36
|
+
* $$\begin{aligned} x &\leftarrow cx + sy \\\\ y &\leftarrow -sx + cy \end{aligned}$$
|
|
39
37
|
*
|
|
40
38
|
* {@includeCode ../../examples/srot/gpu.srot.js}
|
|
41
39
|
*
|
package/src/srot/srot.mjs
CHANGED
|
@@ -11,6 +11,7 @@ import { extractResult } from "../util/result.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
@@ -18,6 +19,7 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
18
19
|
|
|
19
20
|
if (!(device instanceof GPUDevice))
|
|
20
21
|
throw new Error("device must be a GPUDevice.");
|
|
22
|
+
requireSameDevice(device, "srot", { x, y });
|
|
21
23
|
if (
|
|
22
24
|
!Number.isInteger(n) ||
|
|
23
25
|
!Number.isInteger(incx) ||
|
|
@@ -58,9 +60,9 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
58
60
|
let readY = null;
|
|
59
61
|
|
|
60
62
|
try {
|
|
61
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srot-x", true);
|
|
62
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srot-y", true);
|
|
63
|
-
paramsBuffer = createParamsBuffer(
|
|
63
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "srot-x", true);
|
|
64
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "srot-y", true);
|
|
65
|
+
paramsBuffer = createParamsBuffer(device,
|
|
64
66
|
[
|
|
65
67
|
{ value: n, type: "u32" },
|
|
66
68
|
{ value: c, type: "f32" },
|
|
@@ -71,19 +73,19 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
71
73
|
"srot-params",
|
|
72
74
|
);
|
|
73
75
|
|
|
74
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
76
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
75
77
|
xBuffer,
|
|
76
78
|
yBuffer,
|
|
77
79
|
paramsBuffer,
|
|
78
80
|
]);
|
|
79
|
-
const { commandEncoder, ts } = runComputePass(
|
|
81
|
+
const { commandEncoder, ts } = runComputePass(device,
|
|
80
82
|
pipeline,
|
|
81
83
|
bindGroup,
|
|
82
|
-
calcWorkgroups(n),
|
|
84
|
+
calcWorkgroups(device, n),
|
|
83
85
|
);
|
|
84
|
-
readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
85
|
-
readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
86
|
-
submit(commandEncoder);
|
|
86
|
+
readX = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
|
|
87
|
+
readY = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
88
|
+
submit(device, commandEncoder);
|
|
87
89
|
|
|
88
90
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
89
91
|
|
package/src/srotm/srotm.d.mts
CHANGED
|
@@ -2,8 +2,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
|
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
4
|
* Applies a modified Givens plane rotation H to vectors x and y:
|
|
5
|
-
*
|
|
6
|
-
* y = H[1][0]*x + H[1][1]*y
|
|
5
|
+
* $$\begin{pmatrix} x \\\\ y \end{pmatrix} \leftarrow \begin{pmatrix} h_{11} & h_{12} \\\\ h_{21} & h_{22} \end{pmatrix} \begin{pmatrix} x \\\\ y \end{pmatrix}$$
|
|
7
6
|
*
|
|
8
7
|
* {@includeCode ../../examples/srotm/srotm.js}
|
|
9
8
|
*
|
|
@@ -33,8 +32,7 @@ export declare function srotm(
|
|
|
33
32
|
|
|
34
33
|
/**
|
|
35
34
|
* Applies a modified Givens plane rotation H to vectors x and y:
|
|
36
|
-
*
|
|
37
|
-
* y = H[1][0]*x + H[1][1]*y
|
|
35
|
+
* $$\begin{pmatrix} x \\\\ y \end{pmatrix} \leftarrow \begin{pmatrix} h_{11} & h_{12} \\\\ h_{21} & h_{22} \end{pmatrix} \begin{pmatrix} x \\\\ y \end{pmatrix}$$
|
|
38
36
|
*
|
|
39
37
|
* {@includeCode ../../examples/srotm/gpu.srotm.js}
|
|
40
38
|
*
|
package/src/srotm/srotm.mjs
CHANGED
|
@@ -11,6 +11,7 @@ import { extractResult } from "../util/result.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function srotm(device, n, x, incx, y, incy, param) {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
@@ -18,6 +19,7 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
18
19
|
|
|
19
20
|
if (!(device instanceof GPUDevice))
|
|
20
21
|
throw new Error("device must be a GPUDevice.");
|
|
22
|
+
requireSameDevice(device, "srotm", { x, y });
|
|
21
23
|
if (
|
|
22
24
|
!Number.isInteger(n) ||
|
|
23
25
|
!Number.isInteger(incx) ||
|
|
@@ -26,6 +28,13 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
26
28
|
throw new Error("n, incx, and incy must be integers.");
|
|
27
29
|
if (!(param instanceof Float32Array) || param.length !== 5)
|
|
28
30
|
throw new Error("param must be a Float32Array of length 5.");
|
|
31
|
+
if (
|
|
32
|
+
param[0] !== -2 &&
|
|
33
|
+
param[0] !== -1 &&
|
|
34
|
+
param[0] !== 0 &&
|
|
35
|
+
param[0] !== 1
|
|
36
|
+
)
|
|
37
|
+
throw new Error("param[0] (flag) must be one of -2, -1, 0, or 1.");
|
|
29
38
|
if (incx <= 0 || incy <= 0)
|
|
30
39
|
throw new Error("incx and incy must be positive.");
|
|
31
40
|
if (!xIsGpu && !(x instanceof Float32Array))
|
|
@@ -56,10 +65,10 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
56
65
|
let readY = null;
|
|
57
66
|
|
|
58
67
|
try {
|
|
59
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srotm-x", true);
|
|
60
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srotm-y", true);
|
|
61
|
-
paramBuffer = uploadBuffer(param, "srotm-param", false);
|
|
62
|
-
paramsBuffer = createParamsBuffer(
|
|
68
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "srotm-x", true);
|
|
69
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "srotm-y", true);
|
|
70
|
+
paramBuffer = uploadBuffer(device, param, "srotm-param", false);
|
|
71
|
+
paramsBuffer = createParamsBuffer(device,
|
|
63
72
|
[
|
|
64
73
|
{ value: n, type: "u32" },
|
|
65
74
|
{ value: incx, type: "u32" },
|
|
@@ -68,20 +77,20 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
68
77
|
"srotm-params",
|
|
69
78
|
);
|
|
70
79
|
|
|
71
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
80
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
72
81
|
xBuffer,
|
|
73
82
|
yBuffer,
|
|
74
83
|
paramBuffer,
|
|
75
84
|
paramsBuffer,
|
|
76
85
|
]);
|
|
77
|
-
const { commandEncoder, ts } = runComputePass(
|
|
86
|
+
const { commandEncoder, ts } = runComputePass(device,
|
|
78
87
|
pipeline,
|
|
79
88
|
bindGroup,
|
|
80
|
-
calcWorkgroups(n),
|
|
89
|
+
calcWorkgroups(device, n),
|
|
81
90
|
);
|
|
82
|
-
readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
83
|
-
readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
84
|
-
submit(commandEncoder);
|
|
91
|
+
readX = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
|
|
92
|
+
readY = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
93
|
+
submit(device, commandEncoder);
|
|
85
94
|
|
|
86
95
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
87
96
|
|
package/src/sscal/sscal.d.mts
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
|
-
* Scales a single-precision vector by a constant: x
|
|
4
|
+
* Scales a single-precision vector by a constant: $$x \leftarrow \alpha x$$
|
|
5
5
|
*
|
|
6
6
|
* {@includeCode ../../examples/sscal/sscal.js}
|
|
7
7
|
*
|
|
@@ -22,10 +22,10 @@ export declare function sscal(
|
|
|
22
22
|
alpha: number,
|
|
23
23
|
x: Float32Array,
|
|
24
24
|
incx: number,
|
|
25
|
-
): Promise<Float32Array | {
|
|
25
|
+
): Promise<{ x: Float32Array } | { x: Float32Array; gpuTimeMs: number }>;
|
|
26
26
|
|
|
27
27
|
/**
|
|
28
|
-
* Scales a single-precision vector by a constant: x
|
|
28
|
+
* Scales a single-precision vector by a constant: $$x \leftarrow \alpha x$$
|
|
29
29
|
*
|
|
30
30
|
* {@includeCode ../../examples/sscal/gpu.sscal.js}
|
|
31
31
|
*
|
package/src/sscal/sscal.mjs
CHANGED
|
@@ -11,12 +11,14 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function sscal(device, n, alpha, x, incx) {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
17
|
-
|
|
18
|
+
|
|
18
19
|
if (!(device instanceof GPUDevice))
|
|
19
20
|
throw new Error("device must be a GPUDevice.");
|
|
21
|
+
requireSameDevice(device, "sscal", { x });
|
|
20
22
|
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
21
23
|
throw new Error("n and incx must be integers.");
|
|
22
24
|
if (typeof alpha !== "number")
|
|
@@ -26,7 +28,7 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
26
28
|
if (incx <= 0) throw new Error("incx must be positive.");
|
|
27
29
|
if (!(x instanceof Float32Array) && !(x instanceof GpuVector))
|
|
28
30
|
throw new Error("x must be a Float32Array or GpuVector.");
|
|
29
|
-
if (n <= 0) return xIsGpu ? {} : x;
|
|
31
|
+
if (n <= 0) return xIsGpu ? {} : { x };
|
|
30
32
|
if (x.length < (n - 1) * incx + 1)
|
|
31
33
|
throw new Error(
|
|
32
34
|
"x does not have enough elements for the given n and incx.",
|
|
@@ -39,8 +41,8 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
39
41
|
let readBuffer = null;
|
|
40
42
|
|
|
41
43
|
try {
|
|
42
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sscal-x", true);
|
|
43
|
-
paramsBuffer = createParamsBuffer(
|
|
44
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sscal-x", true);
|
|
45
|
+
paramsBuffer = createParamsBuffer(device,
|
|
44
46
|
[
|
|
45
47
|
{ value: n, type: "u32" },
|
|
46
48
|
{ value: alpha, type: "f32" },
|
|
@@ -49,18 +51,18 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
49
51
|
"sscal-params",
|
|
50
52
|
);
|
|
51
53
|
|
|
52
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
54
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
53
55
|
xBuffer,
|
|
54
56
|
paramsBuffer,
|
|
55
57
|
]);
|
|
56
|
-
const { commandEncoder, ts } = runComputePass(
|
|
58
|
+
const { commandEncoder, ts } = runComputePass(device,
|
|
57
59
|
pipeline,
|
|
58
60
|
bindGroup,
|
|
59
|
-
calcWorkgroups(n),
|
|
61
|
+
calcWorkgroups(device, n),
|
|
60
62
|
);
|
|
61
|
-
readBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
63
|
+
readBuffer = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
|
|
62
64
|
|
|
63
|
-
submit(commandEncoder);
|
|
65
|
+
submit(device, commandEncoder);
|
|
64
66
|
|
|
65
67
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
66
68
|
|
|
@@ -72,7 +74,7 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
72
74
|
const result = await extractResult(readBuffer, Float32Array);
|
|
73
75
|
readBuffer = null; // extractResult already destroyed it
|
|
74
76
|
if (gpuTimeMs !== undefined) return { x: result, gpuTimeMs };
|
|
75
|
-
return result;
|
|
77
|
+
return { x: result };
|
|
76
78
|
} finally {
|
|
77
79
|
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
78
80
|
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
package/src/sswap/sswap.d.mts
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
|
-
* Swaps the elements of two single-precision vectors: x
|
|
4
|
+
* Swaps the elements of two single-precision vectors: $$x \leftrightarrow y$$
|
|
5
5
|
*
|
|
6
6
|
* {@includeCode ../../examples/sswap/sswap.js}
|
|
7
7
|
*
|
|
@@ -27,7 +27,7 @@ export declare function sswap(
|
|
|
27
27
|
): Promise<{ x: Float32Array; y: Float32Array } | { x: Float32Array; y: Float32Array; gpuTimeMs: number }>;
|
|
28
28
|
|
|
29
29
|
/**
|
|
30
|
-
* Swaps the elements of two single-precision vectors: x
|
|
30
|
+
* Swaps the elements of two single-precision vectors: $$x \leftrightarrow y$$
|
|
31
31
|
*
|
|
32
32
|
* {@includeCode ../../examples/sswap/gpu.sswap.js}
|
|
33
33
|
*
|
package/src/sswap/sswap.mjs
CHANGED
|
@@ -11,6 +11,7 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function sswap(device, n, x, incx, y, incy) {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
@@ -18,6 +19,7 @@ export async function sswap(device, n, x, incx, y, incy) {
|
|
|
18
19
|
|
|
19
20
|
if (!(device instanceof GPUDevice))
|
|
20
21
|
throw new Error("device must be a GPUDevice.");
|
|
22
|
+
requireSameDevice(device, "sswap", { x, y });
|
|
21
23
|
if (
|
|
22
24
|
!Number.isInteger(n) ||
|
|
23
25
|
!Number.isInteger(incx) ||
|
|
@@ -53,9 +55,9 @@ export async function sswap(device, n, x, incx, y, incy) {
|
|
|
53
55
|
let yReadBuffer = null;
|
|
54
56
|
|
|
55
57
|
try {
|
|
56
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sswap-x", true);
|
|
57
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sswap-y", true);
|
|
58
|
-
paramsBuffer = createParamsBuffer(
|
|
58
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sswap-x", true);
|
|
59
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sswap-y", true);
|
|
60
|
+
paramsBuffer = createParamsBuffer(device,
|
|
59
61
|
[
|
|
60
62
|
{ value: n, type: "u32" },
|
|
61
63
|
{ value: incx, type: "u32" },
|
|
@@ -64,20 +66,20 @@ export async function sswap(device, n, x, incx, y, incy) {
|
|
|
64
66
|
"sswap-params",
|
|
65
67
|
);
|
|
66
68
|
|
|
67
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
69
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
68
70
|
xBuffer,
|
|
69
71
|
yBuffer,
|
|
70
72
|
paramsBuffer,
|
|
71
73
|
]);
|
|
72
|
-
const { commandEncoder, ts } = runComputePass(
|
|
74
|
+
const { commandEncoder, ts } = runComputePass(device,
|
|
73
75
|
pipeline,
|
|
74
76
|
bindGroup,
|
|
75
|
-
calcWorkgroups(n),
|
|
77
|
+
calcWorkgroups(device, n),
|
|
76
78
|
);
|
|
77
|
-
xReadBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
78
|
-
yReadBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
79
|
+
xReadBuffer = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
|
|
80
|
+
yReadBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
79
81
|
|
|
80
|
-
submit(commandEncoder);
|
|
82
|
+
submit(device, commandEncoder);
|
|
81
83
|
|
|
82
84
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
83
85
|
|
package/src/ssymm/ssymm.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 use
|
|
18
|
-
const SYM_WG = 8; // symmetrize.wgsl's @workgroup_size(8, 8)
|
|
19
20
|
|
|
20
21
|
// ssymm: C := alpha*A*B + beta*C (side='left') or alpha*B*A + beta*C
|
|
21
22
|
// (side='right'), A symmetric. No fused kernel — symmetrize then sgemm,
|
|
@@ -29,6 +30,7 @@ export async function ssymm(
|
|
|
29
30
|
|
|
30
31
|
if (!(device instanceof GPUDevice))
|
|
31
32
|
throw new Error("device must be a GPUDevice.");
|
|
33
|
+
requireSameDevice(device, "ssymm", { A, B, C });
|
|
32
34
|
if (side !== "left" && side !== "right")
|
|
33
35
|
throw new Error("side must be 'left' or 'right'.");
|
|
34
36
|
if (uplo !== "lower" && uplo !== "upper")
|
|
@@ -128,22 +130,22 @@ export async function ssymm(
|
|
|
128
130
|
const symPipeline = await getPipeline(device, "symmetrize");
|
|
129
131
|
const gemmWgCount = useLargeTile
|
|
130
132
|
? {
|
|
131
|
-
x:
|
|
132
|
-
y:
|
|
133
|
+
x: requireWorkgroupCount(device, largeWgX, "ssymm", "x"),
|
|
134
|
+
y: requireWorkgroupCount(device, largeWgY, "ssymm", "y"),
|
|
133
135
|
}
|
|
134
136
|
: {
|
|
135
|
-
x:
|
|
136
|
-
y:
|
|
137
|
+
x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "ssymm", "x"),
|
|
138
|
+
y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "ssymm", "y"),
|
|
137
139
|
};
|
|
138
140
|
|
|
139
|
-
const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssymm-A", false);
|
|
140
|
-
const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "ssymm-B", false);
|
|
141
|
-
const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "ssymm-C", true);
|
|
142
|
-
const AdenseBuffer = createStorageBuffer(aOrder * ldDense * 4, "ssymm-Adense");
|
|
141
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssymm-A", false);
|
|
142
|
+
const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "ssymm-B", false);
|
|
143
|
+
const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssymm-C", true);
|
|
144
|
+
const AdenseBuffer = createStorageBuffer(device, aOrder * ldDense * 4, "ssymm-Adense");
|
|
143
145
|
let symParams = null, gemmParams = null;
|
|
144
146
|
|
|
145
147
|
try {
|
|
146
|
-
symParams = createParamsBuffer(
|
|
148
|
+
symParams = createParamsBuffer(device,
|
|
147
149
|
[
|
|
148
150
|
{ value: aOrder, type: "u32" },
|
|
149
151
|
{ value: lda, type: "u32" },
|
|
@@ -152,7 +154,7 @@ export async function ssymm(
|
|
|
152
154
|
],
|
|
153
155
|
"ssymm-sym-params",
|
|
154
156
|
);
|
|
155
|
-
const symBindGroup = createBindGroup(symPipeline.getBindGroupLayout(0), [ABuffer, AdenseBuffer, symParams]);
|
|
157
|
+
const symBindGroup = createBindGroup(device, symPipeline.getBindGroupLayout(0), [ABuffer, AdenseBuffer, symParams]);
|
|
156
158
|
|
|
157
159
|
// X/Y buffers and their own ld, matching swapXY above.
|
|
158
160
|
const XBuffer = swapXY ? BBuffer : AdenseBuffer;
|
|
@@ -160,7 +162,7 @@ export async function ssymm(
|
|
|
160
162
|
const YBuffer = swapXY ? AdenseBuffer : BBuffer;
|
|
161
163
|
const ldY = swapXY ? ldDense : ldb;
|
|
162
164
|
|
|
163
|
-
gemmParams = createParamsBuffer(
|
|
165
|
+
gemmParams = createParamsBuffer(device,
|
|
164
166
|
[
|
|
165
167
|
{ value: mg, type: "u32" },
|
|
166
168
|
{ value: ng, type: "u32" },
|
|
@@ -175,18 +177,25 @@ export async function ssymm(
|
|
|
175
177
|
],
|
|
176
178
|
"ssymm-gemm-params",
|
|
177
179
|
);
|
|
178
|
-
const gemmBindGroup = createBindGroup(gemmPipeline.getBindGroupLayout(0), [
|
|
179
|
-
|
|
180
|
-
|
|
180
|
+
const gemmBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
|
|
181
|
+
XBuffer,
|
|
182
|
+
vec4ViewBinding(device, XBuffer),
|
|
183
|
+
YBuffer,
|
|
184
|
+
vec4ViewBinding(device, YBuffer),
|
|
185
|
+
CBuffer,
|
|
186
|
+
gemmParams,
|
|
187
|
+
]);
|
|
188
|
+
|
|
189
|
+
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
181
190
|
const symDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
|
|
182
191
|
const gemmDesc = querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
|
|
183
|
-
encodePass(commandEncoder, symPipeline, symBindGroup, { x: Math.ceil(aOrder /
|
|
192
|
+
encodePass(commandEncoder, symPipeline, symBindGroup, { x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) }, symDesc);
|
|
184
193
|
encodePass(commandEncoder, gemmPipeline, gemmBindGroup, gemmWgCount, gemmDesc);
|
|
185
194
|
|
|
186
|
-
const ts = resolveTimestamp(commandEncoder, querySet);
|
|
187
|
-
const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
|
|
195
|
+
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
196
|
+
const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
|
|
188
197
|
|
|
189
|
-
submit(commandEncoder);
|
|
198
|
+
submit(device, commandEncoder);
|
|
190
199
|
|
|
191
200
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
192
201
|
|
package/src/ssymv/ssymv.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 ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, incy, layout = "row-major") {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
@@ -19,6 +20,7 @@ export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, in
|
|
|
19
20
|
|
|
20
21
|
if (!(device instanceof GPUDevice))
|
|
21
22
|
throw new Error("device must be a GPUDevice.");
|
|
23
|
+
requireSameDevice(device, "ssymv", { 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")
|
|
@@ -89,10 +91,10 @@ export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, in
|
|
|
89
91
|
let paramsBuffer = null;
|
|
90
92
|
|
|
91
93
|
try {
|
|
92
|
-
ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssymv-A", false);
|
|
93
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "ssymv-x", false);
|
|
94
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "ssymv-y", true);
|
|
95
|
-
paramsBuffer = createParamsBuffer(
|
|
94
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssymv-A", false);
|
|
95
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "ssymv-x", false);
|
|
96
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "ssymv-y", true);
|
|
97
|
+
paramsBuffer = createParamsBuffer(device,
|
|
96
98
|
[
|
|
97
99
|
{ value: n, type: "u32" },
|
|
98
100
|
{ value: alpha, type: "f32" },
|
|
@@ -105,7 +107,7 @@ export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, in
|
|
|
105
107
|
"ssymv-params",
|
|
106
108
|
);
|
|
107
109
|
|
|
108
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
110
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
109
111
|
ABuffer,
|
|
110
112
|
xBuffer,
|
|
111
113
|
yBuffer,
|
|
@@ -113,10 +115,10 @@ export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, in
|
|
|
113
115
|
]);
|
|
114
116
|
|
|
115
117
|
const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
|
|
116
|
-
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
117
|
-
const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
118
|
+
const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
|
|
119
|
+
const readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
118
120
|
|
|
119
|
-
submit(commandEncoder);
|
|
121
|
+
submit(device, commandEncoder);
|
|
120
122
|
|
|
121
123
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
122
124
|
|