wgblas 1.2.1 → 2.0.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 +36 -54
- package/dist/wgblas.browser.js +779 -34
- package/index.d.mts +13 -0
- package/index.mjs +8 -0
- package/package.json +55 -2
- package/src/dasum/dasum.d.mts +2 -2
- package/src/dasum/dasum.mjs +6 -4
- package/src/idamax/idamax.d.mts +51 -0
- package/src/idamax/idamax.mjs +128 -0
- package/src/init.mjs +9 -1
- package/src/isamax/isamax.d.mts +1 -1
- package/src/sasum/sasum.d.mts +1 -1
- package/src/saxpy/saxpy.d.mts +1 -1
- package/src/scopy/scopy.d.mts +1 -1
- package/src/sdot/sdot.d.mts +1 -1
- package/src/sgemm/sgemm.d.mts +102 -0
- package/src/sgemm/sgemm.mjs +195 -0
- package/src/sgemmtr/sgemmtr.d.mts +104 -0
- package/src/sgemmtr/sgemmtr.mjs +203 -0
- package/src/sgemv/sgemv.d.mts +1 -38
- package/src/sgemv/sgemv.mjs +4 -0
- package/src/sger/sger.d.mts +1 -34
- package/src/sger/sger.mjs +2 -0
- package/src/shaders/block_transfer.wgsl +42 -0
- package/src/shaders/browser-shaders.mjs +26 -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/reduction/argmaxF64.wgsl +50 -0
- package/src/shaders/reduction/sumF64.wgsl +2 -2
- package/src/shaders/sgemm_large.wgsl +117 -0
- package/src/shaders/sgemm_small.wgsl +112 -0
- package/src/shaders/sgemmtr_large.wgsl +117 -0
- package/src/shaders/sgemmtr_small.wgsl +110 -0
- package/src/shaders/symmetrize.wgsl +31 -0
- package/src/shaders/triangularize.wgsl +44 -0
- package/src/snrm2/snrm2.d.mts +1 -1
- package/src/srot/srot.d.mts +1 -1
- package/src/srotm/srotm.d.mts +1 -1
- package/src/sscal/sscal.d.mts +1 -1
- package/src/sswap/sswap.d.mts +1 -1
- package/src/ssymm/ssymm.d.mts +103 -0
- package/src/ssymm/ssymm.mjs +209 -0
- package/src/ssymv/ssymv.d.mts +1 -36
- package/src/ssymv/ssymv.mjs +2 -0
- package/src/ssyr/ssyr.d.mts +1 -30
- package/src/ssyr/ssyr.mjs +2 -0
- package/src/ssyr2/ssyr2.d.mts +1 -34
- package/src/ssyr2/ssyr2.mjs +2 -0
- package/src/ssyr2k/ssyr2k.d.mts +100 -0
- package/src/ssyr2k/ssyr2k.mjs +201 -0
- package/src/ssyrk/ssyrk.d.mts +90 -0
- package/src/ssyrk/ssyrk.mjs +176 -0
- package/src/strmm/strmm.d.mts +100 -0
- package/src/strmm/strmm.mjs +211 -0
- package/src/strmv/strmv.d.mts +1 -36
- package/src/strmv/strmv.mjs +2 -0
- package/src/strsm/strsm.d.mts +99 -0
- package/src/strsm/strsm.mjs +342 -0
- package/src/strsv/strsv.d.mts +1 -32
- package/src/strsv/strsv.mjs +2 -0
- package/src/util/buffer.mjs +4 -2
- package/src/util/compute.mjs +6 -3
- package/src/util/f64.mjs +3 -3
|
@@ -1,4 +1,5 @@
|
|
|
1
1
|
import argmax from "./reduction/argmax.wgsl";
|
|
2
|
+
import argmaxF64 from "./reduction/argmaxF64.wgsl";
|
|
2
3
|
import sum from "./reduction/sum.wgsl";
|
|
3
4
|
import sumF64 from "./reduction/sumF64.wgsl";
|
|
4
5
|
import sscal from "./sscal.wgsl";
|
|
@@ -20,13 +21,26 @@ import ssyr from "./ssyr.wgsl";
|
|
|
20
21
|
import ssyr2 from "./ssyr2.wgsl";
|
|
21
22
|
import f64add from "./f64add.wgsl";
|
|
22
23
|
import dekker from "./f64/dekker.wgsl";
|
|
24
|
+
import ddAbs from "./f64/utils/abs.wgsl";
|
|
25
|
+
import ddAddUtil from "./f64/utils/add.wgsl";
|
|
26
|
+
import ddGreater from "./f64/utils/greater.wgsl";
|
|
27
|
+
import ddEqual from "./f64/utils/equal.wgsl";
|
|
23
28
|
import dasum from "./dasum.wgsl";
|
|
29
|
+
import idamax from "./idamax.wgsl";
|
|
24
30
|
import strsv_invert_block from "./strsv_invert_block.wgsl";
|
|
25
31
|
import strsv_apply_inverse from "./strsv_apply_inverse.wgsl";
|
|
26
32
|
import strsv_update from "./strsv_update.wgsl";
|
|
33
|
+
import sgemm_small from "./sgemm_small.wgsl";
|
|
34
|
+
import sgemm_large from "./sgemm_large.wgsl";
|
|
35
|
+
import sgemmtr_small from "./sgemmtr_small.wgsl";
|
|
36
|
+
import sgemmtr_large from "./sgemmtr_large.wgsl";
|
|
37
|
+
import symmetrize from "./symmetrize.wgsl";
|
|
38
|
+
import triangularize from "./triangularize.wgsl";
|
|
39
|
+
import blockTransfer from "./block_transfer.wgsl";
|
|
27
40
|
|
|
28
41
|
export const shaderSources = {
|
|
29
42
|
"reduction/argmax": argmax,
|
|
43
|
+
"reduction/argmaxF64": argmaxF64,
|
|
30
44
|
"reduction/sum": sum,
|
|
31
45
|
"reduction/sumF64": sumF64,
|
|
32
46
|
sscal,
|
|
@@ -48,8 +62,20 @@ export const shaderSources = {
|
|
|
48
62
|
ssyr2,
|
|
49
63
|
f64add,
|
|
50
64
|
"f64/dekker": dekker,
|
|
65
|
+
"f64/utils/abs": ddAbs,
|
|
66
|
+
"f64/utils/add": ddAddUtil,
|
|
67
|
+
"f64/utils/greater": ddGreater,
|
|
68
|
+
"f64/utils/equal": ddEqual,
|
|
51
69
|
dasum,
|
|
70
|
+
idamax,
|
|
52
71
|
strsv_invert_block,
|
|
53
72
|
strsv_apply_inverse,
|
|
54
73
|
strsv_update,
|
|
74
|
+
sgemm_small,
|
|
75
|
+
sgemm_large,
|
|
76
|
+
sgemmtr_small,
|
|
77
|
+
sgemmtr_large,
|
|
78
|
+
symmetrize,
|
|
79
|
+
triangularize,
|
|
80
|
+
block_transfer: blockTransfer,
|
|
55
81
|
};
|
package/src/shaders/dasum.wgsl
CHANGED
|
@@ -1,6 +1,7 @@
|
|
|
1
1
|
// dasum: sum(|x[i]|), double-double (Dekker). Same ILP=4 shape as sasum.wgsl;
|
|
2
|
-
// see f64/
|
|
3
|
-
// GpuVector input isn't pre-abs'd, so ddAbs() applies
|
|
2
|
+
// see f64/utils/add.wgsl for ddAddProtected and why plain ddAdd isn't safe.
|
|
3
|
+
// GpuVector input isn't pre-abs'd, so ddAbs() (f64/utils/abs.wgsl) applies
|
|
4
|
+
// unconditionally below.
|
|
4
5
|
|
|
5
6
|
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
6
7
|
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
@@ -7,93 +7,12 @@
|
|
|
7
7
|
//
|
|
8
8
|
// No bindings, no entry point — a helper library, concatenated with a
|
|
9
9
|
// consumer's own bindings/entry point by getPipeline (WGSL has no #include).
|
|
10
|
+
// The DD struct lives here — abs.wgsl/add.wgsl/greater.wgsl/equal.wgsl all
|
|
11
|
+
// use it but don't redefine it (WGSL errors on duplicate struct definitions
|
|
12
|
+
// once concatenated), so any consumer using those must concatenate this
|
|
13
|
+
// file too, first.
|
|
10
14
|
|
|
11
15
|
struct DD {
|
|
12
16
|
hi: f32,
|
|
13
17
|
lo: f32,
|
|
14
18
|
}
|
|
15
|
-
|
|
16
|
-
// |a| for a double-double pair. Negation is exact (no rounding), so this is
|
|
17
|
-
// just a sign flip on both components — hi alone determines the pair's sign.
|
|
18
|
-
fn ddAbs(a: DD) -> DD {
|
|
19
|
-
if (a.hi < 0.0) {
|
|
20
|
-
return DD(-a.hi, -a.lo);
|
|
21
|
-
}
|
|
22
|
-
return a;
|
|
23
|
-
}
|
|
24
|
-
|
|
25
|
-
// ── A real compiler bug — read before touching anything below ──────────────
|
|
26
|
-
//
|
|
27
|
-
// twoSum/fastTwoSum's error term `e` should be nonzero (that's the point —
|
|
28
|
-
// `s` is rounded). A front-end-optimizer bug zeros it anyway, on both NVIDIA
|
|
29
|
-
// and Mesa (ANV + llvmpipe), via two different mechanisms needing two fixes:
|
|
30
|
-
// bitcast-based subtraction (`fsub`/`negf`, fixes NVIDIA) and materializing
|
|
31
|
-
// the sum through workgroup memory + workgroupBarrier() (fixes Mesa). Only
|
|
32
|
-
// both together (ddAddProtected) is verified correct everywhere — the plain
|
|
33
|
-
// twoSum/fastTwoSum/ddAdd below are reference-only, not safe to use.
|
|
34
|
-
fn negf(x: f32) -> f32 {
|
|
35
|
-
return bitcast<f32>(bitcast<u32>(x) ^ 0x80000000u);
|
|
36
|
-
}
|
|
37
|
-
fn fsub(a: f32, b: f32) -> f32 {
|
|
38
|
-
return a + negf(b);
|
|
39
|
-
}
|
|
40
|
-
|
|
41
|
-
// Knuth/Møller's TwoSum: s = fl(a+b), e = exact rounding error, a+b == s+e.
|
|
42
|
-
// Works for any a, b. UNPROTECTED — see header above.
|
|
43
|
-
fn twoSum(a: f32, b: f32) -> DD {
|
|
44
|
-
let s = a + b;
|
|
45
|
-
let v = s - a;
|
|
46
|
-
let e = (a - (s - v)) + (b - v);
|
|
47
|
-
return DD(s, e);
|
|
48
|
-
}
|
|
49
|
-
|
|
50
|
-
// Dekker's Fast-Two-Sum: same contract, but only correct when |a| >= |b|.
|
|
51
|
-
// UNPROTECTED — see header above.
|
|
52
|
-
fn fastTwoSum(a: f32, b: f32) -> DD {
|
|
53
|
-
let s = a + b;
|
|
54
|
-
let e = b - (s - a);
|
|
55
|
-
return DD(s, e);
|
|
56
|
-
}
|
|
57
|
-
|
|
58
|
-
// Double-double addition (Dekker's Add2). UNPROTECTED — see header above.
|
|
59
|
-
fn ddAdd(a: DD, b: DD) -> DD {
|
|
60
|
-
let s = twoSum(a.hi, b.hi);
|
|
61
|
-
let loSum = a.lo + b.lo;
|
|
62
|
-
return fastTwoSum(s.hi, s.lo + loSum);
|
|
63
|
-
}
|
|
64
|
-
|
|
65
|
-
// ── Protected variants — use these ──────────────────────────────────────────
|
|
66
|
-
//
|
|
67
|
-
// Bitcast subtraction + workgroup-barrier materialization, verified correct
|
|
68
|
-
// on all three backends tested. Costs a real barrier: fine for O(1)-per-
|
|
69
|
-
// thread or O(log n) reduction use, not a long per-element loop. A
|
|
70
|
-
// workgroupBarrier() requires uniform control flow, so:
|
|
71
|
-
// - `threadSlot` must be unique per concurrent caller (e.g. local_invocation_index).
|
|
72
|
-
// - Every thread in the workgroup must call this the same number of times
|
|
73
|
-
// — including ones whose result gets discarded. Compute unconditionally;
|
|
74
|
-
// only the write-back should be conditional.
|
|
75
|
-
var<workgroup> dekkerScratch: array<f32, 64>;
|
|
76
|
-
|
|
77
|
-
fn twoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
|
|
78
|
-
dekkerScratch[threadSlot] = a + b;
|
|
79
|
-
workgroupBarrier();
|
|
80
|
-
let s = dekkerScratch[threadSlot];
|
|
81
|
-
let v = fsub(s, a);
|
|
82
|
-
let e = fsub(a, fsub(s, v)) + fsub(b, v);
|
|
83
|
-
return DD(s, e);
|
|
84
|
-
}
|
|
85
|
-
|
|
86
|
-
fn fastTwoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
|
|
87
|
-
dekkerScratch[threadSlot] = a + b;
|
|
88
|
-
workgroupBarrier();
|
|
89
|
-
let s = dekkerScratch[threadSlot];
|
|
90
|
-
let e = fsub(b, fsub(s, a));
|
|
91
|
-
return DD(s, e);
|
|
92
|
-
}
|
|
93
|
-
|
|
94
|
-
// Protected double-double addition — same contract as ddAdd, but exact.
|
|
95
|
-
fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
96
|
-
let s = twoSumProtected(a.hi, b.hi, threadSlot);
|
|
97
|
-
let loSum = a.lo + b.lo;
|
|
98
|
-
return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
|
|
99
|
-
}
|
|
@@ -0,0 +1,10 @@
|
|
|
1
|
+
// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
2
|
+
|
|
3
|
+
// |a| for a double-double pair. Negation is exact (no rounding), so this is
|
|
4
|
+
// just a sign flip on both components — hi alone determines the pair's sign.
|
|
5
|
+
fn ddAbs(a: DD) -> DD {
|
|
6
|
+
if (a.hi < 0.0) {
|
|
7
|
+
return DD(-a.hi, -a.lo);
|
|
8
|
+
}
|
|
9
|
+
return a;
|
|
10
|
+
}
|
|
@@ -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
|
+
}
|
|
@@ -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
|
+
}
|
|
@@ -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>;
|
|
@@ -0,0 +1,117 @@
|
|
|
1
|
+
// sgemm_large: C = alpha * op(A) * op(B) + beta * C — large-tile half of
|
|
2
|
+
// the two-tier autotuned dispatch (see sgemm.mjs and sgemm_small.wgsl).
|
|
3
|
+
// BM=BN=64, BK=8, TM=8, TN=4 (128 threads/workgroup) — the kernel 9
|
|
4
|
+
// autotuning winner (temp/autotune_sweep.mjs, temp/gen_sweep_kernel.mjs,
|
|
5
|
+
// swept BM/BN/BK/TM/TN and warp-tiled variants), +69% over the old BM=32
|
|
6
|
+
// single-tier baseline at n=512, +84% at n=1024. But BM=64 loses to BM=32
|
|
7
|
+
// below a 6x6=36 workgroup grid (not enough workgroups to fill the GPU at
|
|
8
|
+
// that tile size), hence the two-tier split rather than one global config.
|
|
9
|
+
// Neither vectorized loads (kernel 6) nor warp-tiling (kernel 10) beat this
|
|
10
|
+
// at the sizes tried, including warp-tiled variants in the same sweep at
|
|
11
|
+
// BM=64/128.
|
|
12
|
+
|
|
13
|
+
const BM: u32 = 64u;
|
|
14
|
+
const BN: u32 = 64u;
|
|
15
|
+
const BK: u32 = 8u;
|
|
16
|
+
const TM: u32 = 8u;
|
|
17
|
+
const TN: u32 = 4u;
|
|
18
|
+
const THREADS_X: u32 = BN / TN;
|
|
19
|
+
const THREADS_Y: u32 = BM / TM;
|
|
20
|
+
const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 128
|
|
21
|
+
const STRIDE_A: u32 = NUM_THREADS / BK; // rows of As covered per load-loop step
|
|
22
|
+
const STRIDE_B: u32 = NUM_THREADS / BN; // rows of Bs covered per load-loop step
|
|
23
|
+
|
|
24
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
25
|
+
@group(0) @binding(1) var<storage, read> B: array<f32>;
|
|
26
|
+
@group(0) @binding(2) var<storage, read_write> C: array<f32>;
|
|
27
|
+
|
|
28
|
+
struct Params {
|
|
29
|
+
m: u32,
|
|
30
|
+
n: u32,
|
|
31
|
+
k: u32,
|
|
32
|
+
alpha: f32,
|
|
33
|
+
beta: f32,
|
|
34
|
+
lda: u32,
|
|
35
|
+
ldb: u32,
|
|
36
|
+
ldc: u32,
|
|
37
|
+
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
38
|
+
transB: u32,
|
|
39
|
+
}
|
|
40
|
+
|
|
41
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
42
|
+
|
|
43
|
+
var<workgroup> As: array<f32, BM * BK>;
|
|
44
|
+
var<workgroup> Bs: array<f32, BK * BN>;
|
|
45
|
+
|
|
46
|
+
@compute @workgroup_size(THREADS_X, THREADS_Y)
|
|
47
|
+
fn main(
|
|
48
|
+
@builtin(workgroup_id) wid: vec3u,
|
|
49
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
50
|
+
@builtin(local_invocation_index) tid: u32,
|
|
51
|
+
) {
|
|
52
|
+
let blockRow = wid.y * BM;
|
|
53
|
+
let blockCol = wid.x * BN;
|
|
54
|
+
let threadCol = lid.x;
|
|
55
|
+
let threadRow = lid.y;
|
|
56
|
+
|
|
57
|
+
// Load indices, independent of the compute thread shape — a loop since
|
|
58
|
+
// NUM_THREADS doesn't match the tile size 1:1 at this config.
|
|
59
|
+
let innerRowA = tid / BK;
|
|
60
|
+
let innerColA = tid % BK;
|
|
61
|
+
let innerRowB = tid / BN;
|
|
62
|
+
let innerColB = tid % BN;
|
|
63
|
+
|
|
64
|
+
var threadResults: array<f32, TM * TN>;
|
|
65
|
+
for (var i = 0u; i < TM * TN; i++) {
|
|
66
|
+
threadResults[i] = 0.0;
|
|
67
|
+
}
|
|
68
|
+
var regM: array<f32, TM>;
|
|
69
|
+
var regN: array<f32, TN>;
|
|
70
|
+
|
|
71
|
+
let numTiles = (params.k + BK - 1u) / BK;
|
|
72
|
+
for (var t = 0u; t < numTiles; t++) {
|
|
73
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
74
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
75
|
+
let gColA = t * BK + innerColA;
|
|
76
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
77
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
78
|
+
}
|
|
79
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
80
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
81
|
+
let gColB = blockCol + innerColB;
|
|
82
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
83
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
84
|
+
}
|
|
85
|
+
|
|
86
|
+
workgroupBarrier();
|
|
87
|
+
|
|
88
|
+
for (var dotIdx = 0u; dotIdx < BK; dotIdx++) {
|
|
89
|
+
for (var i = 0u; i < TM; i++) {
|
|
90
|
+
regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
|
|
91
|
+
}
|
|
92
|
+
for (var i = 0u; i < TN; i++) {
|
|
93
|
+
regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
|
|
94
|
+
}
|
|
95
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
96
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
97
|
+
threadResults[resIdxM * TN + resIdxN] += regM[resIdxM] * regN[resIdxN];
|
|
98
|
+
}
|
|
99
|
+
}
|
|
100
|
+
}
|
|
101
|
+
|
|
102
|
+
workgroupBarrier();
|
|
103
|
+
}
|
|
104
|
+
|
|
105
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
106
|
+
let row = blockRow + threadRow * TM + resIdxM;
|
|
107
|
+
if (row < params.m) {
|
|
108
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
109
|
+
let col = blockCol + threadCol * TN + resIdxN;
|
|
110
|
+
if (col < params.n) {
|
|
111
|
+
let cIdx = row * params.ldc + col;
|
|
112
|
+
C[cIdx] = params.alpha * threadResults[resIdxM * TN + resIdxN] + params.beta * C[cIdx];
|
|
113
|
+
}
|
|
114
|
+
}
|
|
115
|
+
}
|
|
116
|
+
}
|
|
117
|
+
}
|