wgblas 0.1.2 → 1.1.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (64) hide show
  1. package/README.md +3 -0
  2. package/dist/wgblas.browser.js +1249 -37
  3. package/index.d.mts +9 -0
  4. package/index.mjs +9 -0
  5. package/package.json +47 -1
  6. package/src/classes/GpuMatrix.d.mts +98 -0
  7. package/src/classes/GpuMatrix.mjs +109 -0
  8. package/src/classes/GpuVector.d.mts +11 -7
  9. package/src/classes/GpuVector.mjs +31 -5
  10. package/src/dasum/dasum.d.mts +48 -0
  11. package/src/dasum/dasum.mjs +121 -0
  12. package/src/init.mjs +3 -1
  13. package/src/isamax/isamax.mjs +66 -67
  14. package/src/random/random.d.mts +37 -0
  15. package/src/random/random.mjs +17 -0
  16. package/src/sasum/sasum.mjs +55 -51
  17. package/src/saxpy/saxpy.mjs +48 -35
  18. package/src/scopy/scopy.mjs +43 -32
  19. package/src/sdot/sdot.mjs +60 -55
  20. package/src/sgemv/sgemv.d.mts +127 -0
  21. package/src/sgemv/sgemv.mjs +148 -0
  22. package/src/sger/sger.d.mts +111 -0
  23. package/src/sger/sger.mjs +136 -0
  24. package/src/shaders/browser-shaders.mjs +26 -0
  25. package/src/shaders/dasum.wgsl +98 -0
  26. package/src/shaders/f64add.wgsl +281 -0
  27. package/src/shaders/isamax.wgsl +32 -9
  28. package/src/shaders/reduction/sumF64.wgsl +49 -0
  29. package/src/shaders/sasum.wgsl +18 -4
  30. package/src/shaders/sdot.wgsl +18 -4
  31. package/src/shaders/sgemv_n.wgsl +75 -0
  32. package/src/shaders/sgemv_t.wgsl +65 -0
  33. package/src/shaders/sger.wgsl +48 -0
  34. package/src/shaders/snrm2.wgsl +22 -4
  35. package/src/shaders/ssymv.wgsl +69 -0
  36. package/src/shaders/ssyr.wgsl +60 -0
  37. package/src/shaders/ssyr2.wgsl +63 -0
  38. package/src/shaders/strmv.wgsl +103 -0
  39. package/src/shaders/strsv_apply_inverse.wgsl +47 -0
  40. package/src/shaders/strsv_invert_block.wgsl +109 -0
  41. package/src/shaders/strsv_update.wgsl +75 -0
  42. package/src/snrm2/snrm2.mjs +56 -52
  43. package/src/srot/srot.mjs +57 -41
  44. package/src/srotm/srotm.mjs +54 -38
  45. package/src/sscal/sscal.mjs +43 -32
  46. package/src/sswap/sswap.mjs +49 -34
  47. package/src/ssymv/ssymv.d.mts +117 -0
  48. package/src/ssymv/ssymv.mjs +135 -0
  49. package/src/ssyr/ssyr.d.mts +100 -0
  50. package/src/ssyr/ssyr.mjs +106 -0
  51. package/src/ssyr2/ssyr2.d.mts +112 -0
  52. package/src/ssyr2/ssyr2.mjs +130 -0
  53. package/src/strmv/strmv.d.mts +117 -0
  54. package/src/strmv/strmv.mjs +138 -0
  55. package/src/strsv/strsv.d.mts +106 -0
  56. package/src/strsv/strsv.mjs +207 -0
  57. package/src/util/benchmark.mjs +1 -1
  58. package/src/util/bindgroup.mjs +14 -10
  59. package/src/util/buffer.mjs +7 -2
  60. package/src/util/compute.mjs +41 -15
  61. package/src/util/f64pack.mjs +152 -0
  62. package/src/util/pipeline.mjs +32 -17
  63. package/src/util/result.mjs +8 -4
  64. package/src/util/workgroup.mjs +10 -10
@@ -0,0 +1,69 @@
1
+ // ssymv: y = alpha * A * x + beta * y
2
+ // A is n×n symmetric, lower (uplo=0) or upper (uplo=1) triangle stored.
3
+ // The logical matrix is fully dense (symmetric), so each row's dot product
4
+ // sums over all n columns; entries on the unstored side of the diagonal are
5
+ // fetched from their mirror position (A[i,j] == A[j,i]).
6
+ // One workgroup per row, grid-stride outer loop.
7
+
8
+ @group(0) @binding(0) var<storage, read> A: array<f32>;
9
+ @group(0) @binding(1) var<storage, read> x: array<f32>;
10
+ @group(0) @binding(2) var<storage, read_write> y: array<f32>;
11
+
12
+ struct Params {
13
+ n: u32,
14
+ alpha: f32,
15
+ beta: f32,
16
+ incx: u32,
17
+ incy: u32,
18
+ lda: u32,
19
+ uplo: u32, // 0 = lower, 1 = upper
20
+ }
21
+
22
+ @group(0) @binding(3) var<uniform> params: Params;
23
+
24
+ const WGS: u32 = 64u;
25
+ var<workgroup> scratch: array<f32, 64>;
26
+
27
+ @compute @workgroup_size(64)
28
+ fn main(
29
+ @builtin(workgroup_id) wgid: vec3u,
30
+ @builtin(local_invocation_id) lid: vec3u,
31
+ @builtin(num_workgroups) nwg: vec3u,
32
+ ) {
33
+ for (var i = wgid.x; i < params.n; i += nwg.x) {
34
+ var acc = 0.0f;
35
+
36
+ // y[i] = Σ_j A[i,j] * x[j]
37
+ for (var j = lid.x; j < params.n; j += WGS) {
38
+ var aVal: f32;
39
+ if params.uplo == 0u {
40
+ // Lower: A[i,j] stored at A[i*lda+j] for j ≤ i, mirrored from A[j*lda+i] otherwise
41
+ if j <= i {
42
+ aVal = A[i * params.lda + j];
43
+ } else {
44
+ aVal = A[j * params.lda + i];
45
+ }
46
+ } else {
47
+ // Upper: A[i,j] stored at A[i*lda+j] for j ≥ i, mirrored from A[j*lda+i] otherwise
48
+ if j >= i {
49
+ aVal = A[i * params.lda + j];
50
+ } else {
51
+ aVal = A[j * params.lda + i];
52
+ }
53
+ }
54
+ acc += aVal * x[j * params.incx];
55
+ }
56
+
57
+ // Parallel reduction: 64 → 1
58
+ scratch[lid.x] = acc;
59
+ workgroupBarrier();
60
+ for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
61
+ if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
62
+ workgroupBarrier();
63
+ }
64
+
65
+ if lid.x == 0u {
66
+ y[i * params.incy] = params.alpha * scratch[0] + params.beta * y[i * params.incy];
67
+ }
68
+ }
69
+ }
@@ -0,0 +1,60 @@
1
+ // ssyr: A := alpha * x * x^T + A (symmetric rank-1 update)
2
+ // A is n×n symmetric; only the triangle specified by uplo is referenced/updated,
3
+ // the other triangle is implied by symmetry (not touched).
4
+
5
+ @group(0) @binding(0) var<storage, read> x: array<f32>;
6
+ @group(0) @binding(1) var<storage, read_write> A: array<f32>;
7
+
8
+ struct Params {
9
+ n: u32,
10
+ alpha: f32,
11
+ incx: u32,
12
+ lda: u32,
13
+ uplo: u32, // 0 = lower, 1 = upper
14
+ }
15
+
16
+ @group(0) @binding(2) var<uniform> params: Params;
17
+
18
+ const WGS: u32 = 64u;
19
+
20
+ @compute @workgroup_size(64)
21
+ fn main(
22
+ @builtin(workgroup_id) wgid: vec3u,
23
+ @builtin(local_invocation_id) lid: vec3u,
24
+ @builtin(num_workgroups) nwg: vec3u,
25
+ ) {
26
+ for (var row = wgid.x; row < params.n; row += nwg.x) {
27
+ let xi = params.alpha * x[row * params.incx];
28
+ let row_base = row * params.lda;
29
+
30
+ // Stored-triangle column range for this row: lower [0,row], upper [row,n).
31
+ var colStart: u32;
32
+ var colEnd: u32;
33
+ if params.uplo == 1u {
34
+ colStart = row;
35
+ colEnd = params.n;
36
+ } else {
37
+ colStart = 0u;
38
+ colEnd = row + 1u;
39
+ }
40
+
41
+ // 4-unrolled loop over the stored range.
42
+ let rangeLen = colEnd - colStart;
43
+ let n4_floor = colStart + (rangeLen / (4u * WGS)) * (4u * WGS);
44
+ for (var col: u32 = colStart + lid.x; col < n4_floor; col += 4u * WGS) {
45
+ let idx0 = row_base + col;
46
+ let idx1 = row_base + col + WGS;
47
+ let idx2 = row_base + col + 2u * WGS;
48
+ let idx3 = row_base + col + 3u * WGS;
49
+ A[idx0] = xi * x[ col * params.incx] + A[idx0];
50
+ A[idx1] = xi * x[(col + WGS) * params.incx] + A[idx1];
51
+ A[idx2] = xi * x[(col + 2u * WGS) * params.incx] + A[idx2];
52
+ A[idx3] = xi * x[(col + 3u * WGS) * params.incx] + A[idx3];
53
+ }
54
+ // Scalar tail: at most 3*WGS elements left after the unrolled block.
55
+ for (var col: u32 = n4_floor + lid.x; col < colEnd; col += WGS) {
56
+ let idx = row_base + col;
57
+ A[idx] = xi * x[col * params.incx] + A[idx];
58
+ }
59
+ }
60
+ }
@@ -0,0 +1,63 @@
1
+ // ssyr2: A := alpha * x * y^T + alpha * y * x^T + A (symmetric rank-2 update)
2
+ // A is n×n symmetric; only the triangle specified by uplo is referenced/updated,
3
+ // the other triangle is implied by symmetry (not touched).
4
+
5
+ @group(0) @binding(0) var<storage, read> x: array<f32>;
6
+ @group(0) @binding(1) var<storage, read> y: array<f32>;
7
+ @group(0) @binding(2) var<storage, read_write> A: array<f32>;
8
+
9
+ struct Params {
10
+ n: u32,
11
+ alpha: f32,
12
+ incx: u32,
13
+ incy: u32,
14
+ lda: u32,
15
+ uplo: u32, // 0 = lower, 1 = upper
16
+ }
17
+
18
+ @group(0) @binding(3) var<uniform> params: Params;
19
+
20
+ const WGS: u32 = 64u;
21
+
22
+ @compute @workgroup_size(64)
23
+ fn main(
24
+ @builtin(workgroup_id) wgid: vec3u,
25
+ @builtin(local_invocation_id) lid: vec3u,
26
+ @builtin(num_workgroups) nwg: vec3u,
27
+ ) {
28
+ for (var row = wgid.x; row < params.n; row += nwg.x) {
29
+ let xi = params.alpha * x[row * params.incx];
30
+ let yi = params.alpha * y[row * params.incy];
31
+ let row_base = row * params.lda;
32
+
33
+ // Stored-triangle column range for this row: lower [0,row], upper [row,n).
34
+ var colStart: u32;
35
+ var colEnd: u32;
36
+ if params.uplo == 1u {
37
+ colStart = row;
38
+ colEnd = params.n;
39
+ } else {
40
+ colStart = 0u;
41
+ colEnd = row + 1u;
42
+ }
43
+
44
+ // 4-unrolled loop over the stored range.
45
+ let rangeLen = colEnd - colStart;
46
+ let n4_floor = colStart + (rangeLen / (4u * WGS)) * (4u * WGS);
47
+ for (var col: u32 = colStart + lid.x; col < n4_floor; col += 4u * WGS) {
48
+ let idx0 = row_base + col;
49
+ let idx1 = row_base + col + WGS;
50
+ let idx2 = row_base + col + 2u * WGS;
51
+ let idx3 = row_base + col + 3u * WGS;
52
+ A[idx0] = xi * y[ col * params.incy] + yi * x[ col * params.incx] + A[idx0];
53
+ A[idx1] = xi * y[(col + WGS) * params.incy] + yi * x[(col + WGS) * params.incx] + A[idx1];
54
+ A[idx2] = xi * y[(col + 2u * WGS) * params.incy] + yi * x[(col + 2u * WGS) * params.incx] + A[idx2];
55
+ A[idx3] = xi * y[(col + 3u * WGS) * params.incy] + yi * x[(col + 3u * WGS) * params.incx] + A[idx3];
56
+ }
57
+ // Scalar tail: at most 3*WGS elements left after the unrolled block.
58
+ for (var col: u32 = n4_floor + lid.x; col < colEnd; col += WGS) {
59
+ let idx = row_base + col;
60
+ A[idx] = xi * y[col * params.incy] + yi * x[col * params.incx] + A[idx];
61
+ }
62
+ }
63
+ }
@@ -0,0 +1,103 @@
1
+ // strmv: y = op(A) * x
2
+ // A is n×n triangular, lower (uplo=0) or upper (uplo=1) triangle stored.
3
+ // op(A) is A (trans=0) or A^T (trans=1).
4
+ // diag=1 (unit) treats the diagonal as 1 without reading A's diagonal values.
5
+ // One workgroup per row, grid-stride outer loop.
6
+
7
+ @group(0) @binding(0) var<storage, read> A: array<f32>;
8
+ @group(0) @binding(1) var<storage, read> x: array<f32>;
9
+ @group(0) @binding(2) var<storage, read_write> y: array<f32>;
10
+
11
+ struct Params {
12
+ n: u32,
13
+ incx: u32,
14
+ incy: u32,
15
+ lda: u32,
16
+ trans: u32, // 0 = no-transpose, 1 = transpose
17
+ uplo: u32, // 0 = lower, 1 = upper
18
+ diag: u32, // 0 = non-unit, 1 = unit
19
+ }
20
+
21
+ @group(0) @binding(3) var<uniform> params: Params;
22
+
23
+ const WGS: u32 = 64u;
24
+ var<workgroup> scratch: array<f32, 64>;
25
+
26
+ @compute @workgroup_size(64)
27
+ fn main(
28
+ @builtin(workgroup_id) wgid: vec3u,
29
+ @builtin(local_invocation_id) lid: vec3u,
30
+ @builtin(num_workgroups) nwg: vec3u,
31
+ ) {
32
+ for (var i = wgid.x; i < params.n; i += nwg.x) {
33
+ var acc = 0.0f;
34
+
35
+ if params.trans == 0u {
36
+ // No-transpose: y[i] = Σ_j A[i,j] * x[j]
37
+ if params.uplo == 0u {
38
+ // Lower: A[i,j] stored at A[i*lda+j] for j ≤ i
39
+ for (var j = lid.x; j <= i; j += WGS) {
40
+ var aVal: f32;
41
+ // unit diagonal: use 1 instead of A's actual diagonal value
42
+ if params.diag == 1u && j == i {
43
+ aVal = 1.0;
44
+ } else if ( j <= i ) {
45
+ aVal = A[i * params.lda + j];
46
+ }
47
+ acc += aVal * x[j * params.incx];
48
+ }
49
+ } else {
50
+ // Upper: A[i,j] stored at A[i*lda+j] for j ≥ i
51
+ for (var j = i + lid.x; j < params.n; j += WGS) {
52
+ var aVal: f32;
53
+ // unit diagonal: use 1 instead of A's actual diagonal value
54
+ if params.diag == 1u && j == i {
55
+ aVal = 1.0;
56
+ } else if ( j >= i ) {
57
+ aVal = A[i * params.lda + j];
58
+ }
59
+ acc += aVal * x[j * params.incx];
60
+ }
61
+ }
62
+ } else {
63
+ // Transpose: y[i] = Σ_j A[j,i] * x[j]
64
+ if params.uplo == 0u {
65
+ // Lower: A[j,i] stored at A[j*lda+i] for j ≥ i
66
+ for (var j = i + lid.x; j < params.n; j += WGS) {
67
+ var aVal: f32;
68
+ // unit diagonal: use 1 instead of A's actual diagonal value
69
+ if params.diag == 1u && j == i {
70
+ aVal = 1.0;
71
+ } else if ( j >= i ) {
72
+ aVal = A[j * params.lda + i];
73
+ }
74
+ acc += aVal * x[j * params.incx];
75
+ }
76
+ } else {
77
+ // Upper: A[j,i] stored at A[j*lda+i] for j ≤ i
78
+ for (var j = lid.x; j <= i; j += WGS) {
79
+ var aVal: f32;
80
+ // unit diagonal: use 1 instead of A's actual diagonal value
81
+ if params.diag == 1u && j == i {
82
+ aVal = 1.0;
83
+ } else if ( j <= i ) {
84
+ aVal = A[j * params.lda + i];
85
+ }
86
+ acc += aVal * x[j * params.incx];
87
+ }
88
+ }
89
+ }
90
+
91
+ // Parallel reduction: 64 → 1
92
+ scratch[lid.x] = acc;
93
+ workgroupBarrier();
94
+ for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
95
+ if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
96
+ workgroupBarrier();
97
+ }
98
+
99
+ if lid.x == 0u {
100
+ y[ i * params.incy ] = scratch[0];
101
+ }
102
+ }
103
+ }
@@ -0,0 +1,47 @@
1
+ // strsv_apply_inverse: given a precomputed block inverse (from
2
+ // strsv_invert_block.wgsl), computes this block's solution as a dense
3
+ // matrix-vector multiply against the block's current remainder in x —
4
+ // replacing what the old strsv_block.wgsl did via a genuinely sequential,
5
+ // barrier-per-row substitution.
6
+ //
7
+ // All blockLen rows are computed in parallel within a single workgroup: the
8
+ // remainder is loaded into workgroup-shared memory once, then each thread
9
+ // independently computes one full row's dot product from that shared copy.
10
+ // No further synchronization is needed after the load — every thread only
11
+ // reads shared memory from then on (never written again within this call)
12
+ // and writes a distinct element of x, so there's no cross-thread hazard to
13
+ // guard against.
14
+
15
+ @group(0) @binding(0) var<storage, read> Ainv: array<f32>;
16
+ @group(0) @binding(1) var<storage, read_write> x: array<f32>;
17
+
18
+ struct Params {
19
+ incx: u32,
20
+ blockIndex: u32,
21
+ blockStart: u32,
22
+ blockEnd: u32,
23
+ }
24
+
25
+ @group(0) @binding(2) var<uniform> params: Params;
26
+
27
+ const BLOCK_SIZE: u32 = 64u;
28
+ var<workgroup> xLocal: array<f32, 64>;
29
+
30
+ @compute @workgroup_size(64)
31
+ fn strsv_apply_inverse_main(@builtin(local_invocation_id) lid: vec3u) {
32
+ let blockLen = params.blockEnd - params.blockStart;
33
+
34
+ if (lid.x < blockLen) {
35
+ xLocal[lid.x] = x[(params.blockStart + lid.x) * params.incx];
36
+ }
37
+ workgroupBarrier();
38
+
39
+ if (lid.x >= blockLen) { return; }
40
+
41
+ let ainvBase = params.blockIndex * BLOCK_SIZE * BLOCK_SIZE;
42
+ var acc = 0.0f;
43
+ for (var j = 0u; j < blockLen; j++) {
44
+ acc += Ainv[ainvBase + lid.x * BLOCK_SIZE + j] * xLocal[j];
45
+ }
46
+ x[(params.blockStart + lid.x) * params.incx] = acc;
47
+ }
@@ -0,0 +1,109 @@
1
+ // strsv_invert_block: computes ONE column (workgroup_id.x) of ONE block's
2
+ // (workgroup_id.y) explicit inverse, via the same one-row-at-a-time
3
+ // substitution as strsv_block.wgsl, but solving against a unit basis vector
4
+ // e_col instead of the real right-hand side, and writing to a dense
5
+ // (BLOCK_SIZE x BLOCK_SIZE, row-major) scratch buffer per block instead of
6
+ // mutating x. Dispatched once for the whole matrix (2D: BLOCK_SIZE columns x
7
+ // numBlocks), fully in parallel -- unlike the sequential per-block main
8
+ // loop in strsv.mjs, no block's inverse depends on any other block or on x.
9
+ //
10
+ // A triangular block's inverse is itself triangular: forward (effectively-
11
+ // lower, e.g. no-trans+lower) blocks have inverse column col nonzero only
12
+ // for row>=col, solved in increasing row order; backward (effectively-
13
+ // upper) blocks have it nonzero only for row<=col, solved in decreasing
14
+ // order. Rows outside a column's nonzero range are written as literal 0 —
15
+ // strsv_apply_inverse.wgsl's dense matvec depends on that, not just on
16
+ // those entries being mathematically implied zero.
17
+
18
+ @group(0) @binding(0) var<storage, read> A: array<f32>;
19
+ @group(0) @binding(1) var<storage, read_write> Ainv: array<f32>;
20
+
21
+ struct Params {
22
+ n: u32,
23
+ lda: u32,
24
+ trans: u32, // 0 = no-transpose, 1 = transpose
25
+ uplo: u32, // 0 = lower, 1 = upper
26
+ diag: u32, // 0 = non-unit, 1 = unit
27
+ }
28
+
29
+ @group(0) @binding(2) var<uniform> params: Params;
30
+
31
+ const WGS: u32 = 64u;
32
+ const BLOCK_SIZE: u32 = 64u;
33
+ var<workgroup> scratch: array<f32, 64>;
34
+
35
+ fn readA(i: u32, j: u32) -> f32 {
36
+ if params.trans == 0u {
37
+ return A[i * params.lda + j];
38
+ } else {
39
+ return A[j * params.lda + i];
40
+ }
41
+ }
42
+
43
+ @compute @workgroup_size(64)
44
+ fn strsv_invert_block_main(
45
+ @builtin(workgroup_id) wgid: vec3u,
46
+ @builtin(local_invocation_id) lid: vec3u,
47
+ ) {
48
+ let col = wgid.x;
49
+ let blockIndex = wgid.y;
50
+ let blockStart = blockIndex * BLOCK_SIZE;
51
+ var blockEnd = blockStart + BLOCK_SIZE;
52
+ if (blockEnd > params.n) { blockEnd = params.n; }
53
+ let blockLen = blockEnd - blockStart;
54
+
55
+ if (col >= blockLen) { return; }
56
+
57
+ let ainvBase = blockIndex * BLOCK_SIZE * BLOCK_SIZE;
58
+ let forward = (params.trans == 0u) == (params.uplo == 0u);
59
+
60
+ if forward {
61
+ for (var r = lid.x; r < col; r += WGS) {
62
+ Ainv[ainvBase + r * BLOCK_SIZE + col] = 0.0;
63
+ }
64
+ } else {
65
+ for (var r = col + 1u + lid.x; r < blockLen; r += WGS) {
66
+ Ainv[ainvBase + r * BLOCK_SIZE + col] = 0.0;
67
+ }
68
+ }
69
+ storageBarrier();
70
+ workgroupBarrier();
71
+
72
+ let numSteps = select(col + 1u, blockLen - col, forward);
73
+ for (var step = 0u; step < numSteps; step++) {
74
+ let localRow = select(col - step, col + step, forward);
75
+ let i = blockStart + localRow;
76
+
77
+ var acc = 0.0f;
78
+ if forward {
79
+ for (var lj = col + lid.x; lj < localRow; lj += WGS) {
80
+ acc += readA(i, blockStart + lj) * Ainv[ainvBase + lj * BLOCK_SIZE + col];
81
+ }
82
+ } else {
83
+ for (var lj = localRow + 1u + lid.x; lj <= col; lj += WGS) {
84
+ acc += readA(i, blockStart + lj) * Ainv[ainvBase + lj * BLOCK_SIZE + col];
85
+ }
86
+ }
87
+
88
+ scratch[lid.x] = acc;
89
+ workgroupBarrier();
90
+ for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
91
+ if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
92
+ workgroupBarrier();
93
+ }
94
+
95
+ if lid.x == 0u {
96
+ let e = select(0.0, 1.0, localRow == col);
97
+ let rhs = e - scratch[0];
98
+ var val: f32;
99
+ if params.diag == 1u {
100
+ val = rhs;
101
+ } else {
102
+ val = rhs / A[i * params.lda + i];
103
+ }
104
+ Ainv[ainvBase + localRow * BLOCK_SIZE + col] = val;
105
+ }
106
+ storageBarrier();
107
+ workgroupBarrier();
108
+ }
109
+ }
@@ -0,0 +1,75 @@
1
+ // strsv_update: subtracts a solved block's contribution from every
2
+ // remaining row in parallel (one workgroup per row, like strmv.wgsl) —
3
+ // this is what turns strsv's O(n) sequential stages into O(n/blockSize).
4
+ // No diag/masking needed: this region never touches the diagonal.
5
+
6
+ @group(0) @binding(0) var<storage, read> A: array<f32>;
7
+ @group(0) @binding(1) var<storage, read_write> x: array<f32>;
8
+
9
+ struct Params {
10
+ n: u32,
11
+ incx: u32,
12
+ lda: u32,
13
+ trans: u32, // 0 = no-transpose, 1 = transpose
14
+ uplo: u32, // 0 = lower, 1 = upper
15
+ blockStart: u32,
16
+ blockEnd: u32, // exclusive
17
+ }
18
+
19
+ @group(0) @binding(2) var<uniform> params: Params;
20
+
21
+ const WGS: u32 = 64u;
22
+ var<workgroup> scratch: array<f32, 64>;
23
+
24
+ @compute @workgroup_size(64)
25
+ fn strsv_update_main(
26
+ @builtin(workgroup_id) wgid: vec3u,
27
+ @builtin(local_invocation_id) lid: vec3u,
28
+ @builtin(num_workgroups) nwg: vec3u,
29
+ ) {
30
+ // forward: remaining rows are [blockEnd,n); backward: [0,blockStart).
31
+ let forward = (params.trans == 0u) == (params.uplo == 0u);
32
+
33
+ var rangeStart: u32;
34
+ var rangeEnd: u32;
35
+
36
+ if forward {
37
+ rangeStart = params.blockEnd;
38
+ rangeEnd = params.n;
39
+ } else {
40
+ rangeStart = 0u;
41
+ rangeEnd = params.blockStart;
42
+ }
43
+
44
+ if (rangeStart >= rangeEnd) { return; }
45
+ let count = rangeEnd - rangeStart;
46
+
47
+ for (var idx = wgid.x; idx < count; idx += nwg.x) {
48
+ let i = rangeStart + idx;
49
+
50
+ // No-trans reads A[i,j]; transpose reads A[j,i] — uplo only sets the range above.
51
+ var acc = 0.0f;
52
+ if params.trans == 0u {
53
+ for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
54
+ acc += A[i * params.lda + j] * x[j * params.incx];
55
+ }
56
+ } else {
57
+ for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
58
+ acc += A[j * params.lda + i] * x[j * params.incx];
59
+ }
60
+ }
61
+
62
+ // Parallel reduction: 64 → 1
63
+ scratch[lid.x] = acc;
64
+ workgroupBarrier();
65
+ for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
66
+ if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
67
+ workgroupBarrier();
68
+ }
69
+
70
+ if lid.x == 0u {
71
+ x[i * params.incx] -= scratch[0];
72
+ }
73
+ workgroupBarrier();
74
+ }
75
+ }
@@ -34,67 +34,71 @@ export async function snrm2(device, n, x, incx) {
34
34
  const pipelineMain = await getPipeline(device, "snrm2");
35
35
  const pipelineReduce = await getPipeline(device, "reduction/sum");
36
36
 
37
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "snrm2-x", false);
38
- const partialsBuffer = createStorageBuffer(2 * WGS * 4, "snrm2-partials"); // 2*WGS partial sums of f32
39
- const resultBuffer = createResultBuffer(4, "snrm2-result"); // final f32 scalar
40
- const paramsBuffer = createParamsBuffer(
41
- [
42
- { value: n, type: "u32" },
43
- { value: incx, type: "u32" },
44
- ],
45
- "snrm2-params",
46
- );
37
+ let xBuffer = null;
38
+ let partialsBuffer = null;
39
+ let resultBuffer = null;
40
+ let paramsBuffer = null;
41
+ let readBuffer = null;
47
42
 
48
- const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
49
- xBuffer,
50
- partialsBuffer,
51
- paramsBuffer,
52
- ]);
53
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(
54
- pipelineMain,
55
- bgMain,
56
- 2 * WGS,
57
- ); // dispatch 2*WGS workgroups
43
+ try {
44
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "snrm2-x", false);
45
+ partialsBuffer = createStorageBuffer(2 * WGS * 4, "snrm2-partials"); // 2*WGS partial sums of f32
46
+ resultBuffer = createResultBuffer(4, "snrm2-result"); // final f32 scalar
47
+ paramsBuffer = createParamsBuffer(
48
+ [
49
+ { value: n, type: "u32" },
50
+ { value: incx, type: "u32" },
51
+ ],
52
+ "snrm2-params",
53
+ );
54
+
55
+ const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
56
+ xBuffer,
57
+ partialsBuffer,
58
+ paramsBuffer,
59
+ ]);
60
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
61
+ pipelineMain,
62
+ bgMain,
63
+ 2 * WGS,
64
+ ); // dispatch 2*WGS workgroups
58
65
 
59
- submit(enc1);
66
+ submit(enc1);
60
67
 
61
- const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
62
- partialsBuffer,
63
- resultBuffer,
64
- ]);
65
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(
66
- pipelineReduce,
67
- bgReduce,
68
- 1,
69
- ); // reduce partials to single result
70
- const readBuffer = stageReadback(enc2, resultBuffer);
68
+ const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
69
+ partialsBuffer,
70
+ resultBuffer,
71
+ ]);
72
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
73
+ pipelineReduce,
74
+ bgReduce,
75
+ 1,
76
+ ); // reduce partials to single result
77
+ readBuffer = stageReadback(enc2, resultBuffer);
71
78
 
72
- submit(enc2);
79
+ submit(enc2);
73
80
 
74
- const [gpuTime1, gpuTime2, sqsumArr] = await Promise.all([
75
- extractTimestamp(ts1),
76
- extractTimestamp(ts2),
77
- extractResult(readBuffer, Float32Array),
78
- ]);
81
+ const resultPromise = extractResult(readBuffer, Float32Array);
82
+ readBuffer = null; // ownership transferred — extractResult's own finally destroys it
79
83
 
80
- // sqrt is taken on CPU after the GPU sum-of-squares reduction
81
- const nrm2 = Math.sqrt(sqsumArr[0]);
84
+ const [gpuTime1, gpuTime2, sqsumArr] = await Promise.all([
85
+ extractTimestamp(ts1),
86
+ extractTimestamp(ts2),
87
+ resultPromise,
88
+ ]);
89
+
90
+ // sqrt is taken on CPU after the GPU sum-of-squares reduction
91
+ const nrm2 = Math.sqrt(sqsumArr[0]);
82
92
 
83
- if (xIsGpu) {
84
- destroyBuffers(partialsBuffer, resultBuffer, paramsBuffer, readBuffer);
85
93
  if (gpuTime1 !== undefined && gpuTime2 !== undefined)
86
94
  return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
87
95
  return { nrm2 };
96
+ } finally {
97
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
98
+ if (partialsBuffer) destroyBuffers(partialsBuffer);
99
+ if (resultBuffer) destroyBuffers(resultBuffer);
100
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
101
+ // Only reached if submit(enc2) threw before ownership was transferred above.
102
+ if (readBuffer) destroyBuffers(readBuffer);
88
103
  }
89
-
90
- destroyBuffers(
91
- xBuffer,
92
- partialsBuffer,
93
- resultBuffer,
94
- paramsBuffer,
95
- readBuffer,
96
- );
97
- if (gpuTime1 !== undefined && gpuTime2 !== undefined)
98
- return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
99
- return { nrm2 };
100
104
  }