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,74 @@
1
+ @group(0) @binding(0) var<storage, read> q_in : array<f32>;
2
+ @group(0) @binding(1) var<storage, read> k_in : array<f32>;
3
+ @group(0) @binding(2) var<storage, read> v_in : array<f32>;
4
+ @group(0) @binding(3) var<storage, read_write> output : array<f32>;
5
+
6
+ struct Params {
7
+ seq_len_q : u32,
8
+ seq_len_kv : u32,
9
+ d_model : u32,
10
+ num_heads : u32,
11
+ head_dim : u32,
12
+ scale : f32,
13
+ }
14
+ @group(0) @binding(4) var<uniform> params : Params;
15
+
16
+ @compute @workgroup_size(64, 1, 1)
17
+ fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
18
+ let q_idx = global_id.x;
19
+ let h_idx = global_id.y;
20
+
21
+ if (q_idx >= params.seq_len_q || h_idx >= params.num_heads) { return; }
22
+
23
+ let head_dim = params.head_dim;
24
+ let d_model = params.d_model;
25
+
26
+ // Cache Q for this head and this q_idx
27
+ var q_cache : array<f32, 64>; // max head_dim 64
28
+ for (var d = 0u; d < head_dim; d = d + 1u) {
29
+ q_cache[d] = q_in[q_idx * d_model + (h_idx * head_dim) + d];
30
+ }
31
+
32
+ var max_logit : f32 = -1e38;
33
+
34
+ // Pass 1: find max
35
+ for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
36
+ var score : f32 = 0.0;
37
+ for (var d = 0u; d < head_dim; d = d + 1u) {
38
+ let k_val = k_in[k_idx * d_model + (h_idx * head_dim) + d];
39
+ score = score + (q_cache[d] * k_val);
40
+ }
41
+ score = score * params.scale;
42
+ if (score > max_logit) { max_logit = score; }
43
+ }
44
+
45
+ // Pass 2: sum exp
46
+ var sum_exp : f32 = 0.0;
47
+ for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
48
+ var score : f32 = 0.0;
49
+ for (var d = 0u; d < head_dim; d = d + 1u) {
50
+ let k_val = k_in[k_idx * d_model + (h_idx * head_dim) + d];
51
+ score = score + (q_cache[d] * k_val);
52
+ }
53
+ score = score * params.scale;
54
+ sum_exp = sum_exp + exp(score - max_logit);
55
+ }
56
+
57
+ // Pass 3: output
58
+ for (var d = 0u; d < head_dim; d = d + 1u) {
59
+ var out_val : f32 = 0.0;
60
+ for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
61
+ var score : f32 = 0.0;
62
+ for (var kd = 0u; kd < head_dim; kd = kd + 1u) {
63
+ let k_val = k_in[k_idx * d_model + (h_idx * head_dim) + kd];
64
+ score = score + (q_cache[kd] * k_val);
65
+ }
66
+ score = score * params.scale;
67
+
68
+ let w = exp(score - max_logit) / sum_exp;
69
+ let v_val = v_in[k_idx * d_model + (h_idx * head_dim) + d];
70
+ out_val = out_val + (w * v_val);
71
+ }
72
+ output[q_idx * d_model + (h_idx * head_dim) + d] = out_val;
73
+ }
74
+ }
@@ -0,0 +1,14 @@
1
+ @group(0) @binding(0) var<storage, read> input : array<f32>;
2
+ @group(0) @binding(1) var<storage, read> scale : array<f32>;
3
+ @group(0) @binding(2) var<storage, read> zero_point : array<f32>;
4
+ @group(0) @binding(3) var<storage, read_write> output : array<f32>;
5
+ struct Params { size : u32, has_zp : 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
+ var zp = 0.0;
12
+ if (params.has_zp == 1u) { zp = zero_point[0]; }
13
+ output[idx] = (input[idx] - zp) * scale[0];
14
+ }
@@ -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
+ undefined
12
+ output[idx] = out_val;
13
+ }
@@ -0,0 +1,22 @@
1
+ @group(0) @binding(0) var<storage, read> tokens : 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
+
5
+ struct Params { seq_len : u32, d_model : 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 token_idx = global_id.x;
11
+ if (token_idx >= params.seq_len) { return; }
12
+
13
+ let d_model = params.d_model;
14
+ let token_id = u32(tokens[token_idx]); // ids arrive as f32 (see CPU/WASM tiers)
15
+
16
+ let in_offset = token_id * d_model;
17
+ let out_offset = token_idx * d_model;
18
+
19
+ for (var i = 0u; i < d_model; i = i + 1u) {
20
+ output[out_offset + i] = weight[in_offset + i];
21
+ }
22
+ }
@@ -0,0 +1,18 @@
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 { in_b: u32, in_h: u32, in_w: u32, in_c: u32, out_b: u32, out_h: u32, out_w: u32, out_c: 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
+ let total = params.out_b * params.out_h * params.out_w * params.out_c;
9
+ if (idx >= total) { return; }
10
+ let oc = idx % params.out_c;
11
+ let ow = (idx / params.out_c) % params.out_w;
12
+ let oh = (idx / (params.out_c * params.out_w)) % params.out_h;
13
+ let ob = idx / (params.out_c * params.out_w * params.out_h);
14
+ let ib = ob % params.in_b;
15
+ let ih = oh % params.in_h; let iw = ow % params.in_w;
16
+ let ic = oc % params.in_c;
17
+ output[idx] = input[((ib * params.in_h + ih) * params.in_w + iw) * params.in_c + ic];
18
+ }
@@ -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
+
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 x = input[idx];
11
+ let cdf = 0.5 * (1.0 + tanh(0.7978845608 * (x + 0.044715 * x * x * x)));
12
+ output[idx] = x * cdf;
13
+ }
@@ -0,0 +1,17 @@
1
+ @group(0) @binding(0) var<storage, read> input : array<f32>;
2
+ @group(0) @binding(1) var<storage, read> indices : array<f32>;
3
+ @group(0) @binding(2) var<storage, read_write> output : array<f32>;
4
+ // Gather along axis 0: output[i, ...] = input[indices[i], ...]. `row_size` is the
5
+ // number of contiguous elements per gathered row (product of input dims after
6
+ // axis 0); `num_idx` is the number of indices. total = num_idx * row_size.
7
+ struct Params { row_size : u32, num_idx : u32, total : u32 }
8
+ @group(0) @binding(3) var<uniform> p : Params;
9
+ @compute @workgroup_size(64)
10
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
11
+ let idx = gid.x;
12
+ if (idx >= p.total) { return; }
13
+ let k = idx % p.row_size;
14
+ let i = idx / p.row_size;
15
+ let row = u32(indices[i]);
16
+ output[idx] = input[row * p.row_size + k];
17
+ }
@@ -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
+ @group(0) @binding(2) var<storage, read> md : array<u32>;
4
+ @compute @workgroup_size(64)
5
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
6
+ let idx = gid.x;
7
+ let total = md[0];
8
+ if (idx >= total) { return; }
9
+ let rank = md[1];
10
+ var rem = idx;
11
+ var in_idx = 0u;
12
+ for (var d = 0u; d < rank; d = d + 1u) {
13
+ let os = md[2u + d];
14
+ let coord = rem / os;
15
+ rem = rem - coord * os;
16
+ in_idx = in_idx + coord * md[2u + rank + d];
17
+ }
18
+ output[idx] = input[in_idx];
19
+ }
@@ -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
+
4
+ struct Params { n : u32, h : u32, w : u32, c : u32 }
5
+ @group(0) @binding(2) var<uniform> params : Params;
6
+
7
+ @compute @workgroup_size(64)
8
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
9
+ let ch = gid.x;
10
+ let nb = gid.y;
11
+ if (nb >= params.n || ch >= params.c) { return; }
12
+ var sum = 0.0;
13
+ for (var y = 0u; y < params.h; y = y + 1u) {
14
+ for (var x = 0u; x < params.w; x = x + 1u) {
15
+ sum = sum + input[((nb * params.h + y) * params.w + x) * params.c + ch];
16
+ }
17
+ }
18
+ output[nb * params.c + ch] = sum / f32(params.h * params.w);
19
+ }
@@ -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
+ var v = x + 3.0; if (v < 0.0) { v = 0.0; } if (v > 6.0) { v = 6.0; } out_val = v / 6.0;
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
+ var v = x + 3.0; if (v < 0.0) { v = 0.0; } if (v > 6.0) { v = 6.0; } out_val = x * v / 6.0;
12
+ output[idx] = out_val;
13
+ }
@@ -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_c : u32, in_l : u32, out_l : 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 i = global_id.x;
10
+ let c = global_id.y;
11
+
12
+ if (i >= params.out_l || c >= params.in_c) { return; }
13
+
14
+ let scale = f32(params.in_l) / f32(params.out_l);
15
+ var src = (f32(i) + 0.5) * scale - 0.5;
16
+ if (src < 0.0) { src = 0.0; }
17
+ if (src > f32(params.in_l - 1u)) { src = f32(params.in_l - 1u); }
18
+
19
+ let lo = u32(src);
20
+ var hi = lo + 1u;
21
+ if (hi >= params.in_l) { hi = params.in_l - 1u; }
22
+ let t = src - f32(lo);
23
+
24
+ let in_base = c * params.in_l;
25
+ let out_base = c * params.out_l;
26
+
27
+ output[out_base + i] = input[in_base + lo] * (1.0 - t) + input[in_base + hi] * t;
28
+ }
@@ -0,0 +1,33 @@
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> bias : array<f32>;
4
+ @group(0) @binding(3) var<storage, read_write> output : array<f32>;
5
+
6
+ struct Params { rows : u32, d_model : u32 }
7
+ @group(0) @binding(4) var<uniform> params : Params;
8
+
9
+ @compute @workgroup_size(64)
10
+ fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
11
+ let row = global_id.x;
12
+ if (row >= params.rows) { return; }
13
+ let d_model = params.d_model;
14
+ let offset = row * d_model;
15
+
16
+ var sum : f32 = 0.0;
17
+ var sq_sum : f32 = 0.0;
18
+
19
+ for (var i = 0u; i < d_model; i = i + 1u) {
20
+ let val = input[offset + i];
21
+ sum = sum + val;
22
+ sq_sum = sq_sum + (val * val);
23
+ }
24
+
25
+ let mean = sum / f32(d_model);
26
+ let variance = (sq_sum / f32(d_model)) - (mean * mean);
27
+ let inv_std = inverseSqrt(variance + 1e-5);
28
+
29
+ for (var i = 0u; i < d_model; i = i + 1u) {
30
+ let norm_val = (input[offset + i] - mean) * inv_std;
31
+ output[offset + i] = norm_val * weight[i] + bias[i];
32
+ }
33
+ }
@@ -0,0 +1,11 @@
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, alpha : 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
+ if (idx >= params.size) { return; }
9
+ let x = input[idx];
10
+ if (x > 0.0) { output[idx] = x; } else { output[idx] = x * params.alpha; }
11
+ }
@@ -0,0 +1,33 @@
1
+ @group(0) @binding(0) var<storage, read> input : array<f32>;
2
+ @group(0) @binding(1) var<storage, read> weight_f32 : array<f32>;
3
+ @group(0) @binding(2) var<storage, read> dummyScale : array<f32>;
4
+ @group(0) @binding(3) var<storage, read> bias : array<f32>;
5
+ @group(0) @binding(4) var<storage, read_write> output : array<f32>;
6
+
7
+ struct Params {
8
+ seq_len : u32,
9
+ d_in : u32,
10
+ d_out : u32,
11
+ }
12
+ @group(0) @binding(5) var<uniform> params : Params;
13
+
14
+ @compute @workgroup_size(64, 1, 1)
15
+ fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
16
+ let row = global_id.y;
17
+ let col = global_id.x;
18
+
19
+ if (row >= params.seq_len || col >= params.d_out) { return; }
20
+
21
+ var sum : f32 = 0.0;
22
+
23
+ for (var k = 0u; k < params.d_in; k = k + 1u) {
24
+ let in_val = input[row * params.d_in + k];
25
+ let w_val = weight_f32[col * params.d_in + k];
26
+ sum = sum + in_val * w_val;
27
+ }
28
+
29
+ let b_val = bias[col];
30
+ // dummyScale is passed just to keep bindings consistent but not used here.
31
+
32
+ output[row * params.d_out + col] = sum + b_val;
33
+ }
@@ -0,0 +1,24 @@
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> bias : array<f32>;
4
+ @group(0) @binding(3) var<storage, read_write> output : array<f32>;
5
+
6
+ struct Params {
7
+ seq_len : u32,
8
+ d_in : u32,
9
+ d_out : u32,
10
+ }
11
+ @group(0) @binding(4) var<uniform> params : Params;
12
+
13
+ @compute @workgroup_size(16, 16, 1)
14
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
15
+ let col = gid.x;
16
+ let row = gid.y;
17
+ if (row >= params.seq_len || col >= params.d_out) { return; }
18
+
19
+ var sum = bias[col];
20
+ for (var k = 0u; k < params.d_in; k = k + 1u) {
21
+ sum = sum + input[row * params.d_in + k] * weight[k * params.d_out + col];
22
+ }
23
+ output[row * params.d_out + col] = sum;
24
+ }
@@ -0,0 +1,42 @@
1
+ @group(0) @binding(0) var<storage, read> input : array<f32>;
2
+ @group(0) @binding(1) var<storage, read> weight_int8_packed : array<u32>;
3
+ @group(0) @binding(2) var<storage, read> weight_scales : array<f32>;
4
+ @group(0) @binding(3) var<storage, read> bias : array<f32>;
5
+ @group(0) @binding(4) var<storage, read_write> output : array<f32>;
6
+
7
+ struct Params {
8
+ seq_len : u32,
9
+ d_in : u32,
10
+ d_out : u32,
11
+ }
12
+ @group(0) @binding(5) var<uniform> params : Params;
13
+
14
+ // Unpack one signed 8-bit integer from a 32-bit packed block
15
+ fn unpack_i8(packed: u32, byte_idx: u32) -> f32 {
16
+ let val_i32 = extractBits(i32(packed), byte_idx * 8u, 8u);
17
+ return f32(val_i32);
18
+ }
19
+
20
+ @compute @workgroup_size(64, 1, 1)
21
+ fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
22
+ let row = global_id.y;
23
+ let col = global_id.x;
24
+
25
+ if (row >= params.seq_len || col >= params.d_out) { return; }
26
+
27
+ var sum : f32 = 0.0;
28
+ let d_in_4 = params.d_in / 4u;
29
+
30
+ for (var i = 0u; i < d_in_4; i = i + 1u) {
31
+ let w_packed = weight_int8_packed[col * d_in_4 + i];
32
+ let in_base = row * params.d_in + i * 4u;
33
+
34
+ sum = sum + input[in_base + 0u] * f32(extractBits(i32(w_packed), 0u, 8u));
35
+ sum = sum + input[in_base + 1u] * f32(extractBits(i32(w_packed), 8u, 8u));
36
+ sum = sum + input[in_base + 2u] * f32(extractBits(i32(w_packed), 16u, 8u));
37
+ sum = sum + input[in_base + 3u] * f32(extractBits(i32(w_packed), 24u, 8u));
38
+ }
39
+
40
+ let scale = weight_scales[col];
41
+ output[row * params.d_out + col] = (sum * scale) + bias[col];
42
+ }
@@ -0,0 +1,22 @@
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
+ sum = sum + exp(input[offset + j] - max_val);
17
+ }
18
+ let logSum = log(sum);
19
+ for (var j = 0u; j < params.d; j = j + 1u) {
20
+ output[offset + j] = (input[offset + j] - max_val) - logSum;
21
+ }
22
+ }
@@ -0,0 +1,37 @@
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
+ h : u32,
6
+ w : u32,
7
+ c : u32,
8
+ out_h : u32,
9
+ out_w : u32,
10
+ ky : u32,
11
+ kx : u32,
12
+ sy : u32,
13
+ sx : u32,
14
+ py : u32,
15
+ px : u32,
16
+ }
17
+ @group(0) @binding(2) var<uniform> params : Params;
18
+
19
+ @compute @workgroup_size(8, 8, 1)
20
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
21
+ let ox = gid.x;
22
+ let oy = gid.y;
23
+ let ch = gid.z;
24
+ if (ox >= params.out_w || oy >= params.out_h || ch >= params.c) { return; }
25
+
26
+ var best = -3.402823466e38;
27
+ for (var yy = 0u; yy < params.ky; yy = yy + 1u) {
28
+ let iy = i32(oy * params.sy + yy) - i32(params.py);
29
+ if (iy < 0 || iy >= i32(params.h)) { continue; }
30
+ for (var xx = 0u; xx < params.kx; xx = xx + 1u) {
31
+ let ix = i32(ox * params.sx + xx) - i32(params.px);
32
+ if (ix < 0 || ix >= i32(params.w)) { continue; }
33
+ best = max(best, input[(u32(iy) * params.w + u32(ix)) * params.c + ch]);
34
+ }
35
+ }
36
+ output[(oy * params.out_w + ox) * params.c + ch] = best;
37
+ }
@@ -0,0 +1,18 @@
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
+ if (x >= params.in_w || c >= params.in_c) { return; }
12
+
13
+ var sum_val = 0.0;
14
+ for (var y = 0u; y < params.in_h; y = y + 1u) {
15
+ sum_val = sum_val + input[(y * params.in_w + x) * params.in_c + c];
16
+ }
17
+ output[c * params.in_w + x] = sum_val / f32(params.in_h);
18
+ }
@@ -0,0 +1,32 @@
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
+ output[idx] = a_val * b_val;
32
+ }