wgblas 0.1.2 → 1.0.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (55) hide show
  1. package/README.md +3 -0
  2. package/dist/wgblas.browser.js +1078 -37
  3. package/index.d.mts +6 -0
  4. package/index.mjs +6 -0
  5. package/package.json +32 -1
  6. package/src/classes/GpuMatrix.d.mts +85 -0
  7. package/src/classes/GpuMatrix.mjs +91 -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 +118 -0
  21. package/src/sgemv/sgemv.mjs +141 -0
  22. package/src/shaders/browser-shaders.mjs +20 -0
  23. package/src/shaders/dasum.wgsl +98 -0
  24. package/src/shaders/f64add.wgsl +281 -0
  25. package/src/shaders/isamax.wgsl +32 -9
  26. package/src/shaders/reduction/sumF64.wgsl +49 -0
  27. package/src/shaders/sasum.wgsl +18 -4
  28. package/src/shaders/sdot.wgsl +18 -4
  29. package/src/shaders/sgemv_n.wgsl +75 -0
  30. package/src/shaders/sgemv_t.wgsl +65 -0
  31. package/src/shaders/snrm2.wgsl +22 -4
  32. package/src/shaders/ssymv.wgsl +69 -0
  33. package/src/shaders/strmv.wgsl +103 -0
  34. package/src/shaders/strsv_apply_inverse.wgsl +47 -0
  35. package/src/shaders/strsv_invert_block.wgsl +109 -0
  36. package/src/shaders/strsv_update.wgsl +75 -0
  37. package/src/snrm2/snrm2.mjs +56 -52
  38. package/src/srot/srot.mjs +57 -41
  39. package/src/srotm/srotm.mjs +54 -38
  40. package/src/sscal/sscal.mjs +43 -32
  41. package/src/sswap/sswap.mjs +49 -34
  42. package/src/ssymv/ssymv.d.mts +109 -0
  43. package/src/ssymv/ssymv.mjs +130 -0
  44. package/src/strmv/strmv.d.mts +109 -0
  45. package/src/strmv/strmv.mjs +132 -0
  46. package/src/strsv/strsv.d.mts +98 -0
  47. package/src/strsv/strsv.mjs +212 -0
  48. package/src/util/benchmark.mjs +1 -1
  49. package/src/util/bindgroup.mjs +14 -10
  50. package/src/util/buffer.mjs +7 -2
  51. package/src/util/compute.mjs +41 -15
  52. package/src/util/f64pack.mjs +152 -0
  53. package/src/util/pipeline.mjs +32 -17
  54. package/src/util/result.mjs +8 -4
  55. package/src/util/workgroup.mjs +10 -10
@@ -25,19 +25,42 @@ fn main(
25
25
  ) {
26
26
  // -1.0 is a safe sentinel: any |x[i]| >= 0 beats it,
27
27
  // so workgroups with no elements lose gracefully in the epilogue.
28
- var best_val: f32 = -1.0;
29
- var best_idx: u32 = 0u;
28
+ var best_val0: f32 = -1.0; var best_idx0: u32 = 0u;
29
+ var best_val1: f32 = -1.0; var best_idx1: u32 = 0u;
30
+ var best_val2: f32 = -1.0; var best_idx2: u32 = 0u;
31
+ var best_val3: f32 = -1.0; var best_idx3: u32 = 0u;
30
32
 
31
- for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
33
+ let stride = num_wg.x * WGS;
34
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
35
+
36
+ for (var id = gid.x; id < n4_floor; id += 4u * stride) {
37
+ let v0 = abs(x[ id * params.x_inc]);
38
+ let v1 = abs(x[(id + stride) * params.x_inc]);
39
+ let v2 = abs(x[(id + 2u * stride) * params.x_inc]);
40
+ let v3 = abs(x[(id + 3u * stride) * params.x_inc]);
41
+ if (v0 > best_val0) { best_val0 = v0; best_idx0 = id; }
42
+ if (v1 > best_val1) { best_val1 = v1; best_idx1 = id + stride; }
43
+ if (v2 > best_val2) { best_val2 = v2; best_idx2 = id + 2u * stride; }
44
+ if (v3 > best_val3) { best_val3 = v3; best_idx3 = id + 3u * stride; }
45
+ }
46
+ for (var id = n4_floor + gid.x; id < params.n; id += stride) {
32
47
  let v = abs(x[id * params.x_inc]);
33
- if (v > best_val) {
34
- best_val = v;
35
- best_idx = id;
36
- }
48
+ if (v > best_val0) { best_val0 = v; best_idx0 = id; }
49
+ }
50
+
51
+ // merge 4 independent lanes; prefer lower index on tie (first occurrence wins)
52
+ if (best_val1 > best_val0 || (best_val1 == best_val0 && best_idx1 < best_idx0)) {
53
+ best_val0 = best_val1; best_idx0 = best_idx1;
54
+ }
55
+ if (best_val2 > best_val0 || (best_val2 == best_val0 && best_idx2 < best_idx0)) {
56
+ best_val0 = best_val2; best_idx0 = best_idx2;
57
+ }
58
+ if (best_val3 > best_val0 || (best_val3 == best_val0 && best_idx3 < best_idx0)) {
59
+ best_val0 = best_val3; best_idx0 = best_idx3;
37
60
  }
38
61
 
39
- tile_val[lid.x] = best_val;
40
- tile_idx[lid.x] = best_idx;
62
+ tile_val[lid.x] = best_val0;
63
+ tile_idx[lid.x] = best_idx0;
41
64
  workgroupBarrier();
42
65
 
43
66
  for (var s = WGS / 2u; s > 0u; s >>= 1u) {
@@ -0,0 +1,49 @@
1
+ // sum reduction (f64): collapses 2*WGS partial [main, aux] pairs into one,
2
+ // using computeSum instead of plain f32 `+` (see reduction/sum.wgsl for the
3
+ // f32 original this mirrors).
4
+ // dispatch: 1 workgroup of WGS threads. partialsMain/partialsAux must have
5
+ // exactly 2*WGS entries each.
6
+ //
7
+ // Concatenated after f64add.wgsl by getPipeline (WGSL has no #include),
8
+ // reusing its decode/encode/computeSum and Packed struct — f64add.wgsl
9
+ // declares no bindings and no entry point of its own (just helper functions),
10
+ // so bindings here start at 0 and the entry point is simply `reduce_f64`.
11
+ //
12
+ // partialsAux/result's aux slot are array<u32>, not array<f32> — aux's bits
13
+ // must never pass through an f32-typed storage slot (NaN-bit-pattern
14
+ // corruption risk, see f64pack.mjs and the Packed struct comment above
15
+ // decode()/encode() in f64add.wgsl).
16
+
17
+ @group(0) @binding(0) var<storage, read> partialsMain: array<f32>;
18
+ @group(0) @binding(1) var<storage, read> partialsAux: array<u32>;
19
+ @group(0) @binding(2) var<storage, read_write> resultMain: array<f32, 1>;
20
+ @group(0) @binding(3) var<storage, read_write> resultAux: array<u32, 1>;
21
+
22
+ const WGS: u32 = 64;
23
+
24
+ var<workgroup> tile: array<Packed, 64>;
25
+
26
+ fn addPair(a: Packed, b: Packed) -> Packed {
27
+ return computeSum(decode(bitcast<u32>(a.main), a.aux), decode(bitcast<u32>(b.main), b.aux));
28
+ }
29
+
30
+ @compute @workgroup_size(64)
31
+ fn reduce_f64(
32
+ @builtin(local_invocation_id) lid: vec3u,
33
+ ) {
34
+ let i = lid.x;
35
+ let a = Packed(partialsMain[i], partialsAux[i]);
36
+ let b = Packed(partialsMain[i + WGS], partialsAux[i + WGS]);
37
+ tile[i] = addPair(a, b);
38
+ workgroupBarrier();
39
+
40
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
41
+ if (i < s) { tile[i] = addPair(tile[i], tile[i + s]); }
42
+ workgroupBarrier();
43
+ }
44
+
45
+ if (i == 0u) {
46
+ resultMain[0] = tile[0].main;
47
+ resultAux[0] = tile[0].aux;
48
+ }
49
+ }
@@ -21,11 +21,25 @@ fn main(
21
21
  @builtin(workgroup_id) wgid: vec3u,
22
22
  @builtin(num_workgroups) num_wg: vec3u,
23
23
  ) {
24
- var acc: f32 = 0.0;
25
- for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
26
- acc += abs(x[id * params.x_inc]);
24
+ var acc0: f32 = 0.0;
25
+ var acc1: f32 = 0.0;
26
+ var acc2: f32 = 0.0;
27
+ var acc3: f32 = 0.0;
28
+
29
+ let stride = num_wg.x * WGS;
30
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
31
+
32
+ for (var id = gid.x; id < n4_floor; id += 4u * stride) {
33
+ acc0 += abs(x[ id * params.x_inc]);
34
+ acc1 += abs(x[(id + stride) * params.x_inc]);
35
+ acc2 += abs(x[(id + 2u * stride) * params.x_inc]);
36
+ acc3 += abs(x[(id + 3u * stride) * params.x_inc]);
27
37
  }
28
- tile[lid.x] = acc;
38
+ for (var id = n4_floor + gid.x; id < params.n; id += stride) {
39
+ acc0 += abs(x[id * params.x_inc]);
40
+ }
41
+
42
+ tile[lid.x] = acc0 + acc1 + acc2 + acc3;
29
43
  workgroupBarrier();
30
44
 
31
45
  for (var s = WGS / 2u; s > 0u; s >>= 1u) {
@@ -23,11 +23,25 @@ fn main(
23
23
  @builtin(workgroup_id) wgid: vec3u,
24
24
  @builtin(num_workgroups) num_wg: vec3u,
25
25
  ) {
26
- var acc: f32 = 0.0;
27
- for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
28
- acc += x[id * params.x_inc] * y[id * params.y_inc];
26
+ var acc0: f32 = 0.0;
27
+ var acc1: f32 = 0.0;
28
+ var acc2: f32 = 0.0;
29
+ var acc3: f32 = 0.0;
30
+
31
+ let stride = num_wg.x * WGS;
32
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
33
+
34
+ for (var id = gid.x; id < n4_floor; id += 4u * stride) {
35
+ acc0 += x[ id * params.x_inc] * y[ id * params.y_inc];
36
+ acc1 += x[(id + stride) * params.x_inc] * y[(id + stride) * params.y_inc];
37
+ acc2 += x[(id + 2u * stride) * params.x_inc] * y[(id + 2u * stride) * params.y_inc];
38
+ acc3 += x[(id + 3u * stride) * params.x_inc] * y[(id + 3u * stride) * params.y_inc];
29
39
  }
30
- tile[lid.x] = acc;
40
+ for (var id = n4_floor + gid.x; id < params.n; id += stride) {
41
+ acc0 += x[id * params.x_inc] * y[id * params.y_inc];
42
+ }
43
+
44
+ tile[lid.x] = acc0 + acc1 + acc2 + acc3;
31
45
  workgroupBarrier();
32
46
 
33
47
  for (var s = WGS / 2u; s > 0u; s >>= 1u) {
@@ -0,0 +1,75 @@
1
+ // sgemv_n: y = alpha * A * x + beta * y (A is m×n row-major, no-transpose)
2
+ //
3
+ // One workgroup per output row, with a grid-stride outer loop so the shader
4
+ // still covers all rows when m exceeds maxComputeWorkgroupsPerDimension.
5
+ // Threads stride through A[row, :] and x with coalesced reads (consecutive
6
+ // threads → consecutive addresses). Four independent accumulators let the GPU
7
+ // pipeline memory requests across iterations (ILP=4), hiding the
8
+ // global-memory latency.
9
+
10
+ @group(0) @binding(0) var<storage, read> A: array<f32>;
11
+ @group(0) @binding(1) var<storage, read> x: array<f32>;
12
+ @group(0) @binding(2) var<storage, read_write> y: array<f32>;
13
+
14
+ struct Params {
15
+ m: u32,
16
+ n: u32,
17
+ alpha: f32,
18
+ beta: f32,
19
+ incx: u32,
20
+ incy: u32,
21
+ lda: u32,
22
+ }
23
+
24
+ @group(0) @binding(3) var<uniform> params: Params;
25
+
26
+ const WGS: u32 = 64u;
27
+ var<workgroup> scratch: array<f32, 64>;
28
+
29
+ @compute @workgroup_size(64)
30
+ fn main(
31
+ @builtin(workgroup_id) wgid: vec3u,
32
+ @builtin(local_invocation_id) lid: vec3u,
33
+ @builtin(num_workgroups) nwg: vec3u,
34
+ ) {
35
+ // Grid-stride loop: each workgroup handles ceil(m / nwg.x) rows.
36
+ for (var row = wgid.x; row < params.m; row += nwg.x) {
37
+ let row_base = row * params.lda;
38
+ var acc0: f32 = 0.0;
39
+ var acc1: f32 = 0.0;
40
+ var acc2: f32 = 0.0;
41
+ var acc3: f32 = 0.0;
42
+
43
+ // 4-unrolled loop: each iteration issues 4 independent loads for A and x.
44
+ // The accumulators are independent so the GPU can overlap the memory
45
+ // requests rather than serialising them behind a dependency chain.
46
+ let n4_floor = (params.n / (4u * WGS)) * (4u * WGS);
47
+ for (var j: u32 = lid.x; j < n4_floor; j += 4u * WGS) {
48
+ acc0 += A[row_base + j ] * x[ j * params.incx];
49
+ acc1 += A[row_base + j + WGS ] * x[(j + WGS) * params.incx];
50
+ acc2 += A[row_base + j + 2u * WGS ] * x[(j + 2u * WGS) * params.incx];
51
+ acc3 += A[row_base + j + 3u * WGS ] * x[(j + 3u * WGS) * params.incx];
52
+ }
53
+ // Scalar tail: at most 3*WGS elements left after the unrolled block.
54
+ for (var j: u32 = n4_floor + lid.x; j < params.n; j += WGS) {
55
+ acc0 += A[row_base + j] * x[j * params.incx];
56
+ }
57
+
58
+ // Parallel reduction: 64 → 32 → 16 → 8 → 4 → 2 → 1
59
+ scratch[lid.x] = acc0 + acc1 + acc2 + acc3;
60
+ workgroupBarrier();
61
+ for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
62
+ if lid.x < stride {
63
+ scratch[lid.x] += scratch[lid.x + stride];
64
+ }
65
+ workgroupBarrier();
66
+ }
67
+
68
+ if lid.x == 0u {
69
+ let yi = row * params.incy;
70
+ y[yi] = params.alpha * scratch[0] + params.beta * y[yi];
71
+ }
72
+ // All 64 threads must agree before the next row reuses scratch[].
73
+ workgroupBarrier();
74
+ }
75
+ }
@@ -0,0 +1,65 @@
1
+ // sgemv_t: y = alpha * A^T * x + beta * y (A is m×n row-major, transposed)
2
+ // each thread owns one column of A → one element of y (length n)
3
+ // tiles over x (length m) using shared memory; four independent accumulators
4
+ // let the GPU pipeline A reads across j within each tile (ILP=4)
5
+
6
+ @group(0) @binding(0) var<storage, read> A: array<f32>;
7
+ @group(0) @binding(1) var<storage, read> x: array<f32>;
8
+ @group(0) @binding(2) var<storage, read_write> y: array<f32>;
9
+
10
+ struct Params {
11
+ m: u32,
12
+ n: u32,
13
+ alpha: f32,
14
+ beta: f32,
15
+ incx: u32,
16
+ incy: u32,
17
+ lda: u32,
18
+ }
19
+
20
+ @group(0) @binding(3) var<uniform> params: Params;
21
+
22
+ const WGS: u32 = 64u;
23
+ var<workgroup> x_tile: array<f32, 64>;
24
+
25
+ @compute @workgroup_size(64)
26
+ fn main(
27
+ @builtin(global_invocation_id) gid: vec3u,
28
+ @builtin(local_invocation_id) lid: vec3u,
29
+ ) {
30
+ // each thread owns column col of A → output y[col]
31
+ let col = gid.x;
32
+ // tile over x (length m, the rows of A)
33
+ let m_floor = (params.m / WGS) * WGS;
34
+ var acc0: f32 = 0.0;
35
+ var acc1: f32 = 0.0;
36
+ var acc2: f32 = 0.0;
37
+ var acc3: f32 = 0.0;
38
+
39
+ for (var base = 0u; base < m_floor; base += WGS) {
40
+ // cooperative load: all 64 threads fill x_tile with x[base..base+WGS]
41
+ x_tile[lid.x] = x[(base + lid.x) * params.incx];
42
+ workgroupBarrier();
43
+
44
+ // 4-unrolled inner loop: 4 independent A reads let the GPU pipeline
45
+ // global-memory requests within each tile. WGS=64 divides by 4 exactly.
46
+ if (col < params.n) {
47
+ for (var j = 0u; j < WGS; j += 4u) {
48
+ acc0 += A[(base + j ) * params.lda + col] * x_tile[j ];
49
+ acc1 += A[(base + j + 1) * params.lda + col] * x_tile[j + 1];
50
+ acc2 += A[(base + j + 2) * params.lda + col] * x_tile[j + 2];
51
+ acc3 += A[(base + j + 3) * params.lda + col] * x_tile[j + 3];
52
+ }
53
+ }
54
+ workgroupBarrier();
55
+ }
56
+
57
+ if (col < params.n) {
58
+ // remainder: m not divisible by WGS — short loop, single accumulator fine
59
+ for (var k = m_floor; k < params.m; k++) {
60
+ acc0 += A[k * params.lda + col] * x[k * params.incx];
61
+ }
62
+ let yi = col * params.incy;
63
+ y[yi] = params.alpha * (acc0 + acc1 + acc2 + acc3) + params.beta * y[yi];
64
+ }
65
+ }
@@ -21,12 +21,30 @@ fn main(
21
21
  @builtin(workgroup_id) wgid: vec3u,
22
22
  @builtin(num_workgroups) num_wg: vec3u,
23
23
  ) {
24
- var acc: f32 = 0.0;
25
- for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
24
+ var acc0: f32 = 0.0;
25
+ var acc1: f32 = 0.0;
26
+ var acc2: f32 = 0.0;
27
+ var acc3: f32 = 0.0;
28
+
29
+ let stride = num_wg.x * WGS;
30
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
31
+
32
+ for (var id = gid.x; id < n4_floor; id += 4u * stride) {
33
+ let v0 = x[ id * params.x_inc];
34
+ let v1 = x[(id + stride) * params.x_inc];
35
+ let v2 = x[(id + 2u * stride) * params.x_inc];
36
+ let v3 = x[(id + 3u * stride) * params.x_inc];
37
+ acc0 += v0 * v0;
38
+ acc1 += v1 * v1;
39
+ acc2 += v2 * v2;
40
+ acc3 += v3 * v3;
41
+ }
42
+ for (var id = n4_floor + gid.x; id < params.n; id += stride) {
26
43
  let v = x[id * params.x_inc];
27
- acc += v * v;
44
+ acc0 += v * v;
28
45
  }
29
- tile[lid.x] = acc;
46
+
47
+ tile[lid.x] = acc0 + acc1 + acc2 + acc3;
30
48
  workgroupBarrier();
31
49
 
32
50
  for (var s = WGS / 2u; s > 0u; s >>= 1u) {
@@ -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,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
+ }