wgblas 2.2.0 → 2.3.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 +5 -2
- package/dist/wgblas.browser.js +1092 -58
- package/index.d.mts +7 -0
- package/index.mjs +7 -0
- package/package.json +40 -4
- package/src/dgemv/dgemv.d.mts +92 -0
- package/src/dgemv/dgemv.mjs +247 -0
- package/src/dger/dger.d.mts +80 -0
- package/src/dger/dger.mjs +213 -0
- package/src/dsymv/dsymv.d.mts +84 -0
- package/src/dsymv/dsymv.mjs +227 -0
- package/src/dsyr/dsyr.d.mts +73 -0
- package/src/dsyr/dsyr.mjs +171 -0
- package/src/dsyr2/dsyr2.d.mts +81 -0
- package/src/dsyr2/dsyr2.mjs +214 -0
- package/src/dtrmv/dtrmv.d.mts +84 -0
- package/src/dtrmv/dtrmv.mjs +219 -0
- package/src/dtrsv/dtrsv.d.mts +79 -0
- package/src/dtrsv/dtrsv.mjs +331 -0
- package/src/shaders/dgemv_n.wgsl +112 -0
- package/src/shaders/dgemv_t.wgsl +102 -0
- package/src/shaders/dger.wgsl +77 -0
- package/src/shaders/dsymv.wgsl +125 -0
- package/src/shaders/dsyr.wgsl +84 -0
- package/src/shaders/dsyr2.wgsl +95 -0
- package/src/shaders/dtrmv.wgsl +111 -0
- package/src/shaders/dtrsv_apply_inverse.wgsl +60 -0
- package/src/shaders/dtrsv_invert_block.wgsl +148 -0
- package/src/shaders/dtrsv_update.wgsl +111 -0
- package/src/shaders/f64/utils/multiply.wgsl +10 -1
- package/src/shaders/index.mjs +63 -0
|
@@ -0,0 +1,111 @@
|
|
|
1
|
+
// dtrmv: y := op(A) * x, double-double (Dekker) f64 emulation of strmv
|
|
2
|
+
// (triangular matrix-vector product). A, x, and y are each split into an f32
|
|
3
|
+
// (hi, lo) pair; WGSL has no f64 type. A is n×n triangular — only the
|
|
4
|
+
// triangle specified by uplo is referenced. Like strmv.wgsl, wgblas's trmv
|
|
5
|
+
// takes a separate output vector y rather than overwriting x in place (the
|
|
6
|
+
// standard BLAS signature is in-place), so there is no aliasing/read-write-
|
|
7
|
+
// ordering concern the way real in-place trmv would have — this is
|
|
8
|
+
// structurally just dsymv.wgsl's per-row DD tree reduction with a triangular
|
|
9
|
+
// (not full-n, not mirrored) column range per row, same shape as dsyr.wgsl's
|
|
10
|
+
// per-row-varying range (main/tail split computed inside the row loop, since
|
|
11
|
+
// range length depends on i).
|
|
12
|
+
|
|
13
|
+
@group(0) @binding(0) var<storage, read> AHi: array<f32>;
|
|
14
|
+
@group(0) @binding(1) var<storage, read> ALo: array<f32>;
|
|
15
|
+
@group(0) @binding(2) var<storage, read> xHi: array<f32>;
|
|
16
|
+
@group(0) @binding(3) var<storage, read> xLo: array<f32>;
|
|
17
|
+
@group(0) @binding(4) var<storage, read_write> yHi: array<f32>;
|
|
18
|
+
@group(0) @binding(5) var<storage, read_write> yLo: array<f32>;
|
|
19
|
+
|
|
20
|
+
struct Params {
|
|
21
|
+
n: u32,
|
|
22
|
+
incx: u32,
|
|
23
|
+
incy: u32,
|
|
24
|
+
lda: u32,
|
|
25
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
26
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
27
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
28
|
+
}
|
|
29
|
+
|
|
30
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
31
|
+
|
|
32
|
+
const WGS: u32 = 64u;
|
|
33
|
+
var<workgroup> tile: array<DD, 64>;
|
|
34
|
+
|
|
35
|
+
@compute @workgroup_size(64)
|
|
36
|
+
fn dtrmv_main(
|
|
37
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
38
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
39
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
40
|
+
) {
|
|
41
|
+
// Effective "upper" range for row i: op(A)[i,j] is stored at A[i,j] for
|
|
42
|
+
// no-transpose+upper or transpose+lower (range [i, n)); at A[j,i] for
|
|
43
|
+
// no-transpose+lower or transpose+upper (range [0, i]).
|
|
44
|
+
let isUpperRange = (params.trans == 0u) == (params.uplo == 1u);
|
|
45
|
+
|
|
46
|
+
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
47
|
+
let rangeStart = select(0u, i, isUpperRange);
|
|
48
|
+
let rangeEnd = select(i + 1u, params.n, isUpperRange);
|
|
49
|
+
|
|
50
|
+
// Range length varies per row, but every thread in the workgroup runs
|
|
51
|
+
// the SAME row at the same time, so the main/tail split computed from
|
|
52
|
+
// rangeStart/rangeEnd is already uniform across threads for this row —
|
|
53
|
+
// just not the same across different rows (see dsyr.wgsl).
|
|
54
|
+
let rangeLen = rangeEnd - rangeStart;
|
|
55
|
+
let range_floor = rangeStart + (rangeLen / WGS) * WGS;
|
|
56
|
+
let mainIters = (range_floor - rangeStart) / WGS;
|
|
57
|
+
let hasTail = select(0u, 1u, range_floor < rangeEnd);
|
|
58
|
+
|
|
59
|
+
var acc = DD(0.0, 0.0);
|
|
60
|
+
|
|
61
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
62
|
+
let j = rangeStart + lid.x + iter * WGS;
|
|
63
|
+
let addr = select(j * params.lda + i, i * params.lda + j, params.trans == 0u);
|
|
64
|
+
let isUnitDiag = params.diag == 1u && j == i;
|
|
65
|
+
let aVal = DD(
|
|
66
|
+
select(AHi[addr], 1.0, isUnitDiag),
|
|
67
|
+
select(ALo[addr], 0.0, isUnitDiag),
|
|
68
|
+
);
|
|
69
|
+
let ix = j * params.incx;
|
|
70
|
+
let prod = ddMulProtected(aVal, DD(xHi[ix], xLo[ix]), lid.x);
|
|
71
|
+
acc = ddAddProtected(acc, prod, lid.x);
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
for (var iter = 0u; iter < hasTail; iter++) {
|
|
75
|
+
let j = range_floor + lid.x;
|
|
76
|
+
let valid = j < rangeEnd;
|
|
77
|
+
let safeJ = select(rangeStart, j, valid); // rangeLen is always >= 1, so rangeStart is a safe in-range fallback
|
|
78
|
+
let addr = select(safeJ * params.lda + i, i * params.lda + safeJ, params.trans == 0u);
|
|
79
|
+
let isUnitDiag = params.diag == 1u && safeJ == i;
|
|
80
|
+
let aVal = DD(
|
|
81
|
+
select(AHi[addr], 1.0, isUnitDiag),
|
|
82
|
+
select(ALo[addr], 0.0, isUnitDiag),
|
|
83
|
+
);
|
|
84
|
+
let ix = safeJ * params.incx;
|
|
85
|
+
let prod = ddMulProtected(aVal, DD(xHi[ix], xLo[ix]), lid.x);
|
|
86
|
+
let contribution = DD(select(0.0, prod.hi, valid), select(0.0, prod.lo, valid));
|
|
87
|
+
acc = ddAddProtected(acc, contribution, lid.x);
|
|
88
|
+
}
|
|
89
|
+
|
|
90
|
+
tile[lid.x] = acc;
|
|
91
|
+
workgroupBarrier();
|
|
92
|
+
|
|
93
|
+
// Inactive threads combine against a throwaway partner and discard it
|
|
94
|
+
// (ddAddProtected must be called unconditionally by every thread).
|
|
95
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
96
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
97
|
+
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
98
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
99
|
+
if (lid.x < s) { tile[lid.x] = combined; }
|
|
100
|
+
workgroupBarrier();
|
|
101
|
+
}
|
|
102
|
+
|
|
103
|
+
if (lid.x == 0u) {
|
|
104
|
+
let iy = i * params.incy;
|
|
105
|
+
yHi[iy] = tile[0].hi;
|
|
106
|
+
yLo[iy] = tile[0].lo;
|
|
107
|
+
}
|
|
108
|
+
// All 64 threads must agree before the next row reuses tile[].
|
|
109
|
+
workgroupBarrier();
|
|
110
|
+
}
|
|
111
|
+
}
|
|
@@ -0,0 +1,60 @@
|
|
|
1
|
+
// dtrsv_apply_inverse: double-double (Dekker) f64 emulation of
|
|
2
|
+
// strsv_apply_inverse.wgsl — given a precomputed block inverse (from
|
|
3
|
+
// dtrsv_invert_block.wgsl), computes this block's solution as a dense
|
|
4
|
+
// matrix-vector multiply against the block's current remainder in x.
|
|
5
|
+
// Ainv, x, and the per-row accumulation are each split into an f32 (hi, lo)
|
|
6
|
+
// pair; WGSL has no f64 type.
|
|
7
|
+
//
|
|
8
|
+
// Unlike dsymv/dtrmv's per-row reduction across 64 threads, each thread here
|
|
9
|
+
// owns one whole row's dot product independently (blockLen <= WGS, so one
|
|
10
|
+
// thread per row already covers it, no tree reduction needed) — same shape
|
|
11
|
+
// as dgemv_t.wgsl's per-thread accumulation. The f32 original early-returns
|
|
12
|
+
// threads with `lid.x >= blockLen` (the last, possibly-short block); DD
|
|
13
|
+
// can't do that, since every thread must call ddMulProtected/ddAddProtected
|
|
14
|
+
// the same number of times. Instead every thread runs the identical
|
|
15
|
+
// `blockLen`-iteration loop (reading past-blockLen rows of Ainv, which the
|
|
16
|
+
// zero-initialized, never-written buffer backing them makes safe garbage —
|
|
17
|
+
// see dgemv_t.wgsl's clamped-dummy-index technique) and only the final
|
|
18
|
+
// write is gated.
|
|
19
|
+
|
|
20
|
+
@group(0) @binding(0) var<storage, read> AinvHi: array<f32>;
|
|
21
|
+
@group(0) @binding(1) var<storage, read> AinvLo: array<f32>;
|
|
22
|
+
@group(0) @binding(2) var<storage, read_write> xHi: array<f32>;
|
|
23
|
+
@group(0) @binding(3) var<storage, read_write> xLo: array<f32>;
|
|
24
|
+
|
|
25
|
+
struct Params {
|
|
26
|
+
incx: u32,
|
|
27
|
+
blockIndex: u32,
|
|
28
|
+
blockStart: u32,
|
|
29
|
+
blockEnd: u32,
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
33
|
+
|
|
34
|
+
const BLOCK_SIZE: u32 = 64u;
|
|
35
|
+
var<workgroup> xLocal: array<DD, 64>;
|
|
36
|
+
|
|
37
|
+
@compute @workgroup_size(64)
|
|
38
|
+
fn dtrsv_apply_inverse_main(@builtin(local_invocation_id) lid: vec3u) {
|
|
39
|
+
let blockLen = params.blockEnd - params.blockStart;
|
|
40
|
+
|
|
41
|
+
if (lid.x < blockLen) {
|
|
42
|
+
let ix = (params.blockStart + lid.x) * params.incx;
|
|
43
|
+
xLocal[lid.x] = DD(xHi[ix], xLo[ix]);
|
|
44
|
+
}
|
|
45
|
+
workgroupBarrier();
|
|
46
|
+
|
|
47
|
+
let ainvBase = params.blockIndex * BLOCK_SIZE * BLOCK_SIZE;
|
|
48
|
+
var acc = DD(0.0, 0.0);
|
|
49
|
+
for (var j = 0u; j < blockLen; j++) {
|
|
50
|
+
let aidx = ainvBase + lid.x * BLOCK_SIZE + j;
|
|
51
|
+
let prod = ddMulProtected(DD(AinvHi[aidx], AinvLo[aidx]), xLocal[j], lid.x);
|
|
52
|
+
acc = ddAddProtected(acc, prod, lid.x);
|
|
53
|
+
}
|
|
54
|
+
|
|
55
|
+
if (lid.x < blockLen) {
|
|
56
|
+
let ix = (params.blockStart + lid.x) * params.incx;
|
|
57
|
+
xHi[ix] = acc.hi;
|
|
58
|
+
xLo[ix] = acc.lo;
|
|
59
|
+
}
|
|
60
|
+
}
|
|
@@ -0,0 +1,148 @@
|
|
|
1
|
+
// dtrsv_invert_block: double-double (Dekker) f64 emulation of
|
|
2
|
+
// strsv_invert_block.wgsl — computes ONE column (workgroup_id.x) of ONE
|
|
3
|
+
// block's (workgroup_id.y) explicit inverse, via the same one-row-at-a-time
|
|
4
|
+
// substitution, solving against a unit basis vector e_col. A, Ainv, and the
|
|
5
|
+
// running accumulation are each split into an f32 (hi, lo) pair; WGSL has no
|
|
6
|
+
// f64 type.
|
|
7
|
+
//
|
|
8
|
+
// This is the first routine in the f64 port order to use ddDivProtected (the
|
|
9
|
+
// non-unit-diagonal division) — deliberately last per TODO.md, since
|
|
10
|
+
// division is where dnrm2 found real double-double edge-case bugs during
|
|
11
|
+
// the L1 port. The division only runs once per (column, step) — same
|
|
12
|
+
// barrier-uniformity requirement as every other protected op here — so
|
|
13
|
+
// every thread in the workgroup computes it redundantly from the
|
|
14
|
+
// already-shared `tile[0]` reduction result (same pattern dgemv_n uses for
|
|
15
|
+
// its O(1) alpha/beta combine after the per-row reduction), with only
|
|
16
|
+
// `lid.x==0` writing the result. `params.diag` is uniform across the whole
|
|
17
|
+
// dispatch (not per-thread), so branching on it to skip the division
|
|
18
|
+
// entirely for a unit diagonal is safe — unlike dtrmv.wgsl's per-element
|
|
19
|
+
// diag check, which varies per thread and needed select() instead.
|
|
20
|
+
|
|
21
|
+
@group(0) @binding(0) var<storage, read> AHi: array<f32>;
|
|
22
|
+
@group(0) @binding(1) var<storage, read> ALo: array<f32>;
|
|
23
|
+
@group(0) @binding(2) var<storage, read_write> AinvHi: array<f32>;
|
|
24
|
+
@group(0) @binding(3) var<storage, read_write> AinvLo: array<f32>;
|
|
25
|
+
|
|
26
|
+
struct Params {
|
|
27
|
+
n: u32,
|
|
28
|
+
lda: u32,
|
|
29
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
30
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
31
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
32
|
+
}
|
|
33
|
+
|
|
34
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
35
|
+
|
|
36
|
+
const WGS: u32 = 64u;
|
|
37
|
+
const BLOCK_SIZE: u32 = 64u;
|
|
38
|
+
var<workgroup> tile: array<DD, 64>;
|
|
39
|
+
|
|
40
|
+
fn readA(i: u32, j: u32) -> DD {
|
|
41
|
+
let idx = select(j * params.lda + i, i * params.lda + j, params.trans == 0u);
|
|
42
|
+
return DD(AHi[idx], ALo[idx]);
|
|
43
|
+
}
|
|
44
|
+
|
|
45
|
+
@compute @workgroup_size(64)
|
|
46
|
+
fn dtrsv_invert_block_main(
|
|
47
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
48
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
49
|
+
) {
|
|
50
|
+
let col = wgid.x;
|
|
51
|
+
let blockIndex = wgid.y;
|
|
52
|
+
let blockStart = blockIndex * BLOCK_SIZE;
|
|
53
|
+
var blockEnd = blockStart + BLOCK_SIZE;
|
|
54
|
+
if (blockEnd > params.n) { blockEnd = params.n; }
|
|
55
|
+
let blockLen = blockEnd - blockStart;
|
|
56
|
+
|
|
57
|
+
if (col >= blockLen) { return; }
|
|
58
|
+
|
|
59
|
+
let ainvBase = blockIndex * BLOCK_SIZE * BLOCK_SIZE;
|
|
60
|
+
let forward = (params.trans == 0u) == (params.uplo == 0u);
|
|
61
|
+
|
|
62
|
+
if forward {
|
|
63
|
+
for (var r = lid.x; r < col; r += WGS) {
|
|
64
|
+
AinvHi[ainvBase + r * BLOCK_SIZE + col] = 0.0;
|
|
65
|
+
AinvLo[ainvBase + r * BLOCK_SIZE + col] = 0.0;
|
|
66
|
+
}
|
|
67
|
+
} else {
|
|
68
|
+
for (var r = col + 1u + lid.x; r < blockLen; r += WGS) {
|
|
69
|
+
AinvHi[ainvBase + r * BLOCK_SIZE + col] = 0.0;
|
|
70
|
+
AinvLo[ainvBase + r * BLOCK_SIZE + col] = 0.0;
|
|
71
|
+
}
|
|
72
|
+
}
|
|
73
|
+
storageBarrier();
|
|
74
|
+
workgroupBarrier();
|
|
75
|
+
|
|
76
|
+
let numSteps = select(col + 1u, blockLen - col, forward);
|
|
77
|
+
for (var step = 0u; step < numSteps; step++) {
|
|
78
|
+
let localRow = select(col - step, col + step, forward);
|
|
79
|
+
let i = blockStart + localRow;
|
|
80
|
+
|
|
81
|
+
// Accumulation range over lj: [col, localRow) forward, [localRow+1, col+1) backward.
|
|
82
|
+
// Range length varies per step, but every thread runs the SAME step at
|
|
83
|
+
// the same time, so the main/tail split is already uniform for this
|
|
84
|
+
// step — just not the same across different steps (see dsyr.wgsl).
|
|
85
|
+
let ljStart = select(localRow + 1u, col, forward);
|
|
86
|
+
let ljEnd = select(col + 1u, localRow, forward);
|
|
87
|
+
let rangeLen = ljEnd - ljStart;
|
|
88
|
+
let range_floor = ljStart + (rangeLen / WGS) * WGS;
|
|
89
|
+
let mainIters = (range_floor - ljStart) / WGS;
|
|
90
|
+
let hasTail = select(0u, 1u, range_floor < ljEnd);
|
|
91
|
+
|
|
92
|
+
var acc = DD(0.0, 0.0);
|
|
93
|
+
|
|
94
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
95
|
+
let lj = ljStart + lid.x + iter * WGS;
|
|
96
|
+
let aij = readA(i, blockStart + lj);
|
|
97
|
+
let aidx = ainvBase + lj * BLOCK_SIZE + col;
|
|
98
|
+
let ainv = DD(AinvHi[aidx], AinvLo[aidx]);
|
|
99
|
+
let prod = ddMulProtected(aij, ainv, lid.x);
|
|
100
|
+
acc = ddAddProtected(acc, prod, lid.x);
|
|
101
|
+
}
|
|
102
|
+
|
|
103
|
+
for (var iter = 0u; iter < hasTail; iter++) {
|
|
104
|
+
let lj = range_floor + lid.x;
|
|
105
|
+
let valid = lj < ljEnd;
|
|
106
|
+
let safeLj = select(ljStart, lj, valid); // ljStart is always in-range when rangeLen > 0, the only case hasTail can be 1
|
|
107
|
+
let aij = readA(i, blockStart + safeLj);
|
|
108
|
+
let aidx = ainvBase + safeLj * BLOCK_SIZE + col;
|
|
109
|
+
let ainv = DD(AinvHi[aidx], AinvLo[aidx]);
|
|
110
|
+
let prod = ddMulProtected(aij, ainv, lid.x);
|
|
111
|
+
let contribution = DD(select(0.0, prod.hi, valid), select(0.0, prod.lo, valid));
|
|
112
|
+
acc = ddAddProtected(acc, contribution, lid.x);
|
|
113
|
+
}
|
|
114
|
+
|
|
115
|
+
tile[lid.x] = acc;
|
|
116
|
+
workgroupBarrier();
|
|
117
|
+
|
|
118
|
+
// Inactive threads combine against a throwaway partner and discard it
|
|
119
|
+
// (ddAddProtected must be called unconditionally by every thread).
|
|
120
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
121
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
122
|
+
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
123
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
124
|
+
if (lid.x < s) { tile[lid.x] = combined; }
|
|
125
|
+
workgroupBarrier();
|
|
126
|
+
}
|
|
127
|
+
|
|
128
|
+
// tile[0] now holds the full accumulation, visible to every lane — every
|
|
129
|
+
// thread redundantly finishes the O(1) rhs/division so the protected
|
|
130
|
+
// ops below stay uniformly called (only the write is gated).
|
|
131
|
+
let e = DD(select(0.0, 1.0, localRow == col), 0.0);
|
|
132
|
+
let rhs = ddSubProtected(e, tile[0], lid.x);
|
|
133
|
+
|
|
134
|
+
var val: DD;
|
|
135
|
+
if params.diag == 1u {
|
|
136
|
+
val = rhs;
|
|
137
|
+
} else {
|
|
138
|
+
val = ddDivProtected(rhs, readA(i, i), lid.x);
|
|
139
|
+
}
|
|
140
|
+
|
|
141
|
+
if (lid.x == 0u) {
|
|
142
|
+
AinvHi[ainvBase + localRow * BLOCK_SIZE + col] = val.hi;
|
|
143
|
+
AinvLo[ainvBase + localRow * BLOCK_SIZE + col] = val.lo;
|
|
144
|
+
}
|
|
145
|
+
storageBarrier();
|
|
146
|
+
workgroupBarrier();
|
|
147
|
+
}
|
|
148
|
+
}
|
|
@@ -0,0 +1,111 @@
|
|
|
1
|
+
// dtrsv_update: double-double (Dekker) f64 emulation of strsv_update.wgsl —
|
|
2
|
+
// subtracts a solved block's contribution from every remaining row in
|
|
3
|
+
// parallel (one workgroup per row, like dtrmv.wgsl). A, x, and the per-row
|
|
4
|
+
// accumulation are each split into an f32 (hi, lo) pair; WGSL has no f64
|
|
5
|
+
// type. No diag/masking needed: this region never touches the diagonal.
|
|
6
|
+
//
|
|
7
|
+
// The column range [blockStart, blockEnd) is the same fixed width for every
|
|
8
|
+
// row in a dispatch (not per-row-varying the way dtrmv's triangular range
|
|
9
|
+
// is), so the main/tail split is computed once, outside the row loop — same
|
|
10
|
+
// shape as dger.wgsl's column loop.
|
|
11
|
+
|
|
12
|
+
@group(0) @binding(0) var<storage, read> AHi: array<f32>;
|
|
13
|
+
@group(0) @binding(1) var<storage, read> ALo: array<f32>;
|
|
14
|
+
@group(0) @binding(2) var<storage, read_write> xHi: array<f32>;
|
|
15
|
+
@group(0) @binding(3) var<storage, read_write> xLo: array<f32>;
|
|
16
|
+
|
|
17
|
+
struct Params {
|
|
18
|
+
n: u32,
|
|
19
|
+
incx: u32,
|
|
20
|
+
lda: u32,
|
|
21
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
22
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
23
|
+
blockStart: u32,
|
|
24
|
+
blockEnd: u32, // exclusive
|
|
25
|
+
}
|
|
26
|
+
|
|
27
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
28
|
+
|
|
29
|
+
const WGS: u32 = 64u;
|
|
30
|
+
var<workgroup> tile: array<DD, 64>;
|
|
31
|
+
|
|
32
|
+
@compute @workgroup_size(64)
|
|
33
|
+
fn dtrsv_update_main(
|
|
34
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
35
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
36
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
37
|
+
) {
|
|
38
|
+
// forward: remaining rows are [blockEnd,n); backward: [0,blockStart).
|
|
39
|
+
// Uniform across the whole dispatch (derived from params only), so an
|
|
40
|
+
// early return here is safe — never a partial-workgroup divergence.
|
|
41
|
+
let forward = (params.trans == 0u) == (params.uplo == 0u);
|
|
42
|
+
|
|
43
|
+
var rangeStart: u32;
|
|
44
|
+
var rangeEnd: u32;
|
|
45
|
+
if forward {
|
|
46
|
+
rangeStart = params.blockEnd;
|
|
47
|
+
rangeEnd = params.n;
|
|
48
|
+
} else {
|
|
49
|
+
rangeStart = 0u;
|
|
50
|
+
rangeEnd = params.blockStart;
|
|
51
|
+
}
|
|
52
|
+
|
|
53
|
+
if (rangeStart >= rangeEnd) { return; }
|
|
54
|
+
let count = rangeEnd - rangeStart;
|
|
55
|
+
|
|
56
|
+
let colCount = params.blockEnd - params.blockStart;
|
|
57
|
+
let col_floor = (colCount / WGS) * WGS;
|
|
58
|
+
let mainIters = col_floor / WGS;
|
|
59
|
+
let hasTail = select(0u, 1u, col_floor < colCount);
|
|
60
|
+
|
|
61
|
+
for (var idx = wgid.x; idx < count; idx += nwg.x) {
|
|
62
|
+
let i = rangeStart + idx;
|
|
63
|
+
var acc = DD(0.0, 0.0);
|
|
64
|
+
|
|
65
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
66
|
+
let j = params.blockStart + lid.x + iter * WGS;
|
|
67
|
+
let aidx = select(j * params.lda + i, i * params.lda + j, params.trans == 0u);
|
|
68
|
+
let ix = j * params.incx;
|
|
69
|
+
let prod = ddMulProtected(DD(AHi[aidx], ALo[aidx]), DD(xHi[ix], xLo[ix]), lid.x);
|
|
70
|
+
acc = ddAddProtected(acc, prod, lid.x);
|
|
71
|
+
}
|
|
72
|
+
|
|
73
|
+
for (var iter = 0u; iter < hasTail; iter++) {
|
|
74
|
+
let j = params.blockStart + col_floor + lid.x;
|
|
75
|
+
let valid = j < params.blockEnd;
|
|
76
|
+
let safeJ = select(params.blockStart, j, valid);
|
|
77
|
+
let aidx = select(safeJ * params.lda + i, i * params.lda + safeJ, params.trans == 0u);
|
|
78
|
+
let ix = safeJ * params.incx;
|
|
79
|
+
let prod = ddMulProtected(DD(AHi[aidx], ALo[aidx]), DD(xHi[ix], xLo[ix]), lid.x);
|
|
80
|
+
let contribution = DD(select(0.0, prod.hi, valid), select(0.0, prod.lo, valid));
|
|
81
|
+
acc = ddAddProtected(acc, contribution, lid.x);
|
|
82
|
+
}
|
|
83
|
+
|
|
84
|
+
tile[lid.x] = acc;
|
|
85
|
+
workgroupBarrier();
|
|
86
|
+
|
|
87
|
+
// Inactive threads combine against a throwaway partner and discard it
|
|
88
|
+
// (ddAddProtected must be called unconditionally by every thread).
|
|
89
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
90
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
91
|
+
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
92
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
93
|
+
if (lid.x < s) { tile[lid.x] = combined; }
|
|
94
|
+
workgroupBarrier();
|
|
95
|
+
}
|
|
96
|
+
|
|
97
|
+
let ix_i = i * params.incx;
|
|
98
|
+
let xVal = DD(xHi[ix_i], xLo[ix_i]);
|
|
99
|
+
// tile[0] now holds the full subtracted term, visible to every lane —
|
|
100
|
+
// every thread redundantly finishes the O(1) subtraction so the
|
|
101
|
+
// protected op below stays uniformly called (only the write is gated).
|
|
102
|
+
let result = ddSubProtected(xVal, tile[0], lid.x);
|
|
103
|
+
|
|
104
|
+
if (lid.x == 0u) {
|
|
105
|
+
xHi[ix_i] = result.hi;
|
|
106
|
+
xLo[ix_i] = result.lo;
|
|
107
|
+
}
|
|
108
|
+
// All 64 threads must agree before the next row (grid-stride) reuses tile[].
|
|
109
|
+
workgroupBarrier();
|
|
110
|
+
}
|
|
111
|
+
}
|
|
@@ -84,7 +84,16 @@ fn ddMulRaw(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
|
84
84
|
return DD(p.hi, crossAndLo);
|
|
85
85
|
}
|
|
86
86
|
|
|
87
|
+
// A third compiler bug (Intel Mesa ANV, via NIR dump): raw.hi never gets
|
|
88
|
+
// materialized as one rounded value — the driver re-fuses a.hi*b.hi with
|
|
89
|
+
// ffma at every use site instead, breaking the a+b == s+e identity
|
|
90
|
+
// TwoSum-style algorithms depend on. Same fix as ddMulRaw's p.lo: force it
|
|
91
|
+
// through workgroup memory + a barrier. (Verified: max forward-error factor
|
|
92
|
+
// over 20000 trials dropped from >1e5 to ~3-4, Intel Mesa Iris Xe + NVIDIA GTX 1650.)
|
|
87
93
|
fn ddMulProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
88
94
|
let raw = ddMulRaw(a, b, threadSlot);
|
|
89
|
-
|
|
95
|
+
dekkerScratch[threadSlot] = raw.hi;
|
|
96
|
+
workgroupBarrier();
|
|
97
|
+
let rawHi = dekkerScratch[threadSlot];
|
|
98
|
+
return fastTwoSumProtected(rawHi, raw.lo, threadSlot);
|
|
90
99
|
}
|
package/src/shaders/index.mjs
CHANGED
|
@@ -200,12 +200,75 @@ routineShaders.strsv = {
|
|
|
200
200
|
import sger from "./sger.wgsl";
|
|
201
201
|
routineShaders.sger = { sger };
|
|
202
202
|
|
|
203
|
+
import dger from "./dger.wgsl"; // f64 sibling of sger — one ddMulProtected + one ddMulProtected + one ddAddProtected per element, no reduction shader needed
|
|
204
|
+
routineShaders.dger = {
|
|
205
|
+
"f64/dekker": dekker,
|
|
206
|
+
"f64/utils/add": ddAddUtil,
|
|
207
|
+
"f64/utils/multiply": ddMulUtil,
|
|
208
|
+
dger,
|
|
209
|
+
};
|
|
210
|
+
|
|
203
211
|
import ssyr from "./ssyr.wgsl";
|
|
204
212
|
routineShaders.ssyr = { ssyr };
|
|
205
213
|
|
|
214
|
+
import dsyr from "./dsyr.wgsl"; // f64 sibling of ssyr — two ddMulProtected + one ddAddProtected per element, no reduction shader needed
|
|
215
|
+
routineShaders.dsyr = {
|
|
216
|
+
"f64/dekker": dekker,
|
|
217
|
+
"f64/utils/add": ddAddUtil,
|
|
218
|
+
"f64/utils/multiply": ddMulUtil,
|
|
219
|
+
dsyr,
|
|
220
|
+
};
|
|
221
|
+
|
|
206
222
|
import ssyr2 from "./ssyr2.wgsl";
|
|
207
223
|
routineShaders.ssyr2 = { ssyr2 };
|
|
208
224
|
|
|
225
|
+
import dsyr2 from "./dsyr2.wgsl"; // f64 sibling of ssyr2 — two ddMulProtected + two ddAddProtected per element, no reduction shader needed
|
|
226
|
+
routineShaders.dsyr2 = {
|
|
227
|
+
"f64/dekker": dekker,
|
|
228
|
+
"f64/utils/add": ddAddUtil,
|
|
229
|
+
"f64/utils/multiply": ddMulUtil,
|
|
230
|
+
dsyr2,
|
|
231
|
+
};
|
|
232
|
+
|
|
233
|
+
import dgemv_n from "./dgemv_n.wgsl"; // f64 sibling of sgemv_n — per-row DD dot product via a full within-workgroup tree reduction (ddot's reduction shape), not sgemv_n's 4-way ILP unroll
|
|
234
|
+
import dgemv_t from "./dgemv_t.wgsl"; // f64 sibling of sgemv_t — per-thread accumulation tiled over x, with out-of-range columns computed against a clamped index (not skipped) to keep protected-op call counts uniform
|
|
235
|
+
routineShaders.dgemv = {
|
|
236
|
+
"f64/dekker": dekker,
|
|
237
|
+
"f64/utils/add": ddAddUtil,
|
|
238
|
+
"f64/utils/multiply": ddMulUtil,
|
|
239
|
+
dgemv_n,
|
|
240
|
+
dgemv_t,
|
|
241
|
+
}; // one or the other, picked by trans
|
|
242
|
+
|
|
243
|
+
import dsymv from "./dsymv.wgsl"; // f64 sibling of ssymv — same per-row DD tree reduction as dgemv_n, with a uplo-mirror address lookup (symIdx) replacing dgemv_n's plain row-major fetch
|
|
244
|
+
routineShaders.dsymv = {
|
|
245
|
+
"f64/dekker": dekker,
|
|
246
|
+
"f64/utils/add": ddAddUtil,
|
|
247
|
+
"f64/utils/multiply": ddMulUtil,
|
|
248
|
+
dsymv,
|
|
249
|
+
};
|
|
250
|
+
|
|
251
|
+
import dtrmv from "./dtrmv.wgsl"; // f64 sibling of strmv — per-row DD tree reduction (dsymv's shape) over a triangular, per-row-varying column range (dsyr's shape), not a mirrored or full-n one; wgblas's trmv already writes to a separate y rather than aliasing x in place, so there's no read/write-ordering concern to port
|
|
252
|
+
routineShaders.dtrmv = {
|
|
253
|
+
"f64/dekker": dekker,
|
|
254
|
+
"f64/utils/add": ddAddUtil,
|
|
255
|
+
"f64/utils/multiply": ddMulUtil,
|
|
256
|
+
dtrmv,
|
|
257
|
+
};
|
|
258
|
+
|
|
259
|
+
import dtrsv_invert_block from "./dtrsv_invert_block.wgsl"; // f64 sibling of strsv_invert_block — same per-column substitution, but the non-unit-diagonal division is the first real ddDivProtected use in this port order (deliberately last per TODO.md)
|
|
260
|
+
import dtrsv_apply_inverse from "./dtrsv_apply_inverse.wgsl"; // f64 sibling of strsv_apply_inverse — per-thread dense matvec against the block inverse, every thread runs the same blockLen-iteration loop (no early return) to keep protected-op call counts uniform
|
|
261
|
+
import dtrsv_update from "./dtrsv_update.wgsl"; // f64 sibling of strsv_update — per-row DD tree reduction (dsymv's shape) over the fixed-width block column range (dger's shape, not per-row-varying)
|
|
262
|
+
routineShaders.dtrsv = {
|
|
263
|
+
"f64/dekker": dekker,
|
|
264
|
+
"f64/utils/add": ddAddUtil,
|
|
265
|
+
"f64/utils/multiply": ddMulUtil,
|
|
266
|
+
"f64/utils/divide": ddDivUtil,
|
|
267
|
+
dtrsv_invert_block,
|
|
268
|
+
dtrsv_apply_inverse,
|
|
269
|
+
dtrsv_update,
|
|
270
|
+
};
|
|
271
|
+
|
|
209
272
|
import sgemm_small from "./sgemm_small.wgsl";
|
|
210
273
|
import sgemm_large from "./sgemm_large.wgsl";
|
|
211
274
|
routineShaders.sgemm = { sgemm_small, sgemm_large }; // one or the other, picked by a tile-size threshold
|