wgblas 0.1.0 → 0.1.1

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/package.json CHANGED
@@ -1,6 +1,6 @@
1
1
  {
2
2
  "name": "wgblas",
3
- "version": "0.1.0",
3
+ "version": "0.1.1",
4
4
  "description": "BLAS on WebGPU",
5
5
  "type": "module",
6
6
  "main": "index.mjs",
@@ -65,6 +65,7 @@
65
65
  "index.d.mts",
66
66
  "src/**/*.mjs",
67
67
  "src/**/*.d.mts",
68
+ "src/**/*.wgsl",
68
69
  "dist/wgblas.browser.js"
69
70
  ],
70
71
  "repository": {
@@ -0,0 +1,59 @@
1
+ // isamax: returns index of element with largest absolute value
2
+ // pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/argmax.wgsl.
3
+
4
+ @group(0) @binding(0) var<storage, read> x: array<f32>;
5
+ @group(0) @binding(1) var<storage, read_write> partials_val: array<f32>;
6
+ @group(0) @binding(2) var<storage, read_write> partials_idx: array<u32>;
7
+ @group(0) @binding(3) var<uniform> params: Params;
8
+
9
+ struct Params {
10
+ n: u32,
11
+ x_inc: u32,
12
+ }
13
+
14
+ const WGS: u32 = 64;
15
+
16
+ var<workgroup> tile_val: array<f32, 64>;
17
+ var<workgroup> tile_idx: array<u32, 64>;
18
+
19
+ @compute @workgroup_size(64)
20
+ fn main(
21
+ @builtin(global_invocation_id) gid: vec3u,
22
+ @builtin(local_invocation_id) lid: vec3u,
23
+ @builtin(workgroup_id) wgid: vec3u,
24
+ @builtin(num_workgroups) num_wg: vec3u,
25
+ ) {
26
+ // -1.0 is a safe sentinel: any |x[i]| >= 0 beats it,
27
+ // so workgroups with no elements lose gracefully in the epilogue.
28
+ var best_val: f32 = -1.0;
29
+ var best_idx: u32 = 0u;
30
+
31
+ for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
32
+ let v = abs(x[id * params.x_inc]);
33
+ if (v > best_val) {
34
+ best_val = v;
35
+ best_idx = id;
36
+ }
37
+ }
38
+
39
+ tile_val[lid.x] = best_val;
40
+ tile_idx[lid.x] = best_idx;
41
+ workgroupBarrier();
42
+
43
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
44
+ if (lid.x < s) {
45
+ let a_val = tile_val[lid.x];
46
+ let b_val = tile_val[lid.x + s];
47
+ if (b_val > a_val || (b_val == a_val && tile_idx[lid.x + s] < tile_idx[lid.x])) {
48
+ tile_val[lid.x] = b_val;
49
+ tile_idx[lid.x] = tile_idx[lid.x + s];
50
+ }
51
+ }
52
+ workgroupBarrier();
53
+ }
54
+
55
+ if (lid.x == 0u) {
56
+ partials_val[wgid.x] = tile_val[0];
57
+ partials_idx[wgid.x] = tile_idx[0];
58
+ }
59
+ }
@@ -0,0 +1,43 @@
1
+ // amax reduction: collapses 2*WGS (value, index) pairs into one index.
2
+ // dispatch: 1 workgroup of WGS threads.
3
+ // partials_val and partials_idx must have exactly 2*WGS entries.
4
+
5
+ @group(0) @binding(0) var<storage, read> partials_val: array<f32>;
6
+ @group(0) @binding(1) var<storage, read> partials_idx: array<u32>;
7
+ @group(0) @binding(2) var<storage, read_write> result: array<u32>;
8
+
9
+ const WGS: u32 = 64;
10
+
11
+ var<workgroup> tile_val: array<f32, 64>;
12
+ var<workgroup> tile_idx: array<u32, 64>;
13
+
14
+ @compute @workgroup_size(64)
15
+ fn reduce(
16
+ @builtin(local_invocation_id) lid: vec3u,
17
+ ) {
18
+ let i = lid.x;
19
+ let a_val = partials_val[i];
20
+ let b_val = partials_val[i + WGS];
21
+ if (b_val > a_val || (b_val == a_val && partials_idx[i + WGS] < partials_idx[i])) {
22
+ tile_val[i] = b_val;
23
+ tile_idx[i] = partials_idx[i + WGS];
24
+ } else {
25
+ tile_val[i] = a_val;
26
+ tile_idx[i] = partials_idx[i];
27
+ }
28
+ workgroupBarrier();
29
+
30
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
31
+ if (i < s) {
32
+ let c_val = tile_val[i];
33
+ let d_val = tile_val[i + s];
34
+ if (d_val > c_val || (d_val == c_val && tile_idx[i + s] < tile_idx[i])) {
35
+ tile_val[i] = d_val;
36
+ tile_idx[i] = tile_idx[i + s];
37
+ }
38
+ }
39
+ workgroupBarrier();
40
+ }
41
+
42
+ if (i == 0u) { result[0] = tile_idx[0]; }
43
+ }
@@ -0,0 +1,26 @@
1
+ // sum reduction: collapses 2*WGS partials into one scalar.
2
+ // dispatch: 1 workgroup of WGS threads.
3
+ // partials must have exactly 2*WGS entries.
4
+
5
+ @group(0) @binding(0) var<storage, read> partials: array<f32>;
6
+ @group(0) @binding(1) var<storage, read_write> result: array<f32>;
7
+
8
+ const WGS: u32 = 64;
9
+
10
+ var<workgroup> tile: array<f32, 64>;
11
+
12
+ @compute @workgroup_size(64)
13
+ fn reduce(
14
+ @builtin(local_invocation_id) lid: vec3u,
15
+ ) {
16
+ let i = lid.x;
17
+ tile[i] = partials[i] + partials[i + WGS];
18
+ workgroupBarrier();
19
+
20
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
21
+ if (i < s) { tile[i] += tile[i + s]; }
22
+ workgroupBarrier();
23
+ }
24
+
25
+ if (i == 0u) { result[0] = tile[0]; }
26
+ }
@@ -0,0 +1,37 @@
1
+ // sasum: result = sum(|x[i]|)
2
+ // pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/abssum.wgsl.
3
+
4
+ @group(0) @binding(0) var<storage, read> x: array<f32>;
5
+ @group(0) @binding(1) var<storage, read_write> partials: array<f32>;
6
+ @group(0) @binding(2) var<uniform> params: Params;
7
+
8
+ struct Params {
9
+ n: u32,
10
+ x_inc: u32,
11
+ }
12
+
13
+ const WGS: u32 = 64;
14
+
15
+ var<workgroup> tile: array<f32, 64>;
16
+
17
+ @compute @workgroup_size(64)
18
+ fn main(
19
+ @builtin(global_invocation_id) gid: vec3u,
20
+ @builtin(local_invocation_id) lid: vec3u,
21
+ @builtin(workgroup_id) wgid: vec3u,
22
+ @builtin(num_workgroups) num_wg: vec3u,
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]);
27
+ }
28
+ tile[lid.x] = acc;
29
+ workgroupBarrier();
30
+
31
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
32
+ if (lid.x < s) { tile[lid.x] += tile[lid.x + s]; }
33
+ workgroupBarrier();
34
+ }
35
+
36
+ if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
37
+ }
@@ -0,0 +1,25 @@
1
+ // saxpy: y = alpha * x + y
2
+
3
+ @group(0) @binding(0) var<storage, read> x: array<f32>;
4
+ @group(0) @binding(1) var<storage, read_write> y: array<f32>;
5
+
6
+ struct Params {
7
+ n: u32,
8
+ alpha: f32,
9
+ x_inc: u32,
10
+ y_inc: u32,
11
+ }
12
+
13
+ @group(0) @binding(2) var<uniform> params: Params;
14
+
15
+ const WGS: u32 = 64;
16
+
17
+ @compute @workgroup_size(64)
18
+ fn main(
19
+ @builtin(global_invocation_id) gid: vec3u,
20
+ @builtin(num_workgroups) num_wg: vec3u,
21
+ ) {
22
+ for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
23
+ y[id * params.y_inc] = params.alpha * x[id * params.x_inc] + y[id * params.y_inc];
24
+ }
25
+ }
@@ -0,0 +1,24 @@
1
+ // scopy: y = x
2
+
3
+ @group(0) @binding(0) var<storage, read> x: array<f32>;
4
+ @group(0) @binding(1) var<storage, read_write> y: array<f32>;
5
+
6
+ struct Params {
7
+ n: u32,
8
+ x_inc: u32,
9
+ y_inc: u32,
10
+ }
11
+
12
+ @group(0) @binding(2) var<uniform> params: Params;
13
+
14
+ const WGS: u32 = 64;
15
+
16
+ @compute @workgroup_size(64)
17
+ fn main(
18
+ @builtin(global_invocation_id) gid: vec3u,
19
+ @builtin(num_workgroups) num_wg: vec3u,
20
+ ) {
21
+ for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
22
+ y[id * params.y_inc] = x[id * params.x_inc];
23
+ }
24
+ }
@@ -0,0 +1,39 @@
1
+ // sdot: result = sum(x[i] * y[i])
2
+ // pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sum.wgsl.
3
+
4
+ @group(0) @binding(0) var<storage, read> x: array<f32>;
5
+ @group(0) @binding(1) var<storage, read> y: array<f32>;
6
+ @group(0) @binding(2) var<storage, read_write> partials: array<f32>;
7
+ @group(0) @binding(3) var<uniform> params: Params;
8
+
9
+ struct Params {
10
+ n: u32,
11
+ x_inc: u32,
12
+ y_inc: u32,
13
+ }
14
+
15
+ const WGS: u32 = 64;
16
+
17
+ var<workgroup> tile: array<f32, 64>;
18
+
19
+ @compute @workgroup_size(64)
20
+ fn main(
21
+ @builtin(global_invocation_id) gid: vec3u,
22
+ @builtin(local_invocation_id) lid: vec3u,
23
+ @builtin(workgroup_id) wgid: vec3u,
24
+ @builtin(num_workgroups) num_wg: vec3u,
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];
29
+ }
30
+ tile[lid.x] = acc;
31
+ workgroupBarrier();
32
+
33
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
34
+ if (lid.x < s) { tile[lid.x] += tile[lid.x + s]; }
35
+ workgroupBarrier();
36
+ }
37
+
38
+ if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
39
+ }
@@ -0,0 +1,38 @@
1
+ // snrm2: result = sqrt(sum(x[i] * x[i]))
2
+ // pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sqsum.wgsl.
3
+
4
+ @group(0) @binding(0) var<storage, read> x: array<f32>;
5
+ @group(0) @binding(1) var<storage, read_write> partials: array<f32>;
6
+ @group(0) @binding(2) var<uniform> params: Params;
7
+
8
+ struct Params {
9
+ n: u32,
10
+ x_inc: u32,
11
+ }
12
+
13
+ const WGS: u32 = 64;
14
+
15
+ var<workgroup> tile: array<f32, 64>;
16
+
17
+ @compute @workgroup_size(64)
18
+ fn main(
19
+ @builtin(global_invocation_id) gid: vec3u,
20
+ @builtin(local_invocation_id) lid: vec3u,
21
+ @builtin(workgroup_id) wgid: vec3u,
22
+ @builtin(num_workgroups) num_wg: vec3u,
23
+ ) {
24
+ var acc: f32 = 0.0;
25
+ for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
26
+ let v = x[id * params.x_inc];
27
+ acc += v * v;
28
+ }
29
+ tile[lid.x] = acc;
30
+ workgroupBarrier();
31
+
32
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
33
+ if (lid.x < s) { tile[lid.x] += tile[lid.x + s]; }
34
+ workgroupBarrier();
35
+ }
36
+
37
+ if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
38
+ }
@@ -0,0 +1,29 @@
1
+ // srot: x = c*x + s*y, y = -s*x + c*y
2
+
3
+ @group(0) @binding(0) var<storage, read_write> x: array<f32>;
4
+ @group(0) @binding(1) var<storage, read_write> y: array<f32>;
5
+
6
+ struct Params {
7
+ n: u32,
8
+ c: f32,
9
+ s: f32,
10
+ x_inc: u32,
11
+ y_inc: u32,
12
+ }
13
+
14
+ @group(0) @binding(2) var<uniform> params: Params;
15
+
16
+ const WGS: u32 = 64;
17
+
18
+ @compute @workgroup_size(64)
19
+ fn main(
20
+ @builtin(global_invocation_id) gid: vec3u,
21
+ @builtin(num_workgroups) num_wg: vec3u,
22
+ ) {
23
+ for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
24
+ let xi = x[id * params.x_inc];
25
+ let yi = y[id * params.y_inc];
26
+ x[id * params.x_inc] = params.c * xi + params.s * yi;
27
+ y[id * params.y_inc] = -params.s * xi + params.c * yi;
28
+ }
29
+ }
@@ -0,0 +1,50 @@
1
+ // srotm: applies modified Givens rotation H to vectors x and y.
2
+ // param[0] = flag: -1 (full H), 0 (unit diagonal), 1 (unit off-diagonal)
3
+ // param = [ flag, h11, h21, h12, h22 ]
4
+ // flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
5
+
6
+ @group(0) @binding(0) var<storage, read_write> x: array<f32>;
7
+ @group(0) @binding(1) var<storage, read_write> y: array<f32>;
8
+ @group(0) @binding(2) var<storage, read> param: array<f32>;
9
+
10
+ struct Params {
11
+ n: u32,
12
+ x_inc: u32,
13
+ y_inc: u32,
14
+ }
15
+
16
+ @group(0) @binding(3) var<uniform> params: Params;
17
+
18
+ const WGS: u32 = 64;
19
+
20
+ @compute @workgroup_size(64)
21
+ fn main(
22
+ @builtin(global_invocation_id) gid: vec3u,
23
+ @builtin(num_workgroups) num_wg: vec3u,
24
+ ) {
25
+ let flag = param[0];
26
+
27
+ var h11: f32; var h12: f32;
28
+ var h21: f32; var h22: f32;
29
+
30
+ if (flag == -1.0) {
31
+ // full 2x2 matrix
32
+ h11 = param[1]; h21 = param[2];
33
+ h12 = param[3]; h22 = param[4];
34
+ } else if (flag == 0.0) {
35
+ // diagonal fixed at 1
36
+ h11 = 1.0; h21 = param[2];
37
+ h12 = param[3]; h22 = 1.0;
38
+ } else if (flag == 1.0) {
39
+ // flag == 1.0: off-diagonal fixed at +1 / -1
40
+ h11 = param[1]; h21 = -1.0;
41
+ h12 = 1.0; h22 = param[4];
42
+ }
43
+
44
+ for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
45
+ let xi = x[id * params.x_inc];
46
+ let yi = y[id * params.y_inc];
47
+ x[id * params.x_inc] = h11 * xi + h12 * yi;
48
+ y[id * params.y_inc] = h21 * xi + h22 * yi;
49
+ }
50
+ }
@@ -0,0 +1,23 @@
1
+ // sscal: x = alpha * x
2
+
3
+ @group(0) @binding(0) var<storage, read_write> x: array<f32>;
4
+
5
+ struct Params {
6
+ n: u32,
7
+ alpha: f32,
8
+ x_inc: u32,
9
+ }
10
+
11
+ @group(0) @binding(1) var<uniform> params: Params;
12
+
13
+ const WGS: u32 = 64;
14
+
15
+ @compute @workgroup_size(64)
16
+ fn main(
17
+ @builtin(global_invocation_id) gid: vec3u,
18
+ @builtin(num_workgroups) num_wg: vec3u,
19
+ ) {
20
+ for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
21
+ x[id * params.x_inc] = params.alpha * x[id * params.x_inc];
22
+ }
23
+ }
@@ -0,0 +1,26 @@
1
+ // sswap: x <-> y
2
+
3
+ @group(0) @binding(0) var<storage, read_write> x: array<f32>;
4
+ @group(0) @binding(1) var<storage, read_write> y: array<f32>;
5
+
6
+ struct Params {
7
+ n: u32,
8
+ x_inc: u32,
9
+ y_inc: u32,
10
+ }
11
+
12
+ @group(0) @binding(2) var<uniform> params: Params;
13
+
14
+ const WGS: u32 = 64;
15
+
16
+ @compute @workgroup_size(64)
17
+ fn main(
18
+ @builtin(global_invocation_id) gid: vec3u,
19
+ @builtin(num_workgroups) num_wg: vec3u,
20
+ ) {
21
+ for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
22
+ let temp = x[id * params.x_inc];
23
+ x[id * params.x_inc] = y[id * params.y_inc];
24
+ y[id * params.y_inc] = temp;
25
+ }
26
+ }