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 +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/package.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
{
|
|
2
2
|
"name": "wgblas",
|
|
3
|
-
"version": "0.1.
|
|
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
|
+
}
|