volvoxai 0.1.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.
- package/LICENSE +21 -0
- package/README.md +145 -0
- package/bin/volvox.js +72 -0
- package/dist/v0.1.0/volvoxai.js +4664 -0
- package/dist/v0.1.0/volvoxai.min.js +1848 -0
- package/dist/v0.1.0/volvoxai.wasm +0 -0
- package/dist/volvoxai.js +4664 -0
- package/dist/volvoxai.min.js +1848 -0
- package/dist/volvoxai.wasm +0 -0
- package/docs/README.md +22 -0
- package/docs/browser-runtime.md +87 -0
- package/docs/efficientdet_tflite_vs_volvoxai.md +445 -0
- package/docs/microkernel_optimization_guide.md +153 -0
- package/docs/model-format.md +108 -0
- package/docs/models.md +103 -0
- package/docs/native-runtime.md +189 -0
- package/docs/operation_list.md +232 -0
- package/docs/operator_fusion_patterns.md +58 -0
- package/docs/quickstart.md +115 -0
- package/docs/roadmap.md +19 -0
- package/docs/testing.md +97 -0
- package/docs/textbook/01-foundations.md +233 -0
- package/docs/textbook/02-tinystories-language-model.md +300 -0
- package/docs/textbook/03-efficientdet-vision-model.md +281 -0
- package/docs/textbook/04-precision-and-quantization.md +208 -0
- package/docs/textbook/05-inside-the-engine.md +155 -0
- package/docs/textbook/06-native-engine-architecture.md +338 -0
- package/docs/textbook/07-glossary-and-next-steps.md +258 -0
- package/docs/textbook/README.md +85 -0
- package/docs/textbook/ko/01-foundations.md +231 -0
- package/docs/textbook/ko/02-tinystories-language-model.md +300 -0
- package/docs/textbook/ko/03-efficientdet-vision-model.md +277 -0
- package/docs/textbook/ko/04-precision-and-quantization.md +206 -0
- package/docs/textbook/ko/05-inside-the-engine.md +154 -0
- package/docs/textbook/ko/06-native-engine-architecture.md +333 -0
- package/docs/textbook/ko/07-glossary-and-next-steps.md +253 -0
- package/docs/textbook/ko/README.md +83 -0
- package/docs/xnnpack_optimization_guide.md +197 -0
- package/js/CPUEngine.js +241 -0
- package/js/Graph.js +49 -0
- package/js/GraphExecutor.js +1020 -0
- package/js/GraphLoader.js +282 -0
- package/js/ShaderLibrary.js +236 -0
- package/js/Tensor.js +25 -0
- package/js/Tokenizer.js +266 -0
- package/js/VolvoxAI.js +130 -0
- package/js/WasmEngine.js +378 -0
- package/js/WebNNEngine.js +169 -0
- package/js/index.js +11 -0
- package/js/ops/add.js +31 -0
- package/js/ops/argMax.js +33 -0
- package/js/ops/averagePool2D.js +38 -0
- package/js/ops/batchNorm2D.js +28 -0
- package/js/ops/cast.js +19 -0
- package/js/ops/clip.js +15 -0
- package/js/ops/concat2.js +18 -0
- package/js/ops/conv1D.js +35 -0
- package/js/ops/conv2D.js +70 -0
- package/js/ops/convTranspose2D.js +45 -0
- package/js/ops/crossAttention.js +69 -0
- package/js/ops/crossSDPA.js +41 -0
- package/js/ops/dequantizeLinear.js +9 -0
- package/js/ops/div.js +15 -0
- package/js/ops/embedding.js +14 -0
- package/js/ops/expand.js +24 -0
- package/js/ops/gELU.js +9 -0
- package/js/ops/gather.js +51 -0
- package/js/ops/gatherElements.js +33 -0
- package/js/ops/globalAveragePool.js +21 -0
- package/js/ops/hardSigmoid.js +12 -0
- package/js/ops/hardSwish.js +12 -0
- package/js/ops/interp1D.js +25 -0
- package/js/ops/layerNorm.js +25 -0
- package/js/ops/leakyReLU.js +10 -0
- package/js/ops/logSoftmax.js +15 -0
- package/js/ops/matMul.js +35 -0
- package/js/ops/maxPool2D.js +36 -0
- package/js/ops/meanHeight.js +17 -0
- package/js/ops/mul.js +31 -0
- package/js/ops/nonMaxSuppression.js +72 -0
- package/js/ops/pReLU.js +11 -0
- package/js/ops/pad.js +35 -0
- package/js/ops/profileX.js +22 -0
- package/js/ops/profileY.js +22 -0
- package/js/ops/rMSNorm.js +14 -0
- package/js/ops/reLU.js +8 -0
- package/js/ops/reduceMean.js +17 -0
- package/js/ops/reduceSum.js +19 -0
- package/js/ops/reshape.js +6 -0
- package/js/ops/resize.js +44 -0
- package/js/ops/sDPA.js +44 -0
- package/js/ops/siLU.js +8 -0
- package/js/ops/sigmoid.js +6 -0
- package/js/ops/slice.js +36 -0
- package/js/ops/softmax.js +18 -0
- package/js/ops/spatialSoftargmaxY.js +28 -0
- package/js/ops/split.js +24 -0
- package/js/ops/sub.js +11 -0
- package/js/ops/tanh.js +7 -0
- package/js/ops/transpose.js +34 -0
- package/js/ops/upsample2x.js +23 -0
- package/js/ops/where.js +15 -0
- package/package.json +33 -0
- package/shaders/add.wgsl +13 -0
- package/shaders/add3Relu.wgsl +23 -0
- package/shaders/addRelu.wgsl +22 -0
- package/shaders/averagePool2D.wgsl +24 -0
- package/shaders/batchNorm2D.wgsl +21 -0
- package/shaders/binaryBroadcast.wgsl +34 -0
- package/shaders/broadcastBinary.wgsl +26 -0
- package/shaders/clip.wgsl +10 -0
- package/shaders/concat2.wgsl +16 -0
- package/shaders/concatCopy.wgsl +10 -0
- package/shaders/concatSigmoidCopy.wgsl +16 -0
- package/shaders/conv1D.wgsl +37 -0
- package/shaders/conv2D.wgsl +80 -0
- package/shaders/conv2DDepthwise4.wgsl +74 -0
- package/shaders/conv2DDepthwise8.wgsl +66 -0
- package/shaders/conv2DPointwise16.wgsl +67 -0
- package/shaders/conv2DPointwise16Tile.wgsl +86 -0
- package/shaders/conv2DPointwise8.wgsl +85 -0
- package/shaders/conv2DPointwise8Vec2.wgsl +70 -0
- package/shaders/conv2DPointwise8Vec4.wgsl +65 -0
- package/shaders/conv2DRegularC3Out16.wgsl +75 -0
- package/shaders/convTranspose2D.wgsl +33 -0
- package/shaders/copy.wgsl +13 -0
- package/shaders/crossAttention.wgsl +140 -0
- package/shaders/crossAttentionF32.wgsl +98 -0
- package/shaders/crossSDPA.wgsl +74 -0
- package/shaders/dequantizeLinear.wgsl +14 -0
- package/shaders/div.wgsl +34 -0
- package/shaders/elementwise.wgsl +13 -0
- package/shaders/embedding.wgsl +22 -0
- package/shaders/expand.wgsl +18 -0
- package/shaders/gELU.wgsl +13 -0
- package/shaders/gather.wgsl +17 -0
- package/shaders/generalTranspose.wgsl +19 -0
- package/shaders/globalAveragePool.wgsl +19 -0
- package/shaders/hardSigmoid.wgsl +13 -0
- package/shaders/hardSwish.wgsl +13 -0
- package/shaders/interp1D.wgsl +28 -0
- package/shaders/layerNorm.wgsl +33 -0
- package/shaders/leakyReLU.wgsl +11 -0
- package/shaders/linearF32.wgsl +33 -0
- package/shaders/linearF32RowMajor.wgsl +24 -0
- package/shaders/linearInt8.wgsl +42 -0
- package/shaders/logSoftmax.wgsl +22 -0
- package/shaders/maxPool2D.wgsl +37 -0
- package/shaders/meanHeight.wgsl +18 -0
- package/shaders/mul.wgsl +32 -0
- package/shaders/nonMaxSuppression.wgsl +92 -0
- package/shaders/pReLU.wgsl +14 -0
- package/shaders/pad.wgsl +19 -0
- package/shaders/profileX.wgsl +28 -0
- package/shaders/profileY.wgsl +28 -0
- package/shaders/quantizeLinear.wgsl +69 -0
- package/shaders/rMSNorm.wgsl +21 -0
- package/shaders/reLU.wgsl +13 -0
- package/shaders/reduce.wgsl +17 -0
- package/shaders/resize.wgsl +52 -0
- package/shaders/sDPA.wgsl +71 -0
- package/shaders/siLU.wgsl +13 -0
- package/shaders/sigmoid.wgsl +13 -0
- package/shaders/slice.wgsl +26 -0
- package/shaders/softmax.wgsl +23 -0
- package/shaders/spatialSoftargmaxY.wgsl +32 -0
- package/shaders/split.wgsl +15 -0
- package/shaders/sub.wgsl +34 -0
- package/shaders/tanh.wgsl +13 -0
- package/shaders/upsample2x.wgsl +24 -0
- package/shaders/where.wgsl +12 -0
- package/volvoxai.wasm +0 -0
|
@@ -0,0 +1,92 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> boxes : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> scores : array<f32>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read_write> output : array<f32>;
|
|
4
|
+
|
|
5
|
+
struct Params {
|
|
6
|
+
batches : u32,
|
|
7
|
+
spatial : u32,
|
|
8
|
+
classes : u32,
|
|
9
|
+
max_output : u32,
|
|
10
|
+
output_rows : u32,
|
|
11
|
+
iou_threshold : f32,
|
|
12
|
+
score_threshold : f32,
|
|
13
|
+
}
|
|
14
|
+
@group(0) @binding(3) var<uniform> params : Params;
|
|
15
|
+
|
|
16
|
+
fn box_iou(b : u32, a_idx : u32, b_idx : u32) -> f32 {
|
|
17
|
+
let base_a = (b * params.spatial + a_idx) * 4u;
|
|
18
|
+
let base_b = (b * params.spatial + b_idx) * 4u;
|
|
19
|
+
let ay1 = boxes[base_a + 0u];
|
|
20
|
+
let ax1 = boxes[base_a + 1u];
|
|
21
|
+
let ay2 = boxes[base_a + 2u];
|
|
22
|
+
let ax2 = boxes[base_a + 3u];
|
|
23
|
+
let by1 = boxes[base_b + 0u];
|
|
24
|
+
let bx1 = boxes[base_b + 1u];
|
|
25
|
+
let by2 = boxes[base_b + 2u];
|
|
26
|
+
let bx2 = boxes[base_b + 3u];
|
|
27
|
+
let xx1 = max(ax1, bx1);
|
|
28
|
+
let yy1 = max(ay1, by1);
|
|
29
|
+
let xx2 = min(ax2, bx2);
|
|
30
|
+
let yy2 = min(ay2, by2);
|
|
31
|
+
let w = max(0.0, xx2 - xx1);
|
|
32
|
+
let h = max(0.0, yy2 - yy1);
|
|
33
|
+
let inter = w * h;
|
|
34
|
+
let area_a = max(0.0, ax2 - ax1) * max(0.0, ay2 - ay1);
|
|
35
|
+
let area_b = max(0.0, bx2 - bx1) * max(0.0, by2 - by1);
|
|
36
|
+
let denom = area_a + area_b - inter;
|
|
37
|
+
if (denom <= 0.0) { return 0.0; }
|
|
38
|
+
return inter / denom;
|
|
39
|
+
}
|
|
40
|
+
|
|
41
|
+
fn selected_suppresses(b : u32, c : u32, candidate : u32, out_start : u32, selected_count : u32) -> bool {
|
|
42
|
+
for (var j = 0u; j < selected_count; j = j + 1u) {
|
|
43
|
+
let row = out_start + j;
|
|
44
|
+
if (row >= params.output_rows) { return false; }
|
|
45
|
+
let prev_b = u32(output[row * 3u + 0u]);
|
|
46
|
+
let prev_c = u32(output[row * 3u + 1u]);
|
|
47
|
+
let prev_s = u32(output[row * 3u + 2u]);
|
|
48
|
+
if (prev_b == b && prev_c == c && box_iou(b, candidate, prev_s) > params.iou_threshold) {
|
|
49
|
+
return true;
|
|
50
|
+
}
|
|
51
|
+
}
|
|
52
|
+
return false;
|
|
53
|
+
}
|
|
54
|
+
|
|
55
|
+
@compute @workgroup_size(1)
|
|
56
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
57
|
+
if (global_id.x != 0u) { return; }
|
|
58
|
+
|
|
59
|
+
for (var i = 0u; i < params.output_rows; i = i + 1u) {
|
|
60
|
+
output[i * 3u + 0u] = -1.0;
|
|
61
|
+
output[i * 3u + 1u] = -1.0;
|
|
62
|
+
output[i * 3u + 2u] = -1.0;
|
|
63
|
+
}
|
|
64
|
+
|
|
65
|
+
var out_idx = 0u;
|
|
66
|
+
for (var b = 0u; b < params.batches; b = b + 1u) {
|
|
67
|
+
for (var c = 0u; c < params.classes; c = c + 1u) {
|
|
68
|
+
let class_start = out_idx;
|
|
69
|
+
var selected = 0u;
|
|
70
|
+
loop {
|
|
71
|
+
if (selected >= params.max_output || out_idx >= params.output_rows) { break; }
|
|
72
|
+
var best_score = params.score_threshold;
|
|
73
|
+
var best_s = params.spatial;
|
|
74
|
+
for (var s = 0u; s < params.spatial; s = s + 1u) {
|
|
75
|
+
let score = scores[b * (params.classes * params.spatial) + c * params.spatial + s];
|
|
76
|
+
if (score < best_score) { continue; }
|
|
77
|
+
if (selected_suppresses(b, c, s, class_start, selected)) { continue; }
|
|
78
|
+
if (best_s == params.spatial || score > best_score) {
|
|
79
|
+
best_score = score;
|
|
80
|
+
best_s = s;
|
|
81
|
+
}
|
|
82
|
+
}
|
|
83
|
+
if (best_s == params.spatial) { break; }
|
|
84
|
+
output[out_idx * 3u + 0u] = f32(b);
|
|
85
|
+
output[out_idx * 3u + 1u] = f32(c);
|
|
86
|
+
output[out_idx * 3u + 2u] = f32(best_s);
|
|
87
|
+
out_idx = out_idx + 1u;
|
|
88
|
+
selected = selected + 1u;
|
|
89
|
+
}
|
|
90
|
+
}
|
|
91
|
+
}
|
|
92
|
+
}
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> weight : array<f32>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read_write> output : array<f32>;
|
|
4
|
+
struct Params { size : u32, c : u32 }
|
|
5
|
+
@group(0) @binding(3) var<uniform> params : Params;
|
|
6
|
+
@compute @workgroup_size(64)
|
|
7
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
8
|
+
let idx = global_id.x;
|
|
9
|
+
if (idx >= params.size) { return; }
|
|
10
|
+
let c_idx = idx % params.c;
|
|
11
|
+
let alpha = weight[c_idx];
|
|
12
|
+
let v = input[idx];
|
|
13
|
+
if (v > 0.0) { output[idx] = v; } else { output[idx] = v * alpha; }
|
|
14
|
+
}
|
package/shaders/pad.wgsl
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
struct Params { b: u32, in_h: u32, in_w: u32, c: u32, out_h: u32, out_w: u32, pt: u32, pl: u32, val: f32 }
|
|
4
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
5
|
+
@compute @workgroup_size(64)
|
|
6
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
7
|
+
let idx = global_id.x;
|
|
8
|
+
let total = params.b * params.out_h * params.out_w * params.c;
|
|
9
|
+
if (idx >= total) { return; }
|
|
10
|
+
let c = idx % params.c;
|
|
11
|
+
let x = (idx / params.c) % params.out_w;
|
|
12
|
+
let y = (idx / (params.c * params.out_w)) % params.out_h;
|
|
13
|
+
let b = idx / (params.c * params.out_w * params.out_h);
|
|
14
|
+
if (y >= params.pt && y < params.pt + params.in_h && x >= params.pl && x < params.pl + params.in_w) {
|
|
15
|
+
output[idx] = input[((b * params.in_h + (y - params.pt)) * params.in_w + (x - params.pl)) * params.c + c];
|
|
16
|
+
} else {
|
|
17
|
+
output[idx] = params.val;
|
|
18
|
+
}
|
|
19
|
+
}
|
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
|
|
4
|
+
struct Params { in_h : u32, in_w : u32, in_c : u32 }
|
|
5
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
6
|
+
|
|
7
|
+
@compute @workgroup_size(64)
|
|
8
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
9
|
+
let x = global_id.x;
|
|
10
|
+
let c = global_id.y;
|
|
11
|
+
|
|
12
|
+
if (x >= params.in_w || c >= params.in_c) { return; }
|
|
13
|
+
|
|
14
|
+
var max_val = -1e38;
|
|
15
|
+
var sum_val = 0.0;
|
|
16
|
+
|
|
17
|
+
for (var y = 0u; y < params.in_h; y = y + 1u) {
|
|
18
|
+
let val = input[(y * params.in_w + x) * params.in_c + c];
|
|
19
|
+
if (val > max_val) { max_val = val; }
|
|
20
|
+
sum_val = sum_val + val;
|
|
21
|
+
}
|
|
22
|
+
|
|
23
|
+
let out_max_idx = c * params.in_w + x;
|
|
24
|
+
let out_mean_idx = (c + params.in_c) * params.in_w + x;
|
|
25
|
+
|
|
26
|
+
output[out_max_idx] = max_val;
|
|
27
|
+
output[out_mean_idx] = sum_val / f32(params.in_h);
|
|
28
|
+
}
|
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
|
|
4
|
+
struct Params { in_h : u32, in_w : u32, in_c : u32 }
|
|
5
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
6
|
+
|
|
7
|
+
@compute @workgroup_size(64)
|
|
8
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
9
|
+
let y = global_id.x;
|
|
10
|
+
let c = global_id.y;
|
|
11
|
+
|
|
12
|
+
if (y >= params.in_h || c >= params.in_c) { return; }
|
|
13
|
+
|
|
14
|
+
var max_val = -1e38;
|
|
15
|
+
var sum_val = 0.0;
|
|
16
|
+
|
|
17
|
+
for (var x = 0u; x < params.in_w; x = x + 1u) {
|
|
18
|
+
let val = input[(y * params.in_w + x) * params.in_c + c];
|
|
19
|
+
if (val > max_val) { max_val = val; }
|
|
20
|
+
sum_val = sum_val + val;
|
|
21
|
+
}
|
|
22
|
+
|
|
23
|
+
let out_max_idx = c * params.in_h + y;
|
|
24
|
+
let out_mean_idx = (c + params.in_c) * params.in_h + y;
|
|
25
|
+
|
|
26
|
+
output[out_max_idx] = max_val;
|
|
27
|
+
output[out_mean_idx] = sum_val / f32(params.in_w);
|
|
28
|
+
}
|
|
@@ -0,0 +1,69 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<u32>;
|
|
3
|
+
|
|
4
|
+
struct Params {
|
|
5
|
+
size : u32,
|
|
6
|
+
input_zp : i32,
|
|
7
|
+
output_zp : i32,
|
|
8
|
+
has_input_scale : u32,
|
|
9
|
+
input_scale : f32,
|
|
10
|
+
output_scale : f32,
|
|
11
|
+
pad0 : u32,
|
|
12
|
+
pad1 : u32,
|
|
13
|
+
}
|
|
14
|
+
|
|
15
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
16
|
+
|
|
17
|
+
fn clamp_i8(v : i32) -> i32 {
|
|
18
|
+
return min(max(v, -128), 127);
|
|
19
|
+
}
|
|
20
|
+
|
|
21
|
+
fn i8_byte(v : i32) -> u32 {
|
|
22
|
+
if (v < 0) {
|
|
23
|
+
return u32(v + 256) & 255u;
|
|
24
|
+
}
|
|
25
|
+
return u32(v) & 255u;
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
fn round_even(v : f32) -> f32 {
|
|
29
|
+
let lo = floor(v);
|
|
30
|
+
let frac = v - lo;
|
|
31
|
+
if (frac < 0.5) {
|
|
32
|
+
return lo;
|
|
33
|
+
}
|
|
34
|
+
if (frac > 0.5) {
|
|
35
|
+
return lo + 1.0;
|
|
36
|
+
}
|
|
37
|
+
if (floor(lo * 0.5) == lo * 0.5) {
|
|
38
|
+
return lo;
|
|
39
|
+
}
|
|
40
|
+
return lo + 1.0;
|
|
41
|
+
}
|
|
42
|
+
|
|
43
|
+
fn quantize_one(idx : u32) -> u32 {
|
|
44
|
+
if (idx >= params.size) {
|
|
45
|
+
return 0u;
|
|
46
|
+
}
|
|
47
|
+
var v = input[idx];
|
|
48
|
+
if (params.has_input_scale == 1u) {
|
|
49
|
+
v = (v - f32(params.input_zp)) * params.input_scale;
|
|
50
|
+
}
|
|
51
|
+
let qf = round_even(v / params.output_scale + f32(params.output_zp));
|
|
52
|
+
let q = clamp_i8(i32(min(max(qf, -128.0), 127.0)));
|
|
53
|
+
return i8_byte(q);
|
|
54
|
+
}
|
|
55
|
+
|
|
56
|
+
@compute @workgroup_size(64)
|
|
57
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
58
|
+
let word_idx = global_id.x;
|
|
59
|
+
let base = word_idx * 4u;
|
|
60
|
+
if (base >= params.size) {
|
|
61
|
+
return;
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
var packed = quantize_one(base);
|
|
65
|
+
packed = packed | (quantize_one(base + 1u) << 8u);
|
|
66
|
+
packed = packed | (quantize_one(base + 2u) << 16u);
|
|
67
|
+
packed = packed | (quantize_one(base + 3u) << 24u);
|
|
68
|
+
output[word_idx] = packed;
|
|
69
|
+
}
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> weight : array<f32>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read_write> output : array<f32>;
|
|
4
|
+
struct Params { seq_len : u32, d_model : u32, eps : f32 }
|
|
5
|
+
@group(0) @binding(3) var<uniform> params : Params;
|
|
6
|
+
@compute @workgroup_size(64)
|
|
7
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
8
|
+
let s = global_id.x;
|
|
9
|
+
if (s >= params.seq_len) { return; }
|
|
10
|
+
var var_val = 0.0;
|
|
11
|
+
let offset = s * params.d_model;
|
|
12
|
+
for (var d = 0u; d < params.d_model; d = d + 1u) {
|
|
13
|
+
let v = input[offset + d];
|
|
14
|
+
var_val = var_val + v * v;
|
|
15
|
+
}
|
|
16
|
+
var_val = var_val / f32(params.d_model);
|
|
17
|
+
let inv_std = 1.0 / sqrt(var_val + params.eps);
|
|
18
|
+
for (var d = 0u; d < params.d_model; d = d + 1u) {
|
|
19
|
+
output[offset + d] = input[offset + d] * inv_std * weight[d];
|
|
20
|
+
}
|
|
21
|
+
}
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
struct Params { size : u32 }
|
|
4
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
5
|
+
@compute @workgroup_size(64)
|
|
6
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
7
|
+
let idx = global_id.x;
|
|
8
|
+
if (idx >= params.size) { return; }
|
|
9
|
+
let x = input[idx];
|
|
10
|
+
var out_val = x;
|
|
11
|
+
if (x > 0.0) { out_val = x; } else { out_val = 0.0; }
|
|
12
|
+
output[idx] = out_val;
|
|
13
|
+
}
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
// Reduce over the last (innermost) axis: `b` rows of length `d` -> `b` outputs.
|
|
4
|
+
// `inv` scales the sum: 1.0 for ReduceSum, 1.0/d for ReduceMean.
|
|
5
|
+
struct Params { b : u32, d : u32, inv : f32 }
|
|
6
|
+
@group(0) @binding(2) var<uniform> p : Params;
|
|
7
|
+
@compute @workgroup_size(64)
|
|
8
|
+
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
|
|
9
|
+
let i = gid.x;
|
|
10
|
+
if (i >= p.b) { return; }
|
|
11
|
+
let offset = i * p.d;
|
|
12
|
+
var sum = 0.0;
|
|
13
|
+
for (var j = 0u; j < p.d; j = j + 1u) {
|
|
14
|
+
sum = sum + input[offset + j];
|
|
15
|
+
}
|
|
16
|
+
output[i] = sum * p.inv;
|
|
17
|
+
}
|
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
|
|
4
|
+
struct Params {
|
|
5
|
+
n : u32,
|
|
6
|
+
h : u32,
|
|
7
|
+
w : u32,
|
|
8
|
+
c : u32,
|
|
9
|
+
out_h : u32,
|
|
10
|
+
out_w : u32,
|
|
11
|
+
mode : u32,
|
|
12
|
+
}
|
|
13
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
14
|
+
|
|
15
|
+
@compute @workgroup_size(8, 8, 1)
|
|
16
|
+
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
|
|
17
|
+
let ox = gid.x;
|
|
18
|
+
let oy = gid.y;
|
|
19
|
+
let z = gid.z;
|
|
20
|
+
let nb = z / params.c;
|
|
21
|
+
let ch = z - nb * params.c;
|
|
22
|
+
if (nb >= params.n || ox >= params.out_w || oy >= params.out_h) { return; }
|
|
23
|
+
|
|
24
|
+
if (params.mode == 0u) {
|
|
25
|
+
let iy = (oy * params.h) / params.out_h;
|
|
26
|
+
let ix = (ox * params.w) / params.out_w;
|
|
27
|
+
output[((nb * params.out_h + oy) * params.out_w + ox) * params.c + ch] =
|
|
28
|
+
input[((nb * params.h + iy) * params.w + ix) * params.c + ch];
|
|
29
|
+
return;
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
let scale_y = f32(params.h) / f32(params.out_h);
|
|
33
|
+
let scale_x = f32(params.w) / f32(params.out_w);
|
|
34
|
+
var fy = (f32(oy) + 0.5) * scale_y - 0.5;
|
|
35
|
+
var fx = (f32(ox) + 0.5) * scale_x - 0.5;
|
|
36
|
+
if (fy < 0.0) { fy = 0.0; }
|
|
37
|
+
if (fx < 0.0) { fx = 0.0; }
|
|
38
|
+
let y0 = u32(fy);
|
|
39
|
+
let x0 = u32(fx);
|
|
40
|
+
let y1 = min(y0 + 1u, params.h - 1u);
|
|
41
|
+
let x1 = min(x0 + 1u, params.w - 1u);
|
|
42
|
+
let dy = fy - f32(y0);
|
|
43
|
+
let dx = fx - f32(x0);
|
|
44
|
+
let base = nb * params.h * params.w * params.c;
|
|
45
|
+
let v00 = input[base + (y0 * params.w + x0) * params.c + ch];
|
|
46
|
+
let v01 = input[base + (y0 * params.w + x1) * params.c + ch];
|
|
47
|
+
let v10 = input[base + (y1 * params.w + x0) * params.c + ch];
|
|
48
|
+
let v11 = input[base + (y1 * params.w + x1) * params.c + ch];
|
|
49
|
+
let val = v00 * (1.0 - dy) * (1.0 - dx) + v01 * (1.0 - dy) * dx +
|
|
50
|
+
v10 * dy * (1.0 - dx) + v11 * dy * dx;
|
|
51
|
+
output[((nb * params.out_h + oy) * params.out_w + ox) * params.c + ch] = val;
|
|
52
|
+
}
|
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> qkv : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
|
|
4
|
+
struct Params {
|
|
5
|
+
seq_len : u32,
|
|
6
|
+
d_model : u32,
|
|
7
|
+
num_heads : u32,
|
|
8
|
+
head_dim : u32,
|
|
9
|
+
scale : f32,
|
|
10
|
+
}
|
|
11
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
12
|
+
|
|
13
|
+
@compute @workgroup_size(64, 1, 1)
|
|
14
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
15
|
+
let q_idx = global_id.x;
|
|
16
|
+
let h_idx = global_id.y;
|
|
17
|
+
|
|
18
|
+
if (q_idx >= params.seq_len || h_idx >= params.num_heads) { return; }
|
|
19
|
+
|
|
20
|
+
let head_dim = params.head_dim;
|
|
21
|
+
let d_model = params.d_model;
|
|
22
|
+
|
|
23
|
+
var max_logit : f32 = -1e38;
|
|
24
|
+
|
|
25
|
+
// Cache Q for this head and this q_idx
|
|
26
|
+
var q_cache : array<f32, 64>; // max head_dim 64
|
|
27
|
+
for (var d = 0u; d < head_dim; d = d + 1u) {
|
|
28
|
+
q_cache[d] = qkv[q_idx * (d_model * 3u) + (h_idx * head_dim) + d];
|
|
29
|
+
}
|
|
30
|
+
|
|
31
|
+
// Pass 1: find max
|
|
32
|
+
for (var k_idx = 0u; k_idx <= q_idx; k_idx = k_idx + 1u) {
|
|
33
|
+
var score : f32 = 0.0;
|
|
34
|
+
for (var d = 0u; d < head_dim; d = d + 1u) {
|
|
35
|
+
let k_val = qkv[k_idx * (d_model * 3u) + d_model + (h_idx * head_dim) + d];
|
|
36
|
+
score = score + (q_cache[d] * k_val);
|
|
37
|
+
}
|
|
38
|
+
score = score * params.scale;
|
|
39
|
+
if (score > max_logit) { max_logit = score; }
|
|
40
|
+
}
|
|
41
|
+
|
|
42
|
+
// Pass 2: sum exp
|
|
43
|
+
var sum_exp : f32 = 0.0;
|
|
44
|
+
for (var k_idx = 0u; k_idx <= q_idx; k_idx = k_idx + 1u) {
|
|
45
|
+
var score : f32 = 0.0;
|
|
46
|
+
for (var d = 0u; d < head_dim; d = d + 1u) {
|
|
47
|
+
let k_val = qkv[k_idx * (d_model * 3u) + d_model + (h_idx * head_dim) + d];
|
|
48
|
+
score = score + (q_cache[d] * k_val);
|
|
49
|
+
}
|
|
50
|
+
score = score * params.scale;
|
|
51
|
+
sum_exp = sum_exp + exp(score - max_logit);
|
|
52
|
+
}
|
|
53
|
+
|
|
54
|
+
// Pass 3: output
|
|
55
|
+
for (var d = 0u; d < head_dim; d = d + 1u) {
|
|
56
|
+
var out_val : f32 = 0.0;
|
|
57
|
+
for (var k_idx = 0u; k_idx <= q_idx; k_idx = k_idx + 1u) {
|
|
58
|
+
var score : f32 = 0.0;
|
|
59
|
+
for (var kd = 0u; kd < head_dim; kd = kd + 1u) {
|
|
60
|
+
let k_val = qkv[k_idx * (d_model * 3u) + d_model + (h_idx * head_dim) + kd];
|
|
61
|
+
score = score + (q_cache[kd] * k_val);
|
|
62
|
+
}
|
|
63
|
+
score = score * params.scale;
|
|
64
|
+
|
|
65
|
+
let w = exp(score - max_logit) / sum_exp;
|
|
66
|
+
let v_val = qkv[k_idx * (d_model * 3u) + d_model * 2u + (h_idx * head_dim) + d];
|
|
67
|
+
out_val = out_val + (w * v_val);
|
|
68
|
+
}
|
|
69
|
+
output[q_idx * d_model + (h_idx * head_dim) + d] = out_val;
|
|
70
|
+
}
|
|
71
|
+
}
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
struct Params { size : u32 }
|
|
4
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
5
|
+
@compute @workgroup_size(64)
|
|
6
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
7
|
+
let idx = global_id.x;
|
|
8
|
+
if (idx >= params.size) { return; }
|
|
9
|
+
let x = input[idx];
|
|
10
|
+
var out_val = x;
|
|
11
|
+
out_val = x * (1.0 / (1.0 + exp(-x)));
|
|
12
|
+
output[idx] = out_val;
|
|
13
|
+
}
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
struct Params { size : u32 }
|
|
4
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
5
|
+
@compute @workgroup_size(64)
|
|
6
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
7
|
+
let idx = global_id.x;
|
|
8
|
+
if (idx >= params.size) { return; }
|
|
9
|
+
let x = input[idx];
|
|
10
|
+
var out_val = x;
|
|
11
|
+
out_val = 1.0 / (1.0 + exp(-x));
|
|
12
|
+
output[idx] = out_val;
|
|
13
|
+
}
|
|
@@ -0,0 +1,26 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
// General strided slice over 4 dims (leading dims are padded to 1). For each
|
|
4
|
+
// output element the source coord along axis k is start[k] + out_coord * step[k].
|
|
5
|
+
struct Params {
|
|
6
|
+
out_b: u32, out_c: u32, out_h: u32, out_w: u32,
|
|
7
|
+
in_c: u32, in_h: u32, in_w: u32,
|
|
8
|
+
s0: u32, s1: u32, s2: u32, s3: u32,
|
|
9
|
+
st0: u32, st1: u32, st2: u32, st3: u32,
|
|
10
|
+
total: u32
|
|
11
|
+
}
|
|
12
|
+
@group(0) @binding(2) var<uniform> p : Params;
|
|
13
|
+
@compute @workgroup_size(64)
|
|
14
|
+
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
|
|
15
|
+
let idx = gid.x;
|
|
16
|
+
if (idx >= p.total) { return; }
|
|
17
|
+
let ow = idx % p.out_w;
|
|
18
|
+
let oh = (idx / p.out_w) % p.out_h;
|
|
19
|
+
let oc = (idx / (p.out_w * p.out_h)) % p.out_c;
|
|
20
|
+
let ob = idx / (p.out_w * p.out_h * p.out_c);
|
|
21
|
+
let ib = p.s0 + ob * p.st0;
|
|
22
|
+
let ic = p.s1 + oc * p.st1;
|
|
23
|
+
let ih = p.s2 + oh * p.st2;
|
|
24
|
+
let iw = p.s3 + ow * p.st3;
|
|
25
|
+
output[idx] = input[((ib * p.in_c + ic) * p.in_h + ih) * p.in_w + iw];
|
|
26
|
+
}
|
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
struct Params { b : u32, d : u32 }
|
|
4
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
5
|
+
@compute @workgroup_size(64)
|
|
6
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
7
|
+
let i = global_id.x;
|
|
8
|
+
if (i >= params.b) { return; }
|
|
9
|
+
let offset = i * params.d;
|
|
10
|
+
var max_val = -100000.0;
|
|
11
|
+
for (var j = 0u; j < params.d; j = j + 1u) {
|
|
12
|
+
if (input[offset + j] > max_val) { max_val = input[offset + j]; }
|
|
13
|
+
}
|
|
14
|
+
var sum = 0.0;
|
|
15
|
+
for (var j = 0u; j < params.d; j = j + 1u) {
|
|
16
|
+
let e = exp(input[offset + j] - max_val);
|
|
17
|
+
output[offset + j] = e;
|
|
18
|
+
sum = sum + e;
|
|
19
|
+
}
|
|
20
|
+
for (var j = 0u; j < params.d; j = j + 1u) {
|
|
21
|
+
output[offset + j] = output[offset + j] / sum;
|
|
22
|
+
}
|
|
23
|
+
}
|
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
|
|
4
|
+
struct Params { in_h : u32, in_w : u32, in_c : u32 }
|
|
5
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
6
|
+
|
|
7
|
+
@compute @workgroup_size(64)
|
|
8
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
9
|
+
let x = global_id.x;
|
|
10
|
+
let c = global_id.y;
|
|
11
|
+
|
|
12
|
+
if (x >= params.in_w || c >= params.in_c) { return; }
|
|
13
|
+
|
|
14
|
+
var max_val = -1e38;
|
|
15
|
+
for (var y = 0u; y < params.in_h; y = y + 1u) {
|
|
16
|
+
let val = input[(y * params.in_w + x) * params.in_c + c];
|
|
17
|
+
if (val > max_val) { max_val = val; }
|
|
18
|
+
}
|
|
19
|
+
|
|
20
|
+
var denom = 0.0;
|
|
21
|
+
var weighted = 0.0;
|
|
22
|
+
for (var y = 0u; y < params.in_h; y = y + 1u) {
|
|
23
|
+
let val = input[(y * params.in_w + x) * params.in_c + c];
|
|
24
|
+
let ev = exp(val - max_val);
|
|
25
|
+
denom = denom + ev;
|
|
26
|
+
weighted = weighted + ev * (f32(y) + 0.5) / f32(params.in_h);
|
|
27
|
+
}
|
|
28
|
+
|
|
29
|
+
var out_val = 0.0;
|
|
30
|
+
if (denom > 0.0) { out_val = weighted / denom; }
|
|
31
|
+
output[c * params.in_w + x] = out_val;
|
|
32
|
+
}
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
struct Params { total : u32, inner : u32, split_size : u32, axis_in : u32, offset : u32 }
|
|
4
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
5
|
+
@compute @workgroup_size(64)
|
|
6
|
+
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
|
|
7
|
+
let i = gid.x;
|
|
8
|
+
if (i >= params.total) { return; }
|
|
9
|
+
let inner_idx = i % params.inner;
|
|
10
|
+
let s = (i / params.inner) % params.split_size;
|
|
11
|
+
let outer_idx = i / (params.split_size * params.inner);
|
|
12
|
+
let in_idx = outer_idx * (params.axis_in * params.inner)
|
|
13
|
+
+ (params.offset + s) * params.inner + inner_idx;
|
|
14
|
+
output[i] = input[in_idx];
|
|
15
|
+
}
|
package/shaders/sub.wgsl
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> a : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> b : array<f32>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read_write> output : array<f32>;
|
|
4
|
+
|
|
5
|
+
struct Params { size : u32, is_b_scalar : u32, b_size : u32, a_size: u32, is_a_scalar: u32 }
|
|
6
|
+
@group(0) @binding(3) var<uniform> params : Params;
|
|
7
|
+
|
|
8
|
+
@compute @workgroup_size(64)
|
|
9
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
10
|
+
let idx = global_id.x;
|
|
11
|
+
if (idx >= params.size) { return; }
|
|
12
|
+
|
|
13
|
+
var a_val : f32 = 0.0;
|
|
14
|
+
if (params.is_a_scalar == 1u) {
|
|
15
|
+
a_val = a[0];
|
|
16
|
+
} else if (params.a_size < params.size && params.a_size > 0u) {
|
|
17
|
+
a_val = a[idx % params.a_size];
|
|
18
|
+
} else {
|
|
19
|
+
a_val = a[idx];
|
|
20
|
+
}
|
|
21
|
+
|
|
22
|
+
var b_val : f32 = 0.0;
|
|
23
|
+
if (params.is_b_scalar == 1u) {
|
|
24
|
+
b_val = b[0];
|
|
25
|
+
} else if (params.b_size < params.size && params.b_size > 0u) {
|
|
26
|
+
b_val = b[idx % params.b_size];
|
|
27
|
+
} else {
|
|
28
|
+
b_val = b[idx];
|
|
29
|
+
}
|
|
30
|
+
|
|
31
|
+
var out_val : f32 = 0.0;
|
|
32
|
+
out_val = a_val - b_val;
|
|
33
|
+
output[idx] = out_val;
|
|
34
|
+
}
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
struct Params { size : u32 }
|
|
4
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
5
|
+
@compute @workgroup_size(64)
|
|
6
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
7
|
+
let idx = global_id.x;
|
|
8
|
+
if (idx >= params.size) { return; }
|
|
9
|
+
let x = input[idx];
|
|
10
|
+
var out_val = x;
|
|
11
|
+
let e2x = exp(2.0 * x); out_val = (e2x - 1.0) / (e2x + 1.0);
|
|
12
|
+
output[idx] = out_val;
|
|
13
|
+
}
|
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
|
|
4
|
+
struct Params {
|
|
5
|
+
n : u32,
|
|
6
|
+
h : u32,
|
|
7
|
+
w : u32,
|
|
8
|
+
c : u32,
|
|
9
|
+
}
|
|
10
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
11
|
+
|
|
12
|
+
@compute @workgroup_size(8, 8, 1)
|
|
13
|
+
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
|
|
14
|
+
let ox = gid.x;
|
|
15
|
+
let oy = gid.y;
|
|
16
|
+
let z = gid.z;
|
|
17
|
+
let nb = z / params.c;
|
|
18
|
+
let ch = z - nb * params.c;
|
|
19
|
+
let out_h = params.h * 2u;
|
|
20
|
+
let out_w = params.w * 2u;
|
|
21
|
+
if (nb >= params.n || ox >= out_w || oy >= out_h) { return; }
|
|
22
|
+
output[((nb * out_h + oy) * out_w + ox) * params.c + ch] =
|
|
23
|
+
input[((nb * params.h + (oy / 2u)) * params.w + (ox / 2u)) * params.c + ch];
|
|
24
|
+
}
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> cond : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> a : array<f32>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read> b : array<f32>;
|
|
4
|
+
@group(0) @binding(3) var<storage, read_write> output : array<f32>;
|
|
5
|
+
struct Params { size : u32 }
|
|
6
|
+
@group(0) @binding(4) var<uniform> params : Params;
|
|
7
|
+
@compute @workgroup_size(64)
|
|
8
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
9
|
+
let idx = global_id.x;
|
|
10
|
+
if (idx >= params.size) { return; }
|
|
11
|
+
if (cond[idx] != 0.0) { output[idx] = a[idx]; } else { output[idx] = b[idx]; }
|
|
12
|
+
}
|
package/volvoxai.wasm
ADDED
|
Binary file
|