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
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
2
|
+
|
|
3
|
+
// ── A real compiler bug — read before touching anything below ──────────────
|
|
4
|
+
//
|
|
5
|
+
// twoSum/fastTwoSum's error term `e` should be nonzero (that's the point —
|
|
6
|
+
// `s` is rounded). A front-end-optimizer bug zeros it anyway, on both NVIDIA
|
|
7
|
+
// and Mesa (ANV + llvmpipe), via two different mechanisms needing two fixes:
|
|
8
|
+
// bitcast-based subtraction (`fsub`/`negf`, fixes NVIDIA) and materializing
|
|
9
|
+
// the sum through workgroup memory + workgroupBarrier() (fixes Mesa). Only
|
|
10
|
+
// both together (ddAddProtected) is verified correct everywhere — the plain
|
|
11
|
+
// twoSum/fastTwoSum/ddAdd below are reference-only, not safe to use.
|
|
12
|
+
fn negf(x: f32) -> f32 {
|
|
13
|
+
return bitcast<f32>(bitcast<u32>(x) ^ 0x80000000u);
|
|
14
|
+
}
|
|
15
|
+
fn fsub(a: f32, b: f32) -> f32 {
|
|
16
|
+
return a + negf(b);
|
|
17
|
+
}
|
|
18
|
+
|
|
19
|
+
// Knuth/Møller's TwoSum: s = fl(a+b), e = exact rounding error, a+b == s+e.
|
|
20
|
+
// Works for any a, b. UNPROTECTED — see header above.
|
|
21
|
+
fn twoSum(a: f32, b: f32) -> DD {
|
|
22
|
+
let s = a + b;
|
|
23
|
+
let v = s - a;
|
|
24
|
+
let e = (a - (s - v)) + (b - v);
|
|
25
|
+
return DD(s, e);
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
// Dekker's Fast-Two-Sum: same contract, but only correct when |a| >= |b|.
|
|
29
|
+
// UNPROTECTED — see header above.
|
|
30
|
+
fn fastTwoSum(a: f32, b: f32) -> DD {
|
|
31
|
+
let s = a + b;
|
|
32
|
+
let e = b - (s - a);
|
|
33
|
+
return DD(s, e);
|
|
34
|
+
}
|
|
35
|
+
|
|
36
|
+
// Double-double addition (Dekker's Add2). UNPROTECTED — see header above.
|
|
37
|
+
fn ddAdd(a: DD, b: DD) -> DD {
|
|
38
|
+
let s = twoSum(a.hi, b.hi);
|
|
39
|
+
let loSum = a.lo + b.lo;
|
|
40
|
+
return fastTwoSum(s.hi, s.lo + loSum);
|
|
41
|
+
}
|
|
42
|
+
|
|
43
|
+
// ── Protected variants — use these ──────────────────────────────────────────
|
|
44
|
+
//
|
|
45
|
+
// Bitcast subtraction + workgroup-barrier materialization, verified correct
|
|
46
|
+
// on all three backends tested. Costs a real barrier: fine for O(1)-per-
|
|
47
|
+
// thread or O(log n) reduction use, not a long per-element loop. A
|
|
48
|
+
// workgroupBarrier() requires uniform control flow, so:
|
|
49
|
+
// - `threadSlot` must be unique per concurrent caller (e.g. local_invocation_index).
|
|
50
|
+
// - Every thread in the workgroup must call this the same number of times
|
|
51
|
+
// — including ones whose result gets discarded. Compute unconditionally;
|
|
52
|
+
// only the write-back should be conditional.
|
|
53
|
+
var<workgroup> dekkerScratch: array<f32, 64>;
|
|
54
|
+
|
|
55
|
+
fn twoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
|
|
56
|
+
dekkerScratch[threadSlot] = a + b;
|
|
57
|
+
workgroupBarrier();
|
|
58
|
+
let s = dekkerScratch[threadSlot];
|
|
59
|
+
let v = fsub(s, a);
|
|
60
|
+
let e = fsub(a, fsub(s, v)) + fsub(b, v);
|
|
61
|
+
return DD(s, e);
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
fn fastTwoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
|
|
65
|
+
dekkerScratch[threadSlot] = a + b;
|
|
66
|
+
workgroupBarrier();
|
|
67
|
+
let s = dekkerScratch[threadSlot];
|
|
68
|
+
let e = fsub(b, fsub(s, a));
|
|
69
|
+
return DD(s, e);
|
|
70
|
+
}
|
|
71
|
+
|
|
72
|
+
// Protected double-double addition — same contract as ddAdd, but exact.
|
|
73
|
+
fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
74
|
+
let s = twoSumProtected(a.hi, b.hi, threadSlot);
|
|
75
|
+
let loSum = a.lo + b.lo;
|
|
76
|
+
return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
|
|
77
|
+
}
|
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
2
|
+
|
|
3
|
+
// a == b for double-double pairs — exact field equality, no rounding
|
|
4
|
+
// involved, so (like ddGreater) this needs no protection.
|
|
5
|
+
fn ddEqual(a: DD, b: DD) -> bool {
|
|
6
|
+
return a.hi == b.hi && a.lo == b.lo;
|
|
7
|
+
}
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
2
|
+
|
|
3
|
+
// a > b for double-double pairs. hi dominates (|lo| <= ulp(hi)/2 always), so
|
|
4
|
+
// comparing hi alone is correct except on an exact hi tie, when lo breaks it.
|
|
5
|
+
// A plain comparison, not a rounding-identity subtraction — no reassociation
|
|
6
|
+
// risk, so unlike twoSum/fastTwoSum this needs no protection.
|
|
7
|
+
fn ddGreater(a: DD, b: DD) -> bool {
|
|
8
|
+
if (a.hi != b.hi) {
|
|
9
|
+
return a.hi > b.hi;
|
|
10
|
+
}
|
|
11
|
+
return a.lo > b.lo;
|
|
12
|
+
}
|
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
// Requires f64/dekker.wgsl concatenated first for the DD struct, and
|
|
2
|
+
// f64/utils/add.wgsl for fsub/negf (bitcast-based subtraction/negation) and,
|
|
3
|
+
// for ddMulProtected at the bottom, fastTwoSumProtected.
|
|
4
|
+
//
|
|
5
|
+
// Use twoProdBit — verified universal (0 corrupting failures across 3000+
|
|
6
|
+
// random trials on NVIDIA/Intel-Mesa-ANV/llvmpipe), no barrier protection
|
|
7
|
+
// needed. The classic approaches below (twoProd, twoProdFma) each fail on
|
|
8
|
+
// one backend in a way barrier materialization doesn't fix; twoProdBit
|
|
9
|
+
// sidesteps the bug instead by deriving the split via bitcast+bitmask
|
|
10
|
+
// rather than an arithmetic identity, leaving nothing for a reassociating
|
|
11
|
+
// compiler to fold. Intel Mesa ANV shows frequent last-bit-only diffs from
|
|
12
|
+
// strict ground truth (never data-corrupting) — consistent with the driver
|
|
13
|
+
// legitimately auto-fusing `x - y*z` into hardware FMA.
|
|
14
|
+
const SPLIT_CONST: f32 = 4097.0;
|
|
15
|
+
|
|
16
|
+
fn bitSplit(a: f32) -> DD {
|
|
17
|
+
let bits = bitcast<u32>(a);
|
|
18
|
+
let hiBits = bits & 0xFFFFF800u; // keep sign+exponent+top 12 mantissa bits
|
|
19
|
+
let hi = bitcast<f32>(hiBits);
|
|
20
|
+
let lo = fsub(a, hi); // exact by Sterbenz's lemma (hi, a share an exponent, are close)
|
|
21
|
+
return DD(hi, lo);
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
fn twoProdBit(a: f32, b: f32) -> DD {
|
|
25
|
+
let s = a * b;
|
|
26
|
+
let aSplit = bitSplit(a);
|
|
27
|
+
let bSplit = bitSplit(b);
|
|
28
|
+
let e = fsub(fsub(fsub(fsub(s, aSplit.hi * bSplit.hi), aSplit.lo * bSplit.hi), aSplit.hi * bSplit.lo), aSplit.lo * bSplit.lo);
|
|
29
|
+
return DD(s, negf(e));
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
// ── Unsafe historical reference — do not use ────────────────────────────
|
|
33
|
+
// Both broken on one backend, confirmed via isolated cross-driver testing,
|
|
34
|
+
// NOT fixed by barrier materialization (unlike addition's bug):
|
|
35
|
+
// - veltkampSplit/twoProd (Dekker's original): fails on NVIDIA — compiler
|
|
36
|
+
// folds `hi = c - (c - a)` to `= a` straight through fsub/negf, even
|
|
37
|
+
// with every intermediate barrier-materialized (11/11 fail, worse than
|
|
38
|
+
// unprotected's 6/11).
|
|
39
|
+
// - twoProdFma (Ogita/Rump/Oishi): fails on llvmpipe — its software fma()
|
|
40
|
+
// likely isn't genuinely fused, making `fma(a,b,-(a*b))` correctly (not
|
|
41
|
+
// buggily) zero. Materializing `s` doesn't change this.
|
|
42
|
+
fn veltkampSplit(a: f32) -> DD {
|
|
43
|
+
let c = SPLIT_CONST * a;
|
|
44
|
+
let big = fsub(c, a);
|
|
45
|
+
let hi = fsub(c, big);
|
|
46
|
+
let lo = fsub(a, hi);
|
|
47
|
+
return DD(hi, lo);
|
|
48
|
+
}
|
|
49
|
+
|
|
50
|
+
fn twoProd(a: f32, b: f32) -> DD {
|
|
51
|
+
let s = a * b;
|
|
52
|
+
let aSplit = veltkampSplit(a);
|
|
53
|
+
let bSplit = veltkampSplit(b);
|
|
54
|
+
let e = fsub(fsub(fsub(fsub(s, aSplit.hi * bSplit.hi), aSplit.lo * bSplit.hi), aSplit.hi * bSplit.lo), aSplit.lo * bSplit.lo);
|
|
55
|
+
return DD(s, negf(e));
|
|
56
|
+
}
|
|
57
|
+
|
|
58
|
+
fn twoProdFma(a: f32, b: f32) -> DD {
|
|
59
|
+
let s = a * b;
|
|
60
|
+
let e = fma(a, b, negf(s));
|
|
61
|
+
return DD(s, e);
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
// DD × DD product (Dekker/Bailey): twoProdBit(a.hi, b.hi) already captures
|
|
65
|
+
// 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
|
+
// only the final renormalization needs barrier protection. Split into
|
|
68
|
+
// ddMulRaw (unprotected) and ddMulProtected (renormalizes via
|
|
69
|
+
// fastTwoSumProtected) so callers with several products can batch them
|
|
70
|
+
// through one shared barrier. ddMulRaw's result isn't a valid DD pair on
|
|
71
|
+
// its own — it must be renormalized before use.
|
|
72
|
+
fn ddMulRaw(a: DD, b: DD) -> DD {
|
|
73
|
+
let p = twoProdBit(a.hi, b.hi);
|
|
74
|
+
let crossAndLo = p.lo + (a.hi * b.lo + a.lo * b.hi);
|
|
75
|
+
return DD(p.hi, crossAndLo);
|
|
76
|
+
}
|
|
77
|
+
|
|
78
|
+
fn ddMulProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
79
|
+
let raw = ddMulRaw(a, b);
|
|
80
|
+
return fastTwoSumProtected(raw.hi, raw.lo, threadSlot);
|
|
81
|
+
}
|
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
// idamax: returns index of element with largest absolute value (f64, double-double)
|
|
2
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/argmaxF64.wgsl.
|
|
3
|
+
// Concatenated after f64/dekker.wgsl (DD struct), f64/utils/abs.wgsl (ddAbs),
|
|
4
|
+
// f64/utils/greater.wgsl (ddGreater), and f64/utils/equal.wgsl (ddEqual).
|
|
5
|
+
|
|
6
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
7
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
8
|
+
@group(0) @binding(2) var<storage, read_write> partialsValHi: array<f32>;
|
|
9
|
+
@group(0) @binding(3) var<storage, read_write> partialsValLo: array<f32>;
|
|
10
|
+
@group(0) @binding(4) var<storage, read_write> partialsIdx: array<u32>;
|
|
11
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
12
|
+
|
|
13
|
+
struct Params {
|
|
14
|
+
n: u32,
|
|
15
|
+
x_inc: u32,
|
|
16
|
+
}
|
|
17
|
+
|
|
18
|
+
const WGS: u32 = 64;
|
|
19
|
+
|
|
20
|
+
var<workgroup> tile_val: array<DD, 64>;
|
|
21
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
22
|
+
|
|
23
|
+
@compute @workgroup_size(64)
|
|
24
|
+
fn idamax_main(
|
|
25
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
26
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
27
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
28
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
29
|
+
) {
|
|
30
|
+
// DD(-1.0, 0.0) is a safe sentinel: any |x[i]| >= 0 beats it,
|
|
31
|
+
// so workgroups with no elements lose gracefully in the epilogue.
|
|
32
|
+
var best_val0 = DD(-1.0, 0.0); var best_idx0: u32 = 0u;
|
|
33
|
+
var best_val1 = DD(-1.0, 0.0); var best_idx1: u32 = 0u;
|
|
34
|
+
var best_val2 = DD(-1.0, 0.0); var best_idx2: u32 = 0u;
|
|
35
|
+
var best_val3 = DD(-1.0, 0.0); var best_idx3: u32 = 0u;
|
|
36
|
+
|
|
37
|
+
let stride = num_wg.x * WGS;
|
|
38
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
39
|
+
|
|
40
|
+
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
41
|
+
let i0 = id * params.x_inc;
|
|
42
|
+
let i1 = (id + stride) * params.x_inc;
|
|
43
|
+
let i2 = (id + 2u * stride) * params.x_inc;
|
|
44
|
+
let i3 = (id + 3u * stride) * params.x_inc;
|
|
45
|
+
let v0 = ddAbs(DD(xHi[i0], xLo[i0]));
|
|
46
|
+
let v1 = ddAbs(DD(xHi[i1], xLo[i1]));
|
|
47
|
+
let v2 = ddAbs(DD(xHi[i2], xLo[i2]));
|
|
48
|
+
let v3 = ddAbs(DD(xHi[i3], xLo[i3]));
|
|
49
|
+
if (ddGreater(v0, best_val0)) { best_val0 = v0; best_idx0 = id; }
|
|
50
|
+
if (ddGreater(v1, best_val1)) { best_val1 = v1; best_idx1 = id + stride; }
|
|
51
|
+
if (ddGreater(v2, best_val2)) { best_val2 = v2; best_idx2 = id + 2u * stride; }
|
|
52
|
+
if (ddGreater(v3, best_val3)) { best_val3 = v3; best_idx3 = id + 3u * stride; }
|
|
53
|
+
}
|
|
54
|
+
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
55
|
+
let i = id * params.x_inc;
|
|
56
|
+
let v = ddAbs(DD(xHi[i], xLo[i]));
|
|
57
|
+
if (ddGreater(v, best_val0)) { best_val0 = v; best_idx0 = id; }
|
|
58
|
+
}
|
|
59
|
+
|
|
60
|
+
// merge 4 independent lanes; prefer lower index on tie (first occurrence wins)
|
|
61
|
+
if (ddGreater(best_val1, best_val0) ||
|
|
62
|
+
(ddEqual(best_val1, best_val0) && best_idx1 < best_idx0)) {
|
|
63
|
+
best_val0 = best_val1; best_idx0 = best_idx1;
|
|
64
|
+
}
|
|
65
|
+
if (ddGreater(best_val2, best_val0) ||
|
|
66
|
+
(ddEqual(best_val2, best_val0) && best_idx2 < best_idx0)) {
|
|
67
|
+
best_val0 = best_val2; best_idx0 = best_idx2;
|
|
68
|
+
}
|
|
69
|
+
if (ddGreater(best_val3, best_val0) ||
|
|
70
|
+
(ddEqual(best_val3, best_val0) && best_idx3 < best_idx0)) {
|
|
71
|
+
best_val0 = best_val3; best_idx0 = best_idx3;
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
tile_val[lid.x] = best_val0;
|
|
75
|
+
tile_idx[lid.x] = best_idx0;
|
|
76
|
+
workgroupBarrier();
|
|
77
|
+
|
|
78
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
79
|
+
if (lid.x < s) {
|
|
80
|
+
let a_val = tile_val[lid.x];
|
|
81
|
+
let b_val = tile_val[lid.x + s];
|
|
82
|
+
if (ddGreater(b_val, a_val) ||
|
|
83
|
+
(ddEqual(b_val, a_val) && tile_idx[lid.x + s] < tile_idx[lid.x])) {
|
|
84
|
+
tile_val[lid.x] = b_val;
|
|
85
|
+
tile_idx[lid.x] = tile_idx[lid.x + s];
|
|
86
|
+
}
|
|
87
|
+
}
|
|
88
|
+
workgroupBarrier();
|
|
89
|
+
}
|
|
90
|
+
|
|
91
|
+
if (lid.x == 0u) {
|
|
92
|
+
partialsValHi[wgid.x] = tile_val[0].hi;
|
|
93
|
+
partialsValLo[wgid.x] = tile_val[0].lo;
|
|
94
|
+
partialsIdx[wgid.x] = tile_idx[0];
|
|
95
|
+
}
|
|
96
|
+
}
|
package/src/shaders/index.mjs
CHANGED
|
@@ -3,25 +3,175 @@
|
|
|
3
3
|
*
|
|
4
4
|
* `shaders/*.wgsl` — one WGSL compute shader per BLAS routine (sscal, saxpy, sdot, …).
|
|
5
5
|
*
|
|
6
|
-
* `
|
|
7
|
-
*
|
|
8
|
-
*
|
|
6
|
+
* `routineShaders` below is the single source of truth: routine name → the WGSL source(s)
|
|
7
|
+
* its `getPipeline()` calls actually reference, verified against every `src/<routine>/<routine>.mjs`
|
|
8
|
+
* rather than inferred from naming convention (see its doc comment for the exceptions). Each
|
|
9
|
+
* shader is imported right above the line that adds it — the import *is* the mapping entry, no
|
|
10
|
+
* separate block to cross-reference. `shaderSources`, the flat name → source registry the
|
|
11
|
+
* browser bundle's runtime lookup needs, is *derived* from `routineShaders` rather than
|
|
12
|
+
* hand-duplicated, so the two can never drift apart. In Node.js neither is read — shaders are
|
|
13
|
+
* `readFileSync` from disk directly; `scripts/build-browser.mjs` inlines this module into the
|
|
14
|
+
* browser's IIFE bundle via esbuild instead.
|
|
9
15
|
*
|
|
10
16
|
* ## Cross-shader patterns
|
|
11
17
|
*
|
|
12
|
-
* **
|
|
13
|
-
*
|
|
14
|
-
*
|
|
18
|
+
* **Single bind group.** Every shader with bindings uses `@group(0)` only — the JS side always
|
|
19
|
+
* calls `pipeline.getBindGroupLayout(0)`, no secondary groups to track. Binding order is
|
|
20
|
+
* consistent too: any read-only storage buffers come before read_write ones, with the
|
|
21
|
+
* `uniform Params` struct always last. `@binding` indices match the position of each resource in
|
|
22
|
+
* the array passed to `createBindGroup`, which appends `resultBuffer` last.
|
|
15
23
|
*
|
|
16
|
-
* **
|
|
17
|
-
* `
|
|
24
|
+
* **All counts and strides are `u32`.** `n`, `x_inc`, `y_inc`, and every other index/count field
|
|
25
|
+
* in a `Params` struct is unsigned, avoiding implicit sign-extension in index expressions like
|
|
26
|
+
* `id * params.x_inc`.
|
|
18
27
|
*
|
|
19
|
-
*
|
|
20
|
-
*
|
|
21
|
-
*
|
|
22
|
-
* **All counts and strides are `u32`.** `n`, `x_inc`, `y_inc`, and any other index fields in the
|
|
23
|
-
* `Params` uniform struct are unsigned. This avoids implicit sign-extension when they appear in
|
|
24
|
-
* index expressions like `id * params.x_inc`.
|
|
28
|
+
* **Entry points don't have to be named `main`.** `loadShader` (`util/pipeline.mjs`)
|
|
29
|
+
* auto-detects the sole `@compute` function in a module instead of requiring a fixed name, so
|
|
30
|
+
* `dasum_main`, `strsv_invert_block_main`, etc. work without renaming.
|
|
25
31
|
*
|
|
26
32
|
* @module devdocs/shaders
|
|
27
33
|
*/
|
|
34
|
+
|
|
35
|
+
/**
|
|
36
|
+
* Routine name → the WGSL source(s) its `getPipeline()` calls reference. Keys are the exact
|
|
37
|
+
* shader names `getPipeline(device, name)` is called with — most routines have one, some pick
|
|
38
|
+
* one of several conditionally (e.g. sgemv's `sgemv_n`/`sgemv_t`, by `trans`), and some have no
|
|
39
|
+
* dedicated shader at all:
|
|
40
|
+
*
|
|
41
|
+
* - `sgemmtr`/`ssyrk`/`ssyr2k` all dispatch through `sgemmtr_small`/`sgemmtr_large`.
|
|
42
|
+
* - `strsm` reuses `strsv_invert_block` and `sscal`, plus its own `block_transfer` and the
|
|
43
|
+
* shared `sgemm_small`/`sgemm_large`.
|
|
44
|
+
* - `dasum`/`idamax` concatenate several f64 utility shaders with their own — see
|
|
45
|
+
* `getPipeline`'s `shaderName: string[]` behaviour.
|
|
46
|
+
* - `random` has no entry — CPU-only, no `getPipeline()` call.
|
|
47
|
+
*
|
|
48
|
+
* Built up entry by entry so each import sits next to the mapping entry that uses it.
|
|
49
|
+
* @public
|
|
50
|
+
*/
|
|
51
|
+
export const routineShaders = {};
|
|
52
|
+
|
|
53
|
+
import sscal from "./sscal.wgsl";
|
|
54
|
+
routineShaders.sscal = { sscal };
|
|
55
|
+
|
|
56
|
+
import sswap from "./sswap.wgsl";
|
|
57
|
+
routineShaders.sswap = { sswap };
|
|
58
|
+
|
|
59
|
+
import saxpy from "./saxpy.wgsl";
|
|
60
|
+
routineShaders.saxpy = { saxpy };
|
|
61
|
+
|
|
62
|
+
import scopy from "./scopy.wgsl";
|
|
63
|
+
routineShaders.scopy = { scopy };
|
|
64
|
+
|
|
65
|
+
import sdot from "./sdot.wgsl";
|
|
66
|
+
import sum from "./reduction/sum.wgsl";
|
|
67
|
+
routineShaders.sdot = { sdot, "reduction/sum": sum };
|
|
68
|
+
|
|
69
|
+
import sasum from "./sasum.wgsl";
|
|
70
|
+
routineShaders.sasum = { sasum, "reduction/sum": sum };
|
|
71
|
+
|
|
72
|
+
import snrm2 from "./snrm2.wgsl";
|
|
73
|
+
import scaledSum from "./reduction/scaledSum.wgsl";
|
|
74
|
+
routineShaders.snrm2 = { snrm2, "reduction/scaledSum": scaledSum };
|
|
75
|
+
|
|
76
|
+
import isamax from "./isamax.wgsl";
|
|
77
|
+
import argmax from "./reduction/argmax.wgsl";
|
|
78
|
+
routineShaders.isamax = { isamax, "reduction/argmax": argmax };
|
|
79
|
+
|
|
80
|
+
import dekker from "./f64/dekker.wgsl";
|
|
81
|
+
import ddAbs from "./f64/utils/abs.wgsl";
|
|
82
|
+
import ddAddUtil from "./f64/utils/add.wgsl";
|
|
83
|
+
import dasum from "./dasum.wgsl";
|
|
84
|
+
import sumF64 from "./reduction/sumF64.wgsl";
|
|
85
|
+
routineShaders.dasum = {
|
|
86
|
+
"f64/dekker": dekker,
|
|
87
|
+
"f64/utils/abs": ddAbs,
|
|
88
|
+
"f64/utils/add": ddAddUtil,
|
|
89
|
+
dasum,
|
|
90
|
+
"reduction/sumF64": sumF64,
|
|
91
|
+
};
|
|
92
|
+
|
|
93
|
+
import ddGreater from "./f64/utils/greater.wgsl";
|
|
94
|
+
import ddEqual from "./f64/utils/equal.wgsl";
|
|
95
|
+
import idamax from "./idamax.wgsl";
|
|
96
|
+
import argmaxF64 from "./reduction/argmaxF64.wgsl";
|
|
97
|
+
routineShaders.idamax = {
|
|
98
|
+
"f64/dekker": dekker,
|
|
99
|
+
"f64/utils/abs": ddAbs,
|
|
100
|
+
"f64/utils/greater": ddGreater,
|
|
101
|
+
"f64/utils/equal": ddEqual,
|
|
102
|
+
idamax,
|
|
103
|
+
"reduction/argmaxF64": argmaxF64,
|
|
104
|
+
};
|
|
105
|
+
|
|
106
|
+
import srot from "./srot.wgsl";
|
|
107
|
+
routineShaders.srot = { srot };
|
|
108
|
+
|
|
109
|
+
import srotm from "./srotm.wgsl";
|
|
110
|
+
routineShaders.srotm = { srotm };
|
|
111
|
+
|
|
112
|
+
import sgemv_n from "./sgemv_n.wgsl";
|
|
113
|
+
import sgemv_t from "./sgemv_t.wgsl";
|
|
114
|
+
routineShaders.sgemv = { sgemv_n, sgemv_t }; // one or the other, picked by trans
|
|
115
|
+
|
|
116
|
+
import ssymv from "./ssymv.wgsl";
|
|
117
|
+
routineShaders.ssymv = { ssymv };
|
|
118
|
+
|
|
119
|
+
import strmv from "./strmv.wgsl";
|
|
120
|
+
routineShaders.strmv = { strmv };
|
|
121
|
+
|
|
122
|
+
import strsv_invert_block from "./strsv_invert_block.wgsl";
|
|
123
|
+
import strsv_apply_inverse from "./strsv_apply_inverse.wgsl";
|
|
124
|
+
import strsv_update from "./strsv_update.wgsl";
|
|
125
|
+
routineShaders.strsv = {
|
|
126
|
+
strsv_invert_block,
|
|
127
|
+
strsv_apply_inverse,
|
|
128
|
+
strsv_update,
|
|
129
|
+
};
|
|
130
|
+
|
|
131
|
+
import sger from "./sger.wgsl";
|
|
132
|
+
routineShaders.sger = { sger };
|
|
133
|
+
|
|
134
|
+
import ssyr from "./ssyr.wgsl";
|
|
135
|
+
routineShaders.ssyr = { ssyr };
|
|
136
|
+
|
|
137
|
+
import ssyr2 from "./ssyr2.wgsl";
|
|
138
|
+
routineShaders.ssyr2 = { ssyr2 };
|
|
139
|
+
|
|
140
|
+
import sgemm_small from "./sgemm_small.wgsl";
|
|
141
|
+
import sgemm_large from "./sgemm_large.wgsl";
|
|
142
|
+
routineShaders.sgemm = { sgemm_small, sgemm_large }; // one or the other, picked by a tile-size threshold
|
|
143
|
+
|
|
144
|
+
import sgemmtr_small from "./sgemmtr_small.wgsl";
|
|
145
|
+
import sgemmtr_large from "./sgemmtr_large.wgsl";
|
|
146
|
+
routineShaders.sgemmtr = { sgemmtr_small, sgemmtr_large };
|
|
147
|
+
|
|
148
|
+
routineShaders.ssyrk = { sgemmtr_small, sgemmtr_large }; // no shader of its own — rides on sgemmtr's
|
|
149
|
+
routineShaders.ssyr2k = { sgemmtr_small, sgemmtr_large }; // no shader of its own — rides on sgemmtr's
|
|
150
|
+
|
|
151
|
+
import symmetrize from "./symmetrize.wgsl";
|
|
152
|
+
routineShaders.ssymm = { sgemm_small, sgemm_large, symmetrize };
|
|
153
|
+
|
|
154
|
+
import triangularize from "./triangularize.wgsl";
|
|
155
|
+
routineShaders.strmm = { sgemm_small, sgemm_large, triangularize };
|
|
156
|
+
|
|
157
|
+
import blockTransfer from "./block_transfer.wgsl";
|
|
158
|
+
routineShaders.strsm = {
|
|
159
|
+
strsv_invert_block,
|
|
160
|
+
block_transfer: blockTransfer,
|
|
161
|
+
sscal,
|
|
162
|
+
sgemm_small,
|
|
163
|
+
sgemm_large,
|
|
164
|
+
};
|
|
165
|
+
|
|
166
|
+
/**
|
|
167
|
+
* Flat shader-name → WGSL source-string registry — what `getPipeline()`/`loadShader()` (see
|
|
168
|
+
* `util/pipeline.mjs`) actually look shaders up in, in the browser. Derived from
|
|
169
|
+
* `routineShaders` by merging every routine's shaders together; shared shaders (e.g.
|
|
170
|
+
* `"reduction/sum"`, used by two different routines above) collapse harmlessly here since
|
|
171
|
+
* every routine's copy is the same imported string, never independently authored text.
|
|
172
|
+
* @public
|
|
173
|
+
*/
|
|
174
|
+
export const shaderSources = Object.assign(
|
|
175
|
+
{},
|
|
176
|
+
...Object.values(routineShaders),
|
|
177
|
+
);
|
|
@@ -0,0 +1,50 @@
|
|
|
1
|
+
// amax reduction (f64, double-double): collapses 2*WGS (value, index) pairs
|
|
2
|
+
// into one index, using ddGreater/ddEqual instead of plain f32 `>`/`==` (see
|
|
3
|
+
// reduction/argmax.wgsl for the f32 original this mirrors).
|
|
4
|
+
// dispatch: 1 workgroup of WGS threads. partialsValHi/partialsValLo/
|
|
5
|
+
// partialsIdx must have exactly 2*WGS entries each. Concatenated after
|
|
6
|
+
// f64/dekker.wgsl (DD struct), f64/utils/greater.wgsl (ddGreater), and
|
|
7
|
+
// f64/utils/equal.wgsl (ddEqual).
|
|
8
|
+
|
|
9
|
+
@group(0) @binding(0) var<storage, read> partialsValHi: array<f32>;
|
|
10
|
+
@group(0) @binding(1) var<storage, read> partialsValLo: array<f32>;
|
|
11
|
+
@group(0) @binding(2) var<storage, read> partialsIdx: array<u32>;
|
|
12
|
+
@group(0) @binding(3) var<storage, read_write> result: array<u32>;
|
|
13
|
+
|
|
14
|
+
const WGS: u32 = 64;
|
|
15
|
+
|
|
16
|
+
var<workgroup> tile_val: array<DD, 64>;
|
|
17
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
18
|
+
|
|
19
|
+
@compute @workgroup_size(64)
|
|
20
|
+
fn reduce_f64(
|
|
21
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
22
|
+
) {
|
|
23
|
+
let i = lid.x;
|
|
24
|
+
let a_val = DD(partialsValHi[i], partialsValLo[i]);
|
|
25
|
+
let b_val = DD(partialsValHi[i + WGS], partialsValLo[i + WGS]);
|
|
26
|
+
if (ddGreater(b_val, a_val) ||
|
|
27
|
+
(ddEqual(b_val, a_val) && partialsIdx[i + WGS] < partialsIdx[i])) {
|
|
28
|
+
tile_val[i] = b_val;
|
|
29
|
+
tile_idx[i] = partialsIdx[i + WGS];
|
|
30
|
+
} else {
|
|
31
|
+
tile_val[i] = a_val;
|
|
32
|
+
tile_idx[i] = partialsIdx[i];
|
|
33
|
+
}
|
|
34
|
+
workgroupBarrier();
|
|
35
|
+
|
|
36
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
37
|
+
if (i < s) {
|
|
38
|
+
let c_val = tile_val[i];
|
|
39
|
+
let d_val = tile_val[i + s];
|
|
40
|
+
if (ddGreater(d_val, c_val) ||
|
|
41
|
+
(ddEqual(d_val, c_val) && tile_idx[i + s] < tile_idx[i])) {
|
|
42
|
+
tile_val[i] = d_val;
|
|
43
|
+
tile_idx[i] = tile_idx[i + s];
|
|
44
|
+
}
|
|
45
|
+
}
|
|
46
|
+
workgroupBarrier();
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
if (i == 0u) { result[0] = tile_idx[0]; }
|
|
50
|
+
}
|
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
// scaledSum reduction: collapses 2*WGS (scale, ssq) partials from
|
|
2
|
+
// snrm2.wgsl into the final norm — sqrt(scale² · ssq) == scale · sqrt(ssq).
|
|
3
|
+
// Mirrors reduction/sum.wgsl's shape exactly, merging via ssqMerge (see
|
|
4
|
+
// snrm2.wgsl for the derivation) instead of plain `+`, and taking the final
|
|
5
|
+
// sqrt here rather than on the CPU — unlike sasum/sdot's plain sum, "sum of
|
|
6
|
+
// squares" isn't a meaningful standalone value to hand back, only
|
|
7
|
+
// scale·sqrt(ssq) is.
|
|
8
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
9
|
+
// partialsScale/partialsSsq must have exactly 2*WGS entries each.
|
|
10
|
+
|
|
11
|
+
@group(0) @binding(0) var<storage, read> partialsScale: array<f32>;
|
|
12
|
+
@group(0) @binding(1) var<storage, read> partialsSsq: array<f32>;
|
|
13
|
+
@group(0) @binding(2) var<storage, read_write> result: array<f32>;
|
|
14
|
+
|
|
15
|
+
const WGS: u32 = 64;
|
|
16
|
+
|
|
17
|
+
// True sum-of-squares represented so far == scale² · ssq — see snrm2.wgsl.
|
|
18
|
+
struct ScaleSsq {
|
|
19
|
+
scale: f32,
|
|
20
|
+
ssq: f32,
|
|
21
|
+
}
|
|
22
|
+
|
|
23
|
+
// Associative merge of two independent (scale, ssq) partials.
|
|
24
|
+
fn ssqMerge(a: ScaleSsq, b: ScaleSsq) -> ScaleSsq {
|
|
25
|
+
if (a.scale == 0.0 && b.scale == 0.0) { return ScaleSsq(0.0, 1.0); }
|
|
26
|
+
if (a.scale >= b.scale) {
|
|
27
|
+
let r = b.scale / a.scale;
|
|
28
|
+
return ScaleSsq(a.scale, a.ssq + b.ssq * r * r);
|
|
29
|
+
}
|
|
30
|
+
let r = a.scale / b.scale;
|
|
31
|
+
return ScaleSsq(b.scale, b.ssq + a.ssq * r * r);
|
|
32
|
+
}
|
|
33
|
+
|
|
34
|
+
var<workgroup> tileScale: array<f32, 64>;
|
|
35
|
+
var<workgroup> tileSsq: array<f32, 64>;
|
|
36
|
+
|
|
37
|
+
@compute @workgroup_size(64)
|
|
38
|
+
fn reduce_scaled(
|
|
39
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
40
|
+
) {
|
|
41
|
+
let i = lid.x;
|
|
42
|
+
let merged0 = ssqMerge(
|
|
43
|
+
ScaleSsq(partialsScale[i], partialsSsq[i]),
|
|
44
|
+
ScaleSsq(partialsScale[i + WGS], partialsSsq[i + WGS]),
|
|
45
|
+
);
|
|
46
|
+
tileScale[i] = merged0.scale;
|
|
47
|
+
tileSsq[i] = merged0.ssq;
|
|
48
|
+
workgroupBarrier();
|
|
49
|
+
|
|
50
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
51
|
+
if (i < s) {
|
|
52
|
+
let merged = ssqMerge(
|
|
53
|
+
ScaleSsq(tileScale[i], tileSsq[i]),
|
|
54
|
+
ScaleSsq(tileScale[i + s], tileSsq[i + s]),
|
|
55
|
+
);
|
|
56
|
+
tileScale[i] = merged.scale;
|
|
57
|
+
tileSsq[i] = merged.ssq;
|
|
58
|
+
}
|
|
59
|
+
workgroupBarrier();
|
|
60
|
+
}
|
|
61
|
+
|
|
62
|
+
if (i == 0u) {
|
|
63
|
+
result[0] = tileScale[0] * sqrt(tileSsq[0]);
|
|
64
|
+
}
|
|
65
|
+
}
|
|
@@ -2,8 +2,8 @@
|
|
|
2
2
|
// into one, using ddAddProtected instead of plain f32 `+` (see
|
|
3
3
|
// reduction/sum.wgsl for the f32 original this mirrors).
|
|
4
4
|
// dispatch: 1 workgroup of WGS threads. partialsHi/partialsLo must have
|
|
5
|
-
// exactly 2*WGS entries each. Concatenated after f64/dekker.wgsl
|
|
6
|
-
//
|
|
5
|
+
// exactly 2*WGS entries each. Concatenated after f64/dekker.wgsl (DD struct)
|
|
6
|
+
// and f64/utils/add.wgsl (ddAddProtected — see it for why plain ddAdd isn't safe).
|
|
7
7
|
|
|
8
8
|
@group(0) @binding(0) var<storage, read> partialsHi: array<f32>;
|
|
9
9
|
@group(0) @binding(1) var<storage, read> partialsLo: array<f32>;
|