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.
@@ -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
+ }