wgblas 2.1.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 +2 -0
- package/dist/wgblas.browser.js +1007 -43
- package/index.d.mts +26 -53
- package/index.mjs +11 -0
- package/package.json +132 -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 +112 -10
- package/src/classes/GpuVector.d.mts +36 -40
- package/src/classes/GpuVector.mjs +39 -2
- package/src/cscal/cscal.d.mts +47 -0
- package/src/cscal/cscal.mjs +98 -0
- package/src/dasum/dasum.mjs +31 -15
- 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/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.mjs +49 -19
- package/src/init.mjs +6 -3
- package/src/isamax/isamax.mjs +17 -14
- package/src/random/random.d.mts +37 -40
- package/src/random/random.mjs +39 -7
- package/src/sasum/sasum.mjs +13 -11
- package/src/saxpy/saxpy.mjs +9 -8
- package/src/scopy/scopy.mjs +8 -6
- package/src/sdot/sdot.mjs +13 -11
- package/src/sgemm/sgemm.d.mts +2 -2
- package/src/sgemm/sgemm.mjs +91 -35
- package/src/sgemmtr/sgemmtr.d.mts +3 -2
- package/src/sgemmtr/sgemmtr.mjs +92 -35
- package/src/sgemv/sgemv.d.mts +2 -2
- package/src/sgemv/sgemv.mjs +41 -25
- package/src/sger/sger.d.mts +2 -2
- package/src/sger/sger.mjs +38 -16
- 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 +69 -0
- package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
- package/src/snrm2/snrm2.mjs +20 -14
- package/src/srot/srot.mjs +10 -7
- package/src/srotm/srotm.mjs +9 -12
- package/src/sscal/sscal.mjs +7 -7
- package/src/sswap/sswap.mjs +14 -8
- package/src/ssymm/ssymm.d.mts +5 -4
- package/src/ssymm/ssymm.mjs +142 -55
- package/src/ssymv/ssymv.d.mts +2 -2
- package/src/ssymv/ssymv.mjs +42 -23
- package/src/ssyr/ssyr.d.mts +2 -2
- package/src/ssyr/ssyr.mjs +34 -15
- package/src/ssyr2/ssyr2.d.mts +2 -2
- package/src/ssyr2/ssyr2.mjs +43 -18
- package/src/ssyr2k/ssyr2k.d.mts +3 -2
- package/src/ssyr2k/ssyr2k.mjs +132 -55
- package/src/ssyrk/ssyrk.d.mts +3 -2
- package/src/ssyrk/ssyrk.mjs +84 -33
- package/src/strmm/strmm.d.mts +5 -4
- package/src/strmm/strmm.mjs +153 -54
- package/src/strmv/strmv.d.mts +2 -2
- package/src/strmv/strmv.mjs +37 -17
- package/src/strsm/strsm.d.mts +6 -4
- package/src/strsm/strsm.mjs +418 -172
- package/src/strsv/strsv.d.mts +5 -3
- package/src/strsv/strsv.mjs +82 -31
- package/src/util/benchmark.mjs +5 -3
- package/src/util/buffer.mjs +33 -12
- package/src/util/complex.mjs +87 -0
- package/src/util/compute.mjs +14 -8
- package/src/util/device.mjs +18 -3
- package/src/util/pipeline.mjs +40 -5
- package/src/util/workgroup.mjs +23 -6
- package/src/shaders/f64add.wgsl +0 -281
package/src/srot/srot.mjs
CHANGED
|
@@ -11,14 +11,13 @@ 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
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
15
15
|
|
|
16
16
|
export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
17
17
|
const xIsGpu = x instanceof GpuVector;
|
|
18
18
|
const yIsGpu = y instanceof GpuVector;
|
|
19
19
|
|
|
20
|
-
|
|
21
|
-
throw new Error("device must be a GPUDevice.");
|
|
20
|
+
requireGpuDevice(device);
|
|
22
21
|
requireSameDevice(device, "srot", { x, y });
|
|
23
22
|
if (
|
|
24
23
|
!Number.isInteger(n) ||
|
|
@@ -28,7 +27,8 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
28
27
|
throw new Error("n, incx, and incy must be integers.");
|
|
29
28
|
if (typeof c !== "number") throw new Error("c must be a number.");
|
|
30
29
|
if (typeof s !== "number") throw new Error("s must be a number.");
|
|
31
|
-
if (Number.isNaN(c) || Number.isNaN(s))
|
|
30
|
+
if (Number.isNaN(c) || Number.isNaN(s))
|
|
31
|
+
throw new Error("c and s must not be NaN.");
|
|
32
32
|
if (!Number.isFinite(c)) throw new Error("c must be finite.");
|
|
33
33
|
if (!Number.isFinite(s)) throw new Error("s must be finite.");
|
|
34
34
|
if (incx <= 0 || incy <= 0)
|
|
@@ -62,7 +62,8 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
62
62
|
try {
|
|
63
63
|
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "srot-x", true);
|
|
64
64
|
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "srot-y", true);
|
|
65
|
-
paramsBuffer = createParamsBuffer(
|
|
65
|
+
paramsBuffer = createParamsBuffer(
|
|
66
|
+
device,
|
|
66
67
|
[
|
|
67
68
|
{ value: n, type: "u32" },
|
|
68
69
|
{ value: c, type: "f32" },
|
|
@@ -78,7 +79,8 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
78
79
|
yBuffer,
|
|
79
80
|
paramsBuffer,
|
|
80
81
|
]);
|
|
81
|
-
const { commandEncoder, ts } = runComputePass(
|
|
82
|
+
const { commandEncoder, ts } = runComputePass(
|
|
83
|
+
device,
|
|
82
84
|
pipeline,
|
|
83
85
|
bindGroup,
|
|
84
86
|
calcWorkgroups(device, n),
|
|
@@ -89,7 +91,8 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
89
91
|
|
|
90
92
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
91
93
|
|
|
92
|
-
if (xIsGpu
|
|
94
|
+
if (xIsGpu) {
|
|
95
|
+
// xIsGpu === yIsGpu, enforced above
|
|
93
96
|
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
94
97
|
return {};
|
|
95
98
|
}
|
package/src/srotm/srotm.mjs
CHANGED
|
@@ -11,14 +11,13 @@ 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
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
15
15
|
|
|
16
16
|
export async function srotm(device, n, x, incx, y, incy, param) {
|
|
17
17
|
const xIsGpu = x instanceof GpuVector;
|
|
18
18
|
const yIsGpu = y instanceof GpuVector;
|
|
19
19
|
|
|
20
|
-
|
|
21
|
-
throw new Error("device must be a GPUDevice.");
|
|
20
|
+
requireGpuDevice(device);
|
|
22
21
|
requireSameDevice(device, "srotm", { x, y });
|
|
23
22
|
if (
|
|
24
23
|
!Number.isInteger(n) ||
|
|
@@ -28,12 +27,7 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
28
27
|
throw new Error("n, incx, and incy must be integers.");
|
|
29
28
|
if (!(param instanceof Float32Array) || param.length !== 5)
|
|
30
29
|
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
|
-
)
|
|
30
|
+
if (param[0] !== -2 && param[0] !== -1 && param[0] !== 0 && param[0] !== 1)
|
|
37
31
|
throw new Error("param[0] (flag) must be one of -2, -1, 0, or 1.");
|
|
38
32
|
if (incx <= 0 || incy <= 0)
|
|
39
33
|
throw new Error("incx and incy must be positive.");
|
|
@@ -68,7 +62,8 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
68
62
|
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "srotm-x", true);
|
|
69
63
|
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "srotm-y", true);
|
|
70
64
|
paramBuffer = uploadBuffer(device, param, "srotm-param", false);
|
|
71
|
-
paramsBuffer = createParamsBuffer(
|
|
65
|
+
paramsBuffer = createParamsBuffer(
|
|
66
|
+
device,
|
|
72
67
|
[
|
|
73
68
|
{ value: n, type: "u32" },
|
|
74
69
|
{ value: incx, type: "u32" },
|
|
@@ -83,7 +78,8 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
83
78
|
paramBuffer,
|
|
84
79
|
paramsBuffer,
|
|
85
80
|
]);
|
|
86
|
-
const { commandEncoder, ts } = runComputePass(
|
|
81
|
+
const { commandEncoder, ts } = runComputePass(
|
|
82
|
+
device,
|
|
87
83
|
pipeline,
|
|
88
84
|
bindGroup,
|
|
89
85
|
calcWorkgroups(device, n),
|
|
@@ -94,7 +90,8 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
94
90
|
|
|
95
91
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
96
92
|
|
|
97
|
-
if (xIsGpu
|
|
93
|
+
if (xIsGpu) {
|
|
94
|
+
// xIsGpu === yIsGpu, enforced above
|
|
98
95
|
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
99
96
|
return {};
|
|
100
97
|
}
|
package/src/sscal/sscal.mjs
CHANGED
|
@@ -11,18 +11,16 @@ 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
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
15
15
|
|
|
16
16
|
export async function sscal(device, n, alpha, x, incx) {
|
|
17
17
|
const xIsGpu = x instanceof GpuVector;
|
|
18
18
|
|
|
19
|
-
|
|
20
|
-
throw new Error("device must be a GPUDevice.");
|
|
19
|
+
requireGpuDevice(device);
|
|
21
20
|
requireSameDevice(device, "sscal", { x });
|
|
22
21
|
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
23
22
|
throw new Error("n and incx must be integers.");
|
|
24
|
-
if (typeof alpha !== "number")
|
|
25
|
-
throw new Error("alpha must be a number.");
|
|
23
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
26
24
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
27
25
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
28
26
|
if (incx <= 0) throw new Error("incx must be positive.");
|
|
@@ -42,7 +40,8 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
42
40
|
|
|
43
41
|
try {
|
|
44
42
|
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sscal-x", true);
|
|
45
|
-
paramsBuffer = createParamsBuffer(
|
|
43
|
+
paramsBuffer = createParamsBuffer(
|
|
44
|
+
device,
|
|
46
45
|
[
|
|
47
46
|
{ value: n, type: "u32" },
|
|
48
47
|
{ value: alpha, type: "f32" },
|
|
@@ -55,7 +54,8 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
55
54
|
xBuffer,
|
|
56
55
|
paramsBuffer,
|
|
57
56
|
]);
|
|
58
|
-
const { commandEncoder, ts } = runComputePass(
|
|
57
|
+
const { commandEncoder, ts } = runComputePass(
|
|
58
|
+
device,
|
|
59
59
|
pipeline,
|
|
60
60
|
bindGroup,
|
|
61
61
|
calcWorkgroups(device, n),
|
package/src/sswap/sswap.mjs
CHANGED
|
@@ -11,14 +11,13 @@ 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
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
15
15
|
|
|
16
16
|
export async function sswap(device, n, x, incx, y, incy) {
|
|
17
17
|
const xIsGpu = x instanceof GpuVector;
|
|
18
18
|
const yIsGpu = y instanceof GpuVector;
|
|
19
19
|
|
|
20
|
-
|
|
21
|
-
throw new Error("device must be a GPUDevice.");
|
|
20
|
+
requireGpuDevice(device);
|
|
22
21
|
requireSameDevice(device, "sswap", { x, y });
|
|
23
22
|
if (
|
|
24
23
|
!Number.isInteger(n) ||
|
|
@@ -57,7 +56,8 @@ export async function sswap(device, n, x, incx, y, incy) {
|
|
|
57
56
|
try {
|
|
58
57
|
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sswap-x", true);
|
|
59
58
|
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sswap-y", true);
|
|
60
|
-
paramsBuffer = createParamsBuffer(
|
|
59
|
+
paramsBuffer = createParamsBuffer(
|
|
60
|
+
device,
|
|
61
61
|
[
|
|
62
62
|
{ value: n, type: "u32" },
|
|
63
63
|
{ value: incx, type: "u32" },
|
|
@@ -71,19 +71,25 @@ export async function sswap(device, n, x, incx, y, incy) {
|
|
|
71
71
|
yBuffer,
|
|
72
72
|
paramsBuffer,
|
|
73
73
|
]);
|
|
74
|
-
const { commandEncoder, ts } = runComputePass(
|
|
74
|
+
const { commandEncoder, ts } = runComputePass(
|
|
75
|
+
device,
|
|
75
76
|
pipeline,
|
|
76
77
|
bindGroup,
|
|
77
78
|
calcWorkgroups(device, n),
|
|
78
79
|
);
|
|
79
|
-
xReadBuffer = xIsGpu
|
|
80
|
-
|
|
80
|
+
xReadBuffer = xIsGpu
|
|
81
|
+
? null
|
|
82
|
+
: stageReadback(device, commandEncoder, xBuffer);
|
|
83
|
+
yReadBuffer = yIsGpu
|
|
84
|
+
? null
|
|
85
|
+
: stageReadback(device, commandEncoder, yBuffer);
|
|
81
86
|
|
|
82
87
|
submit(device, commandEncoder);
|
|
83
88
|
|
|
84
89
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
85
90
|
|
|
86
|
-
if (xIsGpu
|
|
91
|
+
if (xIsGpu) {
|
|
92
|
+
// xIsGpu === yIsGpu, enforced above (x.constructor !== y.constructor throws)
|
|
87
93
|
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
88
94
|
return {};
|
|
89
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
|
@@ -13,23 +13,40 @@ import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
|
|
|
13
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
14
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
15
15
|
import { requireWorkgroupCount } from "../util/workgroup.mjs";
|
|
16
|
-
import {
|
|
16
|
+
import {
|
|
17
|
+
BM_SMALL,
|
|
18
|
+
BN_SMALL,
|
|
19
|
+
BM_LARGE,
|
|
20
|
+
BN_LARGE,
|
|
21
|
+
LARGE_TILE_WORKGROUP_THRESHOLD,
|
|
22
|
+
} from "../util/constants.mjs";
|
|
17
23
|
import { TILE_WG_2D } from "../util/constants.mjs";
|
|
18
|
-
import { requireSameDevice } from "../util/device.mjs";
|
|
19
|
-
|
|
24
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
20
25
|
|
|
21
26
|
// ssymm: C := alpha*A*B + beta*C (side='left') or alpha*B*A + beta*C
|
|
22
27
|
// (side='right'), A symmetric. No fused kernel — symmetrize then sgemm,
|
|
23
28
|
// both on one command encoder. See symmetrize.wgsl.
|
|
24
29
|
export async function ssymm(
|
|
25
|
-
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",
|
|
26
44
|
) {
|
|
27
45
|
const AIsGpu = A instanceof GpuMatrix;
|
|
28
46
|
const BIsGpu = B instanceof GpuMatrix;
|
|
29
47
|
const CIsGpu = C instanceof GpuMatrix;
|
|
30
48
|
|
|
31
|
-
|
|
32
|
-
throw new Error("device must be a GPUDevice.");
|
|
49
|
+
requireGpuDevice(device);
|
|
33
50
|
requireSameDevice(device, "ssymm", { A, B, C });
|
|
34
51
|
if (side !== "left" && side !== "right")
|
|
35
52
|
throw new Error("side must be 'left' or 'right'.");
|
|
@@ -37,17 +54,18 @@ export async function ssymm(
|
|
|
37
54
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
38
55
|
if (layout !== "row-major" && layout !== "column-major")
|
|
39
56
|
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
40
|
-
if (typeof alpha !== "number")
|
|
41
|
-
throw new Error("alpha must be a number.");
|
|
57
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
42
58
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
43
59
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
44
|
-
if (typeof beta !== "number")
|
|
45
|
-
throw new Error("beta must be a number.");
|
|
60
|
+
if (typeof beta !== "number") throw new Error("beta must be a number.");
|
|
46
61
|
if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
|
|
47
62
|
if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
|
|
48
63
|
if (
|
|
49
|
-
!Number.isInteger(m) ||
|
|
50
|
-
!Number.isInteger(
|
|
64
|
+
!Number.isInteger(m) ||
|
|
65
|
+
!Number.isInteger(n) ||
|
|
66
|
+
!Number.isInteger(lda) ||
|
|
67
|
+
!Number.isInteger(ldb) ||
|
|
68
|
+
!Number.isInteger(ldc)
|
|
51
69
|
)
|
|
52
70
|
throw new Error("m, n, lda, ldb, and ldc must be integers.");
|
|
53
71
|
if (!AIsGpu && !(A instanceof Float32Array))
|
|
@@ -69,48 +87,71 @@ export async function ssymm(
|
|
|
69
87
|
|
|
70
88
|
// A: symmetric, order = m (side='left') or n (side='right').
|
|
71
89
|
const aOrder = side === "left" ? m : n;
|
|
72
|
-
if (lda < aOrder)
|
|
90
|
+
if (lda < aOrder)
|
|
91
|
+
throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
|
|
73
92
|
if (AIsGpu) {
|
|
74
|
-
if (lda !== A.lda)
|
|
75
|
-
|
|
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.");
|
|
76
97
|
} else if (A.length < (aOrder - 1) * lda + aOrder) {
|
|
77
|
-
throw new Error(
|
|
98
|
+
throw new Error(
|
|
99
|
+
"A does not have enough elements for the given dimensions and lda.",
|
|
100
|
+
);
|
|
78
101
|
}
|
|
79
102
|
|
|
80
103
|
// B: always m x n, no trans flag — same shape rule as sgemm's C.
|
|
81
104
|
const bOuter = effLayoutB === "column-major" ? n : m;
|
|
82
105
|
const bInner = effLayoutB === "column-major" ? m : n;
|
|
83
106
|
if (ldb < bInner)
|
|
84
|
-
throw new Error(
|
|
107
|
+
throw new Error(
|
|
108
|
+
`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
|
|
109
|
+
);
|
|
85
110
|
if (BIsGpu) {
|
|
86
|
-
if (ldb !== B.lda)
|
|
87
|
-
|
|
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.");
|
|
88
115
|
} else if (B.length < (bOuter - 1) * ldb + bInner) {
|
|
89
|
-
throw new Error(
|
|
116
|
+
throw new Error(
|
|
117
|
+
"B does not have enough elements for the given dimensions and ldb.",
|
|
118
|
+
);
|
|
90
119
|
}
|
|
91
120
|
|
|
92
121
|
// C: always m x n.
|
|
93
122
|
const cOuter = effLayoutC === "column-major" ? n : m;
|
|
94
123
|
const cInner = effLayoutC === "column-major" ? m : n;
|
|
95
124
|
if (ldc < cInner)
|
|
96
|
-
throw new Error(
|
|
125
|
+
throw new Error(
|
|
126
|
+
`ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`,
|
|
127
|
+
);
|
|
97
128
|
if (CIsGpu) {
|
|
98
|
-
if (ldc !== C.lda)
|
|
99
|
-
|
|
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.");
|
|
100
133
|
} else if (C.length < (cOuter - 1) * ldc + cInner) {
|
|
101
|
-
throw new Error(
|
|
134
|
+
throw new Error(
|
|
135
|
+
"C does not have enough elements for the given dimensions and ldc.",
|
|
136
|
+
);
|
|
102
137
|
}
|
|
103
138
|
|
|
104
139
|
// A = A^T, so column-major storage is still A, but the populated
|
|
105
140
|
// triangle swaps — uplo flips (same reasoning ssyr/ssyrk use).
|
|
106
|
-
const uploEffA =
|
|
141
|
+
const uploEffA =
|
|
142
|
+
effLayoutA === "column-major"
|
|
143
|
+
? uplo === "lower"
|
|
144
|
+
? "upper"
|
|
145
|
+
: "lower"
|
|
146
|
+
: uplo;
|
|
107
147
|
|
|
108
148
|
const transB = effLayoutB === "column-major" ? "transpose" : "no-transpose";
|
|
109
149
|
const transDense = "no-transpose"; // Adense is always row-major
|
|
110
150
|
|
|
111
151
|
// X*Y (X=Adense,Y=B for side='left', swapped for 'right'). Column-major
|
|
112
152
|
// C: compute C^T instead (swap X/Y, flip trans, swap m_g/n_g) — sgemm's own trick.
|
|
113
|
-
let mg = m,
|
|
153
|
+
let mg = m,
|
|
154
|
+
ng = n;
|
|
114
155
|
const kg = aOrder;
|
|
115
156
|
let transX = side === "left" ? transDense : transB;
|
|
116
157
|
let transY = side === "left" ? transB : transDense;
|
|
@@ -126,26 +167,45 @@ export async function ssymm(
|
|
|
126
167
|
const largeWgX = Math.ceil(ng / BN_LARGE);
|
|
127
168
|
const largeWgY = Math.ceil(mg / BM_LARGE);
|
|
128
169
|
const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
|
|
129
|
-
const gemmPipeline = await getPipeline(
|
|
170
|
+
const gemmPipeline = await getPipeline(
|
|
171
|
+
device,
|
|
172
|
+
useLargeTile ? "sgemm_large" : "sgemm_small",
|
|
173
|
+
);
|
|
130
174
|
const symPipeline = await getPipeline(device, "symmetrize");
|
|
131
175
|
const gemmWgCount = useLargeTile
|
|
132
176
|
? {
|
|
133
|
-
|
|
134
|
-
|
|
135
|
-
|
|
177
|
+
x: requireWorkgroupCount(device, largeWgX, "ssymm", "x"),
|
|
178
|
+
y: requireWorkgroupCount(device, largeWgY, "ssymm", "y"),
|
|
179
|
+
}
|
|
136
180
|
: {
|
|
137
|
-
|
|
138
|
-
|
|
139
|
-
|
|
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
|
+
};
|
|
140
194
|
|
|
141
195
|
const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssymm-A", false);
|
|
142
196
|
const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "ssymm-B", false);
|
|
143
197
|
const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssymm-C", true);
|
|
144
|
-
const AdenseBuffer = createStorageBuffer(
|
|
145
|
-
|
|
198
|
+
const AdenseBuffer = createStorageBuffer(
|
|
199
|
+
device,
|
|
200
|
+
aOrder * ldDense * 4,
|
|
201
|
+
"ssymm-Adense",
|
|
202
|
+
);
|
|
203
|
+
let symParams = null,
|
|
204
|
+
gemmParams = null;
|
|
146
205
|
|
|
147
206
|
try {
|
|
148
|
-
symParams = createParamsBuffer(
|
|
207
|
+
symParams = createParamsBuffer(
|
|
208
|
+
device,
|
|
149
209
|
[
|
|
150
210
|
{ value: aOrder, type: "u32" },
|
|
151
211
|
{ value: lda, type: "u32" },
|
|
@@ -154,7 +214,11 @@ export async function ssymm(
|
|
|
154
214
|
],
|
|
155
215
|
"ssymm-sym-params",
|
|
156
216
|
);
|
|
157
|
-
const symBindGroup = createBindGroup(
|
|
217
|
+
const symBindGroup = createBindGroup(
|
|
218
|
+
device,
|
|
219
|
+
symPipeline.getBindGroupLayout(0),
|
|
220
|
+
[ABuffer, AdenseBuffer, symParams],
|
|
221
|
+
);
|
|
158
222
|
|
|
159
223
|
// X/Y buffers and their own ld, matching swapXY above.
|
|
160
224
|
const XBuffer = swapXY ? BBuffer : AdenseBuffer;
|
|
@@ -162,13 +226,14 @@ export async function ssymm(
|
|
|
162
226
|
const YBuffer = swapXY ? AdenseBuffer : BBuffer;
|
|
163
227
|
const ldY = swapXY ? ldDense : ldb;
|
|
164
228
|
|
|
165
|
-
gemmParams = createParamsBuffer(
|
|
229
|
+
gemmParams = createParamsBuffer(
|
|
230
|
+
device,
|
|
166
231
|
[
|
|
167
|
-
{ value: mg,
|
|
168
|
-
{ value: ng,
|
|
169
|
-
{ value: kg,
|
|
232
|
+
{ value: mg, type: "u32" },
|
|
233
|
+
{ value: ng, type: "u32" },
|
|
234
|
+
{ value: kg, type: "u32" },
|
|
170
235
|
{ value: alpha, type: "f32" },
|
|
171
|
-
{ value: beta,
|
|
236
|
+
{ value: beta, type: "f32" },
|
|
172
237
|
{ value: ldX, type: "u32" },
|
|
173
238
|
{ value: ldY, type: "u32" },
|
|
174
239
|
{ value: ldc, type: "u32" },
|
|
@@ -177,23 +242,45 @@ export async function ssymm(
|
|
|
177
242
|
],
|
|
178
243
|
"ssymm-gemm-params",
|
|
179
244
|
);
|
|
180
|
-
const gemmBindGroup = createBindGroup(
|
|
181
|
-
|
|
182
|
-
|
|
183
|
-
|
|
184
|
-
|
|
185
|
-
|
|
186
|
-
|
|
187
|
-
|
|
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
|
+
);
|
|
188
257
|
|
|
189
258
|
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
190
|
-
const symDesc = querySet
|
|
191
|
-
|
|
192
|
-
|
|
193
|
-
|
|
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
|
+
);
|
|
194
279
|
|
|
195
280
|
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
196
|
-
const readBuffer = CIsGpu
|
|
281
|
+
const readBuffer = CIsGpu
|
|
282
|
+
? null
|
|
283
|
+
: stageReadback(device, commandEncoder, CBuffer);
|
|
197
284
|
|
|
198
285
|
submit(device, commandEncoder);
|
|
199
286
|
|
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
|