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/sgemv/sgemv.mjs
CHANGED
|
@@ -9,27 +9,40 @@ import { runComputePass, submit } from "../util/compute.mjs";
|
|
|
9
9
|
import { extractResult } from "../util/result.mjs";
|
|
10
10
|
import { extractTimestamp } from "../util/benchmark.mjs";
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
|
-
import {
|
|
12
|
+
import { requireWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
14
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
15
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
15
16
|
|
|
16
|
-
export async function sgemv(
|
|
17
|
+
export async function sgemv(
|
|
18
|
+
device,
|
|
19
|
+
trans,
|
|
20
|
+
m,
|
|
21
|
+
n,
|
|
22
|
+
alpha,
|
|
23
|
+
A,
|
|
24
|
+
lda,
|
|
25
|
+
x,
|
|
26
|
+
incx,
|
|
27
|
+
beta,
|
|
28
|
+
y,
|
|
29
|
+
incy,
|
|
30
|
+
layout = "row-major",
|
|
31
|
+
) {
|
|
17
32
|
const AIsGpu = A instanceof GpuMatrix;
|
|
18
33
|
const xIsGpu = x instanceof GpuVector;
|
|
19
34
|
const yIsGpu = y instanceof GpuVector;
|
|
20
35
|
|
|
21
|
-
|
|
22
|
-
|
|
36
|
+
requireGpuDevice(device);
|
|
37
|
+
requireSameDevice(device, "sgemv", { A, x, y });
|
|
23
38
|
if (trans !== "no-transpose" && trans !== "transpose")
|
|
24
39
|
throw new Error("trans must be 'no-transpose' or 'transpose'.");
|
|
25
40
|
if (layout !== "row-major" && layout !== "column-major")
|
|
26
41
|
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
27
|
-
if (typeof alpha !== "number")
|
|
28
|
-
throw new Error("alpha must be a number.");
|
|
42
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
29
43
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
30
44
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
31
|
-
if (typeof beta !== "number")
|
|
32
|
-
throw new Error("beta must be a number.");
|
|
45
|
+
if (typeof beta !== "number") throw new Error("beta must be a number.");
|
|
33
46
|
if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
|
|
34
47
|
if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
|
|
35
48
|
if (
|
|
@@ -53,15 +66,15 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
|
|
|
53
66
|
"x and y must be the same type (both Float32Array or both GpuVector).",
|
|
54
67
|
);
|
|
55
68
|
if (xIsGpu && !AIsGpu)
|
|
56
|
-
throw new Error(
|
|
57
|
-
"A must be a GpuMatrix when x and y are GpuVectors.",
|
|
58
|
-
);
|
|
69
|
+
throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");
|
|
59
70
|
if (AIsGpu && !xIsGpu)
|
|
71
|
+
throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");
|
|
72
|
+
if (xIsGpu && x._buf === y._buf)
|
|
60
73
|
throw new Error(
|
|
61
|
-
"x and y must
|
|
74
|
+
"x and y must not reference the same GPU buffer when both are GpuVectors.",
|
|
62
75
|
);
|
|
63
|
-
if (
|
|
64
|
-
throw new Error("
|
|
76
|
+
if (AIsGpu && yIsGpu && A._buf === y._buf)
|
|
77
|
+
throw new Error("A and y must not reference the same GPU buffer.");
|
|
65
78
|
if (AIsGpu && lda !== A.lda)
|
|
66
79
|
throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
67
80
|
if (AIsGpu && (A.rows < m || A.cols < n))
|
|
@@ -97,40 +110,56 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
|
|
|
97
110
|
);
|
|
98
111
|
|
|
99
112
|
const shaderName = isNoTrans ? "sgemv_n" : "sgemv_t";
|
|
100
|
-
const pipeline
|
|
113
|
+
const pipeline = await getPipeline(device, shaderName);
|
|
101
114
|
|
|
102
|
-
|
|
103
|
-
|
|
104
|
-
|
|
105
|
-
|
|
106
|
-
[
|
|
107
|
-
{ value: m, type: "u32" },
|
|
108
|
-
{ value: n, type: "u32" },
|
|
109
|
-
{ value: alpha, type: "f32" },
|
|
110
|
-
{ value: beta, type: "f32" },
|
|
111
|
-
{ value: incx, type: "u32" },
|
|
112
|
-
{ value: incy, type: "u32" },
|
|
113
|
-
{ value: lda, type: "u32" },
|
|
114
|
-
],
|
|
115
|
-
"sgemv-params",
|
|
116
|
-
);
|
|
115
|
+
let ABuffer = null;
|
|
116
|
+
let xBuffer = null;
|
|
117
|
+
let yBuffer = null;
|
|
118
|
+
let paramsBuffer = null;
|
|
117
119
|
|
|
118
120
|
try {
|
|
119
|
-
|
|
121
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemv-A", false);
|
|
122
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sgemv-x", false);
|
|
123
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sgemv-y", true);
|
|
124
|
+
paramsBuffer = createParamsBuffer(
|
|
125
|
+
device,
|
|
126
|
+
[
|
|
127
|
+
{ value: m, type: "u32" },
|
|
128
|
+
{ value: n, type: "u32" },
|
|
129
|
+
{ value: alpha, type: "f32" },
|
|
130
|
+
{ value: beta, type: "f32" },
|
|
131
|
+
{ value: incx, type: "u32" },
|
|
132
|
+
{ value: incy, type: "u32" },
|
|
133
|
+
{ value: lda, type: "u32" },
|
|
134
|
+
],
|
|
135
|
+
"sgemv-params",
|
|
136
|
+
);
|
|
137
|
+
|
|
138
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
120
139
|
ABuffer,
|
|
121
140
|
xBuffer,
|
|
122
141
|
yBuffer,
|
|
123
142
|
paramsBuffer,
|
|
124
143
|
]);
|
|
125
144
|
|
|
126
|
-
// NoTrans: one workgroup per row
|
|
145
|
+
// NoTrans: one workgroup per row — sgemv_n.wgsl is a grid-stride loop, so
|
|
146
|
+
// clamping here only costs parallelism. Trans: one thread per output
|
|
147
|
+
// column, and sgemv_t.wgsl indexes straight off global_invocation_id with
|
|
148
|
+
// no fallback, so an over-limit dispatch must be refused, not truncated.
|
|
127
149
|
const wgCount = isNoTrans
|
|
128
150
|
? Math.min(m, device.limits.maxComputeWorkgroupsPerDimension)
|
|
129
|
-
:
|
|
130
|
-
const { commandEncoder, ts } = runComputePass(
|
|
131
|
-
|
|
151
|
+
: requireWorkgroups(device, "sgemv", yLen);
|
|
152
|
+
const { commandEncoder, ts } = runComputePass(
|
|
153
|
+
device,
|
|
154
|
+
pipeline,
|
|
155
|
+
bindGroup,
|
|
156
|
+
wgCount,
|
|
157
|
+
);
|
|
158
|
+
const readBuffer = yIsGpu
|
|
159
|
+
? null
|
|
160
|
+
: stageReadback(device, commandEncoder, yBuffer);
|
|
132
161
|
|
|
133
|
-
submit(commandEncoder);
|
|
162
|
+
submit(device, commandEncoder);
|
|
134
163
|
|
|
135
164
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
136
165
|
|
|
@@ -143,10 +172,9 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
|
|
|
143
172
|
if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
|
|
144
173
|
return { y: result };
|
|
145
174
|
} finally {
|
|
146
|
-
if (!AIsGpu) destroyBuffers(ABuffer);
|
|
147
|
-
if (!xIsGpu) destroyBuffers(xBuffer);
|
|
148
|
-
if (!yIsGpu) destroyBuffers(yBuffer);
|
|
149
|
-
destroyBuffers(paramsBuffer);
|
|
150
|
-
|
|
175
|
+
if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
|
|
176
|
+
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
177
|
+
if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
|
|
178
|
+
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
151
179
|
}
|
|
152
180
|
}
|
package/src/sger/sger.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 rank-1 update A
|
|
5
|
+
* Performs the rank-1 update $$A \leftarrow \alpha x y^{T} + A$$
|
|
6
6
|
*
|
|
7
7
|
* A is an m×n matrix stored in row-major order, updated in place. `lda` is
|
|
8
8
|
* the leading dimension (number of floats between the start of consecutive
|
|
@@ -43,7 +43,7 @@ export declare function sger(
|
|
|
43
43
|
): Promise<{ A: Float32Array; gpuTimeMs?: number }>;
|
|
44
44
|
|
|
45
45
|
/**
|
|
46
|
-
* Performs the rank-1 update A
|
|
46
|
+
* Performs the rank-1 update $$A \leftarrow \alpha x y^{T} + A$$
|
|
47
47
|
*
|
|
48
48
|
* x, y, and A are all kept resident on the GPU. `A`'s own `layout` (set at
|
|
49
49
|
* `GpuMatrix.from` time) determines the operation — there is no separate
|
package/src/sger/sger.mjs
CHANGED
|
@@ -11,16 +11,28 @@ 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 { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
|
-
export async function sger(
|
|
16
|
+
export async function sger(
|
|
17
|
+
device,
|
|
18
|
+
m,
|
|
19
|
+
n,
|
|
20
|
+
alpha,
|
|
21
|
+
x,
|
|
22
|
+
incx,
|
|
23
|
+
y,
|
|
24
|
+
incy,
|
|
25
|
+
A,
|
|
26
|
+
lda,
|
|
27
|
+
layout = "row-major",
|
|
28
|
+
) {
|
|
16
29
|
const AIsGpu = A instanceof GpuMatrix;
|
|
17
30
|
|
|
18
|
-
|
|
19
|
-
|
|
31
|
+
requireGpuDevice(device);
|
|
32
|
+
requireSameDevice(device, "sger", { A, x, y });
|
|
20
33
|
if (layout !== "row-major" && layout !== "column-major")
|
|
21
34
|
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
22
|
-
if (typeof alpha !== "number")
|
|
23
|
-
throw new Error("alpha must be a number.");
|
|
35
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
24
36
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
25
37
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
26
38
|
if (
|
|
@@ -77,9 +89,13 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
|
|
|
77
89
|
"A does not have enough elements for the given m, n, and lda.",
|
|
78
90
|
);
|
|
79
91
|
if (x.length < (m - 1) * incx + 1)
|
|
80
|
-
throw new Error(
|
|
92
|
+
throw new Error(
|
|
93
|
+
"x does not have enough elements for the given m and incx.",
|
|
94
|
+
);
|
|
81
95
|
if (y.length < (n - 1) * incy + 1)
|
|
82
|
-
throw new Error(
|
|
96
|
+
throw new Error(
|
|
97
|
+
"y does not have enough elements for the given n and incy.",
|
|
98
|
+
);
|
|
83
99
|
|
|
84
100
|
const pipeline = await getPipeline(device, "sger");
|
|
85
101
|
|
|
@@ -89,22 +105,23 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
|
|
|
89
105
|
let paramsBuffer = null;
|
|
90
106
|
|
|
91
107
|
try {
|
|
92
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sger-x", false);
|
|
93
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sger-y", false);
|
|
94
|
-
ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sger-A", true);
|
|
108
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sger-x", false);
|
|
109
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sger-y", false);
|
|
110
|
+
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sger-A", true);
|
|
95
111
|
paramsBuffer = createParamsBuffer(
|
|
112
|
+
device,
|
|
96
113
|
[
|
|
97
|
-
{ value: m,
|
|
98
|
-
{ value: n,
|
|
114
|
+
{ value: m, type: "u32" },
|
|
115
|
+
{ value: n, type: "u32" },
|
|
99
116
|
{ value: alpha, type: "f32" },
|
|
100
|
-
{ value: incx,
|
|
101
|
-
{ value: incy,
|
|
102
|
-
{ value: lda,
|
|
117
|
+
{ value: incx, type: "u32" },
|
|
118
|
+
{ value: incy, type: "u32" },
|
|
119
|
+
{ value: lda, type: "u32" },
|
|
103
120
|
],
|
|
104
121
|
"sger-params",
|
|
105
122
|
);
|
|
106
123
|
|
|
107
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
124
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
108
125
|
xBuffer,
|
|
109
126
|
yBuffer,
|
|
110
127
|
ABuffer,
|
|
@@ -114,10 +131,17 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
|
|
|
114
131
|
// One workgroup per row of A; clamped to device limit — the shader's
|
|
115
132
|
// grid-stride loop handles remaining rows when m > dispatch count.
|
|
116
133
|
const wgCount = Math.min(m, device.limits.maxComputeWorkgroupsPerDimension);
|
|
117
|
-
const { commandEncoder, ts } = runComputePass(
|
|
118
|
-
|
|
134
|
+
const { commandEncoder, ts } = runComputePass(
|
|
135
|
+
device,
|
|
136
|
+
pipeline,
|
|
137
|
+
bindGroup,
|
|
138
|
+
wgCount,
|
|
139
|
+
);
|
|
140
|
+
const readBuffer = AIsGpu
|
|
141
|
+
? null
|
|
142
|
+
: stageReadback(device, commandEncoder, ABuffer);
|
|
119
143
|
|
|
120
|
-
submit(commandEncoder);
|
|
144
|
+
submit(device, commandEncoder);
|
|
121
145
|
|
|
122
146
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
123
147
|
|
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
// cscal: x := alpha * x, complex. x is one interleaved f32 array
|
|
2
|
+
// (re0, im0, re1, im1, ...), matching Complex32Array/GpuVector's storage
|
|
3
|
+
// (and cuBLAS's cuComplex / stdlib's Complex64Array) — no repacking needed
|
|
4
|
+
// between JS and GPU.
|
|
5
|
+
// (alphaRe + i*alphaIm)(re + i*im) = (alphaRe*re - alphaIm*im) + i*(alphaRe*im + alphaIm*re)
|
|
6
|
+
|
|
7
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
8
|
+
|
|
9
|
+
struct Params {
|
|
10
|
+
n: u32,
|
|
11
|
+
alphaRe: f32,
|
|
12
|
+
alphaIm: f32,
|
|
13
|
+
x_inc: u32,
|
|
14
|
+
}
|
|
15
|
+
|
|
16
|
+
@group(0) @binding(1) var<uniform> params: Params;
|
|
17
|
+
|
|
18
|
+
const WGS: u32 = 64;
|
|
19
|
+
|
|
20
|
+
@compute @workgroup_size(64)
|
|
21
|
+
fn main(
|
|
22
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
23
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
24
|
+
) {
|
|
25
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
26
|
+
let base = 2u * id * params.x_inc;
|
|
27
|
+
// Both new parts need both old parts, so capture them before either write.
|
|
28
|
+
let re = x[base];
|
|
29
|
+
let im = x[base + 1u];
|
|
30
|
+
x[base] = params.alphaRe * re - params.alphaIm * im;
|
|
31
|
+
x[base + 1u] = params.alphaRe * im + params.alphaIm * re;
|
|
32
|
+
}
|
|
33
|
+
}
|
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
// daxpy: y := alpha * x + y, double-double (Dekker) f64 emulation of saxpy.
|
|
2
|
+
// Each element costs one ddMulProtected (alpha*x[i]) then one ddAddProtected
|
|
3
|
+
// (+ y[i]) — the same two-protected-op shape ddot spends per term, applied
|
|
4
|
+
// straight to the output instead of folded into a reduction. See dscal.wgsl
|
|
5
|
+
// for why this is a uniform main pass plus a ragged, select-masked tail
|
|
6
|
+
// rather than a plain `id < params.n` grid-stride loop.
|
|
7
|
+
|
|
8
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
9
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
10
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
11
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
12
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
13
|
+
|
|
14
|
+
struct Params {
|
|
15
|
+
n: u32,
|
|
16
|
+
alphaHi: f32,
|
|
17
|
+
alphaLo: f32,
|
|
18
|
+
x_inc: u32,
|
|
19
|
+
y_inc: u32,
|
|
20
|
+
}
|
|
21
|
+
|
|
22
|
+
const WGS: u32 = 64;
|
|
23
|
+
|
|
24
|
+
@compute @workgroup_size(64)
|
|
25
|
+
fn daxpy_main(
|
|
26
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
27
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
28
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
29
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
30
|
+
) {
|
|
31
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
32
|
+
let stride = num_wg.x * WGS;
|
|
33
|
+
|
|
34
|
+
let n_floor = (params.n / stride) * stride;
|
|
35
|
+
let mainIters = n_floor / stride;
|
|
36
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
37
|
+
let id = gid.x + iter * stride;
|
|
38
|
+
let ix = id * params.x_inc;
|
|
39
|
+
let iy = id * params.y_inc;
|
|
40
|
+
let prod = ddMulProtected(alpha, DD(xHi[ix], xLo[ix]), lid.x);
|
|
41
|
+
let result = ddAddProtected(prod, DD(yHi[iy], yLo[iy]), lid.x);
|
|
42
|
+
yHi[iy] = result.hi;
|
|
43
|
+
yLo[iy] = result.lo;
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
// Tail is ragged (0 or 1 extra per thread) — pad to this workgroup's worst
|
|
47
|
+
// case so every thread still calls ddMulProtected/ddAddProtected the same
|
|
48
|
+
// number of times (their barriers need that), masking only the write.
|
|
49
|
+
let wgBaseGid = wgid.x * WGS;
|
|
50
|
+
var tailIters = 0u;
|
|
51
|
+
if (n_floor + wgBaseGid < params.n) {
|
|
52
|
+
tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
|
|
53
|
+
}
|
|
54
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
55
|
+
let id = n_floor + gid.x + iter * stride;
|
|
56
|
+
let valid = id < params.n;
|
|
57
|
+
let ix = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
58
|
+
let iy = select(0u, id * params.y_inc, valid);
|
|
59
|
+
let prod = ddMulProtected(alpha, DD(xHi[ix], xLo[ix]), lid.x);
|
|
60
|
+
let result = ddAddProtected(prod, DD(yHi[iy], yLo[iy]), lid.x);
|
|
61
|
+
if (valid) {
|
|
62
|
+
yHi[iy] = result.hi;
|
|
63
|
+
yLo[iy] = result.lo;
|
|
64
|
+
}
|
|
65
|
+
}
|
|
66
|
+
}
|
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
// dcopy: y = x, double-double (Dekker) f64 emulation of scopy. A copy is
|
|
2
|
+
// pure data movement — hi and lo are transferred verbatim, with no
|
|
3
|
+
// arithmetic at all — so (unlike dscal/daxpy/ddot) this needs no
|
|
4
|
+
// ddMulProtected/ddAddProtected renormalizing barrier, and so no
|
|
5
|
+
// ragged-tail-with-barrier split; a plain grid-stride loop is safe, same
|
|
6
|
+
// shape as scopy.wgsl itself.
|
|
7
|
+
|
|
8
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
9
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
10
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
11
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
12
|
+
|
|
13
|
+
struct Params {
|
|
14
|
+
n: u32,
|
|
15
|
+
x_inc: u32,
|
|
16
|
+
y_inc: u32,
|
|
17
|
+
}
|
|
18
|
+
|
|
19
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
20
|
+
|
|
21
|
+
const WGS: u32 = 64;
|
|
22
|
+
|
|
23
|
+
@compute @workgroup_size(64)
|
|
24
|
+
fn main(
|
|
25
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
26
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
27
|
+
) {
|
|
28
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
29
|
+
let ix = id * params.x_inc;
|
|
30
|
+
let iy = id * params.y_inc;
|
|
31
|
+
yHi[iy] = xHi[ix];
|
|
32
|
+
yLo[iy] = xLo[ix];
|
|
33
|
+
}
|
|
34
|
+
}
|
|
@@ -0,0 +1,106 @@
|
|
|
1
|
+
// ddot: sum(x[i] * y[i]), double-double (Dekker). Same ILP=4 shape as
|
|
2
|
+
// dasum.wgsl, which this mirrors closely — the only structural difference is
|
|
3
|
+
// a second input vector and a product where dasum takes an absolute value.
|
|
4
|
+
//
|
|
5
|
+
// See f64/utils/add.wgsl for ddAddProtected and why plain ddAdd isn't safe,
|
|
6
|
+
// and f64/utils/multiply.wgsl for ddMulProtected. The multiply itself
|
|
7
|
+
// (twoProdBit) needs no barrier; only its final renormalisation does, which
|
|
8
|
+
// is why each element costs two protected ops here against dasum's one.
|
|
9
|
+
|
|
10
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
11
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
12
|
+
@group(0) @binding(2) var<storage, read> yHi: array<f32>;
|
|
13
|
+
@group(0) @binding(3) var<storage, read> yLo: array<f32>;
|
|
14
|
+
@group(0) @binding(4) var<storage, read_write> partialsHi: array<f32>;
|
|
15
|
+
@group(0) @binding(5) var<storage, read_write> partialsLo: array<f32>;
|
|
16
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
17
|
+
|
|
18
|
+
struct Params {
|
|
19
|
+
n: u32,
|
|
20
|
+
x_inc: u32,
|
|
21
|
+
y_inc: u32,
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
const WGS: u32 = 64;
|
|
25
|
+
|
|
26
|
+
var<workgroup> tile: array<DD, 64>;
|
|
27
|
+
|
|
28
|
+
@compute @workgroup_size(64)
|
|
29
|
+
fn ddot_main(
|
|
30
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
31
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
32
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
33
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
34
|
+
) {
|
|
35
|
+
var acc0 = DD(0.0, 0.0);
|
|
36
|
+
var acc1 = DD(0.0, 0.0);
|
|
37
|
+
var acc2 = DD(0.0, 0.0);
|
|
38
|
+
var acc3 = DD(0.0, 0.0);
|
|
39
|
+
|
|
40
|
+
let stride = num_wg.x * WGS;
|
|
41
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
42
|
+
|
|
43
|
+
// Same trip count for every thread, but driven by a counter, not `id`
|
|
44
|
+
// itself (the protected ops' barriers need a provably-uniform loop bound).
|
|
45
|
+
let mainIters = n4_floor / (4u * stride);
|
|
46
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
47
|
+
let id = gid.x + iter * 4u * stride;
|
|
48
|
+
let d0 = id;
|
|
49
|
+
let d1 = id + stride;
|
|
50
|
+
let d2 = id + 2u * stride;
|
|
51
|
+
let d3 = id + 3u * stride;
|
|
52
|
+
|
|
53
|
+
let p0 = ddMulProtected(DD(xHi[d0 * params.x_inc], xLo[d0 * params.x_inc]),
|
|
54
|
+
DD(yHi[d0 * params.y_inc], yLo[d0 * params.y_inc]), lid.x);
|
|
55
|
+
let p1 = ddMulProtected(DD(xHi[d1 * params.x_inc], xLo[d1 * params.x_inc]),
|
|
56
|
+
DD(yHi[d1 * params.y_inc], yLo[d1 * params.y_inc]), lid.x);
|
|
57
|
+
let p2 = ddMulProtected(DD(xHi[d2 * params.x_inc], xLo[d2 * params.x_inc]),
|
|
58
|
+
DD(yHi[d2 * params.y_inc], yLo[d2 * params.y_inc]), lid.x);
|
|
59
|
+
let p3 = ddMulProtected(DD(xHi[d3 * params.x_inc], xLo[d3 * params.x_inc]),
|
|
60
|
+
DD(yHi[d3 * params.y_inc], yLo[d3 * params.y_inc]), lid.x);
|
|
61
|
+
|
|
62
|
+
acc0 = ddAddProtected(acc0, p0, lid.x);
|
|
63
|
+
acc1 = ddAddProtected(acc1, p1, lid.x);
|
|
64
|
+
acc2 = ddAddProtected(acc2, p2, lid.x);
|
|
65
|
+
acc3 = ddAddProtected(acc3, p3, lid.x);
|
|
66
|
+
}
|
|
67
|
+
|
|
68
|
+
// Tail is ragged (0-3 extra per thread) — pad to this workgroup's worst case.
|
|
69
|
+
// Out-of-range lanes still run the multiply (it carries a barrier, so every
|
|
70
|
+
// thread must reach it) against index 0, then mask the result to zero.
|
|
71
|
+
let wgBaseGid = wgid.x * WGS;
|
|
72
|
+
var tailIters = 0u;
|
|
73
|
+
if (n4_floor + wgBaseGid < params.n) {
|
|
74
|
+
tailIters = (params.n - 1u - n4_floor - wgBaseGid) / stride + 1u;
|
|
75
|
+
}
|
|
76
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
77
|
+
let id = n4_floor + gid.x + iter * stride;
|
|
78
|
+
let valid = id < params.n;
|
|
79
|
+
let ix = select(0u, id * params.x_inc, valid);
|
|
80
|
+
let iy = select(0u, id * params.y_inc, valid);
|
|
81
|
+
let prod = ddMulProtected(DD(xHi[ix], xLo[ix]), DD(yHi[iy], yLo[iy]), lid.x);
|
|
82
|
+
// select() has no DD overload
|
|
83
|
+
let contribution = DD(select(0.0, prod.hi, valid), select(0.0, prod.lo, valid));
|
|
84
|
+
acc0 = ddAddProtected(acc0, contribution, lid.x);
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
let combined01 = ddAddProtected(acc0, acc1, lid.x);
|
|
88
|
+
let combined23 = ddAddProtected(acc2, acc3, lid.x);
|
|
89
|
+
tile[lid.x] = ddAddProtected(combined01, combined23, lid.x);
|
|
90
|
+
workgroupBarrier();
|
|
91
|
+
|
|
92
|
+
// Inactive threads combine against a throwaway partner and discard it
|
|
93
|
+
// (ddAddProtected must be called unconditionally by every thread).
|
|
94
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
95
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
96
|
+
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
97
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
98
|
+
if (lid.x < s) { tile[lid.x] = combined; }
|
|
99
|
+
workgroupBarrier();
|
|
100
|
+
}
|
|
101
|
+
|
|
102
|
+
if (lid.x == 0u) {
|
|
103
|
+
partialsHi[wgid.x] = tile[0].hi;
|
|
104
|
+
partialsLo[wgid.x] = tile[0].lo;
|
|
105
|
+
}
|
|
106
|
+
}
|