wgblas 1.2.1 → 2.1.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/LICENSE +1 -1
- package/README.md +54 -72
- package/dist/wgblas.browser.js +1637 -858
- package/index.d.mts +51 -6
- package/index.mjs +8 -0
- package/package.json +56 -2
- package/src/classes/GpuMatrix.mjs +17 -10
- package/src/classes/GpuVector.mjs +28 -10
- package/src/dasum/dasum.d.mts +6 -6
- package/src/dasum/dasum.mjs +25 -21
- package/src/devdocs.mjs +13 -0
- package/src/idamax/idamax.d.mts +69 -0
- package/src/idamax/idamax.mjs +130 -0
- package/src/init.mjs +115 -49
- package/src/isamax/isamax.d.mts +21 -3
- package/src/isamax/isamax.mjs +16 -14
- package/src/random/random.d.mts +1 -0
- package/src/sasum/sasum.d.mts +3 -3
- package/src/sasum/sasum.mjs +15 -13
- package/src/saxpy/saxpy.d.mts +3 -3
- package/src/saxpy/saxpy.mjs +10 -8
- package/src/scopy/scopy.d.mts +3 -3
- package/src/scopy/scopy.mjs +10 -8
- package/src/sdot/sdot.d.mts +3 -3
- package/src/sdot/sdot.mjs +16 -14
- package/src/sgemm/sgemm.d.mts +102 -0
- package/src/sgemm/sgemm.mjs +208 -0
- package/src/sgemmtr/sgemmtr.d.mts +104 -0
- package/src/sgemmtr/sgemmtr.mjs +204 -0
- package/src/sgemv/sgemv.d.mts +1 -38
- package/src/sgemv/sgemv.mjs +42 -26
- package/src/sger/sger.d.mts +1 -34
- package/src/sger/sger.mjs +12 -8
- package/src/shaders/block_transfer.wgsl +42 -0
- package/src/shaders/dasum.wgsl +3 -2
- package/src/shaders/f64/dekker.wgsl +4 -85
- package/src/shaders/f64/utils/abs.wgsl +10 -0
- package/src/shaders/f64/utils/add.wgsl +77 -0
- package/src/shaders/f64/utils/equal.wgsl +7 -0
- package/src/shaders/f64/utils/greater.wgsl +12 -0
- package/src/shaders/f64/utils/multiply.wgsl +81 -0
- package/src/shaders/idamax.wgsl +96 -0
- package/src/shaders/index.mjs +164 -14
- package/src/shaders/reduction/argmaxF64.wgsl +50 -0
- package/src/shaders/reduction/scaledSum.wgsl +65 -0
- package/src/shaders/reduction/sumF64.wgsl +2 -2
- package/src/shaders/sgemm_large.wgsl +206 -0
- package/src/shaders/sgemm_small.wgsl +212 -0
- package/src/shaders/sgemmtr_large.wgsl +120 -0
- package/src/shaders/sgemmtr_small.wgsl +113 -0
- 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/shaders/symmetrize.wgsl +31 -0
- package/src/shaders/triangularize.wgsl +44 -0
- package/src/snrm2/snrm2.d.mts +3 -3
- package/src/snrm2/snrm2.mjs +33 -21
- package/src/srot/srot.d.mts +3 -5
- package/src/srot/srot.mjs +11 -9
- package/src/srotm/srotm.d.mts +3 -5
- package/src/srotm/srotm.mjs +19 -10
- package/src/sscal/sscal.d.mts +4 -4
- package/src/sscal/sscal.mjs +12 -10
- package/src/sswap/sswap.d.mts +3 -3
- package/src/sswap/sswap.mjs +11 -9
- package/src/ssymm/ssymm.d.mts +103 -0
- package/src/ssymm/ssymm.mjs +218 -0
- package/src/ssymv/ssymv.d.mts +1 -36
- package/src/ssymv/ssymv.mjs +12 -8
- package/src/ssyr/ssyr.d.mts +1 -30
- package/src/ssyr/ssyr.mjs +11 -7
- package/src/ssyr2/ssyr2.d.mts +1 -34
- package/src/ssyr2/ssyr2.mjs +12 -8
- package/src/ssyr2k/ssyr2k.d.mts +100 -0
- package/src/ssyr2k/ssyr2k.mjs +202 -0
- package/src/ssyrk/ssyrk.d.mts +90 -0
- package/src/ssyrk/ssyrk.mjs +177 -0
- package/src/strmm/strmm.d.mts +100 -0
- package/src/strmm/strmm.mjs +226 -0
- package/src/strmv/strmv.d.mts +1 -36
- package/src/strmv/strmv.mjs +12 -8
- package/src/strsm/strsm.d.mts +99 -0
- package/src/strsm/strsm.mjs +360 -0
- package/src/strsv/strsv.d.mts +1 -32
- package/src/strsv/strsv.mjs +18 -12
- package/src/util/benchmark.mjs +4 -6
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +116 -20
- package/src/util/compute.mjs +12 -12
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +34 -0
- package/src/util/f64.mjs +3 -3
- package/src/util/pipeline.mjs +5 -6
- package/src/util/workgroup.mjs +55 -7
- package/src/shaders/browser-shaders.mjs +0 -55
package/src/shaders/snrm2.wgsl
CHANGED
|
@@ -1,9 +1,22 @@
|
|
|
1
|
-
// snrm2: result = sqrt(sum(x[i] * x[i]))
|
|
2
|
-
//
|
|
1
|
+
// snrm2: result = sqrt(sum(x[i] * x[i])), computed via scaled accumulation
|
|
2
|
+
// (Blue's algorithm / reference BLAS's SLASSQ) rather than naive squaring —
|
|
3
|
+
// naive `sum += x_i * x_i` overflows to inf for |x_i| ≳ 1.8e19 (f32's
|
|
4
|
+
// squaring range is only sqrt(f32_max)) and loses precision on tiny
|
|
5
|
+
// magnitudes squaring into the denormal range. Running state is (scale,
|
|
6
|
+
// ssq) with true-sum-of-squares == scale² · ssq: scale tracks the largest
|
|
7
|
+
// |x_i| seen so far, and every other contribution is expressed *relative
|
|
8
|
+
// to* scale (never squared in absolute terms), so ssq stays near 1
|
|
9
|
+
// regardless of x's magnitude range. Merging two independent partials
|
|
10
|
+
// (ssqMerge) is associative, so this composes with the same 4-way-ILP +
|
|
11
|
+
// tree-reduction shape every other Level 1 reduction here uses — see
|
|
12
|
+
// reduction/scaledSum.wgsl for the pass-2 counterpart, which finishes with
|
|
13
|
+
// scale·sqrt(ssq).
|
|
14
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/scaledSum.wgsl.
|
|
3
15
|
|
|
4
|
-
@group(0) @binding(0) var<storage, read> x:
|
|
5
|
-
@group(0) @binding(1) var<storage, read_write>
|
|
6
|
-
@group(0) @binding(2) var<
|
|
16
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
17
|
+
@group(0) @binding(1) var<storage, read_write> partialsScale: array<f32>;
|
|
18
|
+
@group(0) @binding(2) var<storage, read_write> partialsSsq: array<f32>;
|
|
19
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
7
20
|
|
|
8
21
|
struct Params {
|
|
9
22
|
n: u32,
|
|
@@ -12,7 +25,36 @@ struct Params {
|
|
|
12
25
|
|
|
13
26
|
const WGS: u32 = 64;
|
|
14
27
|
|
|
15
|
-
|
|
28
|
+
struct ScaleSsq {
|
|
29
|
+
scale: f32,
|
|
30
|
+
ssq: f32,
|
|
31
|
+
}
|
|
32
|
+
|
|
33
|
+
// Folds one more |value| into a running (scale, ssq) pair.
|
|
34
|
+
fn ssqAccum(acc: ScaleSsq, absxi: f32) -> ScaleSsq {
|
|
35
|
+
if (absxi == 0.0) { return acc; }
|
|
36
|
+
if (absxi > acc.scale) {
|
|
37
|
+
let r = acc.scale / absxi; // 0/absxi == 0 on the first nonzero value — safe
|
|
38
|
+
return ScaleSsq(absxi, 1.0 + acc.ssq * r * r);
|
|
39
|
+
}
|
|
40
|
+
let r = absxi / acc.scale; // reached only once acc.scale > 0 (absxi <= acc.scale and absxi > 0)
|
|
41
|
+
return ScaleSsq(acc.scale, acc.ssq + r * r);
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
// Associative merge of two independent (scale, ssq) partials — lets this
|
|
45
|
+
// compose with a tree reduction exactly like a plain sum would.
|
|
46
|
+
fn ssqMerge(a: ScaleSsq, b: ScaleSsq) -> ScaleSsq {
|
|
47
|
+
if (a.scale == 0.0 && b.scale == 0.0) { return ScaleSsq(0.0, 1.0); }
|
|
48
|
+
if (a.scale >= b.scale) {
|
|
49
|
+
let r = b.scale / a.scale;
|
|
50
|
+
return ScaleSsq(a.scale, a.ssq + b.ssq * r * r);
|
|
51
|
+
}
|
|
52
|
+
let r = a.scale / b.scale;
|
|
53
|
+
return ScaleSsq(b.scale, b.ssq + a.ssq * r * r);
|
|
54
|
+
}
|
|
55
|
+
|
|
56
|
+
var<workgroup> tileScale: array<f32, 64>;
|
|
57
|
+
var<workgroup> tileSsq: array<f32, 64>;
|
|
16
58
|
|
|
17
59
|
@compute @workgroup_size(64)
|
|
18
60
|
fn main(
|
|
@@ -21,36 +63,43 @@ fn main(
|
|
|
21
63
|
@builtin(workgroup_id) wgid: vec3u,
|
|
22
64
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
23
65
|
) {
|
|
24
|
-
var acc0
|
|
25
|
-
var acc1
|
|
26
|
-
var acc2
|
|
27
|
-
var acc3
|
|
66
|
+
var acc0 = ScaleSsq(0.0, 1.0);
|
|
67
|
+
var acc1 = ScaleSsq(0.0, 1.0);
|
|
68
|
+
var acc2 = ScaleSsq(0.0, 1.0);
|
|
69
|
+
var acc3 = ScaleSsq(0.0, 1.0);
|
|
28
70
|
|
|
29
71
|
let stride = num_wg.x * WGS;
|
|
30
72
|
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
31
73
|
|
|
32
74
|
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
acc0 += v0 * v0;
|
|
38
|
-
acc1 += v1 * v1;
|
|
39
|
-
acc2 += v2 * v2;
|
|
40
|
-
acc3 += v3 * v3;
|
|
75
|
+
acc0 = ssqAccum(acc0, abs(x[ id * params.x_inc]));
|
|
76
|
+
acc1 = ssqAccum(acc1, abs(x[(id + stride) * params.x_inc]));
|
|
77
|
+
acc2 = ssqAccum(acc2, abs(x[(id + 2u * stride) * params.x_inc]));
|
|
78
|
+
acc3 = ssqAccum(acc3, abs(x[(id + 3u * stride) * params.x_inc]));
|
|
41
79
|
}
|
|
42
80
|
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
43
|
-
|
|
44
|
-
acc0 += v * v;
|
|
81
|
+
acc0 = ssqAccum(acc0, abs(x[id * params.x_inc]));
|
|
45
82
|
}
|
|
46
83
|
|
|
47
|
-
|
|
84
|
+
let combined = ssqMerge(ssqMerge(acc0, acc1), ssqMerge(acc2, acc3));
|
|
85
|
+
tileScale[lid.x] = combined.scale;
|
|
86
|
+
tileSsq[lid.x] = combined.ssq;
|
|
48
87
|
workgroupBarrier();
|
|
49
88
|
|
|
50
89
|
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
51
|
-
if (lid.x < s) {
|
|
90
|
+
if (lid.x < s) {
|
|
91
|
+
let merged = ssqMerge(
|
|
92
|
+
ScaleSsq(tileScale[lid.x], tileSsq[lid.x]),
|
|
93
|
+
ScaleSsq(tileScale[lid.x + s], tileSsq[lid.x + s]),
|
|
94
|
+
);
|
|
95
|
+
tileScale[lid.x] = merged.scale;
|
|
96
|
+
tileSsq[lid.x] = merged.ssq;
|
|
97
|
+
}
|
|
52
98
|
workgroupBarrier();
|
|
53
99
|
}
|
|
54
100
|
|
|
55
|
-
if (lid.x == 0u) {
|
|
101
|
+
if (lid.x == 0u) {
|
|
102
|
+
partialsScale[wgid.x] = tileScale[0];
|
|
103
|
+
partialsSsq[wgid.x] = tileSsq[0];
|
|
104
|
+
}
|
|
56
105
|
}
|
package/src/shaders/ssymv.wgsl
CHANGED
|
@@ -63,7 +63,9 @@ fn main(
|
|
|
63
63
|
}
|
|
64
64
|
|
|
65
65
|
if lid.x == 0u {
|
|
66
|
-
|
|
66
|
+
// BLAS beta==0 semantics: y is written, not accumulated — must not read y.
|
|
67
|
+
let acc = params.alpha * scratch[0];
|
|
68
|
+
y[i * params.incy] = select(acc, acc + params.beta * y[i * params.incy], params.beta != 0.0);
|
|
67
69
|
}
|
|
68
70
|
}
|
|
69
71
|
}
|
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
// symmetrize: Adense := full dense expansion of a symmetric matrix stored
|
|
2
|
+
// with only its `uplo` triangle meaningful (the other triangle is implied
|
|
3
|
+
// by symmetry: A[i,j] = A[j,i]). A plain element-wise pass, no tiling or
|
|
4
|
+
// shared memory needed — used to materialize a dense operand for routines
|
|
5
|
+
// that read a symmetric matrix as a normal dense gemm input (e.g. ssymm),
|
|
6
|
+
// rather than teaching the tiled gemm kernel itself to mirror-read.
|
|
7
|
+
|
|
8
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
9
|
+
@group(0) @binding(1) var<storage, read_write> Adense: array<f32>;
|
|
10
|
+
|
|
11
|
+
struct Params {
|
|
12
|
+
n: u32,
|
|
13
|
+
lda: u32,
|
|
14
|
+
ldd: u32, // leading dimension of Adense
|
|
15
|
+
uplo: u32, // 0 = lower (stored where col <= row), 1 = upper (col >= row)
|
|
16
|
+
}
|
|
17
|
+
|
|
18
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
19
|
+
|
|
20
|
+
@compute @workgroup_size(8, 8)
|
|
21
|
+
fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
22
|
+
let row = gid.y;
|
|
23
|
+
let col = gid.x;
|
|
24
|
+
if (row >= params.n || col >= params.n) {
|
|
25
|
+
return;
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
let isStored = select(col >= row, col <= row, params.uplo == 0u);
|
|
29
|
+
let srcIdx = select(col * params.lda + row, row * params.lda + col, isStored);
|
|
30
|
+
Adense[row * params.ldd + col] = A[srcIdx];
|
|
31
|
+
}
|
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
// triangularize: Adense := dense expansion of op(A) (A or A^T per `trans`),
|
|
2
|
+
// zero-filling the unstored triangle (exact for a matmul) so strmm can reuse
|
|
3
|
+
// sgemm's kernel unchanged. `diag=1` substitutes 1.0 on the diagonal.
|
|
4
|
+
|
|
5
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
6
|
+
@group(0) @binding(1) var<storage, read_write> Adense: array<f32>;
|
|
7
|
+
|
|
8
|
+
struct Params {
|
|
9
|
+
n: u32,
|
|
10
|
+
lda: u32,
|
|
11
|
+
ldd: u32, // leading dimension of Adense
|
|
12
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
13
|
+
trans: u32, // 0 = no-transpose (op(A) = A), 1 = transpose (op(A) = A^T)
|
|
14
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
15
|
+
}
|
|
16
|
+
|
|
17
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
18
|
+
|
|
19
|
+
@compute @workgroup_size(8, 8)
|
|
20
|
+
fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
21
|
+
let row = gid.y;
|
|
22
|
+
let col = gid.x;
|
|
23
|
+
if (row >= params.n || col >= params.n) {
|
|
24
|
+
return;
|
|
25
|
+
}
|
|
26
|
+
|
|
27
|
+
if (row == col) {
|
|
28
|
+
Adense[row * params.ldd + col] = select(A[row * params.lda + row], 1.0, params.diag == 1u);
|
|
29
|
+
return;
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
var isMeaningful: bool;
|
|
33
|
+
var srcRow: u32;
|
|
34
|
+
var srcCol: u32;
|
|
35
|
+
if (params.trans == 0u) {
|
|
36
|
+
isMeaningful = select(col >= row, col <= row, params.uplo == 0u);
|
|
37
|
+
srcRow = row; srcCol = col;
|
|
38
|
+
} else {
|
|
39
|
+
isMeaningful = select(col <= row, col >= row, params.uplo == 0u);
|
|
40
|
+
srcRow = col; srcCol = row;
|
|
41
|
+
}
|
|
42
|
+
|
|
43
|
+
Adense[row * params.ldd + col] = select(0.0, A[srcRow * params.lda + srcCol], isMeaningful);
|
|
44
|
+
}
|
package/src/snrm2/snrm2.d.mts
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
|
-
* Computes the Euclidean norm of a vector: result = sqrt
|
|
4
|
+
* Computes the Euclidean norm of a vector: $$\text{result} = \sqrt{\sum_{i} x_i^2}$$
|
|
5
5
|
*
|
|
6
6
|
* {@includeCode ../../examples/snrm2/snrm2.js}
|
|
7
7
|
*
|
|
@@ -24,9 +24,9 @@ export declare function snrm2(
|
|
|
24
24
|
): Promise<{ nrm2: number } | { nrm2: number; gpuTimeMs: number }>;
|
|
25
25
|
|
|
26
26
|
/**
|
|
27
|
-
* Computes the Euclidean norm of a vector: result = sqrt
|
|
27
|
+
* Computes the Euclidean norm of a vector: $$\text{result} = \sqrt{\sum_{i} x_i^2}$$
|
|
28
28
|
*
|
|
29
|
-
* {@includeCode ../../examples/snrm2/
|
|
29
|
+
* {@includeCode ../../examples/snrm2/gpu.snrm2.js}
|
|
30
30
|
*
|
|
31
31
|
* @param device - GPUDevice from `init()`
|
|
32
32
|
* @param n - number of elements (must be a positive integer)
|
package/src/snrm2/snrm2.mjs
CHANGED
|
@@ -12,14 +12,16 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
12
12
|
import { extractResult } from "../util/result.mjs";
|
|
13
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
14
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
15
|
+
import { WGS } from "../util/constants.mjs";
|
|
16
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
15
17
|
|
|
16
|
-
const WGS = 64; // workgroup size
|
|
17
18
|
|
|
18
19
|
export async function snrm2(device, n, x, incx) {
|
|
19
20
|
const xIsGpu = x instanceof GpuVector;
|
|
20
21
|
|
|
21
22
|
if (!(device instanceof GPUDevice))
|
|
22
23
|
throw new Error("device must be a GPUDevice.");
|
|
24
|
+
requireSameDevice(device, "snrm2", { x });
|
|
23
25
|
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
24
26
|
throw new Error("n and incx must be integers.");
|
|
25
27
|
if (incx <= 0) throw new Error("incx must be positive.");
|
|
@@ -32,19 +34,25 @@ export async function snrm2(device, n, x, incx) {
|
|
|
32
34
|
);
|
|
33
35
|
|
|
34
36
|
const pipelineMain = await getPipeline(device, "snrm2");
|
|
35
|
-
const pipelineReduce = await getPipeline(device, "reduction/
|
|
37
|
+
const pipelineReduce = await getPipeline(device, "reduction/scaledSum");
|
|
36
38
|
|
|
37
39
|
let xBuffer = null;
|
|
38
|
-
let
|
|
40
|
+
let partialsScaleBuffer = null;
|
|
41
|
+
let partialsSsqBuffer = null;
|
|
39
42
|
let resultBuffer = null;
|
|
40
43
|
let paramsBuffer = null;
|
|
41
44
|
let readBuffer = null;
|
|
42
45
|
|
|
43
46
|
try {
|
|
44
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "snrm2-x", false);
|
|
45
|
-
|
|
46
|
-
|
|
47
|
-
|
|
47
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "snrm2-x", false);
|
|
48
|
+
// 2*WGS partial (scale, ssq) pairs — see snrm2.wgsl for what they represent.
|
|
49
|
+
partialsScaleBuffer = createStorageBuffer(device,
|
|
50
|
+
2 * WGS * 4,
|
|
51
|
+
"snrm2-partials-scale",
|
|
52
|
+
);
|
|
53
|
+
partialsSsqBuffer = createStorageBuffer(device, 2 * WGS * 4, "snrm2-partials-ssq");
|
|
54
|
+
resultBuffer = createResultBuffer(device, 4, "snrm2-result"); // final f32 scalar
|
|
55
|
+
paramsBuffer = createParamsBuffer(device,
|
|
48
56
|
[
|
|
49
57
|
{ value: n, type: "u32" },
|
|
50
58
|
{ value: incx, type: "u32" },
|
|
@@ -52,53 +60,57 @@ export async function snrm2(device, n, x, incx) {
|
|
|
52
60
|
"snrm2-params",
|
|
53
61
|
);
|
|
54
62
|
|
|
55
|
-
const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
|
|
63
|
+
const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
|
|
56
64
|
xBuffer,
|
|
57
|
-
|
|
65
|
+
partialsScaleBuffer,
|
|
66
|
+
partialsSsqBuffer,
|
|
58
67
|
paramsBuffer,
|
|
59
68
|
]);
|
|
60
|
-
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
69
|
+
const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
|
|
61
70
|
pipelineMain,
|
|
62
71
|
bgMain,
|
|
63
72
|
2 * WGS,
|
|
64
73
|
); // dispatch 2*WGS workgroups
|
|
65
74
|
|
|
66
|
-
submit(enc1);
|
|
75
|
+
submit(device, enc1);
|
|
67
76
|
|
|
68
|
-
const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
|
|
69
|
-
|
|
77
|
+
const bgReduce = createBindGroup(device, pipelineReduce.getBindGroupLayout(0), [
|
|
78
|
+
partialsScaleBuffer,
|
|
79
|
+
partialsSsqBuffer,
|
|
70
80
|
resultBuffer,
|
|
71
81
|
]);
|
|
72
|
-
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
82
|
+
const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
|
|
73
83
|
pipelineReduce,
|
|
74
84
|
bgReduce,
|
|
75
85
|
1,
|
|
76
86
|
); // reduce partials to single result
|
|
77
|
-
readBuffer = stageReadback(enc2, resultBuffer);
|
|
87
|
+
readBuffer = stageReadback(device, enc2, resultBuffer);
|
|
78
88
|
|
|
79
|
-
submit(enc2);
|
|
89
|
+
submit(device, enc2);
|
|
80
90
|
|
|
81
91
|
const resultPromise = extractResult(readBuffer, Float32Array);
|
|
82
92
|
readBuffer = null; // ownership transferred — extractResult's own finally destroys it
|
|
83
93
|
|
|
84
|
-
const [gpuTime1, gpuTime2,
|
|
94
|
+
const [gpuTime1, gpuTime2, resultArr] = await Promise.all([
|
|
85
95
|
extractTimestamp(ts1),
|
|
86
96
|
extractTimestamp(ts2),
|
|
87
97
|
resultPromise,
|
|
88
98
|
]);
|
|
89
99
|
|
|
90
|
-
//
|
|
91
|
-
|
|
100
|
+
// reduction/scaledSum.wgsl already computes scale·sqrt(ssq) on the GPU —
|
|
101
|
+
// unlike the old naive-sum version, there's no separate sqrt step here.
|
|
102
|
+
const nrm2 = resultArr[0];
|
|
92
103
|
|
|
93
104
|
if (gpuTime1 !== undefined && gpuTime2 !== undefined)
|
|
94
105
|
return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
|
|
95
106
|
return { nrm2 };
|
|
96
107
|
} finally {
|
|
97
108
|
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
98
|
-
if (
|
|
109
|
+
if (partialsScaleBuffer) destroyBuffers(partialsScaleBuffer);
|
|
110
|
+
if (partialsSsqBuffer) destroyBuffers(partialsSsqBuffer);
|
|
99
111
|
if (resultBuffer) destroyBuffers(resultBuffer);
|
|
100
112
|
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
101
|
-
// Only reached if submit(enc2) threw before ownership was transferred above.
|
|
113
|
+
// Only reached if submit(device, enc2) threw before ownership was transferred above.
|
|
102
114
|
if (readBuffer) destroyBuffers(readBuffer);
|
|
103
115
|
}
|
|
104
116
|
}
|
package/src/srot/srot.d.mts
CHANGED
|
@@ -2,8 +2,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
|
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
4
|
* Applies a Givens plane rotation to vectors x and y:
|
|
5
|
-
*
|
|
6
|
-
* y = -s*x + c*y
|
|
5
|
+
* $$\begin{aligned} x &\leftarrow cx + sy \\\\ y &\leftarrow -sx + cy \end{aligned}$$
|
|
7
6
|
*
|
|
8
7
|
* {@includeCode ../../examples/srot/srot.js}
|
|
9
8
|
*
|
|
@@ -34,10 +33,9 @@ export declare function srot(
|
|
|
34
33
|
|
|
35
34
|
/**
|
|
36
35
|
* Applies a Givens plane rotation to vectors x and y:
|
|
37
|
-
*
|
|
38
|
-
* y = -s*x + c*y
|
|
36
|
+
* $$\begin{aligned} x &\leftarrow cx + sy \\\\ y &\leftarrow -sx + cy \end{aligned}$$
|
|
39
37
|
*
|
|
40
|
-
* {@includeCode ../../examples/srot/
|
|
38
|
+
* {@includeCode ../../examples/srot/gpu.srot.js}
|
|
41
39
|
*
|
|
42
40
|
* @param device - GPUDevice from `init()`
|
|
43
41
|
* @param n - number of elements (must be a positive integer)
|
package/src/srot/srot.mjs
CHANGED
|
@@ -11,6 +11,7 @@ import { extractResult } from "../util/result.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
@@ -18,6 +19,7 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
18
19
|
|
|
19
20
|
if (!(device instanceof GPUDevice))
|
|
20
21
|
throw new Error("device must be a GPUDevice.");
|
|
22
|
+
requireSameDevice(device, "srot", { x, y });
|
|
21
23
|
if (
|
|
22
24
|
!Number.isInteger(n) ||
|
|
23
25
|
!Number.isInteger(incx) ||
|
|
@@ -58,9 +60,9 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
58
60
|
let readY = null;
|
|
59
61
|
|
|
60
62
|
try {
|
|
61
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srot-x", true);
|
|
62
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srot-y", true);
|
|
63
|
-
paramsBuffer = createParamsBuffer(
|
|
63
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "srot-x", true);
|
|
64
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "srot-y", true);
|
|
65
|
+
paramsBuffer = createParamsBuffer(device,
|
|
64
66
|
[
|
|
65
67
|
{ value: n, type: "u32" },
|
|
66
68
|
{ value: c, type: "f32" },
|
|
@@ -71,19 +73,19 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
71
73
|
"srot-params",
|
|
72
74
|
);
|
|
73
75
|
|
|
74
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
76
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
75
77
|
xBuffer,
|
|
76
78
|
yBuffer,
|
|
77
79
|
paramsBuffer,
|
|
78
80
|
]);
|
|
79
|
-
const { commandEncoder, ts } = runComputePass(
|
|
81
|
+
const { commandEncoder, ts } = runComputePass(device,
|
|
80
82
|
pipeline,
|
|
81
83
|
bindGroup,
|
|
82
|
-
calcWorkgroups(n),
|
|
84
|
+
calcWorkgroups(device, n),
|
|
83
85
|
);
|
|
84
|
-
readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
85
|
-
readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
86
|
-
submit(commandEncoder);
|
|
86
|
+
readX = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
|
|
87
|
+
readY = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
88
|
+
submit(device, commandEncoder);
|
|
87
89
|
|
|
88
90
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
89
91
|
|
package/src/srotm/srotm.d.mts
CHANGED
|
@@ -2,8 +2,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
|
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
4
|
* Applies a modified Givens plane rotation H to vectors x and y:
|
|
5
|
-
*
|
|
6
|
-
* y = H[1][0]*x + H[1][1]*y
|
|
5
|
+
* $$\begin{pmatrix} x \\\\ y \end{pmatrix} \leftarrow \begin{pmatrix} h_{11} & h_{12} \\\\ h_{21} & h_{22} \end{pmatrix} \begin{pmatrix} x \\\\ y \end{pmatrix}$$
|
|
7
6
|
*
|
|
8
7
|
* {@includeCode ../../examples/srotm/srotm.js}
|
|
9
8
|
*
|
|
@@ -33,10 +32,9 @@ export declare function srotm(
|
|
|
33
32
|
|
|
34
33
|
/**
|
|
35
34
|
* Applies a modified Givens plane rotation H to vectors x and y:
|
|
36
|
-
*
|
|
37
|
-
* y = H[1][0]*x + H[1][1]*y
|
|
35
|
+
* $$\begin{pmatrix} x \\\\ y \end{pmatrix} \leftarrow \begin{pmatrix} h_{11} & h_{12} \\\\ h_{21} & h_{22} \end{pmatrix} \begin{pmatrix} x \\\\ y \end{pmatrix}$$
|
|
38
36
|
*
|
|
39
|
-
* {@includeCode ../../examples/srotm/
|
|
37
|
+
* {@includeCode ../../examples/srotm/gpu.srotm.js}
|
|
40
38
|
*
|
|
41
39
|
* @param device - GPUDevice from `init()`
|
|
42
40
|
* @param n - number of elements (must be a positive integer)
|
package/src/srotm/srotm.mjs
CHANGED
|
@@ -11,6 +11,7 @@ import { extractResult } from "../util/result.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function srotm(device, n, x, incx, y, incy, param) {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
@@ -18,6 +19,7 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
18
19
|
|
|
19
20
|
if (!(device instanceof GPUDevice))
|
|
20
21
|
throw new Error("device must be a GPUDevice.");
|
|
22
|
+
requireSameDevice(device, "srotm", { x, y });
|
|
21
23
|
if (
|
|
22
24
|
!Number.isInteger(n) ||
|
|
23
25
|
!Number.isInteger(incx) ||
|
|
@@ -26,6 +28,13 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
26
28
|
throw new Error("n, incx, and incy must be integers.");
|
|
27
29
|
if (!(param instanceof Float32Array) || param.length !== 5)
|
|
28
30
|
throw new Error("param must be a Float32Array of length 5.");
|
|
31
|
+
if (
|
|
32
|
+
param[0] !== -2 &&
|
|
33
|
+
param[0] !== -1 &&
|
|
34
|
+
param[0] !== 0 &&
|
|
35
|
+
param[0] !== 1
|
|
36
|
+
)
|
|
37
|
+
throw new Error("param[0] (flag) must be one of -2, -1, 0, or 1.");
|
|
29
38
|
if (incx <= 0 || incy <= 0)
|
|
30
39
|
throw new Error("incx and incy must be positive.");
|
|
31
40
|
if (!xIsGpu && !(x instanceof Float32Array))
|
|
@@ -56,10 +65,10 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
56
65
|
let readY = null;
|
|
57
66
|
|
|
58
67
|
try {
|
|
59
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srotm-x", true);
|
|
60
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srotm-y", true);
|
|
61
|
-
paramBuffer = uploadBuffer(param, "srotm-param", false);
|
|
62
|
-
paramsBuffer = createParamsBuffer(
|
|
68
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "srotm-x", true);
|
|
69
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "srotm-y", true);
|
|
70
|
+
paramBuffer = uploadBuffer(device, param, "srotm-param", false);
|
|
71
|
+
paramsBuffer = createParamsBuffer(device,
|
|
63
72
|
[
|
|
64
73
|
{ value: n, type: "u32" },
|
|
65
74
|
{ value: incx, type: "u32" },
|
|
@@ -68,20 +77,20 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
68
77
|
"srotm-params",
|
|
69
78
|
);
|
|
70
79
|
|
|
71
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
80
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
72
81
|
xBuffer,
|
|
73
82
|
yBuffer,
|
|
74
83
|
paramBuffer,
|
|
75
84
|
paramsBuffer,
|
|
76
85
|
]);
|
|
77
|
-
const { commandEncoder, ts } = runComputePass(
|
|
86
|
+
const { commandEncoder, ts } = runComputePass(device,
|
|
78
87
|
pipeline,
|
|
79
88
|
bindGroup,
|
|
80
|
-
calcWorkgroups(n),
|
|
89
|
+
calcWorkgroups(device, n),
|
|
81
90
|
);
|
|
82
|
-
readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
83
|
-
readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
84
|
-
submit(commandEncoder);
|
|
91
|
+
readX = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
|
|
92
|
+
readY = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
93
|
+
submit(device, commandEncoder);
|
|
85
94
|
|
|
86
95
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
87
96
|
|
package/src/sscal/sscal.d.mts
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
|
-
* Scales a single-precision vector by a constant: x
|
|
4
|
+
* Scales a single-precision vector by a constant: $$x \leftarrow \alpha x$$
|
|
5
5
|
*
|
|
6
6
|
* {@includeCode ../../examples/sscal/sscal.js}
|
|
7
7
|
*
|
|
@@ -22,12 +22,12 @@ 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
|
-
* {@includeCode ../../examples/sscal/
|
|
30
|
+
* {@includeCode ../../examples/sscal/gpu.sscal.js}
|
|
31
31
|
*
|
|
32
32
|
* @param device - GPUDevice from `init()`
|
|
33
33
|
* @param n - number of elements to scale (must be a positive integer)
|
package/src/sscal/sscal.mjs
CHANGED
|
@@ -11,12 +11,14 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function sscal(device, n, alpha, x, incx) {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
17
|
-
|
|
18
|
+
|
|
18
19
|
if (!(device instanceof GPUDevice))
|
|
19
20
|
throw new Error("device must be a GPUDevice.");
|
|
21
|
+
requireSameDevice(device, "sscal", { x });
|
|
20
22
|
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
21
23
|
throw new Error("n and incx must be integers.");
|
|
22
24
|
if (typeof alpha !== "number")
|
|
@@ -26,7 +28,7 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
26
28
|
if (incx <= 0) throw new Error("incx must be positive.");
|
|
27
29
|
if (!(x instanceof Float32Array) && !(x instanceof GpuVector))
|
|
28
30
|
throw new Error("x must be a Float32Array or GpuVector.");
|
|
29
|
-
if (n <= 0) return xIsGpu ? {} : x;
|
|
31
|
+
if (n <= 0) return xIsGpu ? {} : { x };
|
|
30
32
|
if (x.length < (n - 1) * incx + 1)
|
|
31
33
|
throw new Error(
|
|
32
34
|
"x does not have enough elements for the given n and incx.",
|
|
@@ -39,8 +41,8 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
39
41
|
let readBuffer = null;
|
|
40
42
|
|
|
41
43
|
try {
|
|
42
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sscal-x", true);
|
|
43
|
-
paramsBuffer = createParamsBuffer(
|
|
44
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sscal-x", true);
|
|
45
|
+
paramsBuffer = createParamsBuffer(device,
|
|
44
46
|
[
|
|
45
47
|
{ value: n, type: "u32" },
|
|
46
48
|
{ value: alpha, type: "f32" },
|
|
@@ -49,18 +51,18 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
49
51
|
"sscal-params",
|
|
50
52
|
);
|
|
51
53
|
|
|
52
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
54
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
53
55
|
xBuffer,
|
|
54
56
|
paramsBuffer,
|
|
55
57
|
]);
|
|
56
|
-
const { commandEncoder, ts } = runComputePass(
|
|
58
|
+
const { commandEncoder, ts } = runComputePass(device,
|
|
57
59
|
pipeline,
|
|
58
60
|
bindGroup,
|
|
59
|
-
calcWorkgroups(n),
|
|
61
|
+
calcWorkgroups(device, n),
|
|
60
62
|
);
|
|
61
|
-
readBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
63
|
+
readBuffer = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
|
|
62
64
|
|
|
63
|
-
submit(commandEncoder);
|
|
65
|
+
submit(device, commandEncoder);
|
|
64
66
|
|
|
65
67
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
66
68
|
|
|
@@ -72,7 +74,7 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
72
74
|
const result = await extractResult(readBuffer, Float32Array);
|
|
73
75
|
readBuffer = null; // extractResult already destroyed it
|
|
74
76
|
if (gpuTimeMs !== undefined) return { x: result, gpuTimeMs };
|
|
75
|
-
return result;
|
|
77
|
+
return { x: result };
|
|
76
78
|
} finally {
|
|
77
79
|
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
78
80
|
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
package/src/sswap/sswap.d.mts
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
|
-
* Swaps the elements of two single-precision vectors: x
|
|
4
|
+
* Swaps the elements of two single-precision vectors: $$x \leftrightarrow y$$
|
|
5
5
|
*
|
|
6
6
|
* {@includeCode ../../examples/sswap/sswap.js}
|
|
7
7
|
*
|
|
@@ -27,9 +27,9 @@ 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
|
-
* {@includeCode ../../examples/sswap/
|
|
32
|
+
* {@includeCode ../../examples/sswap/gpu.sswap.js}
|
|
33
33
|
*
|
|
34
34
|
* @param device - GPUDevice from `init()`
|
|
35
35
|
* @param n - number of elements to swap (must be a positive integer)
|