wgblas 2.1.0 → 2.2.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +2 -0
- package/dist/wgblas.browser.js +1007 -43
- package/index.d.mts +26 -53
- package/index.mjs +11 -0
- package/package.json +132 -63
- package/src/classes/Complex32.d.mts +43 -0
- package/src/classes/Complex32.mjs +82 -0
- package/src/classes/Complex64.d.mts +44 -0
- package/src/classes/Complex64.mjs +76 -0
- package/src/classes/GpuMatrix.d.mts +41 -41
- package/src/classes/GpuMatrix.mjs +112 -10
- package/src/classes/GpuVector.d.mts +36 -40
- package/src/classes/GpuVector.mjs +39 -2
- package/src/cscal/cscal.d.mts +47 -0
- package/src/cscal/cscal.mjs +98 -0
- package/src/dasum/dasum.mjs +31 -15
- package/src/daxpy/daxpy.d.mts +56 -0
- package/src/daxpy/daxpy.mjs +150 -0
- package/src/dcopy/dcopy.d.mts +52 -0
- package/src/dcopy/dcopy.mjs +140 -0
- package/src/ddot/ddot.d.mts +62 -0
- package/src/ddot/ddot.mjs +184 -0
- package/src/dnrm2/dnrm2.d.mts +50 -0
- package/src/dnrm2/dnrm2.mjs +189 -0
- package/src/drot/drot.d.mts +67 -0
- package/src/drot/drot.mjs +170 -0
- package/src/drotm/drotm.d.mts +67 -0
- package/src/drotm/drotm.mjs +171 -0
- package/src/dscal/dscal.d.mts +52 -0
- package/src/dscal/dscal.mjs +119 -0
- package/src/dswap/dswap.d.mts +57 -0
- package/src/dswap/dswap.mjs +155 -0
- package/src/idamax/idamax.mjs +49 -19
- package/src/init.mjs +6 -3
- package/src/isamax/isamax.mjs +17 -14
- package/src/random/random.d.mts +37 -40
- package/src/random/random.mjs +39 -7
- package/src/sasum/sasum.mjs +13 -11
- package/src/saxpy/saxpy.mjs +9 -8
- package/src/scopy/scopy.mjs +8 -6
- package/src/sdot/sdot.mjs +13 -11
- package/src/sgemm/sgemm.d.mts +2 -2
- package/src/sgemm/sgemm.mjs +91 -35
- package/src/sgemmtr/sgemmtr.d.mts +3 -2
- package/src/sgemmtr/sgemmtr.mjs +92 -35
- package/src/sgemv/sgemv.d.mts +2 -2
- package/src/sgemv/sgemv.mjs +41 -25
- package/src/sger/sger.d.mts +2 -2
- package/src/sger/sger.mjs +38 -16
- package/src/shaders/__test_pipeline_a.wgsl +3 -0
- package/src/shaders/__test_pipeline_b.wgsl +2 -0
- package/src/shaders/cscal.wgsl +33 -0
- package/src/shaders/daxpy.wgsl +66 -0
- package/src/shaders/dcopy.wgsl +34 -0
- package/src/shaders/ddot.wgsl +106 -0
- package/src/shaders/dnrm2.wgsl +167 -0
- package/src/shaders/drot.wgsl +81 -0
- package/src/shaders/drotm.wgsl +99 -0
- package/src/shaders/dscal.wgsl +60 -0
- package/src/shaders/dswap.wgsl +38 -0
- package/src/shaders/f64/utils/add.wgsl +6 -0
- package/src/shaders/f64/utils/divide.wgsl +45 -0
- package/src/shaders/f64/utils/multiply.wgsl +19 -10
- package/src/shaders/f64/utils/sqrt.wgsl +45 -0
- package/src/shaders/index.mjs +69 -0
- package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
- package/src/snrm2/snrm2.mjs +20 -14
- package/src/srot/srot.mjs +10 -7
- package/src/srotm/srotm.mjs +9 -12
- package/src/sscal/sscal.mjs +7 -7
- package/src/sswap/sswap.mjs +14 -8
- package/src/ssymm/ssymm.d.mts +5 -4
- package/src/ssymm/ssymm.mjs +142 -55
- package/src/ssymv/ssymv.d.mts +2 -2
- package/src/ssymv/ssymv.mjs +42 -23
- package/src/ssyr/ssyr.d.mts +2 -2
- package/src/ssyr/ssyr.mjs +34 -15
- package/src/ssyr2/ssyr2.d.mts +2 -2
- package/src/ssyr2/ssyr2.mjs +43 -18
- package/src/ssyr2k/ssyr2k.d.mts +3 -2
- package/src/ssyr2k/ssyr2k.mjs +132 -55
- package/src/ssyrk/ssyrk.d.mts +3 -2
- package/src/ssyrk/ssyrk.mjs +84 -33
- package/src/strmm/strmm.d.mts +5 -4
- package/src/strmm/strmm.mjs +153 -54
- package/src/strmv/strmv.d.mts +2 -2
- package/src/strmv/strmv.mjs +37 -17
- package/src/strsm/strsm.d.mts +6 -4
- package/src/strsm/strsm.mjs +418 -172
- package/src/strsv/strsv.d.mts +5 -3
- package/src/strsv/strsv.mjs +82 -31
- package/src/util/benchmark.mjs +5 -3
- package/src/util/buffer.mjs +33 -12
- package/src/util/complex.mjs +87 -0
- package/src/util/compute.mjs +14 -8
- package/src/util/device.mjs +18 -3
- package/src/util/pipeline.mjs +40 -5
- package/src/util/workgroup.mjs +23 -6
- package/src/shaders/f64add.wgsl +0 -281
|
@@ -0,0 +1,99 @@
|
|
|
1
|
+
// drotm: applies a modified Givens rotation H to vectors x and y — double-
|
|
2
|
+
// double (Dekker) f64 emulation of srotm. paramHi/paramLo[0] = flag: -1
|
|
3
|
+
// (full H), 0 (unit diagonal), 1 (unit off-diagonal). param = [ flag, h11,
|
|
4
|
+
// h21, h12, h22 ], each entry an f32 (hi, lo) pair.
|
|
5
|
+
// flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
|
|
6
|
+
//
|
|
7
|
+
// h11/h12/h21/h22 are resolved once, outside the loop, from the (uniform
|
|
8
|
+
// across every thread) flag — same shape as srot's c/s, so no barrier is
|
|
9
|
+
// needed for that selection itself. Each element then costs four
|
|
10
|
+
// ddMulProtected + two ddAddProtected, same as drot.
|
|
11
|
+
|
|
12
|
+
@group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
|
|
13
|
+
@group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
|
|
14
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
15
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
16
|
+
@group(0) @binding(4) var<storage, read> paramHi: array<f32>;
|
|
17
|
+
@group(0) @binding(5) var<storage, read> paramLo: array<f32>;
|
|
18
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
19
|
+
|
|
20
|
+
struct Params {
|
|
21
|
+
n: u32,
|
|
22
|
+
x_inc: u32,
|
|
23
|
+
y_inc: u32,
|
|
24
|
+
}
|
|
25
|
+
|
|
26
|
+
const WGS: u32 = 64;
|
|
27
|
+
const ONE: DD = DD(1.0, 0.0);
|
|
28
|
+
const NEG_ONE: DD = DD(-1.0, 0.0);
|
|
29
|
+
|
|
30
|
+
@compute @workgroup_size(64)
|
|
31
|
+
fn drotm_main(
|
|
32
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
33
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
34
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
35
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
36
|
+
) {
|
|
37
|
+
let flag = paramHi[0]; // exact small integer (-1, 0, or 1) — lo is always 0
|
|
38
|
+
|
|
39
|
+
var h11: DD; var h12: DD;
|
|
40
|
+
var h21: DD; var h22: DD;
|
|
41
|
+
|
|
42
|
+
if (flag == -1.0) {
|
|
43
|
+
// full 2x2 matrix
|
|
44
|
+
h11 = DD(paramHi[1], paramLo[1]); h21 = DD(paramHi[2], paramLo[2]);
|
|
45
|
+
h12 = DD(paramHi[3], paramLo[3]); h22 = DD(paramHi[4], paramLo[4]);
|
|
46
|
+
} else if (flag == 0.0) {
|
|
47
|
+
// diagonal fixed at 1
|
|
48
|
+
h11 = ONE; h21 = DD(paramHi[2], paramLo[2]);
|
|
49
|
+
h12 = DD(paramHi[3], paramLo[3]); h22 = ONE;
|
|
50
|
+
} else {
|
|
51
|
+
// flag == 1.0: off-diagonal fixed at +1 / -1
|
|
52
|
+
h11 = DD(paramHi[1], paramLo[1]); h21 = NEG_ONE;
|
|
53
|
+
h12 = ONE; h22 = DD(paramHi[4], paramLo[4]);
|
|
54
|
+
}
|
|
55
|
+
|
|
56
|
+
let stride = num_wg.x * WGS;
|
|
57
|
+
|
|
58
|
+
let n_floor = (params.n / stride) * stride;
|
|
59
|
+
let mainIters = n_floor / stride;
|
|
60
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
61
|
+
let id = gid.x + iter * stride;
|
|
62
|
+
let ix = id * params.x_inc;
|
|
63
|
+
let iy = id * params.y_inc;
|
|
64
|
+
let xi = DD(xHi[ix], xLo[ix]);
|
|
65
|
+
let yi = DD(yHi[iy], yLo[iy]);
|
|
66
|
+
let xNew = ddAddProtected(ddMulProtected(h11, xi, lid.x), ddMulProtected(h12, yi, lid.x), lid.x);
|
|
67
|
+
let yNew = ddAddProtected(ddMulProtected(h21, xi, lid.x), ddMulProtected(h22, yi, lid.x), lid.x);
|
|
68
|
+
xHi[ix] = xNew.hi;
|
|
69
|
+
xLo[ix] = xNew.lo;
|
|
70
|
+
yHi[iy] = yNew.hi;
|
|
71
|
+
yLo[iy] = yNew.lo;
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
// Tail is ragged (0 or 1 extra per thread) — pad to this workgroup's worst
|
|
75
|
+
// case so every thread in the workgroup still calls ddMulProtected/
|
|
76
|
+
// ddAddProtected the same number of times (their barriers need that),
|
|
77
|
+
// masking only the write.
|
|
78
|
+
let wgBaseGid = wgid.x * WGS;
|
|
79
|
+
var tailIters = 0u;
|
|
80
|
+
if (n_floor + wgBaseGid < params.n) {
|
|
81
|
+
tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
|
|
82
|
+
}
|
|
83
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
84
|
+
let id = n_floor + gid.x + iter * stride;
|
|
85
|
+
let valid = id < params.n;
|
|
86
|
+
let ix = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
87
|
+
let iy = select(0u, id * params.y_inc, valid);
|
|
88
|
+
let xi = DD(xHi[ix], xLo[ix]);
|
|
89
|
+
let yi = DD(yHi[iy], yLo[iy]);
|
|
90
|
+
let xNew = ddAddProtected(ddMulProtected(h11, xi, lid.x), ddMulProtected(h12, yi, lid.x), lid.x);
|
|
91
|
+
let yNew = ddAddProtected(ddMulProtected(h21, xi, lid.x), ddMulProtected(h22, yi, lid.x), lid.x);
|
|
92
|
+
if (valid) {
|
|
93
|
+
xHi[ix] = xNew.hi;
|
|
94
|
+
xLo[ix] = xNew.lo;
|
|
95
|
+
yHi[iy] = yNew.hi;
|
|
96
|
+
yLo[iy] = yNew.lo;
|
|
97
|
+
}
|
|
98
|
+
}
|
|
99
|
+
}
|
|
@@ -0,0 +1,60 @@
|
|
|
1
|
+
// dscal: x := alpha * x, double-double (Dekker) f64 emulation of sscal.
|
|
2
|
+
// alpha and x are each an f32 (hi, lo) pair. See f64/utils/multiply.wgsl for
|
|
3
|
+
// ddMulProtected and why plain ddMulRaw isn't safe without a renormalizing
|
|
4
|
+
// barrier — that barrier needs a provably uniform loop trip count across
|
|
5
|
+
// every thread in the workgroup, so (like dasum.wgsl's reduction loop) this
|
|
6
|
+
// splits into a uniform main pass plus a ragged, select-masked tail rather
|
|
7
|
+
// than a plain `id < params.n` grid-stride loop.
|
|
8
|
+
|
|
9
|
+
@group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
|
|
10
|
+
@group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
|
|
11
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
12
|
+
|
|
13
|
+
struct Params {
|
|
14
|
+
n: u32,
|
|
15
|
+
alphaHi: f32,
|
|
16
|
+
alphaLo: f32,
|
|
17
|
+
x_inc: u32,
|
|
18
|
+
}
|
|
19
|
+
|
|
20
|
+
const WGS: u32 = 64;
|
|
21
|
+
|
|
22
|
+
@compute @workgroup_size(64)
|
|
23
|
+
fn dscal_main(
|
|
24
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
25
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
26
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
27
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
28
|
+
) {
|
|
29
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
30
|
+
let stride = num_wg.x * WGS;
|
|
31
|
+
|
|
32
|
+
let n_floor = (params.n / stride) * stride;
|
|
33
|
+
let mainIters = n_floor / stride;
|
|
34
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
35
|
+
let id = gid.x + iter * stride;
|
|
36
|
+
let i = id * params.x_inc;
|
|
37
|
+
let result = ddMulProtected(alpha, DD(xHi[i], xLo[i]), lid.x);
|
|
38
|
+
xHi[i] = result.hi;
|
|
39
|
+
xLo[i] = result.lo;
|
|
40
|
+
}
|
|
41
|
+
|
|
42
|
+
// Tail is ragged (0 or 1 extra per thread) — pad to this workgroup's worst
|
|
43
|
+
// case so every thread in the workgroup still calls ddMulProtected the
|
|
44
|
+
// same number of times (its barrier needs that), masking only the write.
|
|
45
|
+
let wgBaseGid = wgid.x * WGS;
|
|
46
|
+
var tailIters = 0u;
|
|
47
|
+
if (n_floor + wgBaseGid < params.n) {
|
|
48
|
+
tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
|
|
49
|
+
}
|
|
50
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
51
|
+
let id = n_floor + gid.x + iter * stride;
|
|
52
|
+
let valid = id < params.n;
|
|
53
|
+
let i = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
54
|
+
let result = ddMulProtected(alpha, DD(xHi[i], xLo[i]), lid.x);
|
|
55
|
+
if (valid) {
|
|
56
|
+
xHi[i] = result.hi;
|
|
57
|
+
xLo[i] = result.lo;
|
|
58
|
+
}
|
|
59
|
+
}
|
|
60
|
+
}
|
|
@@ -0,0 +1,38 @@
|
|
|
1
|
+
// dswap: x <-> y, double-double (Dekker) f64 emulation of sswap. A swap is
|
|
2
|
+
// pure data movement — hi and lo are exchanged verbatim, with no arithmetic
|
|
3
|
+
// 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 sswap.wgsl itself.
|
|
7
|
+
|
|
8
|
+
@group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
|
|
9
|
+
@group(0) @binding(1) var<storage, read_write> 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
|
+
let tempHi = xHi[ix];
|
|
32
|
+
let tempLo = xLo[ix];
|
|
33
|
+
xHi[ix] = yHi[iy];
|
|
34
|
+
xLo[ix] = yLo[iy];
|
|
35
|
+
yHi[iy] = tempHi;
|
|
36
|
+
yLo[iy] = tempLo;
|
|
37
|
+
}
|
|
38
|
+
}
|
|
@@ -75,3 +75,9 @@ fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
|
75
75
|
let loSum = a.lo + b.lo;
|
|
76
76
|
return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
|
|
77
77
|
}
|
|
78
|
+
|
|
79
|
+
// Double-double subtraction — a - b, via exact negation (a sign-bit flip,
|
|
80
|
+
// no rounding) then ddAddProtected. Same protection contract.
|
|
81
|
+
fn ddSubProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
82
|
+
return ddAddProtected(a, DD(negf(b.hi), negf(b.lo)), threadSlot);
|
|
83
|
+
}
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
// Requires f64/dekker.wgsl concatenated first for the DD struct, and
|
|
2
|
+
// f64/utils/add.wgsl (ddSubProtected/ddAddProtected/negf) and
|
|
3
|
+
// f64/utils/multiply.wgsl (ddMulProtected).
|
|
4
|
+
//
|
|
5
|
+
// Verified empirically on real hardware (NVIDIA GTX 1650 + Intel Mesa
|
|
6
|
+
// integrated GPU, ~8000 random trials each plus explicit edge cases): no
|
|
7
|
+
// compiler-reassociation-style corruption of the kind that broke twoSum/
|
|
8
|
+
// twoProd (see add.wgsl's/multiply.wgsl's own headers) — every failure
|
|
9
|
+
// found was a genuine algorithm gap, not a driver miscompile, and both are
|
|
10
|
+
// fixed below (the b.hi==0.0 guard). One real, expected-shape difference
|
|
11
|
+
// from every other protected op here: the low-power backend's observed
|
|
12
|
+
// forward-error factor for this op specifically runs noticeably higher
|
|
13
|
+
// (~7-14000x eps, vs ~3-9x on high-performance) than ddSqrtProtected's
|
|
14
|
+
// (~3-4x on both) — division inherently amplifies input imprecision more
|
|
15
|
+
// than a sum/product does, so a real routine built on this needs its own
|
|
16
|
+
// backend-calibrated threshold, same as every other f64 arithmetic routine
|
|
17
|
+
// in this codebase (see e.g. tests/drot/src/test.drot.js's THRESHOLDS).
|
|
18
|
+
//
|
|
19
|
+
// One Newton-style long-division refinement (Bailey/QD-style): q1 = a.hi /
|
|
20
|
+
// b.hi is a plain f32 quotient, accurate to ~24 bits. Computing the residual
|
|
21
|
+
// a - q1*b in DD arithmetic (not f32) recovers the bits q1 lost, and a
|
|
22
|
+
// second plain division of that residual resolves them into a correction
|
|
23
|
+
// term — combining q1 + q2 gives roughly double a lone f32 divide's
|
|
24
|
+
// precision, matching this scheme's ~48-bit double-double target (already
|
|
25
|
+
// short of real f64's 52 bits, so a second refinement step would chase
|
|
26
|
+
// precision this representation has no room for).
|
|
27
|
+
// b.hi == 0.0 makes q1 = a.hi/0.0 already the IEEE-754-correct answer
|
|
28
|
+
// (±Infinity, or NaN for 0/0) via plain float division, but the refinement
|
|
29
|
+
// below would corrupt it: p1 = q1*b multiplies that Infinity by a zero
|
|
30
|
+
// divisor, and Infinity*0 is NaN by definition, poisoning everything after.
|
|
31
|
+
// Substituting a safe non-zero denominator via select() — rather than
|
|
32
|
+
// branching/returning early — keeps every thread calling ddMulProtected/
|
|
33
|
+
// ddSubProtected/ddAddProtected unconditionally, which their internal
|
|
34
|
+
// workgroupBarrier() requires; only the final result is selected between
|
|
35
|
+
// the refined value and q1's own already-correct answer.
|
|
36
|
+
fn ddDivProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
37
|
+
let bIsZero = b.hi == 0.0;
|
|
38
|
+
let q1 = a.hi / b.hi;
|
|
39
|
+
let safeB = DD(select(b.hi, 1.0, bIsZero), select(b.lo, 0.0, bIsZero));
|
|
40
|
+
let p1 = ddMulProtected(DD(q1, 0.0), safeB, threadSlot);
|
|
41
|
+
let r1 = ddSubProtected(a, p1, threadSlot);
|
|
42
|
+
let q2 = r1.hi / safeB.hi;
|
|
43
|
+
let refined = ddAddProtected(DD(q1, 0.0), DD(q2, 0.0), threadSlot);
|
|
44
|
+
return DD(select(refined.hi, q1, bIsZero), select(refined.lo, 0.0, bIsZero));
|
|
45
|
+
}
|
|
@@ -15,7 +15,11 @@ const SPLIT_CONST: f32 = 4097.0;
|
|
|
15
15
|
|
|
16
16
|
fn bitSplit(a: f32) -> DD {
|
|
17
17
|
let bits = bitcast<u32>(a);
|
|
18
|
-
|
|
18
|
+
// Top 11 mantissa bits, so hi carries 12 significant bits with the implicit
|
|
19
|
+
// leading 1 — the halves are multiplied pairwise and f32 holds 24, so a
|
|
20
|
+
// wider split rounds those products and the "exact" error term goes wrong.
|
|
21
|
+
// Matches SPLIT_CONST = 2^12+1 used by the Veltkamp path below.
|
|
22
|
+
let hiBits = bits & 0xFFFFF000u;
|
|
19
23
|
let hi = bitcast<f32>(hiBits);
|
|
20
24
|
let lo = fsub(a, hi); // exact by Sterbenz's lemma (hi, a share an exponent, are close)
|
|
21
25
|
return DD(hi, lo);
|
|
@@ -63,19 +67,24 @@ fn twoProdFma(a: f32, b: f32) -> DD {
|
|
|
63
67
|
|
|
64
68
|
// DD × DD product (Dekker/Bailey): twoProdBit(a.hi, b.hi) already captures
|
|
65
69
|
// the dominant term to full DD precision, and the cross terms are below the
|
|
66
|
-
// ~48-bit floor anyway, so folding them in with plain f32 loses nothing
|
|
67
|
-
//
|
|
68
|
-
//
|
|
69
|
-
//
|
|
70
|
-
//
|
|
71
|
-
//
|
|
72
|
-
|
|
70
|
+
// ~48-bit floor anyway, so folding them in with plain f32 loses nothing.
|
|
71
|
+
//
|
|
72
|
+
// Another real compiler bug, distinct from add.wgsl's twoSum one — confirmed
|
|
73
|
+
// on Intel Mesa ANV: when p.lo feeds straight into `crossAndLo` unobserved,
|
|
74
|
+
// the compiler folds it away entirely. Materializing p.lo itself through
|
|
75
|
+
// workgroup memory + workgroupBarrier() (like twoSumProtected does for its
|
|
76
|
+
// sum) is what fixes it, so ddMulRaw now takes threadSlot and always pays
|
|
77
|
+
// that barrier — no longer a plain unprotected batchable helper.
|
|
78
|
+
fn ddMulRaw(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
73
79
|
let p = twoProdBit(a.hi, b.hi);
|
|
74
|
-
|
|
80
|
+
dekkerScratch[threadSlot] = p.lo;
|
|
81
|
+
workgroupBarrier();
|
|
82
|
+
let pLo = dekkerScratch[threadSlot];
|
|
83
|
+
let crossAndLo = pLo + (a.hi * b.lo + a.lo * b.hi);
|
|
75
84
|
return DD(p.hi, crossAndLo);
|
|
76
85
|
}
|
|
77
86
|
|
|
78
87
|
fn ddMulProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
79
|
-
let raw = ddMulRaw(a, b);
|
|
88
|
+
let raw = ddMulRaw(a, b, threadSlot);
|
|
80
89
|
return fastTwoSumProtected(raw.hi, raw.lo, threadSlot);
|
|
81
90
|
}
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
// Requires f64/dekker.wgsl concatenated first for the DD struct, and
|
|
2
|
+
// f64/utils/add.wgsl (ddSubProtected/ddAddProtected) and
|
|
3
|
+
// f64/utils/multiply.wgsl (twoProdBit — squaring a plain f32 needs no
|
|
4
|
+
// barrier, per multiply.wgsl's own note that twoProdBit is universally safe
|
|
5
|
+
// unprotected).
|
|
6
|
+
//
|
|
7
|
+
// Verified empirically on real hardware (NVIDIA GTX 1650 + Intel Mesa
|
|
8
|
+
// integrated GPU, ~8000 random trials each plus explicit edge cases): no
|
|
9
|
+
// compiler-reassociation-style corruption of the kind that broke twoSum/
|
|
10
|
+
// twoProd (see add.wgsl's/multiply.wgsl's own headers) — every failure
|
|
11
|
+
// found was a genuine algorithm gap (the a.hi==0.0 case below), not a
|
|
12
|
+
// driver miscompile, and is fixed. Observed forward-error factor against a
|
|
13
|
+
// true f64 reference stayed ~3-4x eps on both backends across every random
|
|
14
|
+
// trial — noticeably tighter than ddDivProtected's own low-power spread
|
|
15
|
+
// (see divide.wgsl's header), since sqrt has no denominator to be unlucky
|
|
16
|
+
// about.
|
|
17
|
+
//
|
|
18
|
+
// One Newton refinement step (the classic extended-precision sqrt trick):
|
|
19
|
+
// x0 = sqrt(a.hi) is a plain f32 approximation; the residual a - x0^2,
|
|
20
|
+
// computed in DD arithmetic, recovers what x0 lost, and linearizing sqrt
|
|
21
|
+
// around x0 (dividing that residual by 2*x0) gives a correction term
|
|
22
|
+
// roughly doubling the precision — same ~48-bit target as ddDivProtected,
|
|
23
|
+
// so one step is enough.
|
|
24
|
+
//
|
|
25
|
+
// Undefined for a.hi < 0.0, same as plain sqrt() — callers must guard
|
|
26
|
+
// themselves; this never checks.
|
|
27
|
+
//
|
|
28
|
+
// a.hi == 0.0 (a genuinely zero input, not an underflowed one — zero is
|
|
29
|
+
// exactly representable in f32, unlike this scheme's real range limits;
|
|
30
|
+
// see splitDoubleDouble's own doc comment) makes x0 = sqrt(0) = 0, and the
|
|
31
|
+
// correction step would divide by 2*x0 = 0. Substituting a safe non-zero
|
|
32
|
+
// denominator via select() — rather than branching/returning early — keeps
|
|
33
|
+
// every thread calling ddSubProtected/ddAddProtected unconditionally, which
|
|
34
|
+
// their internal workgroupBarrier() requires; only the final result is
|
|
35
|
+
// selected between the computed value and the exact DD(0,0) answer.
|
|
36
|
+
fn ddSqrtProtected(a: DD, threadSlot: u32) -> DD {
|
|
37
|
+
let isZero = a.hi == 0.0;
|
|
38
|
+
let x0 = sqrt(a.hi);
|
|
39
|
+
let x0sq = twoProdBit(x0, x0);
|
|
40
|
+
let r = ddSubProtected(a, x0sq, threadSlot);
|
|
41
|
+
let safeDenom = select(2.0 * x0, 1.0, isZero);
|
|
42
|
+
let correction = r.hi / safeDenom;
|
|
43
|
+
let result = ddAddProtected(DD(x0, 0.0), DD(correction, 0.0), threadSlot);
|
|
44
|
+
return DD(select(result.hi, 0.0, isZero), select(result.lo, 0.0, isZero));
|
|
45
|
+
}
|
package/src/shaders/index.mjs
CHANGED
|
@@ -53,15 +53,24 @@ export const routineShaders = {};
|
|
|
53
53
|
import sscal from "./sscal.wgsl";
|
|
54
54
|
routineShaders.sscal = { sscal };
|
|
55
55
|
|
|
56
|
+
import cscal from "./cscal.wgsl";
|
|
57
|
+
routineShaders.cscal = { cscal };
|
|
58
|
+
|
|
56
59
|
import sswap from "./sswap.wgsl";
|
|
57
60
|
routineShaders.sswap = { sswap };
|
|
58
61
|
|
|
62
|
+
import dswap from "./dswap.wgsl"; // f64 sibling of sswap — pure data movement, no arithmetic, no reduction/barrier shader needed
|
|
63
|
+
routineShaders.dswap = { dswap };
|
|
64
|
+
|
|
59
65
|
import saxpy from "./saxpy.wgsl";
|
|
60
66
|
routineShaders.saxpy = { saxpy };
|
|
61
67
|
|
|
62
68
|
import scopy from "./scopy.wgsl";
|
|
63
69
|
routineShaders.scopy = { scopy };
|
|
64
70
|
|
|
71
|
+
import dcopy from "./dcopy.wgsl"; // f64 sibling of scopy — pure data movement, no arithmetic, no reduction/barrier shader needed
|
|
72
|
+
routineShaders.dcopy = { dcopy };
|
|
73
|
+
|
|
65
74
|
import sdot from "./sdot.wgsl";
|
|
66
75
|
import sum from "./reduction/sum.wgsl";
|
|
67
76
|
routineShaders.sdot = { sdot, "reduction/sum": sum };
|
|
@@ -90,6 +99,34 @@ routineShaders.dasum = {
|
|
|
90
99
|
"reduction/sumF64": sumF64,
|
|
91
100
|
};
|
|
92
101
|
|
|
102
|
+
import ddMulUtil from "./f64/utils/multiply.wgsl";
|
|
103
|
+
import ddot from "./ddot.wgsl";
|
|
104
|
+
// multiply.wgsl needs dekker's DD struct and add.wgsl's fsub/negf and
|
|
105
|
+
// fastTwoSumProtected, so those two precede it here.
|
|
106
|
+
routineShaders.ddot = {
|
|
107
|
+
"f64/dekker": dekker,
|
|
108
|
+
"f64/utils/add": ddAddUtil,
|
|
109
|
+
"f64/utils/multiply": ddMulUtil,
|
|
110
|
+
ddot,
|
|
111
|
+
"reduction/sumF64": sumF64,
|
|
112
|
+
};
|
|
113
|
+
|
|
114
|
+
import dscal from "./dscal.wgsl"; // f64 sibling of sscal — no reduction shader needed, unlike dasum/ddot
|
|
115
|
+
routineShaders.dscal = {
|
|
116
|
+
"f64/dekker": dekker,
|
|
117
|
+
"f64/utils/add": ddAddUtil,
|
|
118
|
+
"f64/utils/multiply": ddMulUtil,
|
|
119
|
+
dscal,
|
|
120
|
+
};
|
|
121
|
+
|
|
122
|
+
import daxpy from "./daxpy.wgsl"; // f64 sibling of saxpy — one ddMulProtected + one ddAddProtected per element, no reduction shader needed
|
|
123
|
+
routineShaders.daxpy = {
|
|
124
|
+
"f64/dekker": dekker,
|
|
125
|
+
"f64/utils/add": ddAddUtil,
|
|
126
|
+
"f64/utils/multiply": ddMulUtil,
|
|
127
|
+
daxpy,
|
|
128
|
+
};
|
|
129
|
+
|
|
93
130
|
import ddGreater from "./f64/utils/greater.wgsl";
|
|
94
131
|
import ddEqual from "./f64/utils/equal.wgsl";
|
|
95
132
|
import idamax from "./idamax.wgsl";
|
|
@@ -106,9 +143,41 @@ routineShaders.idamax = {
|
|
|
106
143
|
import srot from "./srot.wgsl";
|
|
107
144
|
routineShaders.srot = { srot };
|
|
108
145
|
|
|
146
|
+
import drot from "./drot.wgsl"; // f64 sibling of srot — four ddMulProtected + two ddAddProtected per element, no reduction shader needed
|
|
147
|
+
routineShaders.drot = {
|
|
148
|
+
"f64/dekker": dekker,
|
|
149
|
+
"f64/utils/add": ddAddUtil,
|
|
150
|
+
"f64/utils/multiply": ddMulUtil,
|
|
151
|
+
drot,
|
|
152
|
+
};
|
|
153
|
+
|
|
109
154
|
import srotm from "./srotm.wgsl";
|
|
110
155
|
routineShaders.srotm = { srotm };
|
|
111
156
|
|
|
157
|
+
import drotm from "./drotm.wgsl"; // f64 sibling of srotm — four ddMulProtected + two ddAddProtected per element, no reduction shader needed
|
|
158
|
+
routineShaders.drotm = {
|
|
159
|
+
"f64/dekker": dekker,
|
|
160
|
+
"f64/utils/add": ddAddUtil,
|
|
161
|
+
"f64/utils/multiply": ddMulUtil,
|
|
162
|
+
drotm,
|
|
163
|
+
};
|
|
164
|
+
|
|
165
|
+
import ddDivUtil from "./f64/utils/divide.wgsl";
|
|
166
|
+
import ddSqrtUtil from "./f64/utils/sqrt.wgsl";
|
|
167
|
+
import dnrm2 from "./dnrm2.wgsl"; // f64 sibling of snrm2 — scaled accumulation (Blue's algorithm) ported to double-double, branch-free (select()) since ddDivProtected/ddMulProtected/ddAddProtected's barriers need every thread to take the same path
|
|
168
|
+
import scaledSumF64 from "./reduction/scaledSumF64.wgsl";
|
|
169
|
+
routineShaders.dnrm2 = {
|
|
170
|
+
"f64/dekker": dekker,
|
|
171
|
+
"f64/utils/abs": ddAbs,
|
|
172
|
+
"f64/utils/greater": ddGreater,
|
|
173
|
+
"f64/utils/add": ddAddUtil,
|
|
174
|
+
"f64/utils/multiply": ddMulUtil,
|
|
175
|
+
"f64/utils/divide": ddDivUtil,
|
|
176
|
+
"f64/utils/sqrt": ddSqrtUtil,
|
|
177
|
+
dnrm2,
|
|
178
|
+
"reduction/scaledSumF64": scaledSumF64,
|
|
179
|
+
};
|
|
180
|
+
|
|
112
181
|
import sgemv_n from "./sgemv_n.wgsl";
|
|
113
182
|
import sgemv_t from "./sgemv_t.wgsl";
|
|
114
183
|
routineShaders.sgemv = { sgemv_n, sgemv_t }; // one or the other, picked by trans
|
|
@@ -0,0 +1,93 @@
|
|
|
1
|
+
// scaledSum reduction (f64, double-double): collapses 2*WGS (scale, ssq) DD
|
|
2
|
+
// partials from dnrm2.wgsl into the final norm — sqrt(scale² · ssq) ==
|
|
3
|
+
// scale · sqrt(ssq), via ddMulProtected/ddSqrtProtected. Mirrors
|
|
4
|
+
// reduction/scaledSum.wgsl's shape exactly; ssqMergeProtected is duplicated
|
|
5
|
+
// from dnrm2.wgsl rather than shared via f64/utils/ — see that file's own
|
|
6
|
+
// header for why (same convention the f32 pair already uses).
|
|
7
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
8
|
+
// partialsScale*/partialsSsq* must have exactly 2*WGS entries each.
|
|
9
|
+
|
|
10
|
+
@group(0) @binding(0) var<storage, read> partialsScaleHi: array<f32>;
|
|
11
|
+
@group(0) @binding(1) var<storage, read> partialsScaleLo: array<f32>;
|
|
12
|
+
@group(0) @binding(2) var<storage, read> partialsSsqHi: array<f32>;
|
|
13
|
+
@group(0) @binding(3) var<storage, read> partialsSsqLo: array<f32>;
|
|
14
|
+
@group(0) @binding(4) var<storage, read_write> resultHi: array<f32, 1>;
|
|
15
|
+
@group(0) @binding(5) var<storage, read_write> resultLo: array<f32, 1>;
|
|
16
|
+
|
|
17
|
+
const WGS: u32 = 64;
|
|
18
|
+
|
|
19
|
+
struct ScaleSsq {
|
|
20
|
+
scale: DD,
|
|
21
|
+
ssq: DD,
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
fn ddSelect(a: DD, b: DD, cond: bool) -> DD {
|
|
25
|
+
return DD(select(a.hi, b.hi, cond), select(a.lo, b.lo, cond));
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
// Associative merge of two independent (scale, ssq) partials — see
|
|
29
|
+
// dnrm2.wgsl for the derivation and why this is branch-free.
|
|
30
|
+
fn ssqMergeProtected(a: ScaleSsq, b: ScaleSsq, threadSlot: u32) -> ScaleSsq {
|
|
31
|
+
let isBigger = !ddGreater(b.scale, a.scale); // a.scale >= b.scale
|
|
32
|
+
let bigger = ddSelect(b.scale, a.scale, isBigger);
|
|
33
|
+
let smaller = ddSelect(a.scale, b.scale, isBigger);
|
|
34
|
+
let biggerSsq = ddSelect(b.ssq, a.ssq, isBigger);
|
|
35
|
+
let smallerSsq = ddSelect(a.ssq, b.ssq, isBigger);
|
|
36
|
+
let biggerIsZero = bigger.hi == 0.0;
|
|
37
|
+
let safeBigger = ddSelect(bigger, DD(1.0, 0.0), biggerIsZero);
|
|
38
|
+
let r = ddDivProtected(smaller, safeBigger, threadSlot);
|
|
39
|
+
let rsq = ddMulProtected(r, r, threadSlot);
|
|
40
|
+
let smallerSsqTimesRsq = ddMulProtected(smallerSsq, rsq, threadSlot);
|
|
41
|
+
let newSsq = ddAddProtected(biggerSsq, smallerSsqTimesRsq, threadSlot);
|
|
42
|
+
return ScaleSsq(bigger, newSsq);
|
|
43
|
+
}
|
|
44
|
+
|
|
45
|
+
var<workgroup> tileScaleHi: array<f32, 64>;
|
|
46
|
+
var<workgroup> tileScaleLo: array<f32, 64>;
|
|
47
|
+
var<workgroup> tileSsqHi: array<f32, 64>;
|
|
48
|
+
var<workgroup> tileSsqLo: array<f32, 64>;
|
|
49
|
+
|
|
50
|
+
@compute @workgroup_size(64)
|
|
51
|
+
fn reduce_scaled_f64(
|
|
52
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
53
|
+
) {
|
|
54
|
+
let i = lid.x;
|
|
55
|
+
let a = ScaleSsq(DD(partialsScaleHi[i], partialsScaleLo[i]), DD(partialsSsqHi[i], partialsSsqLo[i]));
|
|
56
|
+
let b = ScaleSsq(DD(partialsScaleHi[i + WGS], partialsScaleLo[i + WGS]), DD(partialsSsqHi[i + WGS], partialsSsqLo[i + WGS]));
|
|
57
|
+
let merged0 = ssqMergeProtected(a, b, i);
|
|
58
|
+
tileScaleHi[i] = merged0.scale.hi;
|
|
59
|
+
tileScaleLo[i] = merged0.scale.lo;
|
|
60
|
+
tileSsqHi[i] = merged0.ssq.hi;
|
|
61
|
+
tileSsqLo[i] = merged0.ssq.lo;
|
|
62
|
+
workgroupBarrier();
|
|
63
|
+
|
|
64
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
65
|
+
let partner = select(i, i + s, i < s);
|
|
66
|
+
let ai = ScaleSsq(DD(tileScaleHi[i], tileScaleLo[i]), DD(tileSsqHi[i], tileSsqLo[i]));
|
|
67
|
+
let bi = ScaleSsq(DD(tileScaleHi[partner], tileScaleLo[partner]), DD(tileSsqHi[partner], tileSsqLo[partner]));
|
|
68
|
+
let merged = ssqMergeProtected(ai, bi, i);
|
|
69
|
+
workgroupBarrier();
|
|
70
|
+
if (i < s) {
|
|
71
|
+
tileScaleHi[i] = merged.scale.hi;
|
|
72
|
+
tileScaleLo[i] = merged.scale.lo;
|
|
73
|
+
tileSsqHi[i] = merged.ssq.hi;
|
|
74
|
+
tileSsqLo[i] = merged.ssq.lo;
|
|
75
|
+
}
|
|
76
|
+
workgroupBarrier();
|
|
77
|
+
}
|
|
78
|
+
|
|
79
|
+
// ddSqrtProtected/ddMulProtected's own workgroupBarrier()s need every
|
|
80
|
+
// thread to call them — every thread redundantly computes the same final
|
|
81
|
+
// scale·sqrt(ssq) from tile[0] (still visible to all after the reduction
|
|
82
|
+
// above), and only the write-back is conditional. Guarding the calls
|
|
83
|
+
// themselves behind `if (i == 0u)` (as the plain-f32 original safely
|
|
84
|
+
// does with its unprotected `sqrt()`) would leave 63 threads never
|
|
85
|
+
// reaching a barrier the one remaining thread still needs.
|
|
86
|
+
let scale = DD(tileScaleHi[0], tileScaleLo[0]);
|
|
87
|
+
let ssq = DD(tileSsqHi[0], tileSsqLo[0]);
|
|
88
|
+
let result = ddMulProtected(scale, ddSqrtProtected(ssq, i), i);
|
|
89
|
+
if (i == 0u) {
|
|
90
|
+
resultHi[0] = result.hi;
|
|
91
|
+
resultLo[0] = result.lo;
|
|
92
|
+
}
|
|
93
|
+
}
|
package/src/snrm2/snrm2.mjs
CHANGED
|
@@ -13,14 +13,12 @@ import { extractResult } from "../util/result.mjs";
|
|
|
13
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
14
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
15
15
|
import { WGS } from "../util/constants.mjs";
|
|
16
|
-
import { requireSameDevice } from "../util/device.mjs";
|
|
17
|
-
|
|
16
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
18
17
|
|
|
19
18
|
export async function snrm2(device, n, x, incx) {
|
|
20
19
|
const xIsGpu = x instanceof GpuVector;
|
|
21
20
|
|
|
22
|
-
|
|
23
|
-
throw new Error("device must be a GPUDevice.");
|
|
21
|
+
requireGpuDevice(device);
|
|
24
22
|
requireSameDevice(device, "snrm2", { x });
|
|
25
23
|
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
26
24
|
throw new Error("n and incx must be integers.");
|
|
@@ -46,13 +44,19 @@ export async function snrm2(device, n, x, incx) {
|
|
|
46
44
|
try {
|
|
47
45
|
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "snrm2-x", false);
|
|
48
46
|
// 2*WGS partial (scale, ssq) pairs — see snrm2.wgsl for what they represent.
|
|
49
|
-
partialsScaleBuffer = createStorageBuffer(
|
|
47
|
+
partialsScaleBuffer = createStorageBuffer(
|
|
48
|
+
device,
|
|
50
49
|
2 * WGS * 4,
|
|
51
50
|
"snrm2-partials-scale",
|
|
52
51
|
);
|
|
53
|
-
partialsSsqBuffer = createStorageBuffer(
|
|
52
|
+
partialsSsqBuffer = createStorageBuffer(
|
|
53
|
+
device,
|
|
54
|
+
2 * WGS * 4,
|
|
55
|
+
"snrm2-partials-ssq",
|
|
56
|
+
);
|
|
54
57
|
resultBuffer = createResultBuffer(device, 4, "snrm2-result"); // final f32 scalar
|
|
55
|
-
paramsBuffer = createParamsBuffer(
|
|
58
|
+
paramsBuffer = createParamsBuffer(
|
|
59
|
+
device,
|
|
56
60
|
[
|
|
57
61
|
{ value: n, type: "u32" },
|
|
58
62
|
{ value: incx, type: "u32" },
|
|
@@ -66,7 +70,8 @@ export async function snrm2(device, n, x, incx) {
|
|
|
66
70
|
partialsSsqBuffer,
|
|
67
71
|
paramsBuffer,
|
|
68
72
|
]);
|
|
69
|
-
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
73
|
+
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
74
|
+
device,
|
|
70
75
|
pipelineMain,
|
|
71
76
|
bgMain,
|
|
72
77
|
2 * WGS,
|
|
@@ -74,12 +79,13 @@ export async function snrm2(device, n, x, incx) {
|
|
|
74
79
|
|
|
75
80
|
submit(device, enc1);
|
|
76
81
|
|
|
77
|
-
const bgReduce = createBindGroup(
|
|
78
|
-
|
|
79
|
-
|
|
80
|
-
resultBuffer,
|
|
81
|
-
|
|
82
|
-
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
82
|
+
const bgReduce = createBindGroup(
|
|
83
|
+
device,
|
|
84
|
+
pipelineReduce.getBindGroupLayout(0),
|
|
85
|
+
[partialsScaleBuffer, partialsSsqBuffer, resultBuffer],
|
|
86
|
+
);
|
|
87
|
+
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
88
|
+
device,
|
|
83
89
|
pipelineReduce,
|
|
84
90
|
bgReduce,
|
|
85
91
|
1,
|