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.
Files changed (172) hide show
  1. package/LICENSE +21 -0
  2. package/README.md +145 -0
  3. package/bin/volvox.js +72 -0
  4. package/dist/v0.1.0/volvoxai.js +4664 -0
  5. package/dist/v0.1.0/volvoxai.min.js +1848 -0
  6. package/dist/v0.1.0/volvoxai.wasm +0 -0
  7. package/dist/volvoxai.js +4664 -0
  8. package/dist/volvoxai.min.js +1848 -0
  9. package/dist/volvoxai.wasm +0 -0
  10. package/docs/README.md +22 -0
  11. package/docs/browser-runtime.md +87 -0
  12. package/docs/efficientdet_tflite_vs_volvoxai.md +445 -0
  13. package/docs/microkernel_optimization_guide.md +153 -0
  14. package/docs/model-format.md +108 -0
  15. package/docs/models.md +103 -0
  16. package/docs/native-runtime.md +189 -0
  17. package/docs/operation_list.md +232 -0
  18. package/docs/operator_fusion_patterns.md +58 -0
  19. package/docs/quickstart.md +115 -0
  20. package/docs/roadmap.md +19 -0
  21. package/docs/testing.md +97 -0
  22. package/docs/textbook/01-foundations.md +233 -0
  23. package/docs/textbook/02-tinystories-language-model.md +300 -0
  24. package/docs/textbook/03-efficientdet-vision-model.md +281 -0
  25. package/docs/textbook/04-precision-and-quantization.md +208 -0
  26. package/docs/textbook/05-inside-the-engine.md +155 -0
  27. package/docs/textbook/06-native-engine-architecture.md +338 -0
  28. package/docs/textbook/07-glossary-and-next-steps.md +258 -0
  29. package/docs/textbook/README.md +85 -0
  30. package/docs/textbook/ko/01-foundations.md +231 -0
  31. package/docs/textbook/ko/02-tinystories-language-model.md +300 -0
  32. package/docs/textbook/ko/03-efficientdet-vision-model.md +277 -0
  33. package/docs/textbook/ko/04-precision-and-quantization.md +206 -0
  34. package/docs/textbook/ko/05-inside-the-engine.md +154 -0
  35. package/docs/textbook/ko/06-native-engine-architecture.md +333 -0
  36. package/docs/textbook/ko/07-glossary-and-next-steps.md +253 -0
  37. package/docs/textbook/ko/README.md +83 -0
  38. package/docs/xnnpack_optimization_guide.md +197 -0
  39. package/js/CPUEngine.js +241 -0
  40. package/js/Graph.js +49 -0
  41. package/js/GraphExecutor.js +1020 -0
  42. package/js/GraphLoader.js +282 -0
  43. package/js/ShaderLibrary.js +236 -0
  44. package/js/Tensor.js +25 -0
  45. package/js/Tokenizer.js +266 -0
  46. package/js/VolvoxAI.js +130 -0
  47. package/js/WasmEngine.js +378 -0
  48. package/js/WebNNEngine.js +169 -0
  49. package/js/index.js +11 -0
  50. package/js/ops/add.js +31 -0
  51. package/js/ops/argMax.js +33 -0
  52. package/js/ops/averagePool2D.js +38 -0
  53. package/js/ops/batchNorm2D.js +28 -0
  54. package/js/ops/cast.js +19 -0
  55. package/js/ops/clip.js +15 -0
  56. package/js/ops/concat2.js +18 -0
  57. package/js/ops/conv1D.js +35 -0
  58. package/js/ops/conv2D.js +70 -0
  59. package/js/ops/convTranspose2D.js +45 -0
  60. package/js/ops/crossAttention.js +69 -0
  61. package/js/ops/crossSDPA.js +41 -0
  62. package/js/ops/dequantizeLinear.js +9 -0
  63. package/js/ops/div.js +15 -0
  64. package/js/ops/embedding.js +14 -0
  65. package/js/ops/expand.js +24 -0
  66. package/js/ops/gELU.js +9 -0
  67. package/js/ops/gather.js +51 -0
  68. package/js/ops/gatherElements.js +33 -0
  69. package/js/ops/globalAveragePool.js +21 -0
  70. package/js/ops/hardSigmoid.js +12 -0
  71. package/js/ops/hardSwish.js +12 -0
  72. package/js/ops/interp1D.js +25 -0
  73. package/js/ops/layerNorm.js +25 -0
  74. package/js/ops/leakyReLU.js +10 -0
  75. package/js/ops/logSoftmax.js +15 -0
  76. package/js/ops/matMul.js +35 -0
  77. package/js/ops/maxPool2D.js +36 -0
  78. package/js/ops/meanHeight.js +17 -0
  79. package/js/ops/mul.js +31 -0
  80. package/js/ops/nonMaxSuppression.js +72 -0
  81. package/js/ops/pReLU.js +11 -0
  82. package/js/ops/pad.js +35 -0
  83. package/js/ops/profileX.js +22 -0
  84. package/js/ops/profileY.js +22 -0
  85. package/js/ops/rMSNorm.js +14 -0
  86. package/js/ops/reLU.js +8 -0
  87. package/js/ops/reduceMean.js +17 -0
  88. package/js/ops/reduceSum.js +19 -0
  89. package/js/ops/reshape.js +6 -0
  90. package/js/ops/resize.js +44 -0
  91. package/js/ops/sDPA.js +44 -0
  92. package/js/ops/siLU.js +8 -0
  93. package/js/ops/sigmoid.js +6 -0
  94. package/js/ops/slice.js +36 -0
  95. package/js/ops/softmax.js +18 -0
  96. package/js/ops/spatialSoftargmaxY.js +28 -0
  97. package/js/ops/split.js +24 -0
  98. package/js/ops/sub.js +11 -0
  99. package/js/ops/tanh.js +7 -0
  100. package/js/ops/transpose.js +34 -0
  101. package/js/ops/upsample2x.js +23 -0
  102. package/js/ops/where.js +15 -0
  103. package/package.json +33 -0
  104. package/shaders/add.wgsl +13 -0
  105. package/shaders/add3Relu.wgsl +23 -0
  106. package/shaders/addRelu.wgsl +22 -0
  107. package/shaders/averagePool2D.wgsl +24 -0
  108. package/shaders/batchNorm2D.wgsl +21 -0
  109. package/shaders/binaryBroadcast.wgsl +34 -0
  110. package/shaders/broadcastBinary.wgsl +26 -0
  111. package/shaders/clip.wgsl +10 -0
  112. package/shaders/concat2.wgsl +16 -0
  113. package/shaders/concatCopy.wgsl +10 -0
  114. package/shaders/concatSigmoidCopy.wgsl +16 -0
  115. package/shaders/conv1D.wgsl +37 -0
  116. package/shaders/conv2D.wgsl +80 -0
  117. package/shaders/conv2DDepthwise4.wgsl +74 -0
  118. package/shaders/conv2DDepthwise8.wgsl +66 -0
  119. package/shaders/conv2DPointwise16.wgsl +67 -0
  120. package/shaders/conv2DPointwise16Tile.wgsl +86 -0
  121. package/shaders/conv2DPointwise8.wgsl +85 -0
  122. package/shaders/conv2DPointwise8Vec2.wgsl +70 -0
  123. package/shaders/conv2DPointwise8Vec4.wgsl +65 -0
  124. package/shaders/conv2DRegularC3Out16.wgsl +75 -0
  125. package/shaders/convTranspose2D.wgsl +33 -0
  126. package/shaders/copy.wgsl +13 -0
  127. package/shaders/crossAttention.wgsl +140 -0
  128. package/shaders/crossAttentionF32.wgsl +98 -0
  129. package/shaders/crossSDPA.wgsl +74 -0
  130. package/shaders/dequantizeLinear.wgsl +14 -0
  131. package/shaders/div.wgsl +34 -0
  132. package/shaders/elementwise.wgsl +13 -0
  133. package/shaders/embedding.wgsl +22 -0
  134. package/shaders/expand.wgsl +18 -0
  135. package/shaders/gELU.wgsl +13 -0
  136. package/shaders/gather.wgsl +17 -0
  137. package/shaders/generalTranspose.wgsl +19 -0
  138. package/shaders/globalAveragePool.wgsl +19 -0
  139. package/shaders/hardSigmoid.wgsl +13 -0
  140. package/shaders/hardSwish.wgsl +13 -0
  141. package/shaders/interp1D.wgsl +28 -0
  142. package/shaders/layerNorm.wgsl +33 -0
  143. package/shaders/leakyReLU.wgsl +11 -0
  144. package/shaders/linearF32.wgsl +33 -0
  145. package/shaders/linearF32RowMajor.wgsl +24 -0
  146. package/shaders/linearInt8.wgsl +42 -0
  147. package/shaders/logSoftmax.wgsl +22 -0
  148. package/shaders/maxPool2D.wgsl +37 -0
  149. package/shaders/meanHeight.wgsl +18 -0
  150. package/shaders/mul.wgsl +32 -0
  151. package/shaders/nonMaxSuppression.wgsl +92 -0
  152. package/shaders/pReLU.wgsl +14 -0
  153. package/shaders/pad.wgsl +19 -0
  154. package/shaders/profileX.wgsl +28 -0
  155. package/shaders/profileY.wgsl +28 -0
  156. package/shaders/quantizeLinear.wgsl +69 -0
  157. package/shaders/rMSNorm.wgsl +21 -0
  158. package/shaders/reLU.wgsl +13 -0
  159. package/shaders/reduce.wgsl +17 -0
  160. package/shaders/resize.wgsl +52 -0
  161. package/shaders/sDPA.wgsl +71 -0
  162. package/shaders/siLU.wgsl +13 -0
  163. package/shaders/sigmoid.wgsl +13 -0
  164. package/shaders/slice.wgsl +26 -0
  165. package/shaders/softmax.wgsl +23 -0
  166. package/shaders/spatialSoftargmaxY.wgsl +32 -0
  167. package/shaders/split.wgsl +15 -0
  168. package/shaders/sub.wgsl +34 -0
  169. package/shaders/tanh.wgsl +13 -0
  170. package/shaders/upsample2x.wgsl +24 -0
  171. package/shaders/where.wgsl +12 -0
  172. 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
+ }
@@ -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
+ }
@@ -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