wgblas 2.2.1 → 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 +1082 -57
- 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/index.mjs +63 -0
|
@@ -0,0 +1,112 @@
|
|
|
1
|
+
// dgemv_n: y := alpha * A * x + beta * y, double-double (Dekker) f64
|
|
2
|
+
// emulation of sgemv_n (matrix-vector product, no transpose). A, x, y,
|
|
3
|
+
// alpha, and beta are each split into an f32 (hi, lo) pair; WGSL has no f64
|
|
4
|
+
// type. Same one-workgroup-per-output-row shape as sgemv_n.wgsl (grid-stride
|
|
5
|
+
// over rows), but the per-row dot product is reduced via a full
|
|
6
|
+
// within-workgroup DD tree reduction — the same reduction ddot.wgsl uses to
|
|
7
|
+
// combine 64 lanes down to one — rather than sgemv_n.wgsl's 4-way ILP
|
|
8
|
+
// unroll + shared-memory tree (ddMulProtected/ddAddProtected each carry a
|
|
9
|
+
// workgroupBarrier(), so unrolling would pay that barrier four times over
|
|
10
|
+
// per iteration for no benefit). Since each row is handled by exactly one
|
|
11
|
+
// workgroup, no separate cross-workgroup reduction pass (ddot's reduce_f64)
|
|
12
|
+
// is needed here — the tree reduction alone finishes the row.
|
|
13
|
+
|
|
14
|
+
@group(0) @binding(0) var<storage, read> AHi: array<f32>;
|
|
15
|
+
@group(0) @binding(1) var<storage, read> ALo: array<f32>;
|
|
16
|
+
@group(0) @binding(2) var<storage, read> xHi: array<f32>;
|
|
17
|
+
@group(0) @binding(3) var<storage, read> xLo: array<f32>;
|
|
18
|
+
@group(0) @binding(4) var<storage, read_write> yHi: array<f32>;
|
|
19
|
+
@group(0) @binding(5) var<storage, read_write> yLo: array<f32>;
|
|
20
|
+
|
|
21
|
+
struct Params {
|
|
22
|
+
m: u32,
|
|
23
|
+
n: u32,
|
|
24
|
+
alphaHi: f32,
|
|
25
|
+
alphaLo: f32,
|
|
26
|
+
betaHi: f32,
|
|
27
|
+
betaLo: f32,
|
|
28
|
+
incx: u32,
|
|
29
|
+
incy: u32,
|
|
30
|
+
lda: u32,
|
|
31
|
+
}
|
|
32
|
+
|
|
33
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
34
|
+
|
|
35
|
+
const WGS: u32 = 64u;
|
|
36
|
+
var<workgroup> tile: array<DD, 64>;
|
|
37
|
+
|
|
38
|
+
@compute @workgroup_size(64)
|
|
39
|
+
fn dgemv_n_main(
|
|
40
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
41
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
42
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
43
|
+
) {
|
|
44
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
45
|
+
let beta = DD(params.betaHi, params.betaLo);
|
|
46
|
+
let isBetaNonzero = params.betaHi != 0.0 || params.betaLo != 0.0;
|
|
47
|
+
|
|
48
|
+
// Column loop's trip count depends only on n and WGS, not on row — same
|
|
49
|
+
// main/tail split applies uniformly to every row (see dger.wgsl).
|
|
50
|
+
let n_floor = (params.n / WGS) * WGS;
|
|
51
|
+
let mainIters = n_floor / WGS;
|
|
52
|
+
let hasTail = select(0u, 1u, n_floor < params.n);
|
|
53
|
+
|
|
54
|
+
for (var row = wgid.x; row < params.m; row += nwg.x) {
|
|
55
|
+
let row_base = row * params.lda;
|
|
56
|
+
var acc = DD(0.0, 0.0);
|
|
57
|
+
|
|
58
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
59
|
+
let j = lid.x + iter * WGS;
|
|
60
|
+
let ia = row_base + j;
|
|
61
|
+
let ix = j * params.incx;
|
|
62
|
+
let prod = ddMulProtected(DD(AHi[ia], ALo[ia]), DD(xHi[ix], xLo[ix]), lid.x);
|
|
63
|
+
acc = ddAddProtected(acc, prod, lid.x);
|
|
64
|
+
}
|
|
65
|
+
|
|
66
|
+
for (var iter = 0u; iter < hasTail; iter++) {
|
|
67
|
+
let j = n_floor + lid.x;
|
|
68
|
+
let valid = j < params.n;
|
|
69
|
+
let ia = select(0u, row_base + j, valid);
|
|
70
|
+
let ix = select(0u, j * params.incx, valid);
|
|
71
|
+
let prod = ddMulProtected(DD(AHi[ia], ALo[ia]), DD(xHi[ix], xLo[ix]), lid.x);
|
|
72
|
+
let contribution = DD(select(0.0, prod.hi, valid), select(0.0, prod.lo, valid));
|
|
73
|
+
acc = ddAddProtected(acc, contribution, lid.x);
|
|
74
|
+
}
|
|
75
|
+
|
|
76
|
+
tile[lid.x] = acc;
|
|
77
|
+
workgroupBarrier();
|
|
78
|
+
|
|
79
|
+
// Inactive threads combine against a throwaway partner and discard it
|
|
80
|
+
// (ddAddProtected must be called unconditionally by every thread).
|
|
81
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
82
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
83
|
+
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
84
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
85
|
+
if (lid.x < s) { tile[lid.x] = combined; }
|
|
86
|
+
workgroupBarrier();
|
|
87
|
+
}
|
|
88
|
+
|
|
89
|
+
// tile[0] now holds the full row dot product, visible to every lane —
|
|
90
|
+
// every thread redundantly finishes the O(1) alpha/beta combine so the
|
|
91
|
+
// protected ops below stay uniformly called (only the write is gated).
|
|
92
|
+
let dot = tile[0];
|
|
93
|
+
let scaled = ddMulProtected(alpha, dot, lid.x);
|
|
94
|
+
let iy = row * params.incy;
|
|
95
|
+
let yVal = DD(yHi[iy], yLo[iy]);
|
|
96
|
+
let betaTimesY = ddMulProtected(beta, yVal, lid.x);
|
|
97
|
+
let withBeta = ddAddProtected(scaled, betaTimesY, lid.x);
|
|
98
|
+
// BLAS beta==0 semantics: y is written, not accumulated — the result
|
|
99
|
+
// must not depend on y's prior value (though it is still read above).
|
|
100
|
+
let result = DD(
|
|
101
|
+
select(scaled.hi, withBeta.hi, isBetaNonzero),
|
|
102
|
+
select(scaled.lo, withBeta.lo, isBetaNonzero),
|
|
103
|
+
);
|
|
104
|
+
|
|
105
|
+
if (lid.x == 0u) {
|
|
106
|
+
yHi[iy] = result.hi;
|
|
107
|
+
yLo[iy] = result.lo;
|
|
108
|
+
}
|
|
109
|
+
// All 64 threads must agree before the next row reuses tile[].
|
|
110
|
+
workgroupBarrier();
|
|
111
|
+
}
|
|
112
|
+
}
|
|
@@ -0,0 +1,102 @@
|
|
|
1
|
+
// dgemv_t: y := alpha * A^T * x + beta * y, double-double (Dekker) f64
|
|
2
|
+
// emulation of sgemv_t (matrix-vector product, transposed). A, x, y, alpha,
|
|
3
|
+
// and beta are each split into an f32 (hi, lo) pair; WGSL has no f64 type.
|
|
4
|
+
// Same one-thread-per-output-column shape as sgemv_t.wgsl, tiling over x
|
|
5
|
+
// (length m) via shared memory. Because ddMulProtected/ddAddProtected each
|
|
6
|
+
// carry a workgroupBarrier(), every thread in the workgroup — including ones
|
|
7
|
+
// whose column is out of range (n not a multiple of WGS) — must call them
|
|
8
|
+
// the same number of times: unlike sgemv_t.wgsl's `if (col < n)` guard
|
|
9
|
+
// around the whole accumulation, out-of-range threads here compute against
|
|
10
|
+
// a clamped dummy column unconditionally and simply never write their
|
|
11
|
+
// result (same technique dger.wgsl/ddot.wgsl use for their ragged tails).
|
|
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
|
+
m: u32,
|
|
22
|
+
n: u32,
|
|
23
|
+
alphaHi: f32,
|
|
24
|
+
alphaLo: f32,
|
|
25
|
+
betaHi: f32,
|
|
26
|
+
betaLo: f32,
|
|
27
|
+
incx: u32,
|
|
28
|
+
incy: u32,
|
|
29
|
+
lda: u32,
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
33
|
+
|
|
34
|
+
const WGS: u32 = 64u;
|
|
35
|
+
var<workgroup> xTile: array<DD, 64>;
|
|
36
|
+
|
|
37
|
+
@compute @workgroup_size(64)
|
|
38
|
+
fn dgemv_t_main(
|
|
39
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
40
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
41
|
+
) {
|
|
42
|
+
let col = gid.x;
|
|
43
|
+
let colValid = col < params.n;
|
|
44
|
+
// Clamp so out-of-range threads still index safely — their contribution
|
|
45
|
+
// is computed but never written back.
|
|
46
|
+
let safeCol = select(0u, col, colValid);
|
|
47
|
+
|
|
48
|
+
var acc = DD(0.0, 0.0);
|
|
49
|
+
|
|
50
|
+
let m_floor = (params.m / WGS) * WGS;
|
|
51
|
+
let mainIters = m_floor / WGS;
|
|
52
|
+
|
|
53
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
54
|
+
let base = iter * WGS;
|
|
55
|
+
// Cooperative load: all 64 threads fill xTile with x[base..base+WGS),
|
|
56
|
+
// independent of col — safe regardless of whether this thread's column
|
|
57
|
+
// is in range.
|
|
58
|
+
let ix = (base + lid.x) * params.incx;
|
|
59
|
+
xTile[lid.x] = DD(xHi[ix], xLo[ix]);
|
|
60
|
+
workgroupBarrier();
|
|
61
|
+
|
|
62
|
+
for (var j = 0u; j < WGS; j++) {
|
|
63
|
+
let ia = (base + j) * params.lda + safeCol;
|
|
64
|
+
let prod = ddMulProtected(DD(AHi[ia], ALo[ia]), xTile[j], lid.x);
|
|
65
|
+
acc = ddAddProtected(acc, prod, lid.x);
|
|
66
|
+
}
|
|
67
|
+
workgroupBarrier();
|
|
68
|
+
}
|
|
69
|
+
|
|
70
|
+
// Remainder rows (m not a multiple of WGS): the count is the same for
|
|
71
|
+
// every thread regardless of col, so this loop is already barrier-safe
|
|
72
|
+
// without needing a second tiling pass.
|
|
73
|
+
let remCount = params.m - m_floor;
|
|
74
|
+
for (var k = 0u; k < remCount; k++) {
|
|
75
|
+
let row = m_floor + k;
|
|
76
|
+
let ia = row * params.lda + safeCol;
|
|
77
|
+
let ix = row * params.incx;
|
|
78
|
+
let prod = ddMulProtected(DD(AHi[ia], ALo[ia]), DD(xHi[ix], xLo[ix]), lid.x);
|
|
79
|
+
acc = ddAddProtected(acc, prod, lid.x);
|
|
80
|
+
}
|
|
81
|
+
|
|
82
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
83
|
+
let beta = DD(params.betaHi, params.betaLo);
|
|
84
|
+
let scaled = ddMulProtected(alpha, acc, lid.x);
|
|
85
|
+
|
|
86
|
+
let iy = safeCol * params.incy;
|
|
87
|
+
let yVal = DD(yHi[iy], yLo[iy]);
|
|
88
|
+
let betaTimesY = ddMulProtected(beta, yVal, lid.x);
|
|
89
|
+
let withBeta = ddAddProtected(scaled, betaTimesY, lid.x);
|
|
90
|
+
let isBetaNonzero = params.betaHi != 0.0 || params.betaLo != 0.0;
|
|
91
|
+
// BLAS beta==0 semantics: y is written, not accumulated — the result
|
|
92
|
+
// must not depend on y's prior value (though it is still read above).
|
|
93
|
+
let result = DD(
|
|
94
|
+
select(scaled.hi, withBeta.hi, isBetaNonzero),
|
|
95
|
+
select(scaled.lo, withBeta.lo, isBetaNonzero),
|
|
96
|
+
);
|
|
97
|
+
|
|
98
|
+
if (colValid) {
|
|
99
|
+
yHi[iy] = result.hi;
|
|
100
|
+
yLo[iy] = result.lo;
|
|
101
|
+
}
|
|
102
|
+
}
|
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
// dger: A := alpha * x * y^T + A, double-double (Dekker) f64 emulation of
|
|
2
|
+
// sger (rank-1 update). x, y, A, and alpha are each split into an f32
|
|
3
|
+
// (hi, lo) pair; WGSL has no f64 type. One workgroup per row of A
|
|
4
|
+
// (grid-stride over rows); within a row, ddMulProtected/ddAddProtected each
|
|
5
|
+
// carry a workgroupBarrier(), so the column loop uses the same
|
|
6
|
+
// uniform-main + ragged-tail split dscal.wgsl uses, rather than sger.wgsl's
|
|
7
|
+
// 4-way ILP unroll — unrolling would pay the barrier round-trip four times
|
|
8
|
+
// over per iteration for no benefit here.
|
|
9
|
+
|
|
10
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
11
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
12
|
+
@group(0) @binding(2) var<storage, read> yHi: array<f32>;
|
|
13
|
+
@group(0) @binding(3) var<storage, read> yLo: array<f32>;
|
|
14
|
+
@group(0) @binding(4) var<storage, read_write> AHi: array<f32>;
|
|
15
|
+
@group(0) @binding(5) var<storage, read_write> ALo: array<f32>;
|
|
16
|
+
|
|
17
|
+
struct Params {
|
|
18
|
+
m: u32,
|
|
19
|
+
n: u32,
|
|
20
|
+
alphaHi: f32,
|
|
21
|
+
alphaLo: f32,
|
|
22
|
+
x_inc: u32,
|
|
23
|
+
y_inc: u32,
|
|
24
|
+
lda: u32,
|
|
25
|
+
}
|
|
26
|
+
|
|
27
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
28
|
+
|
|
29
|
+
const WGS: u32 = 64u;
|
|
30
|
+
|
|
31
|
+
@compute @workgroup_size(64)
|
|
32
|
+
fn dger_main(
|
|
33
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
34
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
35
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
36
|
+
) {
|
|
37
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
38
|
+
|
|
39
|
+
// Column loop's trip count depends only on n and WGS, not on row — the
|
|
40
|
+
// same main/tail split applies uniformly to every row, so every thread
|
|
41
|
+
// reaches each protected call the same number of times overall.
|
|
42
|
+
let n_floor = (params.n / WGS) * WGS;
|
|
43
|
+
let mainIters = n_floor / WGS;
|
|
44
|
+
let hasTail = select(0u, 1u, n_floor < params.n);
|
|
45
|
+
|
|
46
|
+
for (var row = wgid.x; row < params.m; row += nwg.x) {
|
|
47
|
+
// Every thread in the workgroup redundantly computes the same alpha*x[row]
|
|
48
|
+
// — wasteful but harmless, and keeps the per-row protected-call count
|
|
49
|
+
// trivially identical across threads.
|
|
50
|
+
let ix = row * params.x_inc;
|
|
51
|
+
let xi = ddMulProtected(alpha, DD(xHi[ix], xLo[ix]), lid.x);
|
|
52
|
+
let row_base = row * params.lda;
|
|
53
|
+
|
|
54
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
55
|
+
let col = lid.x + iter * WGS;
|
|
56
|
+
let idx = row_base + col;
|
|
57
|
+
let iy = col * params.y_inc;
|
|
58
|
+
let prod = ddMulProtected(xi, DD(yHi[iy], yLo[iy]), lid.x);
|
|
59
|
+
let result = ddAddProtected(prod, DD(AHi[idx], ALo[idx]), lid.x);
|
|
60
|
+
AHi[idx] = result.hi;
|
|
61
|
+
ALo[idx] = result.lo;
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
for (var iter = 0u; iter < hasTail; iter++) {
|
|
65
|
+
let col = n_floor + lid.x;
|
|
66
|
+
let valid = col < params.n;
|
|
67
|
+
let idx = select(0u, row_base + col, valid);
|
|
68
|
+
let iy = select(0u, col * params.y_inc, valid);
|
|
69
|
+
let prod = ddMulProtected(xi, DD(yHi[iy], yLo[iy]), lid.x);
|
|
70
|
+
let result = ddAddProtected(prod, DD(AHi[idx], ALo[idx]), lid.x);
|
|
71
|
+
if (valid) {
|
|
72
|
+
AHi[idx] = result.hi;
|
|
73
|
+
ALo[idx] = result.lo;
|
|
74
|
+
}
|
|
75
|
+
}
|
|
76
|
+
}
|
|
77
|
+
}
|
|
@@ -0,0 +1,125 @@
|
|
|
1
|
+
// dsymv: y := alpha * A * x + beta * y, double-double (Dekker) f64 emulation
|
|
2
|
+
// of ssymv (symmetric matrix-vector product). A, x, y, alpha, and beta are
|
|
3
|
+
// each split into an f32 (hi, lo) pair; WGSL has no f64 type. A is n×n
|
|
4
|
+
// symmetric — only the triangle specified by uplo is physically stored, so
|
|
5
|
+
// entries on the unstored side of the diagonal are fetched from their
|
|
6
|
+
// mirror position (A[i,j] == A[j,i]), same as ssymv.wgsl. Same
|
|
7
|
+
// one-workgroup-per-row + full within-workgroup DD tree reduction shape as
|
|
8
|
+
// dgemv_n.wgsl (every row sums over all n columns here, unlike dsyr's
|
|
9
|
+
// triangle-restricted range, since the logical matrix is fully dense).
|
|
10
|
+
|
|
11
|
+
@group(0) @binding(0) var<storage, read> AHi: array<f32>;
|
|
12
|
+
@group(0) @binding(1) var<storage, read> ALo: array<f32>;
|
|
13
|
+
@group(0) @binding(2) var<storage, read> xHi: array<f32>;
|
|
14
|
+
@group(0) @binding(3) var<storage, read> xLo: array<f32>;
|
|
15
|
+
@group(0) @binding(4) var<storage, read_write> yHi: array<f32>;
|
|
16
|
+
@group(0) @binding(5) var<storage, read_write> yLo: array<f32>;
|
|
17
|
+
|
|
18
|
+
struct Params {
|
|
19
|
+
n: u32,
|
|
20
|
+
alphaHi: f32,
|
|
21
|
+
alphaLo: f32,
|
|
22
|
+
betaHi: f32,
|
|
23
|
+
betaLo: f32,
|
|
24
|
+
incx: u32,
|
|
25
|
+
incy: u32,
|
|
26
|
+
lda: u32,
|
|
27
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
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
|
+
// A[i,j] flat index for the symmetric matrix — stored position if (i,j) is
|
|
36
|
+
// on the uplo side of the diagonal, mirrored from (j,i) otherwise. Pure
|
|
37
|
+
// addressing (no protected op inside), so branching here is safe.
|
|
38
|
+
fn symIdx(i: u32, j: u32) -> u32 {
|
|
39
|
+
var row: u32;
|
|
40
|
+
var col: u32;
|
|
41
|
+
if params.uplo == 0u {
|
|
42
|
+
// Lower: stored at (i,j) for j <= i, mirrored from (j,i) otherwise.
|
|
43
|
+
if j <= i { row = i; col = j; } else { row = j; col = i; }
|
|
44
|
+
} else {
|
|
45
|
+
// Upper: stored at (i,j) for j >= i, mirrored from (j,i) otherwise.
|
|
46
|
+
if j >= i { row = i; col = j; } else { row = j; col = i; }
|
|
47
|
+
}
|
|
48
|
+
return row * params.lda + col;
|
|
49
|
+
}
|
|
50
|
+
|
|
51
|
+
@compute @workgroup_size(64)
|
|
52
|
+
fn dsymv_main(
|
|
53
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
54
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
55
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
56
|
+
) {
|
|
57
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
58
|
+
let beta = DD(params.betaHi, params.betaLo);
|
|
59
|
+
let isBetaNonzero = params.betaHi != 0.0 || params.betaLo != 0.0;
|
|
60
|
+
|
|
61
|
+
// Column loop's trip count depends only on n and WGS, not on row — same
|
|
62
|
+
// main/tail split applies uniformly to every row (see dger.wgsl).
|
|
63
|
+
let n_floor = (params.n / WGS) * WGS;
|
|
64
|
+
let mainIters = n_floor / WGS;
|
|
65
|
+
let hasTail = select(0u, 1u, n_floor < params.n);
|
|
66
|
+
|
|
67
|
+
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
68
|
+
var acc = DD(0.0, 0.0);
|
|
69
|
+
|
|
70
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
71
|
+
let j = lid.x + iter * WGS;
|
|
72
|
+
let idx = symIdx(i, j);
|
|
73
|
+
let ix = j * params.incx;
|
|
74
|
+
let prod = ddMulProtected(DD(AHi[idx], ALo[idx]), DD(xHi[ix], xLo[ix]), lid.x);
|
|
75
|
+
acc = ddAddProtected(acc, prod, lid.x);
|
|
76
|
+
}
|
|
77
|
+
|
|
78
|
+
for (var iter = 0u; iter < hasTail; iter++) {
|
|
79
|
+
let j = n_floor + lid.x;
|
|
80
|
+
let valid = j < params.n;
|
|
81
|
+
let safeJ = select(0u, j, valid);
|
|
82
|
+
let idx = symIdx(i, safeJ);
|
|
83
|
+
let ix = safeJ * params.incx;
|
|
84
|
+
let prod = ddMulProtected(DD(AHi[idx], ALo[idx]), DD(xHi[ix], xLo[ix]), lid.x);
|
|
85
|
+
let contribution = DD(select(0.0, prod.hi, valid), select(0.0, prod.lo, valid));
|
|
86
|
+
acc = ddAddProtected(acc, contribution, lid.x);
|
|
87
|
+
}
|
|
88
|
+
|
|
89
|
+
tile[lid.x] = acc;
|
|
90
|
+
workgroupBarrier();
|
|
91
|
+
|
|
92
|
+
// Inactive threads combine against a throwaway partner and discard it
|
|
93
|
+
// (ddAddProtected must be called unconditionally by every thread).
|
|
94
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
95
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
96
|
+
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
97
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
98
|
+
if (lid.x < s) { tile[lid.x] = combined; }
|
|
99
|
+
workgroupBarrier();
|
|
100
|
+
}
|
|
101
|
+
|
|
102
|
+
// tile[0] now holds the full row dot product, visible to every lane —
|
|
103
|
+
// every thread redundantly finishes the O(1) alpha/beta combine so the
|
|
104
|
+
// protected ops below stay uniformly called (only the write is gated).
|
|
105
|
+
let dot = tile[0];
|
|
106
|
+
let scaled = ddMulProtected(alpha, dot, lid.x);
|
|
107
|
+
let iy = i * params.incy;
|
|
108
|
+
let yVal = DD(yHi[iy], yLo[iy]);
|
|
109
|
+
let betaTimesY = ddMulProtected(beta, yVal, lid.x);
|
|
110
|
+
let withBeta = ddAddProtected(scaled, betaTimesY, lid.x);
|
|
111
|
+
// BLAS beta==0 semantics: y is written, not accumulated — the result
|
|
112
|
+
// must not depend on y's prior value (though it is still read above).
|
|
113
|
+
let result = DD(
|
|
114
|
+
select(scaled.hi, withBeta.hi, isBetaNonzero),
|
|
115
|
+
select(scaled.lo, withBeta.lo, isBetaNonzero),
|
|
116
|
+
);
|
|
117
|
+
|
|
118
|
+
if (lid.x == 0u) {
|
|
119
|
+
yHi[iy] = result.hi;
|
|
120
|
+
yLo[iy] = result.lo;
|
|
121
|
+
}
|
|
122
|
+
// All 64 threads must agree before the next row reuses tile[].
|
|
123
|
+
workgroupBarrier();
|
|
124
|
+
}
|
|
125
|
+
}
|
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
// dsyr: A := alpha * x * x^T + A, double-double (Dekker) f64 emulation of
|
|
2
|
+
// ssyr (symmetric rank-1 update). x, A, and alpha are each split into an f32
|
|
3
|
+
// (hi, lo) pair; WGSL has no f64 type. Only the triangle specified by uplo
|
|
4
|
+
// is referenced/updated. One workgroup per row of A (grid-stride over
|
|
5
|
+
// rows); within a row, ddMulProtected/ddAddProtected each carry a
|
|
6
|
+
// workgroupBarrier(), so the column loop uses the same uniform-main +
|
|
7
|
+
// ragged-tail split dger.wgsl uses over its *stored* range, rather than
|
|
8
|
+
// ssyr.wgsl's 4-way ILP unroll.
|
|
9
|
+
|
|
10
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
11
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
12
|
+
@group(0) @binding(2) var<storage, read_write> AHi: array<f32>;
|
|
13
|
+
@group(0) @binding(3) var<storage, read_write> ALo: array<f32>;
|
|
14
|
+
|
|
15
|
+
struct Params {
|
|
16
|
+
n: u32,
|
|
17
|
+
alphaHi: f32,
|
|
18
|
+
alphaLo: f32,
|
|
19
|
+
incx: u32,
|
|
20
|
+
lda: u32,
|
|
21
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
25
|
+
|
|
26
|
+
const WGS: u32 = 64u;
|
|
27
|
+
|
|
28
|
+
@compute @workgroup_size(64)
|
|
29
|
+
fn dsyr_main(
|
|
30
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
31
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
32
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
33
|
+
) {
|
|
34
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
35
|
+
|
|
36
|
+
for (var row = wgid.x; row < params.n; row += nwg.x) {
|
|
37
|
+
let ix = row * params.incx;
|
|
38
|
+
let xi = ddMulProtected(alpha, DD(xHi[ix], xLo[ix]), lid.x);
|
|
39
|
+
let row_base = row * params.lda;
|
|
40
|
+
|
|
41
|
+
// Stored-triangle column range for this row: lower [0,row], upper [row,n).
|
|
42
|
+
var colStart: u32;
|
|
43
|
+
var colEnd: u32;
|
|
44
|
+
if params.uplo == 1u {
|
|
45
|
+
colStart = row;
|
|
46
|
+
colEnd = params.n;
|
|
47
|
+
} else {
|
|
48
|
+
colStart = 0u;
|
|
49
|
+
colEnd = row + 1u;
|
|
50
|
+
}
|
|
51
|
+
|
|
52
|
+
// Range length varies per row, but every thread in the workgroup runs
|
|
53
|
+
// the SAME row at the same time (the outer loop is shared), so the
|
|
54
|
+
// main/tail split computed from colStart/colEnd is already uniform
|
|
55
|
+
// across threads for this row — just not the same across different rows.
|
|
56
|
+
let rangeLen = colEnd - colStart;
|
|
57
|
+
let range_floor = colStart + (rangeLen / WGS) * WGS;
|
|
58
|
+
let mainIters = (range_floor - colStart) / WGS;
|
|
59
|
+
let hasTail = select(0u, 1u, range_floor < colEnd);
|
|
60
|
+
|
|
61
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
62
|
+
let col = colStart + lid.x + iter * WGS;
|
|
63
|
+
let idx = row_base + col;
|
|
64
|
+
let ic = col * params.incx;
|
|
65
|
+
let prod = ddMulProtected(xi, DD(xHi[ic], xLo[ic]), lid.x);
|
|
66
|
+
let result = ddAddProtected(prod, DD(AHi[idx], ALo[idx]), lid.x);
|
|
67
|
+
AHi[idx] = result.hi;
|
|
68
|
+
ALo[idx] = result.lo;
|
|
69
|
+
}
|
|
70
|
+
|
|
71
|
+
for (var iter = 0u; iter < hasTail; iter++) {
|
|
72
|
+
let col = range_floor + lid.x;
|
|
73
|
+
let valid = col < colEnd;
|
|
74
|
+
let idx = select(0u, row_base + col, valid);
|
|
75
|
+
let ic = select(0u, col * params.incx, valid);
|
|
76
|
+
let prod = ddMulProtected(xi, DD(xHi[ic], xLo[ic]), lid.x);
|
|
77
|
+
let result = ddAddProtected(prod, DD(AHi[idx], ALo[idx]), lid.x);
|
|
78
|
+
if (valid) {
|
|
79
|
+
AHi[idx] = result.hi;
|
|
80
|
+
ALo[idx] = result.lo;
|
|
81
|
+
}
|
|
82
|
+
}
|
|
83
|
+
}
|
|
84
|
+
}
|
|
@@ -0,0 +1,95 @@
|
|
|
1
|
+
// dsyr2: A := alpha * x * y^T + alpha * y * x^T + A, double-double (Dekker)
|
|
2
|
+
// f64 emulation of ssyr2 (symmetric rank-2 update). x, y, A, and alpha are
|
|
3
|
+
// each split into an f32 (hi, lo) pair; WGSL has no f64 type. Only the
|
|
4
|
+
// triangle specified by uplo is referenced/updated. One workgroup per row
|
|
5
|
+
// of A (grid-stride over rows); within a row, ddMulProtected/ddAddProtected
|
|
6
|
+
// each carry a workgroupBarrier(), so the column loop uses the same
|
|
7
|
+
// uniform-main + ragged-tail split dsyr.wgsl uses over its *stored* range,
|
|
8
|
+
// rather than ssyr2.wgsl's 4-way ILP unroll.
|
|
9
|
+
|
|
10
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
11
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
12
|
+
@group(0) @binding(2) var<storage, read> yHi: array<f32>;
|
|
13
|
+
@group(0) @binding(3) var<storage, read> yLo: array<f32>;
|
|
14
|
+
@group(0) @binding(4) var<storage, read_write> AHi: array<f32>;
|
|
15
|
+
@group(0) @binding(5) var<storage, read_write> ALo: array<f32>;
|
|
16
|
+
|
|
17
|
+
struct Params {
|
|
18
|
+
n: u32,
|
|
19
|
+
alphaHi: f32,
|
|
20
|
+
alphaLo: f32,
|
|
21
|
+
incx: u32,
|
|
22
|
+
incy: u32,
|
|
23
|
+
lda: u32,
|
|
24
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
25
|
+
}
|
|
26
|
+
|
|
27
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
28
|
+
|
|
29
|
+
const WGS: u32 = 64u;
|
|
30
|
+
|
|
31
|
+
@compute @workgroup_size(64)
|
|
32
|
+
fn dsyr2_main(
|
|
33
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
34
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
35
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
36
|
+
) {
|
|
37
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
38
|
+
|
|
39
|
+
for (var row = wgid.x; row < params.n; row += nwg.x) {
|
|
40
|
+
let ix = row * params.incx;
|
|
41
|
+
let iy = row * params.incy;
|
|
42
|
+
let xi = ddMulProtected(alpha, DD(xHi[ix], xLo[ix]), lid.x);
|
|
43
|
+
let yi = ddMulProtected(alpha, DD(yHi[iy], yLo[iy]), lid.x);
|
|
44
|
+
let row_base = row * params.lda;
|
|
45
|
+
|
|
46
|
+
// Stored-triangle column range for this row: lower [0,row], upper [row,n).
|
|
47
|
+
var colStart: u32;
|
|
48
|
+
var colEnd: u32;
|
|
49
|
+
if params.uplo == 1u {
|
|
50
|
+
colStart = row;
|
|
51
|
+
colEnd = params.n;
|
|
52
|
+
} else {
|
|
53
|
+
colStart = 0u;
|
|
54
|
+
colEnd = row + 1u;
|
|
55
|
+
}
|
|
56
|
+
|
|
57
|
+
// Range length varies per row, but every thread in the workgroup runs
|
|
58
|
+
// the SAME row at the same time (the outer loop is shared), so the
|
|
59
|
+
// main/tail split computed from colStart/colEnd is already uniform
|
|
60
|
+
// across threads for this row — just not the same across different rows.
|
|
61
|
+
let rangeLen = colEnd - colStart;
|
|
62
|
+
let range_floor = colStart + (rangeLen / WGS) * WGS;
|
|
63
|
+
let mainIters = (range_floor - colStart) / WGS;
|
|
64
|
+
let hasTail = select(0u, 1u, range_floor < colEnd);
|
|
65
|
+
|
|
66
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
67
|
+
let col = colStart + lid.x + iter * WGS;
|
|
68
|
+
let idx = row_base + col;
|
|
69
|
+
let jc_x = col * params.incx;
|
|
70
|
+
let jc_y = col * params.incy;
|
|
71
|
+
let prod1 = ddMulProtected(xi, DD(yHi[jc_y], yLo[jc_y]), lid.x);
|
|
72
|
+
let prod2 = ddMulProtected(yi, DD(xHi[jc_x], xLo[jc_x]), lid.x);
|
|
73
|
+
let sum1 = ddAddProtected(prod1, prod2, lid.x);
|
|
74
|
+
let result = ddAddProtected(sum1, DD(AHi[idx], ALo[idx]), lid.x);
|
|
75
|
+
AHi[idx] = result.hi;
|
|
76
|
+
ALo[idx] = result.lo;
|
|
77
|
+
}
|
|
78
|
+
|
|
79
|
+
for (var iter = 0u; iter < hasTail; iter++) {
|
|
80
|
+
let col = range_floor + lid.x;
|
|
81
|
+
let valid = col < colEnd;
|
|
82
|
+
let idx = select(0u, row_base + col, valid);
|
|
83
|
+
let jc_x = select(0u, col * params.incx, valid);
|
|
84
|
+
let jc_y = select(0u, col * params.incy, valid);
|
|
85
|
+
let prod1 = ddMulProtected(xi, DD(yHi[jc_y], yLo[jc_y]), lid.x);
|
|
86
|
+
let prod2 = ddMulProtected(yi, DD(xHi[jc_x], xLo[jc_x]), lid.x);
|
|
87
|
+
let sum1 = ddAddProtected(prod1, prod2, lid.x);
|
|
88
|
+
let result = ddAddProtected(sum1, DD(AHi[idx], ALo[idx]), lid.x);
|
|
89
|
+
if (valid) {
|
|
90
|
+
AHi[idx] = result.hi;
|
|
91
|
+
ALo[idx] = result.lo;
|
|
92
|
+
}
|
|
93
|
+
}
|
|
94
|
+
}
|
|
95
|
+
}
|