wgblas 2.0.0 → 2.2.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +20 -18
- package/dist/wgblas.browser.js +2172 -1174
- package/index.d.mts +49 -44
- package/index.mjs +11 -0
- package/package.json +133 -63
- package/src/classes/Complex32.d.mts +43 -0
- package/src/classes/Complex32.mjs +82 -0
- package/src/classes/Complex64.d.mts +44 -0
- package/src/classes/Complex64.mjs +76 -0
- package/src/classes/GpuMatrix.d.mts +41 -41
- package/src/classes/GpuMatrix.mjs +126 -17
- package/src/classes/GpuVector.d.mts +36 -40
- package/src/classes/GpuVector.mjs +66 -11
- package/src/cscal/cscal.d.mts +47 -0
- package/src/cscal/cscal.mjs +98 -0
- package/src/dasum/dasum.d.mts +4 -4
- package/src/dasum/dasum.mjs +38 -20
- package/src/daxpy/daxpy.d.mts +56 -0
- package/src/daxpy/daxpy.mjs +150 -0
- package/src/dcopy/dcopy.d.mts +52 -0
- package/src/dcopy/dcopy.mjs +140 -0
- package/src/ddot/ddot.d.mts +62 -0
- package/src/ddot/ddot.mjs +184 -0
- package/src/devdocs.mjs +13 -0
- package/src/dnrm2/dnrm2.d.mts +50 -0
- package/src/dnrm2/dnrm2.mjs +189 -0
- package/src/drot/drot.d.mts +67 -0
- package/src/drot/drot.mjs +170 -0
- package/src/drotm/drotm.d.mts +67 -0
- package/src/drotm/drotm.mjs +171 -0
- package/src/dscal/dscal.d.mts +52 -0
- package/src/dscal/dscal.mjs +119 -0
- package/src/dswap/dswap.d.mts +57 -0
- package/src/dswap/dswap.mjs +155 -0
- package/src/idamax/idamax.d.mts +20 -2
- package/src/idamax/idamax.mjs +56 -24
- package/src/init.mjs +117 -56
- package/src/isamax/isamax.d.mts +20 -2
- package/src/isamax/isamax.mjs +21 -16
- package/src/random/random.d.mts +37 -39
- package/src/random/random.mjs +39 -7
- package/src/sasum/sasum.d.mts +2 -2
- package/src/sasum/sasum.mjs +20 -16
- package/src/saxpy/saxpy.d.mts +2 -2
- package/src/saxpy/saxpy.mjs +14 -11
- package/src/scopy/scopy.d.mts +2 -2
- package/src/scopy/scopy.mjs +13 -9
- package/src/sdot/sdot.d.mts +2 -2
- package/src/sdot/sdot.mjs +21 -17
- package/src/sgemm/sgemm.d.mts +2 -2
- package/src/sgemm/sgemm.mjs +109 -40
- package/src/sgemmtr/sgemmtr.d.mts +3 -2
- package/src/sgemmtr/sgemmtr.mjs +98 -40
- package/src/sgemv/sgemv.d.mts +2 -2
- package/src/sgemv/sgemv.mjs +69 -41
- package/src/sger/sger.d.mts +2 -2
- package/src/sger/sger.mjs +43 -19
- package/src/shaders/__test_pipeline_a.wgsl +3 -0
- package/src/shaders/__test_pipeline_b.wgsl +2 -0
- package/src/shaders/cscal.wgsl +33 -0
- package/src/shaders/daxpy.wgsl +66 -0
- package/src/shaders/dcopy.wgsl +34 -0
- package/src/shaders/ddot.wgsl +106 -0
- package/src/shaders/dnrm2.wgsl +167 -0
- package/src/shaders/drot.wgsl +81 -0
- package/src/shaders/drotm.wgsl +99 -0
- package/src/shaders/dscal.wgsl +60 -0
- package/src/shaders/dswap.wgsl +38 -0
- package/src/shaders/f64/utils/add.wgsl +6 -0
- package/src/shaders/f64/utils/divide.wgsl +45 -0
- package/src/shaders/f64/utils/multiply.wgsl +19 -10
- package/src/shaders/f64/utils/sqrt.wgsl +45 -0
- package/src/shaders/index.mjs +233 -14
- package/src/shaders/reduction/scaledSum.wgsl +65 -0
- package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
- package/src/shaders/sgemm_large.wgsl +107 -18
- package/src/shaders/sgemm_small.wgsl +115 -15
- package/src/shaders/sgemmtr_large.wgsl +4 -1
- package/src/shaders/sgemmtr_small.wgsl +4 -1
- package/src/shaders/sgemv_n.wgsl +3 -1
- package/src/shaders/sgemv_t.wgsl +3 -1
- package/src/shaders/snrm2.wgsl +72 -23
- package/src/shaders/ssymv.wgsl +3 -1
- package/src/snrm2/snrm2.d.mts +2 -2
- package/src/snrm2/snrm2.mjs +41 -23
- package/src/srot/srot.d.mts +2 -4
- package/src/srot/srot.mjs +16 -11
- package/src/srotm/srotm.d.mts +2 -4
- package/src/srotm/srotm.mjs +17 -11
- package/src/sscal/sscal.d.mts +3 -3
- package/src/sscal/sscal.mjs +14 -12
- package/src/sswap/sswap.d.mts +2 -2
- package/src/sswap/sswap.mjs +18 -10
- package/src/ssymm/ssymm.d.mts +5 -4
- package/src/ssymm/ssymm.mjs +150 -54
- package/src/ssymv/ssymv.d.mts +2 -2
- package/src/ssymv/ssymv.mjs +47 -26
- package/src/ssyr/ssyr.d.mts +2 -2
- package/src/ssyr/ssyr.mjs +38 -17
- package/src/ssyr2/ssyr2.d.mts +2 -2
- package/src/ssyr2/ssyr2.mjs +48 -21
- package/src/ssyr2k/ssyr2k.d.mts +3 -2
- package/src/ssyr2k/ssyr2k.mjs +140 -62
- package/src/ssyrk/ssyrk.d.mts +3 -2
- package/src/ssyrk/ssyrk.mjs +91 -39
- package/src/strmm/strmm.d.mts +5 -4
- package/src/strmm/strmm.mjs +174 -60
- package/src/strmv/strmv.d.mts +2 -2
- package/src/strmv/strmv.mjs +42 -20
- package/src/strsm/strsm.d.mts +6 -4
- package/src/strsm/strsm.mjs +438 -174
- package/src/strsv/strsv.d.mts +5 -3
- package/src/strsv/strsv.mjs +89 -34
- package/src/util/benchmark.mjs +9 -9
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +139 -24
- package/src/util/complex.mjs +87 -0
- package/src/util/compute.mjs +19 -16
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +49 -0
- package/src/util/pipeline.mjs +44 -10
- package/src/util/workgroup.mjs +72 -7
- package/src/shaders/browser-shaders.mjs +0 -81
- package/src/shaders/f64add.wgsl +0 -281
package/src/srotm/srotm.mjs
CHANGED
|
@@ -11,13 +11,14 @@ 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 { requireGpuDevice, 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;
|
|
17
18
|
const yIsGpu = y instanceof GpuVector;
|
|
18
19
|
|
|
19
|
-
|
|
20
|
-
|
|
20
|
+
requireGpuDevice(device);
|
|
21
|
+
requireSameDevice(device, "srotm", { x, y });
|
|
21
22
|
if (
|
|
22
23
|
!Number.isInteger(n) ||
|
|
23
24
|
!Number.isInteger(incx) ||
|
|
@@ -26,6 +27,8 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
26
27
|
throw new Error("n, incx, and incy must be integers.");
|
|
27
28
|
if (!(param instanceof Float32Array) || param.length !== 5)
|
|
28
29
|
throw new Error("param must be a Float32Array of length 5.");
|
|
30
|
+
if (param[0] !== -2 && param[0] !== -1 && param[0] !== 0 && param[0] !== 1)
|
|
31
|
+
throw new Error("param[0] (flag) must be one of -2, -1, 0, or 1.");
|
|
29
32
|
if (incx <= 0 || incy <= 0)
|
|
30
33
|
throw new Error("incx and incy must be positive.");
|
|
31
34
|
if (!xIsGpu && !(x instanceof Float32Array))
|
|
@@ -56,10 +59,11 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
56
59
|
let readY = null;
|
|
57
60
|
|
|
58
61
|
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
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "srotm-x", true);
|
|
63
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "srotm-y", true);
|
|
64
|
+
paramBuffer = uploadBuffer(device, param, "srotm-param", false);
|
|
62
65
|
paramsBuffer = createParamsBuffer(
|
|
66
|
+
device,
|
|
63
67
|
[
|
|
64
68
|
{ value: n, type: "u32" },
|
|
65
69
|
{ value: incx, type: "u32" },
|
|
@@ -68,24 +72,26 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
68
72
|
"srotm-params",
|
|
69
73
|
);
|
|
70
74
|
|
|
71
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
75
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
72
76
|
xBuffer,
|
|
73
77
|
yBuffer,
|
|
74
78
|
paramBuffer,
|
|
75
79
|
paramsBuffer,
|
|
76
80
|
]);
|
|
77
81
|
const { commandEncoder, ts } = runComputePass(
|
|
82
|
+
device,
|
|
78
83
|
pipeline,
|
|
79
84
|
bindGroup,
|
|
80
|
-
calcWorkgroups(n),
|
|
85
|
+
calcWorkgroups(device, n),
|
|
81
86
|
);
|
|
82
|
-
readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
83
|
-
readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
84
|
-
submit(commandEncoder);
|
|
87
|
+
readX = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
|
|
88
|
+
readY = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
89
|
+
submit(device, commandEncoder);
|
|
85
90
|
|
|
86
91
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
87
92
|
|
|
88
|
-
if (xIsGpu
|
|
93
|
+
if (xIsGpu) {
|
|
94
|
+
// xIsGpu === yIsGpu, enforced above
|
|
89
95
|
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
90
96
|
return {};
|
|
91
97
|
}
|
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,22 +11,22 @@ 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 { requireGpuDevice, 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
|
-
|
|
19
|
-
|
|
18
|
+
|
|
19
|
+
requireGpuDevice(device);
|
|
20
|
+
requireSameDevice(device, "sscal", { x });
|
|
20
21
|
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
21
22
|
throw new Error("n and incx must be integers.");
|
|
22
|
-
if (typeof alpha !== "number")
|
|
23
|
-
throw new Error("alpha must be a number.");
|
|
23
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
24
24
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
25
25
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
26
26
|
if (incx <= 0) throw new Error("incx must be positive.");
|
|
27
27
|
if (!(x instanceof Float32Array) && !(x instanceof GpuVector))
|
|
28
28
|
throw new Error("x must be a Float32Array or GpuVector.");
|
|
29
|
-
if (n <= 0) return xIsGpu ? {} : x;
|
|
29
|
+
if (n <= 0) return xIsGpu ? {} : { x };
|
|
30
30
|
if (x.length < (n - 1) * incx + 1)
|
|
31
31
|
throw new Error(
|
|
32
32
|
"x does not have enough elements for the given n and incx.",
|
|
@@ -39,8 +39,9 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
39
39
|
let readBuffer = null;
|
|
40
40
|
|
|
41
41
|
try {
|
|
42
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sscal-x", true);
|
|
42
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sscal-x", true);
|
|
43
43
|
paramsBuffer = createParamsBuffer(
|
|
44
|
+
device,
|
|
44
45
|
[
|
|
45
46
|
{ value: n, type: "u32" },
|
|
46
47
|
{ value: alpha, type: "f32" },
|
|
@@ -49,18 +50,19 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
49
50
|
"sscal-params",
|
|
50
51
|
);
|
|
51
52
|
|
|
52
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
53
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
53
54
|
xBuffer,
|
|
54
55
|
paramsBuffer,
|
|
55
56
|
]);
|
|
56
57
|
const { commandEncoder, ts } = runComputePass(
|
|
58
|
+
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,13 +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 { requireGpuDevice, 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;
|
|
17
18
|
const yIsGpu = y instanceof GpuVector;
|
|
18
19
|
|
|
19
|
-
|
|
20
|
-
|
|
20
|
+
requireGpuDevice(device);
|
|
21
|
+
requireSameDevice(device, "sswap", { x, y });
|
|
21
22
|
if (
|
|
22
23
|
!Number.isInteger(n) ||
|
|
23
24
|
!Number.isInteger(incx) ||
|
|
@@ -53,9 +54,10 @@ export async function sswap(device, n, x, incx, y, incy) {
|
|
|
53
54
|
let yReadBuffer = null;
|
|
54
55
|
|
|
55
56
|
try {
|
|
56
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sswap-x", true);
|
|
57
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sswap-y", true);
|
|
57
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sswap-x", true);
|
|
58
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sswap-y", true);
|
|
58
59
|
paramsBuffer = createParamsBuffer(
|
|
60
|
+
device,
|
|
59
61
|
[
|
|
60
62
|
{ value: n, type: "u32" },
|
|
61
63
|
{ value: incx, type: "u32" },
|
|
@@ -64,24 +66,30 @@ 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
74
|
const { commandEncoder, ts } = runComputePass(
|
|
75
|
+
device,
|
|
73
76
|
pipeline,
|
|
74
77
|
bindGroup,
|
|
75
|
-
calcWorkgroups(n),
|
|
78
|
+
calcWorkgroups(device, n),
|
|
76
79
|
);
|
|
77
|
-
xReadBuffer = xIsGpu
|
|
78
|
-
|
|
80
|
+
xReadBuffer = xIsGpu
|
|
81
|
+
? null
|
|
82
|
+
: stageReadback(device, commandEncoder, xBuffer);
|
|
83
|
+
yReadBuffer = yIsGpu
|
|
84
|
+
? null
|
|
85
|
+
: stageReadback(device, commandEncoder, yBuffer);
|
|
79
86
|
|
|
80
|
-
submit(commandEncoder);
|
|
87
|
+
submit(device, commandEncoder);
|
|
81
88
|
|
|
82
89
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
83
90
|
|
|
84
|
-
if (xIsGpu
|
|
91
|
+
if (xIsGpu) {
|
|
92
|
+
// xIsGpu === yIsGpu, enforced above (x.constructor !== y.constructor throws)
|
|
85
93
|
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
86
94
|
return {};
|
|
87
95
|
}
|
package/src/ssymm/ssymm.d.mts
CHANGED
|
@@ -2,8 +2,9 @@ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
4
|
* Performs the symmetric matrix-matrix operation
|
|
5
|
-
* C
|
|
6
|
-
* C
|
|
5
|
+
* $$C \leftarrow \alpha A B + \beta C \quad (\texttt{side='left'})$$
|
|
6
|
+
* $$C \leftarrow \alpha B A + \beta C \quad (\texttt{side='right'})$$
|
|
7
|
+
* `A` is symmetric, only
|
|
7
8
|
* its `uplo` triangle stored; `B` and `C` are general m×n matrices.
|
|
8
9
|
*
|
|
9
10
|
* - `side='left'`: `A` is m×m — `A` premultiplies `B`
|
|
@@ -59,8 +60,8 @@ export declare function ssymm(
|
|
|
59
60
|
|
|
60
61
|
/**
|
|
61
62
|
* Performs the symmetric matrix-matrix operation
|
|
62
|
-
* C
|
|
63
|
-
* C
|
|
63
|
+
* $$C \leftarrow \alpha A B + \beta C \quad (\texttt{side='left'})$$
|
|
64
|
+
* $$C \leftarrow \alpha B A + \beta C \quad (\texttt{side='right'})$$
|
|
64
65
|
*
|
|
65
66
|
* A, B, and C are all kept GPU-resident. Each matrix's own `layout` (set at
|
|
66
67
|
* `GpuMatrix.from` time) determines the operation — there is no separate
|
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,41 +12,60 @@ 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";
|
|
14
|
-
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
|
|
18
|
-
|
|
15
|
+
import { requireWorkgroupCount } from "../util/workgroup.mjs";
|
|
16
|
+
import {
|
|
17
|
+
BM_SMALL,
|
|
18
|
+
BN_SMALL,
|
|
19
|
+
BM_LARGE,
|
|
20
|
+
BN_LARGE,
|
|
21
|
+
LARGE_TILE_WORKGROUP_THRESHOLD,
|
|
22
|
+
} from "../util/constants.mjs";
|
|
23
|
+
import { TILE_WG_2D } from "../util/constants.mjs";
|
|
24
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
19
25
|
|
|
20
26
|
// ssymm: C := alpha*A*B + beta*C (side='left') or alpha*B*A + beta*C
|
|
21
27
|
// (side='right'), A symmetric. No fused kernel — symmetrize then sgemm,
|
|
22
28
|
// both on one command encoder. See symmetrize.wgsl.
|
|
23
29
|
export async function ssymm(
|
|
24
|
-
device,
|
|
30
|
+
device,
|
|
31
|
+
side,
|
|
32
|
+
uplo,
|
|
33
|
+
m,
|
|
34
|
+
n,
|
|
35
|
+
alpha,
|
|
36
|
+
A,
|
|
37
|
+
lda,
|
|
38
|
+
B,
|
|
39
|
+
ldb,
|
|
40
|
+
beta,
|
|
41
|
+
C,
|
|
42
|
+
ldc,
|
|
43
|
+
layout = "row-major",
|
|
25
44
|
) {
|
|
26
45
|
const AIsGpu = A instanceof GpuMatrix;
|
|
27
46
|
const BIsGpu = B instanceof GpuMatrix;
|
|
28
47
|
const CIsGpu = C instanceof GpuMatrix;
|
|
29
48
|
|
|
30
|
-
|
|
31
|
-
|
|
49
|
+
requireGpuDevice(device);
|
|
50
|
+
requireSameDevice(device, "ssymm", { A, B, C });
|
|
32
51
|
if (side !== "left" && side !== "right")
|
|
33
52
|
throw new Error("side must be 'left' or 'right'.");
|
|
34
53
|
if (uplo !== "lower" && uplo !== "upper")
|
|
35
54
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
36
55
|
if (layout !== "row-major" && layout !== "column-major")
|
|
37
56
|
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
38
|
-
if (typeof alpha !== "number")
|
|
39
|
-
throw new Error("alpha must be a number.");
|
|
57
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
40
58
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
41
59
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
42
|
-
if (typeof beta !== "number")
|
|
43
|
-
throw new Error("beta must be a number.");
|
|
60
|
+
if (typeof beta !== "number") throw new Error("beta must be a number.");
|
|
44
61
|
if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
|
|
45
62
|
if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
|
|
46
63
|
if (
|
|
47
|
-
!Number.isInteger(m) ||
|
|
48
|
-
!Number.isInteger(
|
|
64
|
+
!Number.isInteger(m) ||
|
|
65
|
+
!Number.isInteger(n) ||
|
|
66
|
+
!Number.isInteger(lda) ||
|
|
67
|
+
!Number.isInteger(ldb) ||
|
|
68
|
+
!Number.isInteger(ldc)
|
|
49
69
|
)
|
|
50
70
|
throw new Error("m, n, lda, ldb, and ldc must be integers.");
|
|
51
71
|
if (!AIsGpu && !(A instanceof Float32Array))
|
|
@@ -67,48 +87,71 @@ export async function ssymm(
|
|
|
67
87
|
|
|
68
88
|
// A: symmetric, order = m (side='left') or n (side='right').
|
|
69
89
|
const aOrder = side === "left" ? m : n;
|
|
70
|
-
if (lda < aOrder)
|
|
90
|
+
if (lda < aOrder)
|
|
91
|
+
throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
|
|
71
92
|
if (AIsGpu) {
|
|
72
|
-
if (lda !== A.lda)
|
|
73
|
-
|
|
93
|
+
if (lda !== A.lda)
|
|
94
|
+
throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
95
|
+
if (A.rows < aOrder || A.cols < aOrder)
|
|
96
|
+
throw new Error("A is too small for the given m/n and side.");
|
|
74
97
|
} else if (A.length < (aOrder - 1) * lda + aOrder) {
|
|
75
|
-
throw new Error(
|
|
98
|
+
throw new Error(
|
|
99
|
+
"A does not have enough elements for the given dimensions and lda.",
|
|
100
|
+
);
|
|
76
101
|
}
|
|
77
102
|
|
|
78
103
|
// B: always m x n, no trans flag — same shape rule as sgemm's C.
|
|
79
104
|
const bOuter = effLayoutB === "column-major" ? n : m;
|
|
80
105
|
const bInner = effLayoutB === "column-major" ? m : n;
|
|
81
106
|
if (ldb < bInner)
|
|
82
|
-
throw new Error(
|
|
107
|
+
throw new Error(
|
|
108
|
+
`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
|
|
109
|
+
);
|
|
83
110
|
if (BIsGpu) {
|
|
84
|
-
if (ldb !== B.lda)
|
|
85
|
-
|
|
111
|
+
if (ldb !== B.lda)
|
|
112
|
+
throw new Error("ldb must match B.lda when B is a GpuMatrix.");
|
|
113
|
+
if (B.rows < m || B.cols < n)
|
|
114
|
+
throw new Error("B is too small for the given m and n.");
|
|
86
115
|
} else if (B.length < (bOuter - 1) * ldb + bInner) {
|
|
87
|
-
throw new Error(
|
|
116
|
+
throw new Error(
|
|
117
|
+
"B does not have enough elements for the given dimensions and ldb.",
|
|
118
|
+
);
|
|
88
119
|
}
|
|
89
120
|
|
|
90
121
|
// C: always m x n.
|
|
91
122
|
const cOuter = effLayoutC === "column-major" ? n : m;
|
|
92
123
|
const cInner = effLayoutC === "column-major" ? m : n;
|
|
93
124
|
if (ldc < cInner)
|
|
94
|
-
throw new Error(
|
|
125
|
+
throw new Error(
|
|
126
|
+
`ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`,
|
|
127
|
+
);
|
|
95
128
|
if (CIsGpu) {
|
|
96
|
-
if (ldc !== C.lda)
|
|
97
|
-
|
|
129
|
+
if (ldc !== C.lda)
|
|
130
|
+
throw new Error("ldc must match C.lda when C is a GpuMatrix.");
|
|
131
|
+
if (C.rows < m || C.cols < n)
|
|
132
|
+
throw new Error("C is too small for the given m and n.");
|
|
98
133
|
} else if (C.length < (cOuter - 1) * ldc + cInner) {
|
|
99
|
-
throw new Error(
|
|
134
|
+
throw new Error(
|
|
135
|
+
"C does not have enough elements for the given dimensions and ldc.",
|
|
136
|
+
);
|
|
100
137
|
}
|
|
101
138
|
|
|
102
139
|
// A = A^T, so column-major storage is still A, but the populated
|
|
103
140
|
// triangle swaps — uplo flips (same reasoning ssyr/ssyrk use).
|
|
104
|
-
const uploEffA =
|
|
141
|
+
const uploEffA =
|
|
142
|
+
effLayoutA === "column-major"
|
|
143
|
+
? uplo === "lower"
|
|
144
|
+
? "upper"
|
|
145
|
+
: "lower"
|
|
146
|
+
: uplo;
|
|
105
147
|
|
|
106
148
|
const transB = effLayoutB === "column-major" ? "transpose" : "no-transpose";
|
|
107
149
|
const transDense = "no-transpose"; // Adense is always row-major
|
|
108
150
|
|
|
109
151
|
// X*Y (X=Adense,Y=B for side='left', swapped for 'right'). Column-major
|
|
110
152
|
// C: compute C^T instead (swap X/Y, flip trans, swap m_g/n_g) — sgemm's own trick.
|
|
111
|
-
let mg = m,
|
|
153
|
+
let mg = m,
|
|
154
|
+
ng = n;
|
|
112
155
|
const kg = aOrder;
|
|
113
156
|
let transX = side === "left" ? transDense : transB;
|
|
114
157
|
let transY = side === "left" ? transB : transDense;
|
|
@@ -124,26 +167,45 @@ export async function ssymm(
|
|
|
124
167
|
const largeWgX = Math.ceil(ng / BN_LARGE);
|
|
125
168
|
const largeWgY = Math.ceil(mg / BM_LARGE);
|
|
126
169
|
const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
127
|
-
const gemmPipeline = await getPipeline(
|
|
170
|
+
const gemmPipeline = await getPipeline(
|
|
171
|
+
device,
|
|
172
|
+
useLargeTile ? "sgemm_large" : "sgemm_small",
|
|
173
|
+
);
|
|
128
174
|
const symPipeline = await getPipeline(device, "symmetrize");
|
|
129
175
|
const gemmWgCount = useLargeTile
|
|
130
176
|
? {
|
|
131
|
-
|
|
132
|
-
|
|
133
|
-
|
|
177
|
+
x: requireWorkgroupCount(device, largeWgX, "ssymm", "x"),
|
|
178
|
+
y: requireWorkgroupCount(device, largeWgY, "ssymm", "y"),
|
|
179
|
+
}
|
|
134
180
|
: {
|
|
135
|
-
|
|
136
|
-
|
|
137
|
-
|
|
181
|
+
x: requireWorkgroupCount(
|
|
182
|
+
device,
|
|
183
|
+
Math.ceil(ng / BN_SMALL),
|
|
184
|
+
"ssymm",
|
|
185
|
+
"x",
|
|
186
|
+
),
|
|
187
|
+
y: requireWorkgroupCount(
|
|
188
|
+
device,
|
|
189
|
+
Math.ceil(mg / BM_SMALL),
|
|
190
|
+
"ssymm",
|
|
191
|
+
"y",
|
|
192
|
+
),
|
|
193
|
+
};
|
|
138
194
|
|
|
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(
|
|
143
|
-
|
|
195
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssymm-A", false);
|
|
196
|
+
const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "ssymm-B", false);
|
|
197
|
+
const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssymm-C", true);
|
|
198
|
+
const AdenseBuffer = createStorageBuffer(
|
|
199
|
+
device,
|
|
200
|
+
aOrder * ldDense * 4,
|
|
201
|
+
"ssymm-Adense",
|
|
202
|
+
);
|
|
203
|
+
let symParams = null,
|
|
204
|
+
gemmParams = null;
|
|
144
205
|
|
|
145
206
|
try {
|
|
146
207
|
symParams = createParamsBuffer(
|
|
208
|
+
device,
|
|
147
209
|
[
|
|
148
210
|
{ value: aOrder, type: "u32" },
|
|
149
211
|
{ value: lda, type: "u32" },
|
|
@@ -152,7 +214,11 @@ export async function ssymm(
|
|
|
152
214
|
],
|
|
153
215
|
"ssymm-sym-params",
|
|
154
216
|
);
|
|
155
|
-
const symBindGroup = createBindGroup(
|
|
217
|
+
const symBindGroup = createBindGroup(
|
|
218
|
+
device,
|
|
219
|
+
symPipeline.getBindGroupLayout(0),
|
|
220
|
+
[ABuffer, AdenseBuffer, symParams],
|
|
221
|
+
);
|
|
156
222
|
|
|
157
223
|
// X/Y buffers and their own ld, matching swapXY above.
|
|
158
224
|
const XBuffer = swapXY ? BBuffer : AdenseBuffer;
|
|
@@ -161,12 +227,13 @@ export async function ssymm(
|
|
|
161
227
|
const ldY = swapXY ? ldDense : ldb;
|
|
162
228
|
|
|
163
229
|
gemmParams = createParamsBuffer(
|
|
230
|
+
device,
|
|
164
231
|
[
|
|
165
|
-
{ value: mg,
|
|
166
|
-
{ value: ng,
|
|
167
|
-
{ value: kg,
|
|
232
|
+
{ value: mg, type: "u32" },
|
|
233
|
+
{ value: ng, type: "u32" },
|
|
234
|
+
{ value: kg, type: "u32" },
|
|
168
235
|
{ value: alpha, type: "f32" },
|
|
169
|
-
{ value: beta,
|
|
236
|
+
{ value: beta, type: "f32" },
|
|
170
237
|
{ value: ldX, type: "u32" },
|
|
171
238
|
{ value: ldY, type: "u32" },
|
|
172
239
|
{ value: ldc, type: "u32" },
|
|
@@ -175,18 +242,47 @@ export async function ssymm(
|
|
|
175
242
|
],
|
|
176
243
|
"ssymm-gemm-params",
|
|
177
244
|
);
|
|
178
|
-
const gemmBindGroup = createBindGroup(
|
|
245
|
+
const gemmBindGroup = createBindGroup(
|
|
246
|
+
device,
|
|
247
|
+
gemmPipeline.getBindGroupLayout(0),
|
|
248
|
+
[
|
|
249
|
+
XBuffer,
|
|
250
|
+
vec4ViewBinding(device, XBuffer),
|
|
251
|
+
YBuffer,
|
|
252
|
+
vec4ViewBinding(device, YBuffer),
|
|
253
|
+
CBuffer,
|
|
254
|
+
gemmParams,
|
|
255
|
+
],
|
|
256
|
+
);
|
|
179
257
|
|
|
180
|
-
const { commandEncoder, querySet } = beginTimedEncoder();
|
|
181
|
-
const symDesc = querySet
|
|
182
|
-
|
|
183
|
-
|
|
184
|
-
|
|
258
|
+
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
259
|
+
const symDesc = querySet
|
|
260
|
+
? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
|
|
261
|
+
: undefined;
|
|
262
|
+
const gemmDesc = querySet
|
|
263
|
+
? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
|
|
264
|
+
: undefined;
|
|
265
|
+
encodePass(
|
|
266
|
+
commandEncoder,
|
|
267
|
+
symPipeline,
|
|
268
|
+
symBindGroup,
|
|
269
|
+
{ x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) },
|
|
270
|
+
symDesc,
|
|
271
|
+
);
|
|
272
|
+
encodePass(
|
|
273
|
+
commandEncoder,
|
|
274
|
+
gemmPipeline,
|
|
275
|
+
gemmBindGroup,
|
|
276
|
+
gemmWgCount,
|
|
277
|
+
gemmDesc,
|
|
278
|
+
);
|
|
185
279
|
|
|
186
|
-
const ts = resolveTimestamp(commandEncoder, querySet);
|
|
187
|
-
const readBuffer = CIsGpu
|
|
280
|
+
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
281
|
+
const readBuffer = CIsGpu
|
|
282
|
+
? null
|
|
283
|
+
: stageReadback(device, commandEncoder, CBuffer);
|
|
188
284
|
|
|
189
|
-
submit(commandEncoder);
|
|
285
|
+
submit(device, commandEncoder);
|
|
190
286
|
|
|
191
287
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
192
288
|
|
package/src/ssymv/ssymv.d.mts
CHANGED
|
@@ -2,7 +2,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
|
|
|
2
2
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
3
3
|
|
|
4
4
|
/**
|
|
5
|
-
* Performs the symmetric matrix-vector operation y
|
|
5
|
+
* Performs the symmetric matrix-vector operation $$y \leftarrow \alpha A x + \beta y$$
|
|
6
6
|
*
|
|
7
7
|
* A is an n×n symmetric matrix stored in row-major order. Only the triangle
|
|
8
8
|
* specified by `uplo` is referenced; the other triangle is inferred by symmetry.
|
|
@@ -45,7 +45,7 @@ export declare function ssymv(
|
|
|
45
45
|
): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
|
|
46
46
|
|
|
47
47
|
/**
|
|
48
|
-
* Performs the symmetric matrix-vector operation y
|
|
48
|
+
* Performs the symmetric matrix-vector operation $$y \leftarrow \alpha A x + \beta y$$
|
|
49
49
|
*
|
|
50
50
|
* x and y are kept resident on the GPU. A must be a GpuMatrix; its own
|
|
51
51
|
* `layout` (set at `GpuMatrix.from` time) determines the operation — there is
|