wgblas 2.1.0 → 2.2.1
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 +1016 -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 +29 -11
- 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
|
@@ -0,0 +1,184 @@
|
|
|
1
|
+
import {
|
|
2
|
+
createStorageBuffer,
|
|
3
|
+
createParamsBuffer,
|
|
4
|
+
createResultBuffer,
|
|
5
|
+
stageReadback,
|
|
6
|
+
destroyBuffers,
|
|
7
|
+
uploadBuffer,
|
|
8
|
+
} from "../util/buffer.mjs";
|
|
9
|
+
import { createBindGroup } from "../util/bindgroup.mjs";
|
|
10
|
+
import { runComputePass, submit } from "../util/compute.mjs";
|
|
11
|
+
import { extractTimestamp } from "../util/benchmark.mjs";
|
|
12
|
+
import { extractResult } from "../util/result.mjs";
|
|
13
|
+
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
|
+
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
15
|
+
import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
|
|
16
|
+
import { WGS } from "../util/constants.mjs";
|
|
17
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
18
|
+
|
|
19
|
+
export async function ddot(device, n, x, incx, y, incy) {
|
|
20
|
+
const xIsGpu = x instanceof GpuVector;
|
|
21
|
+
const yIsGpu = y instanceof GpuVector;
|
|
22
|
+
|
|
23
|
+
requireGpuDevice(device);
|
|
24
|
+
requireSameDevice(device, "ddot", { x, y });
|
|
25
|
+
if (
|
|
26
|
+
!Number.isInteger(n) ||
|
|
27
|
+
!Number.isInteger(incx) ||
|
|
28
|
+
!Number.isInteger(incy)
|
|
29
|
+
)
|
|
30
|
+
throw new Error("n, incx, and incy must be integers.");
|
|
31
|
+
if (incx <= 0 || incy <= 0)
|
|
32
|
+
throw new Error("incx and incy must be positive.");
|
|
33
|
+
if (!xIsGpu && !(x instanceof Float64Array))
|
|
34
|
+
throw new Error("x must be a Float64Array or GpuVector.");
|
|
35
|
+
if (!yIsGpu && !(y instanceof Float64Array))
|
|
36
|
+
throw new Error("y must be a Float64Array or GpuVector.");
|
|
37
|
+
if (xIsGpu && x.dtype !== Float64Array)
|
|
38
|
+
throw new Error("x must be a Float64Array-backed GpuVector.");
|
|
39
|
+
if (yIsGpu && y.dtype !== Float64Array)
|
|
40
|
+
throw new Error("y must be a Float64Array-backed GpuVector.");
|
|
41
|
+
if (xIsGpu !== yIsGpu)
|
|
42
|
+
throw new Error(
|
|
43
|
+
"x and y must be the same type (both Float64Array or both GpuVector).",
|
|
44
|
+
);
|
|
45
|
+
if (n <= 0) return { dot: 0 };
|
|
46
|
+
if (x.length < (n - 1) * incx + 1)
|
|
47
|
+
throw new Error(
|
|
48
|
+
"x does not have enough elements for the given n and incx.",
|
|
49
|
+
);
|
|
50
|
+
if (y.length < (n - 1) * incy + 1)
|
|
51
|
+
throw new Error(
|
|
52
|
+
"y does not have enough elements for the given n and incy.",
|
|
53
|
+
);
|
|
54
|
+
|
|
55
|
+
// Concatenated with f64/dekker.wgsl (DD struct) and f64/utils/add.wgsl
|
|
56
|
+
// (ddAddProtected) — WGSL has no #include; entryPoint omitted since each
|
|
57
|
+
// module has only one @compute. Only the first pass multiplies, so
|
|
58
|
+
// f64/utils/multiply.wgsl is left out of the reduction's module.
|
|
59
|
+
const ddCore = ["f64/dekker", "f64/utils/add"];
|
|
60
|
+
const pipelineMain = await getPipeline(device, [
|
|
61
|
+
...ddCore,
|
|
62
|
+
"f64/utils/multiply",
|
|
63
|
+
"ddot",
|
|
64
|
+
]);
|
|
65
|
+
const pipelineReduce = await getPipeline(device, [
|
|
66
|
+
...ddCore,
|
|
67
|
+
"reduction/sumF64",
|
|
68
|
+
]);
|
|
69
|
+
|
|
70
|
+
let xHiBuffer = null;
|
|
71
|
+
let xLoBuffer = null;
|
|
72
|
+
let yHiBuffer = null;
|
|
73
|
+
let yLoBuffer = null;
|
|
74
|
+
let partialsHiBuffer = null;
|
|
75
|
+
let partialsLoBuffer = null;
|
|
76
|
+
let resultHiBuffer = null;
|
|
77
|
+
let resultLoBuffer = null;
|
|
78
|
+
let paramsBuffer = null;
|
|
79
|
+
let readHiBuffer = null;
|
|
80
|
+
let readLoBuffer = null;
|
|
81
|
+
|
|
82
|
+
try {
|
|
83
|
+
if (xIsGpu) {
|
|
84
|
+
xHiBuffer = x._buf;
|
|
85
|
+
xLoBuffer = x._loBuf;
|
|
86
|
+
yHiBuffer = y._buf;
|
|
87
|
+
yLoBuffer = y._loBuf;
|
|
88
|
+
} else {
|
|
89
|
+
// No abs here (unlike dasum) — the product carries its own sign.
|
|
90
|
+
const xSplit = splitDoubleDouble(x);
|
|
91
|
+
const ySplit = splitDoubleDouble(y);
|
|
92
|
+
xHiBuffer = uploadBuffer(device, xSplit.hi, "ddot-xHi", false);
|
|
93
|
+
xLoBuffer = uploadBuffer(device, xSplit.lo, "ddot-xLo", false);
|
|
94
|
+
yHiBuffer = uploadBuffer(device, ySplit.hi, "ddot-yHi", false);
|
|
95
|
+
yLoBuffer = uploadBuffer(device, ySplit.lo, "ddot-yLo", false);
|
|
96
|
+
}
|
|
97
|
+
partialsHiBuffer = createStorageBuffer(
|
|
98
|
+
device,
|
|
99
|
+
2 * WGS * 4,
|
|
100
|
+
"ddot-partialsHi",
|
|
101
|
+
);
|
|
102
|
+
partialsLoBuffer = createStorageBuffer(
|
|
103
|
+
device,
|
|
104
|
+
2 * WGS * 4,
|
|
105
|
+
"ddot-partialsLo",
|
|
106
|
+
);
|
|
107
|
+
resultHiBuffer = createResultBuffer(device, 4, "ddot-result-hi");
|
|
108
|
+
resultLoBuffer = createResultBuffer(device, 4, "ddot-result-lo");
|
|
109
|
+
paramsBuffer = createParamsBuffer(
|
|
110
|
+
device,
|
|
111
|
+
[
|
|
112
|
+
{ value: n, type: "u32" },
|
|
113
|
+
{ value: incx, type: "u32" },
|
|
114
|
+
{ value: incy, type: "u32" },
|
|
115
|
+
],
|
|
116
|
+
"ddot-params",
|
|
117
|
+
);
|
|
118
|
+
|
|
119
|
+
const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
|
|
120
|
+
xHiBuffer,
|
|
121
|
+
xLoBuffer,
|
|
122
|
+
yHiBuffer,
|
|
123
|
+
yLoBuffer,
|
|
124
|
+
partialsHiBuffer,
|
|
125
|
+
partialsLoBuffer,
|
|
126
|
+
paramsBuffer,
|
|
127
|
+
]);
|
|
128
|
+
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
129
|
+
device,
|
|
130
|
+
pipelineMain,
|
|
131
|
+
bgMain,
|
|
132
|
+
2 * WGS,
|
|
133
|
+
); // dispatch 2*WGS workgroups — one partial per workgroup
|
|
134
|
+
|
|
135
|
+
submit(device, enc1);
|
|
136
|
+
|
|
137
|
+
const bgReduce = createBindGroup(
|
|
138
|
+
device,
|
|
139
|
+
pipelineReduce.getBindGroupLayout(0),
|
|
140
|
+
[partialsHiBuffer, partialsLoBuffer, resultHiBuffer, resultLoBuffer],
|
|
141
|
+
);
|
|
142
|
+
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
143
|
+
device,
|
|
144
|
+
pipelineReduce,
|
|
145
|
+
bgReduce,
|
|
146
|
+
1,
|
|
147
|
+
); // dispatch 1 workgroup to reduce the partial sums to a single result
|
|
148
|
+
readHiBuffer = stageReadback(device, enc2, resultHiBuffer);
|
|
149
|
+
readLoBuffer = stageReadback(device, enc2, resultLoBuffer);
|
|
150
|
+
|
|
151
|
+
submit(device, enc2);
|
|
152
|
+
|
|
153
|
+
const hiPromise = extractResult(readHiBuffer, Float32Array);
|
|
154
|
+
const loPromise = extractResult(readLoBuffer, Float32Array);
|
|
155
|
+
readHiBuffer = null; // ownership transferred — extractResult's own finally destroys it
|
|
156
|
+
readLoBuffer = null;
|
|
157
|
+
|
|
158
|
+
const [gpuTime1, gpuTime2, hiArr, loArr] = await Promise.all([
|
|
159
|
+
extractTimestamp(ts1),
|
|
160
|
+
extractTimestamp(ts2),
|
|
161
|
+
hiPromise,
|
|
162
|
+
loPromise,
|
|
163
|
+
]);
|
|
164
|
+
|
|
165
|
+
// dot is always a scalar readback — both paths return { dot }
|
|
166
|
+
const dot = mergeDoubleDouble(hiArr, loArr)[0];
|
|
167
|
+
if (gpuTime1 !== undefined && gpuTime2 !== undefined)
|
|
168
|
+
return { dot, gpuTimeMs: gpuTime1 + gpuTime2 };
|
|
169
|
+
return { dot };
|
|
170
|
+
} finally {
|
|
171
|
+
if (!xIsGpu && xHiBuffer) destroyBuffers(xHiBuffer);
|
|
172
|
+
if (!xIsGpu && xLoBuffer) destroyBuffers(xLoBuffer);
|
|
173
|
+
if (!yIsGpu && yHiBuffer) destroyBuffers(yHiBuffer);
|
|
174
|
+
if (!yIsGpu && yLoBuffer) destroyBuffers(yLoBuffer);
|
|
175
|
+
if (partialsHiBuffer) destroyBuffers(partialsHiBuffer);
|
|
176
|
+
if (partialsLoBuffer) destroyBuffers(partialsLoBuffer);
|
|
177
|
+
if (resultHiBuffer) destroyBuffers(resultHiBuffer);
|
|
178
|
+
if (resultLoBuffer) destroyBuffers(resultLoBuffer);
|
|
179
|
+
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
180
|
+
// Only reached if submit(device, enc2) threw before ownership was transferred above.
|
|
181
|
+
if (readHiBuffer) destroyBuffers(readHiBuffer);
|
|
182
|
+
if (readLoBuffer) destroyBuffers(readLoBuffer);
|
|
183
|
+
}
|
|
184
|
+
}
|
|
@@ -0,0 +1,50 @@
|
|
|
1
|
+
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
|
+
|
|
3
|
+
/**
|
|
4
|
+
* Computes the Euclidean norm of a double-precision vector:
|
|
5
|
+
* $$\text{result} = \sqrt{\sum_{i} x_i^2}$$
|
|
6
|
+
* — double-double (Dekker) f64 emulation of {@link snrm2}, since WGSL has
|
|
7
|
+
* no native f64 type.
|
|
8
|
+
*
|
|
9
|
+
* {@includeCode ../../examples/dnrm2/dnrm2.js}
|
|
10
|
+
*
|
|
11
|
+
* **Browser (standalone HTML):**
|
|
12
|
+
* {@includeCode ../../examples/dnrm2/web/dnrm2.html}
|
|
13
|
+
*
|
|
14
|
+
* @param device - GPUDevice from `init()`
|
|
15
|
+
* @param n - number of elements (must be a positive integer)
|
|
16
|
+
* @param x - Float64Array input vector
|
|
17
|
+
* @param incx - stride for x (must be a positive integer)
|
|
18
|
+
* @returns Euclidean norm scalar — always a CPU readback, even for GpuVector inputs
|
|
19
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/dnrm2/dnrm2.mjs#L21">Source code: dnrm2.mjs (L21)</a>
|
|
20
|
+
* @category BLAS Level 1
|
|
21
|
+
*/
|
|
22
|
+
export declare function dnrm2(
|
|
23
|
+
device: GPUDevice,
|
|
24
|
+
n: number,
|
|
25
|
+
x: Float64Array,
|
|
26
|
+
incx: number,
|
|
27
|
+
): Promise<{ nrm2: number } | { nrm2: number; gpuTimeMs: number }>;
|
|
28
|
+
|
|
29
|
+
/**
|
|
30
|
+
* Computes the Euclidean norm of a double-precision vector:
|
|
31
|
+
* $$\text{result} = \sqrt{\sum_{i} x_i^2}$$
|
|
32
|
+
* — GPU-resident overload; see the Float64Array overload above for the
|
|
33
|
+
* routine itself.
|
|
34
|
+
*
|
|
35
|
+
* {@includeCode ../../examples/dnrm2/gpu.dnrm2.js}
|
|
36
|
+
*
|
|
37
|
+
* @param device - GPUDevice from `init()`
|
|
38
|
+
* @param n - number of elements (must be a positive integer)
|
|
39
|
+
* @param x - GpuVector input vector (must be Float64Array-backed)
|
|
40
|
+
* @param incx - stride for x (must be a positive integer)
|
|
41
|
+
* @returns Euclidean norm scalar — always a CPU readback, even for GpuVector inputs
|
|
42
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/dnrm2/dnrm2.mjs#L21">Source code: dnrm2.mjs (L21)</a>
|
|
43
|
+
* @category BLAS Level 1
|
|
44
|
+
*/
|
|
45
|
+
export declare function dnrm2(
|
|
46
|
+
device: GPUDevice,
|
|
47
|
+
n: number,
|
|
48
|
+
x: GpuVector,
|
|
49
|
+
incx: number,
|
|
50
|
+
): Promise<{ nrm2: number } | { nrm2: number; gpuTimeMs: number }>;
|
|
@@ -0,0 +1,189 @@
|
|
|
1
|
+
import {
|
|
2
|
+
uploadBuffer,
|
|
3
|
+
createStorageBuffer,
|
|
4
|
+
createParamsBuffer,
|
|
5
|
+
createResultBuffer,
|
|
6
|
+
stageReadback,
|
|
7
|
+
destroyBuffers,
|
|
8
|
+
} from "../util/buffer.mjs";
|
|
9
|
+
import { createBindGroup } from "../util/bindgroup.mjs";
|
|
10
|
+
import { runComputePass, submit } from "../util/compute.mjs";
|
|
11
|
+
import { extractTimestamp } from "../util/benchmark.mjs";
|
|
12
|
+
import { extractResult } from "../util/result.mjs";
|
|
13
|
+
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
|
+
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
15
|
+
import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
|
|
16
|
+
import { WGS } from "../util/constants.mjs";
|
|
17
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
18
|
+
|
|
19
|
+
// dnrm2: sqrt(sum(x[i]*x[i])), double-double (Dekker) f64 emulation of
|
|
20
|
+
// snrm2 — same scaled accumulation (Blue's algorithm) snrm2 uses to avoid
|
|
21
|
+
// overflow, with (scale, ssq) as double-double pairs instead of plain f32.
|
|
22
|
+
// See dnrm2.wgsl for why the per-element branch that algorithm needs had to
|
|
23
|
+
// become branch-free (select()-based) here.
|
|
24
|
+
export async function dnrm2(device, n, x, incx) {
|
|
25
|
+
const xIsGpu = x instanceof GpuVector;
|
|
26
|
+
|
|
27
|
+
requireGpuDevice(device);
|
|
28
|
+
requireSameDevice(device, "dnrm2", { x });
|
|
29
|
+
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
30
|
+
throw new Error("n and incx must be integers.");
|
|
31
|
+
if (incx <= 0) throw new Error("incx must be positive.");
|
|
32
|
+
if (!xIsGpu && !(x instanceof Float64Array))
|
|
33
|
+
throw new Error("x must be a Float64Array or GpuVector.");
|
|
34
|
+
if (xIsGpu && x.dtype !== Float64Array)
|
|
35
|
+
throw new Error("x must be a Float64Array-backed GpuVector.");
|
|
36
|
+
if (n <= 0) return { nrm2: 0 };
|
|
37
|
+
if (x.length < (n - 1) * incx + 1)
|
|
38
|
+
throw new Error(
|
|
39
|
+
"x does not have enough elements for the given n and incx.",
|
|
40
|
+
);
|
|
41
|
+
|
|
42
|
+
// Concatenated with f64/dekker.wgsl (DD struct), f64/utils/abs.wgsl
|
|
43
|
+
// (ddAbs), f64/utils/greater.wgsl (ddGreater), f64/utils/add.wgsl
|
|
44
|
+
// (ddAddProtected), f64/utils/multiply.wgsl (ddMulProtected),
|
|
45
|
+
// f64/utils/divide.wgsl (ddDivProtected), and f64/utils/sqrt.wgsl
|
|
46
|
+
// (ddSqrtProtected) — WGSL has no #include.
|
|
47
|
+
const f64Deps = [
|
|
48
|
+
"f64/dekker",
|
|
49
|
+
"f64/utils/abs",
|
|
50
|
+
"f64/utils/greater",
|
|
51
|
+
"f64/utils/add",
|
|
52
|
+
"f64/utils/multiply",
|
|
53
|
+
"f64/utils/divide",
|
|
54
|
+
"f64/utils/sqrt",
|
|
55
|
+
];
|
|
56
|
+
const pipelineMain = await getPipeline(device, [...f64Deps, "dnrm2"]);
|
|
57
|
+
const pipelineReduce = await getPipeline(device, [
|
|
58
|
+
...f64Deps,
|
|
59
|
+
"reduction/scaledSumF64",
|
|
60
|
+
]);
|
|
61
|
+
|
|
62
|
+
let xHiBuffer = null;
|
|
63
|
+
let xLoBuffer = null;
|
|
64
|
+
let partialsScaleHiBuffer = null;
|
|
65
|
+
let partialsScaleLoBuffer = null;
|
|
66
|
+
let partialsSsqHiBuffer = null;
|
|
67
|
+
let partialsSsqLoBuffer = null;
|
|
68
|
+
let resultHiBuffer = null;
|
|
69
|
+
let resultLoBuffer = null;
|
|
70
|
+
let paramsBuffer = null;
|
|
71
|
+
let readHiBuffer = null;
|
|
72
|
+
let readLoBuffer = null;
|
|
73
|
+
|
|
74
|
+
try {
|
|
75
|
+
if (xIsGpu) {
|
|
76
|
+
xHiBuffer = x._buf;
|
|
77
|
+
xLoBuffer = x._loBuf;
|
|
78
|
+
} else {
|
|
79
|
+
const { hi, lo } = splitDoubleDouble(x);
|
|
80
|
+
xHiBuffer = uploadBuffer(device, hi, "dnrm2-xHi", false);
|
|
81
|
+
xLoBuffer = uploadBuffer(device, lo, "dnrm2-xLo", false);
|
|
82
|
+
}
|
|
83
|
+
// 2*WGS partial (scale, ssq) DD pairs — see dnrm2.wgsl for what they represent.
|
|
84
|
+
partialsScaleHiBuffer = createStorageBuffer(
|
|
85
|
+
device,
|
|
86
|
+
2 * WGS * 4,
|
|
87
|
+
"dnrm2-partials-scaleHi",
|
|
88
|
+
);
|
|
89
|
+
partialsScaleLoBuffer = createStorageBuffer(
|
|
90
|
+
device,
|
|
91
|
+
2 * WGS * 4,
|
|
92
|
+
"dnrm2-partials-scaleLo",
|
|
93
|
+
);
|
|
94
|
+
partialsSsqHiBuffer = createStorageBuffer(
|
|
95
|
+
device,
|
|
96
|
+
2 * WGS * 4,
|
|
97
|
+
"dnrm2-partials-ssqHi",
|
|
98
|
+
);
|
|
99
|
+
partialsSsqLoBuffer = createStorageBuffer(
|
|
100
|
+
device,
|
|
101
|
+
2 * WGS * 4,
|
|
102
|
+
"dnrm2-partials-ssqLo",
|
|
103
|
+
);
|
|
104
|
+
resultHiBuffer = createResultBuffer(device, 4, "dnrm2-result-hi");
|
|
105
|
+
resultLoBuffer = createResultBuffer(device, 4, "dnrm2-result-lo");
|
|
106
|
+
paramsBuffer = createParamsBuffer(
|
|
107
|
+
device,
|
|
108
|
+
[
|
|
109
|
+
{ value: n, type: "u32" },
|
|
110
|
+
{ value: incx, type: "u32" },
|
|
111
|
+
],
|
|
112
|
+
"dnrm2-params",
|
|
113
|
+
);
|
|
114
|
+
|
|
115
|
+
const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
|
|
116
|
+
xHiBuffer,
|
|
117
|
+
xLoBuffer,
|
|
118
|
+
partialsScaleHiBuffer,
|
|
119
|
+
partialsScaleLoBuffer,
|
|
120
|
+
partialsSsqHiBuffer,
|
|
121
|
+
partialsSsqLoBuffer,
|
|
122
|
+
paramsBuffer,
|
|
123
|
+
]);
|
|
124
|
+
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
125
|
+
device,
|
|
126
|
+
pipelineMain,
|
|
127
|
+
bgMain,
|
|
128
|
+
2 * WGS,
|
|
129
|
+
); // dispatch 2*WGS workgroups
|
|
130
|
+
|
|
131
|
+
submit(device, enc1);
|
|
132
|
+
|
|
133
|
+
const bgReduce = createBindGroup(
|
|
134
|
+
device,
|
|
135
|
+
pipelineReduce.getBindGroupLayout(0),
|
|
136
|
+
[
|
|
137
|
+
partialsScaleHiBuffer,
|
|
138
|
+
partialsScaleLoBuffer,
|
|
139
|
+
partialsSsqHiBuffer,
|
|
140
|
+
partialsSsqLoBuffer,
|
|
141
|
+
resultHiBuffer,
|
|
142
|
+
resultLoBuffer,
|
|
143
|
+
],
|
|
144
|
+
);
|
|
145
|
+
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
146
|
+
device,
|
|
147
|
+
pipelineReduce,
|
|
148
|
+
bgReduce,
|
|
149
|
+
1,
|
|
150
|
+
); // reduce partials to a single result
|
|
151
|
+
readHiBuffer = stageReadback(device, enc2, resultHiBuffer);
|
|
152
|
+
readLoBuffer = stageReadback(device, enc2, resultLoBuffer);
|
|
153
|
+
|
|
154
|
+
submit(device, enc2);
|
|
155
|
+
|
|
156
|
+
const hiPromise = extractResult(readHiBuffer, Float32Array);
|
|
157
|
+
const loPromise = extractResult(readLoBuffer, Float32Array);
|
|
158
|
+
readHiBuffer = null; // ownership transferred — extractResult's own finally destroys it
|
|
159
|
+
readLoBuffer = null;
|
|
160
|
+
|
|
161
|
+
const [gpuTime1, gpuTime2, hiArr, loArr] = await Promise.all([
|
|
162
|
+
extractTimestamp(ts1),
|
|
163
|
+
extractTimestamp(ts2),
|
|
164
|
+
hiPromise,
|
|
165
|
+
loPromise,
|
|
166
|
+
]);
|
|
167
|
+
|
|
168
|
+
// reduction/scaledSumF64.wgsl already computes scale·sqrt(ssq) on the
|
|
169
|
+
// GPU — no separate sqrt step here.
|
|
170
|
+
const nrm2 = mergeDoubleDouble(hiArr, loArr)[0];
|
|
171
|
+
|
|
172
|
+
if (gpuTime1 !== undefined && gpuTime2 !== undefined)
|
|
173
|
+
return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
|
|
174
|
+
return { nrm2 };
|
|
175
|
+
} finally {
|
|
176
|
+
if (!xIsGpu && xHiBuffer) destroyBuffers(xHiBuffer);
|
|
177
|
+
if (!xIsGpu && xLoBuffer) destroyBuffers(xLoBuffer);
|
|
178
|
+
if (partialsScaleHiBuffer) destroyBuffers(partialsScaleHiBuffer);
|
|
179
|
+
if (partialsScaleLoBuffer) destroyBuffers(partialsScaleLoBuffer);
|
|
180
|
+
if (partialsSsqHiBuffer) destroyBuffers(partialsSsqHiBuffer);
|
|
181
|
+
if (partialsSsqLoBuffer) destroyBuffers(partialsSsqLoBuffer);
|
|
182
|
+
if (resultHiBuffer) destroyBuffers(resultHiBuffer);
|
|
183
|
+
if (resultLoBuffer) destroyBuffers(resultLoBuffer);
|
|
184
|
+
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
185
|
+
// Only reached if submit(device, enc2) threw before ownership was transferred above.
|
|
186
|
+
if (readHiBuffer) destroyBuffers(readHiBuffer);
|
|
187
|
+
if (readLoBuffer) destroyBuffers(readLoBuffer);
|
|
188
|
+
}
|
|
189
|
+
}
|
|
@@ -0,0 +1,67 @@
|
|
|
1
|
+
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
|
+
|
|
3
|
+
/**
|
|
4
|
+
* Applies a Givens plane rotation to double-precision vectors x and y:
|
|
5
|
+
* $$\begin{aligned} x &\leftarrow cx + sy \\\\ y &\leftarrow -sx + cy \end{aligned}$$
|
|
6
|
+
* — double-double (Dekker) f64 emulation of {@link srot}, since WGSL has no
|
|
7
|
+
* native f64 type.
|
|
8
|
+
*
|
|
9
|
+
* {@includeCode ../../examples/drot/drot.js}
|
|
10
|
+
*
|
|
11
|
+
* **Browser (standalone HTML):**
|
|
12
|
+
* {@includeCode ../../examples/drot/web/drot.html}
|
|
13
|
+
*
|
|
14
|
+
* @param device - GPUDevice from `init()`
|
|
15
|
+
* @param n - number of elements (must be a positive integer)
|
|
16
|
+
* @param x - Float64Array input/output vector
|
|
17
|
+
* @param incx - stride for x (must be a positive integer)
|
|
18
|
+
* @param y - Float64Array input/output vector
|
|
19
|
+
* @param incy - stride for y (must be a positive integer)
|
|
20
|
+
* @param c - cosine of rotation angle
|
|
21
|
+
* @param s - sine of rotation angle
|
|
22
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/drot/drot.mjs#L19">Source code: drot.mjs (L19)</a>
|
|
23
|
+
* @category BLAS Level 1
|
|
24
|
+
*/
|
|
25
|
+
export declare function drot(
|
|
26
|
+
device: GPUDevice,
|
|
27
|
+
n: number,
|
|
28
|
+
x: Float64Array,
|
|
29
|
+
incx: number,
|
|
30
|
+
y: Float64Array,
|
|
31
|
+
incy: number,
|
|
32
|
+
c: number,
|
|
33
|
+
s: number,
|
|
34
|
+
): Promise<
|
|
35
|
+
| { x: Float64Array; y: Float64Array }
|
|
36
|
+
| { x: Float64Array; y: Float64Array; gpuTimeMs: number }
|
|
37
|
+
>;
|
|
38
|
+
|
|
39
|
+
/**
|
|
40
|
+
* Applies a Givens plane rotation to double-precision vectors x and y:
|
|
41
|
+
* $$\begin{aligned} x &\leftarrow cx + sy \\\\ y &\leftarrow -sx + cy \end{aligned}$$
|
|
42
|
+
* — GPU-resident overload; see the Float64Array overload above for the
|
|
43
|
+
* routine itself.
|
|
44
|
+
*
|
|
45
|
+
* {@includeCode ../../examples/drot/gpu.drot.js}
|
|
46
|
+
*
|
|
47
|
+
* @param device - GPUDevice from `init()`
|
|
48
|
+
* @param n - number of elements (must be a positive integer)
|
|
49
|
+
* @param x - GpuVector input/output vector (must be Float64Array-backed, mutated in place)
|
|
50
|
+
* @param incx - stride for x (must be a positive integer)
|
|
51
|
+
* @param y - GpuVector input/output vector (must be Float64Array-backed, mutated in place)
|
|
52
|
+
* @param incy - stride for y (must be a positive integer)
|
|
53
|
+
* @param c - cosine of rotation angle
|
|
54
|
+
* @param s - sine of rotation angle
|
|
55
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/drot/drot.mjs#L19">Source code: drot.mjs (L19)</a>
|
|
56
|
+
* @category BLAS Level 1
|
|
57
|
+
*/
|
|
58
|
+
export declare function drot(
|
|
59
|
+
device: GPUDevice,
|
|
60
|
+
n: number,
|
|
61
|
+
x: GpuVector,
|
|
62
|
+
incx: number,
|
|
63
|
+
y: GpuVector,
|
|
64
|
+
incy: number,
|
|
65
|
+
c: number,
|
|
66
|
+
s: number,
|
|
67
|
+
): Promise<{} | { gpuTimeMs: number }>;
|
|
@@ -0,0 +1,170 @@
|
|
|
1
|
+
import {
|
|
2
|
+
uploadBuffer,
|
|
3
|
+
createParamsBuffer,
|
|
4
|
+
stageReadback,
|
|
5
|
+
destroyBuffers,
|
|
6
|
+
} from "../util/buffer.mjs";
|
|
7
|
+
import { createBindGroup } from "../util/bindgroup.mjs";
|
|
8
|
+
import { runComputePass, submit } from "../util/compute.mjs";
|
|
9
|
+
import { extractResult } from "../util/result.mjs";
|
|
10
|
+
import { extractTimestamp } from "../util/benchmark.mjs";
|
|
11
|
+
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
|
+
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
|
+
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
|
|
15
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
16
|
+
|
|
17
|
+
// drot: x := c*x + s*y, y := -s*x + c*y, double-double (Dekker) f64
|
|
18
|
+
// emulation of srot — x, y, c, and s are each split into an f32 (hi, lo)
|
|
19
|
+
// pair; WGSL has no f64 type.
|
|
20
|
+
export async function drot(device, n, x, incx, y, incy, c, s) {
|
|
21
|
+
const xIsGpu = x instanceof GpuVector;
|
|
22
|
+
const yIsGpu = y instanceof GpuVector;
|
|
23
|
+
|
|
24
|
+
requireGpuDevice(device);
|
|
25
|
+
requireSameDevice(device, "drot", { x, y });
|
|
26
|
+
if (
|
|
27
|
+
!Number.isInteger(n) ||
|
|
28
|
+
!Number.isInteger(incx) ||
|
|
29
|
+
!Number.isInteger(incy)
|
|
30
|
+
)
|
|
31
|
+
throw new Error("n, incx, and incy must be integers.");
|
|
32
|
+
if (typeof c !== "number") throw new Error("c must be a number.");
|
|
33
|
+
if (typeof s !== "number") throw new Error("s must be a number.");
|
|
34
|
+
if (Number.isNaN(c) || Number.isNaN(s))
|
|
35
|
+
throw new Error("c and s must not be NaN.");
|
|
36
|
+
if (!Number.isFinite(c)) throw new Error("c must be finite.");
|
|
37
|
+
if (!Number.isFinite(s)) throw new Error("s must be finite.");
|
|
38
|
+
if (incx <= 0 || incy <= 0)
|
|
39
|
+
throw new Error("incx and incy must be positive.");
|
|
40
|
+
if (!(x instanceof Float64Array) && !xIsGpu)
|
|
41
|
+
throw new Error("x must be a Float64Array or GpuVector.");
|
|
42
|
+
if (!(y instanceof Float64Array) && !yIsGpu)
|
|
43
|
+
throw new Error("y must be a Float64Array or GpuVector.");
|
|
44
|
+
if (xIsGpu && x.dtype !== Float64Array)
|
|
45
|
+
throw new Error("x must be a Float64Array-backed GpuVector.");
|
|
46
|
+
if (yIsGpu && y.dtype !== Float64Array)
|
|
47
|
+
throw new Error("y must be a Float64Array-backed GpuVector.");
|
|
48
|
+
if (xIsGpu !== yIsGpu)
|
|
49
|
+
throw new Error(
|
|
50
|
+
"x and y must be the same type (both Float64Array or both GpuVector).",
|
|
51
|
+
);
|
|
52
|
+
if (n <= 0) return xIsGpu ? {} : { x, y };
|
|
53
|
+
if (x.length < (n - 1) * incx + 1)
|
|
54
|
+
throw new Error(
|
|
55
|
+
"x does not have enough elements for the given n and incx.",
|
|
56
|
+
);
|
|
57
|
+
if (y.length < (n - 1) * incy + 1)
|
|
58
|
+
throw new Error(
|
|
59
|
+
"y does not have enough elements for the given n and incy.",
|
|
60
|
+
);
|
|
61
|
+
|
|
62
|
+
// Concatenated with f64/dekker.wgsl (DD struct), f64/utils/add.wgsl
|
|
63
|
+
// (fsub/negf/fastTwoSumProtected/ddAddProtected), and f64/utils/multiply.wgsl
|
|
64
|
+
// (ddMulProtected) — WGSL has no #include.
|
|
65
|
+
const f64Deps = ["f64/dekker", "f64/utils/add", "f64/utils/multiply"];
|
|
66
|
+
const pipeline = await getPipeline(device, [...f64Deps, "drot"]);
|
|
67
|
+
|
|
68
|
+
const { hi: cHi, lo: cLo } = splitDoubleDouble(new Float64Array([c]));
|
|
69
|
+
const { hi: sHi, lo: sLo } = splitDoubleDouble(new Float64Array([s]));
|
|
70
|
+
|
|
71
|
+
let xHiBuffer = null;
|
|
72
|
+
let xLoBuffer = null;
|
|
73
|
+
let yHiBuffer = null;
|
|
74
|
+
let yLoBuffer = null;
|
|
75
|
+
let paramsBuffer = null;
|
|
76
|
+
let xReadHiBuffer = null;
|
|
77
|
+
let xReadLoBuffer = null;
|
|
78
|
+
let yReadHiBuffer = null;
|
|
79
|
+
let yReadLoBuffer = null;
|
|
80
|
+
|
|
81
|
+
try {
|
|
82
|
+
if (xIsGpu) {
|
|
83
|
+
xHiBuffer = x._buf;
|
|
84
|
+
xLoBuffer = x._loBuf;
|
|
85
|
+
yHiBuffer = y._buf;
|
|
86
|
+
yLoBuffer = y._loBuf;
|
|
87
|
+
} else {
|
|
88
|
+
const xSplit = splitDoubleDouble(x);
|
|
89
|
+
const ySplit = splitDoubleDouble(y);
|
|
90
|
+
xHiBuffer = uploadBuffer(device, xSplit.hi, "drot-xHi", true);
|
|
91
|
+
xLoBuffer = uploadBuffer(device, xSplit.lo, "drot-xLo", true);
|
|
92
|
+
yHiBuffer = uploadBuffer(device, ySplit.hi, "drot-yHi", true);
|
|
93
|
+
yLoBuffer = uploadBuffer(device, ySplit.lo, "drot-yLo", true);
|
|
94
|
+
}
|
|
95
|
+
paramsBuffer = createParamsBuffer(
|
|
96
|
+
device,
|
|
97
|
+
[
|
|
98
|
+
{ value: n, type: "u32" },
|
|
99
|
+
{ value: cHi[0], type: "f32" },
|
|
100
|
+
{ value: cLo[0], type: "f32" },
|
|
101
|
+
{ value: sHi[0], type: "f32" },
|
|
102
|
+
{ value: sLo[0], type: "f32" },
|
|
103
|
+
{ value: incx, type: "u32" },
|
|
104
|
+
{ value: incy, type: "u32" },
|
|
105
|
+
],
|
|
106
|
+
"drot-params",
|
|
107
|
+
);
|
|
108
|
+
|
|
109
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
110
|
+
xHiBuffer,
|
|
111
|
+
xLoBuffer,
|
|
112
|
+
yHiBuffer,
|
|
113
|
+
yLoBuffer,
|
|
114
|
+
paramsBuffer,
|
|
115
|
+
]);
|
|
116
|
+
const { commandEncoder, ts } = runComputePass(
|
|
117
|
+
device,
|
|
118
|
+
pipeline,
|
|
119
|
+
bindGroup,
|
|
120
|
+
calcWorkgroups(device, n),
|
|
121
|
+
);
|
|
122
|
+
xReadHiBuffer = xIsGpu
|
|
123
|
+
? null
|
|
124
|
+
: stageReadback(device, commandEncoder, xHiBuffer);
|
|
125
|
+
xReadLoBuffer = xIsGpu
|
|
126
|
+
? null
|
|
127
|
+
: stageReadback(device, commandEncoder, xLoBuffer);
|
|
128
|
+
yReadHiBuffer = yIsGpu
|
|
129
|
+
? null
|
|
130
|
+
: stageReadback(device, commandEncoder, yHiBuffer);
|
|
131
|
+
yReadLoBuffer = yIsGpu
|
|
132
|
+
? null
|
|
133
|
+
: stageReadback(device, commandEncoder, yLoBuffer);
|
|
134
|
+
|
|
135
|
+
submit(device, commandEncoder);
|
|
136
|
+
|
|
137
|
+
const gpuTimeMs = await extractTimestamp(ts);
|
|
138
|
+
|
|
139
|
+
if (xIsGpu) {
|
|
140
|
+
// xIsGpu === yIsGpu, enforced above
|
|
141
|
+
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
142
|
+
return {};
|
|
143
|
+
}
|
|
144
|
+
|
|
145
|
+
const xHi = await extractResult(xReadHiBuffer, Float32Array);
|
|
146
|
+
xReadHiBuffer = null; // extractResult already destroyed it
|
|
147
|
+
const xLo = await extractResult(xReadLoBuffer, Float32Array);
|
|
148
|
+
xReadLoBuffer = null;
|
|
149
|
+
const yHi = await extractResult(yReadHiBuffer, Float32Array);
|
|
150
|
+
yReadHiBuffer = null;
|
|
151
|
+
const yLo = await extractResult(yReadLoBuffer, Float32Array);
|
|
152
|
+
yReadLoBuffer = null;
|
|
153
|
+
const resultX = mergeDoubleDouble(xHi, xLo);
|
|
154
|
+
const resultY = mergeDoubleDouble(yHi, yLo);
|
|
155
|
+
if (gpuTimeMs !== undefined) return { x: resultX, y: resultY, gpuTimeMs };
|
|
156
|
+
return { x: resultX, y: resultY };
|
|
157
|
+
} finally {
|
|
158
|
+
if (!xIsGpu && xHiBuffer) destroyBuffers(xHiBuffer);
|
|
159
|
+
if (!xIsGpu && xLoBuffer) destroyBuffers(xLoBuffer);
|
|
160
|
+
if (!yIsGpu && yHiBuffer) destroyBuffers(yHiBuffer);
|
|
161
|
+
if (!yIsGpu && yLoBuffer) destroyBuffers(yLoBuffer);
|
|
162
|
+
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
163
|
+
// Only reached if extractTimestamp or extractResult threw before
|
|
164
|
+
// clearing these — on the success path they're already null.
|
|
165
|
+
if (xReadHiBuffer) destroyBuffers(xReadHiBuffer);
|
|
166
|
+
if (xReadLoBuffer) destroyBuffers(xReadLoBuffer);
|
|
167
|
+
if (yReadHiBuffer) destroyBuffers(yReadHiBuffer);
|
|
168
|
+
if (yReadLoBuffer) destroyBuffers(yReadLoBuffer);
|
|
169
|
+
}
|
|
170
|
+
}
|