wgblas 0.1.0 → 0.1.2

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 CHANGED
@@ -51,24 +51,21 @@ Restart Firefox after making changes.
51
51
 
52
52
  ## Requirements
53
53
 
54
- - Node.js 18+
54
+ - Node.js 22+
55
55
 
56
- ## Example usage
57
-
58
- This package is not yet published to npm. Clone the repo and install it locally in your project:
56
+ ## Installation
59
57
 
60
58
  ```sh
61
- git clone https://github.com/manit2004/wgblas.git
62
- cd your-project
63
- npm install /path/to/wgblas
59
+ npm install wgblas
64
60
  ```
65
61
 
62
+ ## Example usage
63
+
66
64
  ### Example Code Snippet
67
65
 
68
66
  ```js
69
- import { init, cleanup } from "wgblas";
67
+ import { init, cleanup, randomFloat32Array } from "wgblas";
70
68
  import { sscal } from "wgblas/sscal";
71
- import { randomFloat32Array } from "wgblas/util/random";
72
69
 
73
70
  const device = await init();
74
71
 
@@ -92,7 +89,7 @@ No bundler needed. Load the pre-built browser bundle from the CDN and use `windo
92
89
  <head>
93
90
  <meta charset="UTF-8" />
94
91
  <title>sscal — wgblas browser example</title>
95
- <script src="https://manit2004.github.io/wgblas/wgblas.browser.js"></script>
92
+ <script src="https://unpkg.com/wgblas/dist/wgblas.browser.js"></script>
96
93
  </head>
97
94
  <body>
98
95
  <pre id="out">Running…</pre>
@@ -126,11 +123,10 @@ No bundler needed. Load the pre-built browser bundle from the CDN and use `windo
126
123
  `GpuVector` keeps data resident on the GPU between operations — upload once, chain any number of operations, read back once. This eliminates the redundant uploads and readbacks between steps, which are often more expensive than the compute itself.
127
124
 
128
125
  ```js
129
- import { init, cleanup } from "wgblas";
126
+ import { init, cleanup, randomFloat32Array } from "wgblas";
130
127
  import { saxpy } from "wgblas/saxpy";
131
128
  import { sscal } from "wgblas/sscal";
132
129
  import { GpuVector } from "wgblas/classes/GpuVector";
133
- import { randomFloat32Array } from "wgblas/util/random";
134
130
 
135
131
  const device = await init();
136
132
 
package/package.json CHANGED
@@ -1,6 +1,6 @@
1
1
  {
2
2
  "name": "wgblas",
3
- "version": "0.1.0",
3
+ "version": "0.1.2",
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
+ }