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 +8 -12
- package/package.json +2 -1
- package/src/shaders/isamax.wgsl +59 -0
- package/src/shaders/reduction/argmax.wgsl +43 -0
- package/src/shaders/reduction/sum.wgsl +26 -0
- package/src/shaders/sasum.wgsl +37 -0
- package/src/shaders/saxpy.wgsl +25 -0
- package/src/shaders/scopy.wgsl +24 -0
- package/src/shaders/sdot.wgsl +39 -0
- package/src/shaders/snrm2.wgsl +38 -0
- package/src/shaders/srot.wgsl +29 -0
- package/src/shaders/srotm.wgsl +50 -0
- package/src/shaders/sscal.wgsl +23 -0
- package/src/shaders/sswap.wgsl +26 -0
package/README.md
CHANGED
|
@@ -51,24 +51,21 @@ Restart Firefox after making changes.
|
|
|
51
51
|
|
|
52
52
|
## Requirements
|
|
53
53
|
|
|
54
|
-
- Node.js
|
|
54
|
+
- Node.js 22+
|
|
55
55
|
|
|
56
|
-
##
|
|
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
|
-
|
|
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://
|
|
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.
|
|
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
|
+
}
|