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,86 @@
1
+ @group(0) @binding(0) var<storage, read> input : array<f32>;
2
+ @group(0) @binding(1) var<storage, read> weight : array<vec4<f32>>;
3
+ @group(0) @binding(2) var<storage, read> bias : array<vec4<f32>>;
4
+ @group(0) @binding(3) var<storage, read_write> output : array<vec4<f32>>;
5
+
6
+ struct Params {
7
+ n : u32,
8
+ in_h : u32,
9
+ in_w : u32,
10
+ in_c : u32,
11
+ out_c : u32,
12
+ out_h : u32,
13
+ out_w : u32,
14
+ kh : u32,
15
+ kw : u32,
16
+ sy : u32,
17
+ sx : u32,
18
+ pt : u32,
19
+ pl : u32,
20
+ groups : u32,
21
+ relu : u32,
22
+ dy : u32,
23
+ dx : u32,
24
+ }
25
+ @group(0) @binding(4) var<uniform> params : Params;
26
+
27
+ const TILE_C : u32 = 64u;
28
+ var<workgroup> tile_weight : array<vec4<f32>, 256>;
29
+
30
+ fn apply_relu4(v_in : vec4<f32>) -> vec4<f32> {
31
+ var v = v_in;
32
+ if (params.relu == 1u) {
33
+ v = max(v, vec4<f32>(0.0));
34
+ } else if (params.relu >= 2u) {
35
+ v = min(max(v, vec4<f32>(0.0)), vec4<f32>(6.0));
36
+ }
37
+ return v;
38
+ }
39
+
40
+ @compute @workgroup_size(8, 8, 1)
41
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>,
42
+ @builtin(local_invocation_id) lid3 : vec3<u32>) {
43
+ let ox = gid.x;
44
+ let oy = gid.y;
45
+ let oc16_count = params.out_c / 16u;
46
+ let nb = gid.z / oc16_count;
47
+ let oc0 = (gid.z - nb * oc16_count) * 16u;
48
+ let lid = lid3.y * 8u + lid3.x;
49
+ let in_bounds = nb < params.n && ox < params.out_w && oy < params.out_h;
50
+
51
+ var sum0 = vec4<f32>(0.0);
52
+ var sum1 = vec4<f32>(0.0);
53
+ var sum2 = vec4<f32>(0.0);
54
+ var sum3 = vec4<f32>(0.0);
55
+ let input_base = ((nb * params.in_h + oy) * params.in_w + ox) * params.in_c;
56
+
57
+ for (var tile_start = 0u; tile_start < params.in_c; tile_start = tile_start + TILE_C) {
58
+ let tile_count = min(TILE_C, params.in_c - tile_start);
59
+ for (var wi = lid; wi < tile_count * 4u; wi = wi + 64u) {
60
+ let ic = tile_start + wi / 4u;
61
+ let oc_vec = wi - (wi / 4u) * 4u;
62
+ tile_weight[wi] = weight[(ic * params.out_c + oc0) / 4u + oc_vec];
63
+ }
64
+ workgroupBarrier();
65
+
66
+ if (in_bounds) {
67
+ for (var ti = 0u; ti < tile_count; ti = ti + 1u) {
68
+ let x = input[input_base + tile_start + ti];
69
+ let wbase = ti * 4u;
70
+ sum0 = sum0 + x * tile_weight[wbase];
71
+ sum1 = sum1 + x * tile_weight[wbase + 1u];
72
+ sum2 = sum2 + x * tile_weight[wbase + 2u];
73
+ sum3 = sum3 + x * tile_weight[wbase + 3u];
74
+ }
75
+ }
76
+ workgroupBarrier();
77
+ }
78
+
79
+ if (!in_bounds) { return; }
80
+ let output_base = (((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc0) / 4u;
81
+ let bias_base = oc0 / 4u;
82
+ output[output_base] = apply_relu4(sum0 + bias[bias_base]);
83
+ output[output_base + 1u] = apply_relu4(sum1 + bias[bias_base + 1u]);
84
+ output[output_base + 2u] = apply_relu4(sum2 + bias[bias_base + 2u]);
85
+ output[output_base + 3u] = apply_relu4(sum3 + bias[bias_base + 3u]);
86
+ }
@@ -0,0 +1,85 @@
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
+ n : u32,
8
+ in_h : u32,
9
+ in_w : u32,
10
+ in_c : u32,
11
+ out_c : u32,
12
+ out_h : u32,
13
+ out_w : u32,
14
+ kh : u32,
15
+ kw : u32,
16
+ sy : u32,
17
+ sx : u32,
18
+ pt : u32,
19
+ pl : u32,
20
+ groups : u32,
21
+ relu : u32,
22
+ dy : u32,
23
+ dx : u32,
24
+ }
25
+ @group(0) @binding(4) var<uniform> params : Params;
26
+
27
+ fn apply_relu(v_in : f32) -> f32 {
28
+ var v = v_in;
29
+ if (params.relu == 1u) {
30
+ v = max(v, 0.0);
31
+ } else if (params.relu >= 2u) {
32
+ v = min(max(v, 0.0), 6.0);
33
+ }
34
+ return v;
35
+ }
36
+
37
+ @compute @workgroup_size(8, 8, 1)
38
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
39
+ let ox = gid.x;
40
+ let oy = gid.y;
41
+ let oc8_count = (params.out_c + 7u) / 8u;
42
+ let nb = gid.z / oc8_count;
43
+ let oc0 = (gid.z - nb * oc8_count) * 8u;
44
+ if (nb >= params.n || ox >= params.out_w || oy >= params.out_h || oc0 >= params.out_c) { return; }
45
+
46
+ let has1 = oc0 + 1u < params.out_c;
47
+ let has2 = oc0 + 2u < params.out_c;
48
+ let has3 = oc0 + 3u < params.out_c;
49
+ let has4 = oc0 + 4u < params.out_c;
50
+ let has5 = oc0 + 5u < params.out_c;
51
+ let has6 = oc0 + 6u < params.out_c;
52
+ let has7 = oc0 + 7u < params.out_c;
53
+ var sum0 = 0.0;
54
+ var sum1 = 0.0;
55
+ var sum2 = 0.0;
56
+ var sum3 = 0.0;
57
+ var sum4 = 0.0;
58
+ var sum5 = 0.0;
59
+ var sum6 = 0.0;
60
+ var sum7 = 0.0;
61
+
62
+ let input_base = ((nb * params.in_h + oy) * params.in_w + ox) * params.in_c;
63
+ for (var ic = 0u; ic < params.in_c; ic = ic + 1u) {
64
+ let x = input[input_base + ic];
65
+ let wbase = ic * params.out_c + oc0;
66
+ sum0 = sum0 + x * weight[wbase];
67
+ if (has1) { sum1 = sum1 + x * weight[wbase + 1u]; }
68
+ if (has2) { sum2 = sum2 + x * weight[wbase + 2u]; }
69
+ if (has3) { sum3 = sum3 + x * weight[wbase + 3u]; }
70
+ if (has4) { sum4 = sum4 + x * weight[wbase + 4u]; }
71
+ if (has5) { sum5 = sum5 + x * weight[wbase + 5u]; }
72
+ if (has6) { sum6 = sum6 + x * weight[wbase + 6u]; }
73
+ if (has7) { sum7 = sum7 + x * weight[wbase + 7u]; }
74
+ }
75
+
76
+ let output_base = ((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc0;
77
+ output[output_base] = apply_relu(sum0 + bias[oc0]);
78
+ if (has1) { output[output_base + 1u] = apply_relu(sum1 + bias[oc0 + 1u]); }
79
+ if (has2) { output[output_base + 2u] = apply_relu(sum2 + bias[oc0 + 2u]); }
80
+ if (has3) { output[output_base + 3u] = apply_relu(sum3 + bias[oc0 + 3u]); }
81
+ if (has4) { output[output_base + 4u] = apply_relu(sum4 + bias[oc0 + 4u]); }
82
+ if (has5) { output[output_base + 5u] = apply_relu(sum5 + bias[oc0 + 5u]); }
83
+ if (has6) { output[output_base + 6u] = apply_relu(sum6 + bias[oc0 + 6u]); }
84
+ if (has7) { output[output_base + 7u] = apply_relu(sum7 + bias[oc0 + 7u]); }
85
+ }
@@ -0,0 +1,70 @@
1
+ @group(0) @binding(0) var<storage, read> input : array<f32>;
2
+ @group(0) @binding(1) var<storage, read> weight : array<vec2<f32>>;
3
+ @group(0) @binding(2) var<storage, read> bias : array<vec2<f32>>;
4
+ @group(0) @binding(3) var<storage, read_write> output : array<vec2<f32>>;
5
+
6
+ struct Params {
7
+ n : u32,
8
+ in_h : u32,
9
+ in_w : u32,
10
+ in_c : u32,
11
+ out_c : u32,
12
+ out_h : u32,
13
+ out_w : u32,
14
+ kh : u32,
15
+ kw : u32,
16
+ sy : u32,
17
+ sx : u32,
18
+ pt : u32,
19
+ pl : u32,
20
+ groups : u32,
21
+ relu : u32,
22
+ dy : u32,
23
+ dx : u32,
24
+ }
25
+ @group(0) @binding(4) var<uniform> params : Params;
26
+
27
+ fn apply_relu2(v_in : vec2<f32>) -> vec2<f32> {
28
+ var v = v_in;
29
+ if (params.relu == 1u) {
30
+ v = max(v, vec2<f32>(0.0));
31
+ } else if (params.relu >= 2u) {
32
+ v = min(max(v, vec2<f32>(0.0)), vec2<f32>(6.0));
33
+ }
34
+ return v;
35
+ }
36
+
37
+ @compute @workgroup_size(8, 8, 1)
38
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
39
+ let ox = gid.x;
40
+ let oy = gid.y;
41
+ let oc8_count = (params.out_c + 7u) / 8u;
42
+ let nb = gid.z / oc8_count;
43
+ let oc0 = (gid.z - nb * oc8_count) * 8u;
44
+ if (nb >= params.n || ox >= params.out_w || oy >= params.out_h || oc0 >= params.out_c) { return; }
45
+
46
+ let has1 = oc0 + 2u < params.out_c;
47
+ let has2 = oc0 + 4u < params.out_c;
48
+ let has3 = oc0 + 6u < params.out_c;
49
+ var sum0 = vec2<f32>(0.0);
50
+ var sum1 = vec2<f32>(0.0);
51
+ var sum2 = vec2<f32>(0.0);
52
+ var sum3 = vec2<f32>(0.0);
53
+
54
+ let input_base = ((nb * params.in_h + oy) * params.in_w + ox) * params.in_c;
55
+ for (var ic = 0u; ic < params.in_c; ic = ic + 1u) {
56
+ let x = input[input_base + ic];
57
+ let wbase = (ic * params.out_c + oc0) / 2u;
58
+ sum0 = sum0 + x * weight[wbase];
59
+ if (has1) { sum1 = sum1 + x * weight[wbase + 1u]; }
60
+ if (has2) { sum2 = sum2 + x * weight[wbase + 2u]; }
61
+ if (has3) { sum3 = sum3 + x * weight[wbase + 3u]; }
62
+ }
63
+
64
+ let output_base = (((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc0) / 2u;
65
+ let bias_base = oc0 / 2u;
66
+ output[output_base] = apply_relu2(sum0 + bias[bias_base]);
67
+ if (has1) { output[output_base + 1u] = apply_relu2(sum1 + bias[bias_base + 1u]); }
68
+ if (has2) { output[output_base + 2u] = apply_relu2(sum2 + bias[bias_base + 2u]); }
69
+ if (has3) { output[output_base + 3u] = apply_relu2(sum3 + bias[bias_base + 3u]); }
70
+ }
@@ -0,0 +1,65 @@
1
+ @group(0) @binding(0) var<storage, read> input : array<f32>;
2
+ @group(0) @binding(1) var<storage, read> weight : array<vec4<f32>>;
3
+ @group(0) @binding(2) var<storage, read> bias : array<vec4<f32>>;
4
+ @group(0) @binding(3) var<storage, read_write> output : array<vec4<f32>>;
5
+
6
+ struct Params {
7
+ n : u32,
8
+ in_h : u32,
9
+ in_w : u32,
10
+ in_c : u32,
11
+ out_c : u32,
12
+ out_h : u32,
13
+ out_w : u32,
14
+ kh : u32,
15
+ kw : u32,
16
+ sy : u32,
17
+ sx : u32,
18
+ pt : u32,
19
+ pl : u32,
20
+ groups : u32,
21
+ relu : u32,
22
+ dy : u32,
23
+ dx : u32,
24
+ }
25
+ @group(0) @binding(4) var<uniform> params : Params;
26
+
27
+ fn apply_relu4(v_in : vec4<f32>) -> vec4<f32> {
28
+ var v = v_in;
29
+ if (params.relu == 1u) {
30
+ v = max(v, vec4<f32>(0.0));
31
+ } else if (params.relu >= 2u) {
32
+ v = min(max(v, vec4<f32>(0.0)), vec4<f32>(6.0));
33
+ }
34
+ return v;
35
+ }
36
+
37
+ @compute @workgroup_size(8, 8, 1)
38
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
39
+ let ox = gid.x;
40
+ let oy = gid.y;
41
+ let oc8_count = (params.out_c + 7u) / 8u;
42
+ let nb = gid.z / oc8_count;
43
+ let oc0 = (gid.z - nb * oc8_count) * 8u;
44
+ if (nb >= params.n || ox >= params.out_w || oy >= params.out_h || oc0 >= params.out_c) { return; }
45
+
46
+ var sum0 = vec4<f32>(0.0);
47
+ var sum1 = vec4<f32>(0.0);
48
+ let input_base = ((nb * params.in_h + oy) * params.in_w + ox) * params.in_c;
49
+ let has_second = oc0 + 4u < params.out_c;
50
+ for (var ic = 0u; ic < params.in_c; ic = ic + 1u) {
51
+ let x = input[input_base + ic];
52
+ let wbase = (ic * params.out_c + oc0) / 4u;
53
+ sum0 = sum0 + x * weight[wbase];
54
+ if (has_second) {
55
+ sum1 = sum1 + x * weight[wbase + 1u];
56
+ }
57
+ }
58
+
59
+ let output_base = (((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc0) / 4u;
60
+ let bias_base = oc0 / 4u;
61
+ output[output_base] = apply_relu4(sum0 + bias[bias_base]);
62
+ if (has_second) {
63
+ output[output_base + 1u] = apply_relu4(sum1 + bias[bias_base + 1u]);
64
+ }
65
+ }
@@ -0,0 +1,75 @@
1
+ @group(0) @binding(0) var<storage, read> input : array<f32>;
2
+ @group(0) @binding(1) var<storage, read> weight : array<vec4<f32>>;
3
+ @group(0) @binding(2) var<storage, read> bias : array<vec4<f32>>;
4
+ @group(0) @binding(3) var<storage, read_write> output : array<vec4<f32>>;
5
+
6
+ struct Params {
7
+ n : u32,
8
+ in_h : u32,
9
+ in_w : u32,
10
+ in_c : u32,
11
+ out_c : u32,
12
+ out_h : u32,
13
+ out_w : u32,
14
+ kh : u32,
15
+ kw : u32,
16
+ sy : u32,
17
+ sx : u32,
18
+ pt : u32,
19
+ pl : u32,
20
+ groups : u32,
21
+ relu : u32,
22
+ dy : u32,
23
+ dx : u32,
24
+ }
25
+ @group(0) @binding(4) var<uniform> params : Params;
26
+
27
+ fn apply_relu4(v_in : vec4<f32>) -> vec4<f32> {
28
+ var v = v_in;
29
+ if (params.relu == 1u) {
30
+ v = max(v, vec4<f32>(0.0));
31
+ } else if (params.relu >= 2u) {
32
+ v = min(max(v, vec4<f32>(0.0)), vec4<f32>(6.0));
33
+ }
34
+ return v;
35
+ }
36
+
37
+ @compute @workgroup_size(8, 8, 1)
38
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
39
+ let ox = gid.x;
40
+ let oy = gid.y;
41
+ let oc16_count = params.out_c / 16u;
42
+ let nb = gid.z / oc16_count;
43
+ let oc0 = (gid.z - nb * oc16_count) * 16u;
44
+ if (nb >= params.n || ox >= params.out_w || oy >= params.out_h) { return; }
45
+
46
+ var sum0 = vec4<f32>(0.0);
47
+ var sum1 = vec4<f32>(0.0);
48
+ var sum2 = vec4<f32>(0.0);
49
+ var sum3 = vec4<f32>(0.0);
50
+ for (var yy = 0u; yy < params.kh; yy = yy + 1u) {
51
+ let iy = i32(oy * params.sy + yy * params.dy) - i32(params.pt);
52
+ if (iy < 0 || iy >= i32(params.in_h)) { continue; }
53
+ for (var xx = 0u; xx < params.kw; xx = xx + 1u) {
54
+ let ix = i32(ox * params.sx + xx * params.dx) - i32(params.pl);
55
+ if (ix < 0 || ix >= i32(params.in_w)) { continue; }
56
+ let input_base = ((nb * params.in_h + u32(iy)) * params.in_w + u32(ix)) * 3u;
57
+ let weight_base = (((yy * params.kw + xx) * 3u) * params.out_c + oc0) / 4u;
58
+ for (var ic = 0u; ic < 3u; ic = ic + 1u) {
59
+ let x = input[input_base + ic];
60
+ let wbase = weight_base + ic * (params.out_c / 4u);
61
+ sum0 = sum0 + x * weight[wbase];
62
+ sum1 = sum1 + x * weight[wbase + 1u];
63
+ sum2 = sum2 + x * weight[wbase + 2u];
64
+ sum3 = sum3 + x * weight[wbase + 3u];
65
+ }
66
+ }
67
+ }
68
+
69
+ let output_base = (((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc0) / 4u;
70
+ let bias_base = oc0 / 4u;
71
+ output[output_base] = apply_relu4(sum0 + bias[bias_base]);
72
+ output[output_base + 1u] = apply_relu4(sum1 + bias[bias_base + 1u]);
73
+ output[output_base + 2u] = apply_relu4(sum2 + bias[bias_base + 2u]);
74
+ output[output_base + 3u] = apply_relu4(sum3 + bias[bias_base + 3u]);
75
+ }
@@ -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
+ struct Params { b: u32, in_h: u32, in_w: u32, in_c: u32, out_h: u32, out_w: u32, out_c: u32, kh: u32, kw: u32, sh: u32, sw: u32, ph: u32, pw: u32, has_bias: u32 }
6
+ @group(0) @binding(4) var<uniform> params : Params;
7
+ @compute @workgroup_size(8, 8, 1)
8
+ fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
9
+ let x = global_id.x; let y = global_id.y; let batch_c = global_id.z;
10
+ if (x >= params.out_w || y >= params.out_h || batch_c >= (params.b * params.out_c)) { return; }
11
+ let ob = batch_c / params.out_c;
12
+ let oc = batch_c % params.out_c;
13
+ var sum = 0.0;
14
+ if (params.has_bias == 1u) { sum = bias[oc]; }
15
+ for (var ic = 0u; ic < params.in_c; ic = ic + 1u) {
16
+ for (var ky = 0u; ky < params.kh; ky = ky + 1u) {
17
+ for (var kx = 0u; kx < params.kw; kx = kx + 1u) {
18
+ let oy_shifted = i32(y) + i32(params.ph) - i32(ky);
19
+ let ox_shifted = i32(x) + i32(params.pw) - i32(kx);
20
+ if (oy_shifted % i32(params.sh) == 0 && ox_shifted % i32(params.sw) == 0) {
21
+ let iy = oy_shifted / i32(params.sh);
22
+ let ix = ox_shifted / i32(params.sw);
23
+ if (iy >= 0 && iy < i32(params.in_h) && ix >= 0 && ix < i32(params.in_w)) {
24
+ let in_val = input[((ob * params.in_h + u32(iy)) * params.in_w + u32(ix)) * params.in_c + ic];
25
+ let w_val = weight[ic * (params.out_c * params.kh * params.kw) + oc * (params.kh * params.kw) + ky * params.kw + kx];
26
+ sum = sum + in_val * w_val;
27
+ }
28
+ }
29
+ }
30
+ }
31
+ }
32
+ output[((ob * params.out_h + y) * params.out_w + x) * params.out_c + oc] = sum;
33
+ }
@@ -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;
12
+ output[idx] = out_val;
13
+ }
@@ -0,0 +1,140 @@
1
+ @group(0) @binding(0) var<storage, read> q_in : array<f32>;
2
+ @group(0) @binding(1) var<storage, read> kv_in : array<f32>;
3
+ @group(0) @binding(2) var<storage, read> weight_int8_packed : array<u32>;
4
+ @group(0) @binding(3) var<storage, read> w_scale : array<f32>;
5
+ @group(0) @binding(4) var<storage, read> bias : array<f32>;
6
+ @group(0) @binding(5) var<storage, read_write> output : array<f32>;
7
+
8
+ struct Params {
9
+ seq_len_q : u32,
10
+ seq_len_kv : u32,
11
+ d_model : u32,
12
+ num_heads : u32,
13
+ head_dim : u32,
14
+ scale_factor : f32,
15
+ has_scale : u32,
16
+ has_bias : u32,
17
+ }
18
+ @group(0) @binding(6) var<uniform> params : Params;
19
+
20
+ @compute @workgroup_size(64, 1, 1)
21
+ fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
22
+ let q_idx = global_id.x;
23
+ let h_idx = global_id.y;
24
+
25
+ if (q_idx >= params.seq_len_q || h_idx >= params.num_heads) { return; }
26
+
27
+ let head_dim = params.head_dim;
28
+ let d_model = params.d_model;
29
+ let d_model_4 = d_model / 4u;
30
+
31
+ // Cache Q for this head and this q_idx
32
+ var q_cache : array<f32, 64>; // max head_dim 64
33
+ for (var d = 0u; d < head_dim; d = d + 1u) {
34
+ var sum = 0.0;
35
+ let out_col = h_idx * head_dim + d;
36
+ for (var i = 0u; i < d_model_4; i = i + 1u) {
37
+ let w_packed = weight_int8_packed[out_col * d_model_4 + i];
38
+ let in_base = q_idx * d_model + i * 4u;
39
+ sum = sum + q_in[in_base + 0u] * f32(extractBits(i32(w_packed), 0u, 8u));
40
+ sum = sum + q_in[in_base + 1u] * f32(extractBits(i32(w_packed), 8u, 8u));
41
+ sum = sum + q_in[in_base + 2u] * f32(extractBits(i32(w_packed), 16u, 8u));
42
+ sum = sum + q_in[in_base + 3u] * f32(extractBits(i32(w_packed), 24u, 8u));
43
+ }
44
+ if (params.has_scale > 0u) { sum = sum * w_scale[out_col]; }
45
+ if (params.has_bias > 0u) { sum = sum + bias[out_col]; }
46
+ q_cache[d] = sum;
47
+ }
48
+
49
+ var max_logit : f32 = -1e38;
50
+
51
+ // Pass 1: find max
52
+ for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
53
+ var score : f32 = 0.0;
54
+ for (var d = 0u; d < head_dim; d = d + 1u) {
55
+ var k_val = 0.0;
56
+ let out_col = d_model + h_idx * head_dim + d;
57
+ for (var i = 0u; i < d_model_4; i = i + 1u) {
58
+ let w_packed = weight_int8_packed[out_col * d_model_4 + i];
59
+ let in_base = k_idx * d_model + i * 4u;
60
+ k_val = k_val + kv_in[in_base + 0u] * f32(extractBits(i32(w_packed), 0u, 8u));
61
+ k_val = k_val + kv_in[in_base + 1u] * f32(extractBits(i32(w_packed), 8u, 8u));
62
+ k_val = k_val + kv_in[in_base + 2u] * f32(extractBits(i32(w_packed), 16u, 8u));
63
+ k_val = k_val + kv_in[in_base + 3u] * f32(extractBits(i32(w_packed), 24u, 8u));
64
+ }
65
+ if (params.has_scale > 0u) { k_val = k_val * w_scale[out_col]; }
66
+ if (params.has_bias > 0u) { k_val = k_val + bias[out_col]; }
67
+
68
+ score = score + (q_cache[d] * k_val);
69
+ }
70
+ score = score * params.scale_factor;
71
+ if (score > max_logit) { max_logit = score; }
72
+ }
73
+
74
+ // Pass 2: sum exp
75
+ var sum_exp : f32 = 0.0;
76
+ for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
77
+ var score : f32 = 0.0;
78
+ for (var d = 0u; d < head_dim; d = d + 1u) {
79
+ var k_val = 0.0;
80
+ let out_col = d_model + h_idx * head_dim + d;
81
+ for (var i = 0u; i < d_model_4; i = i + 1u) {
82
+ let w_packed = weight_int8_packed[out_col * d_model_4 + i];
83
+ let in_base = k_idx * d_model + i * 4u;
84
+ k_val = k_val + kv_in[in_base + 0u] * f32(extractBits(i32(w_packed), 0u, 8u));
85
+ k_val = k_val + kv_in[in_base + 1u] * f32(extractBits(i32(w_packed), 8u, 8u));
86
+ k_val = k_val + kv_in[in_base + 2u] * f32(extractBits(i32(w_packed), 16u, 8u));
87
+ k_val = k_val + kv_in[in_base + 3u] * f32(extractBits(i32(w_packed), 24u, 8u));
88
+ }
89
+ if (params.has_scale > 0u) { k_val = k_val * w_scale[out_col]; }
90
+ if (params.has_bias > 0u) { k_val = k_val + bias[out_col]; }
91
+
92
+ score = score + (q_cache[d] * k_val);
93
+ }
94
+ score = score * params.scale_factor;
95
+ sum_exp = sum_exp + exp(score - max_logit);
96
+ }
97
+
98
+ // Pass 3: output
99
+ for (var d = 0u; d < head_dim; d = d + 1u) {
100
+ var out_val : f32 = 0.0;
101
+ for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
102
+ var score : f32 = 0.0;
103
+ for (var kd = 0u; kd < head_dim; kd = kd + 1u) {
104
+ var k_val = 0.0;
105
+ let out_col = d_model + h_idx * head_dim + kd;
106
+ for (var i = 0u; i < d_model_4; i = i + 1u) {
107
+ let w_packed = weight_int8_packed[out_col * d_model_4 + i];
108
+ let in_base = k_idx * d_model + i * 4u;
109
+ k_val = k_val + kv_in[in_base + 0u] * f32(extractBits(i32(w_packed), 0u, 8u));
110
+ k_val = k_val + kv_in[in_base + 1u] * f32(extractBits(i32(w_packed), 8u, 8u));
111
+ k_val = k_val + kv_in[in_base + 2u] * f32(extractBits(i32(w_packed), 16u, 8u));
112
+ k_val = k_val + kv_in[in_base + 3u] * f32(extractBits(i32(w_packed), 24u, 8u));
113
+ }
114
+ if (params.has_scale > 0u) { k_val = k_val * w_scale[out_col]; }
115
+ if (params.has_bias > 0u) { k_val = k_val + bias[out_col]; }
116
+
117
+ score = score + (q_cache[kd] * k_val);
118
+ }
119
+ score = score * params.scale_factor;
120
+
121
+ let w = exp(score - max_logit) / sum_exp;
122
+
123
+ var v_val = 0.0;
124
+ let v_col = d_model * 2u + h_idx * head_dim + d;
125
+ for (var i = 0u; i < d_model_4; i = i + 1u) {
126
+ let w_packed = weight_int8_packed[v_col * d_model_4 + i];
127
+ let in_base = k_idx * d_model + i * 4u;
128
+ v_val = v_val + kv_in[in_base + 0u] * f32(extractBits(i32(w_packed), 0u, 8u));
129
+ v_val = v_val + kv_in[in_base + 1u] * f32(extractBits(i32(w_packed), 8u, 8u));
130
+ v_val = v_val + kv_in[in_base + 2u] * f32(extractBits(i32(w_packed), 16u, 8u));
131
+ v_val = v_val + kv_in[in_base + 3u] * f32(extractBits(i32(w_packed), 24u, 8u));
132
+ }
133
+ if (params.has_scale > 0u) { v_val = v_val * w_scale[v_col]; }
134
+ if (params.has_bias > 0u) { v_val = v_val + bias[v_col]; }
135
+
136
+ out_val = out_val + (w * v_val);
137
+ }
138
+ output[q_idx * d_model + (h_idx * head_dim) + d] = out_val;
139
+ }
140
+ }
@@ -0,0 +1,98 @@
1
+ @group(0) @binding(0) var<storage, read> q_in : array<f32>;
2
+ @group(0) @binding(1) var<storage, read> kv_in : array<f32>;
3
+ @group(0) @binding(2) var<storage, read> weight : array<f32>;
4
+ @group(0) @binding(3) var<storage, read> w_scale : array<f32>;
5
+ @group(0) @binding(4) var<storage, read> bias : array<f32>;
6
+ @group(0) @binding(5) var<storage, read_write> output : array<f32>;
7
+
8
+ struct Params {
9
+ seq_len_q : u32,
10
+ seq_len_kv : u32,
11
+ d_model : u32,
12
+ num_heads : u32,
13
+ head_dim : u32,
14
+ scale_factor : f32,
15
+ has_scale : u32,
16
+ has_bias : u32,
17
+ }
18
+ @group(0) @binding(6) var<uniform> params : Params;
19
+
20
+ fn project(src : ptr<function, array<f32, 64>>, row : u32, out_col : u32) -> f32 {
21
+ var sum = 0.0;
22
+ let w_base = out_col * params.d_model;
23
+ for (var i = 0u; i < params.d_model; i = i + 1u) {
24
+ sum = sum + (*src)[i] * weight[w_base + i];
25
+ }
26
+ if (params.has_scale > 0u) { sum = sum * w_scale[out_col]; }
27
+ if (params.has_bias > 0u) { sum = sum + bias[out_col]; }
28
+ return sum;
29
+ }
30
+
31
+ @compute @workgroup_size(64, 1, 1)
32
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
33
+ let q_idx = gid.x;
34
+ let h_idx = gid.y;
35
+ if (q_idx >= params.seq_len_q || h_idx >= params.num_heads) { return; }
36
+ if (params.d_model > 64u || params.head_dim > 64u) { return; }
37
+
38
+ var q_src : array<f32, 64>;
39
+ for (var i = 0u; i < params.d_model; i = i + 1u) {
40
+ q_src[i] = q_in[q_idx * params.d_model + i];
41
+ }
42
+
43
+ var q_proj : array<f32, 64>;
44
+ for (var d = 0u; d < params.head_dim; d = d + 1u) {
45
+ let out_col = h_idx * params.head_dim + d;
46
+ q_proj[d] = project(&q_src, q_idx, out_col);
47
+ }
48
+
49
+ var max_logit = -1e38;
50
+ for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
51
+ var kv_src : array<f32, 64>;
52
+ for (var i = 0u; i < params.d_model; i = i + 1u) {
53
+ kv_src[i] = kv_in[k_idx * params.d_model + i];
54
+ }
55
+ var score = 0.0;
56
+ for (var d = 0u; d < params.head_dim; d = d + 1u) {
57
+ let out_col = params.d_model + h_idx * params.head_dim + d;
58
+ score = score + q_proj[d] * project(&kv_src, k_idx, out_col);
59
+ }
60
+ score = score * params.scale_factor;
61
+ if (score > max_logit) { max_logit = score; }
62
+ }
63
+
64
+ var sum_exp = 0.0;
65
+ for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
66
+ var kv_src : array<f32, 64>;
67
+ for (var i = 0u; i < params.d_model; i = i + 1u) {
68
+ kv_src[i] = kv_in[k_idx * params.d_model + i];
69
+ }
70
+ var score = 0.0;
71
+ for (var d = 0u; d < params.head_dim; d = d + 1u) {
72
+ let out_col = params.d_model + h_idx * params.head_dim + d;
73
+ score = score + q_proj[d] * project(&kv_src, k_idx, out_col);
74
+ }
75
+ score = score * params.scale_factor;
76
+ sum_exp = sum_exp + exp(score - max_logit);
77
+ }
78
+
79
+ for (var d = 0u; d < params.head_dim; d = d + 1u) {
80
+ var out_val = 0.0;
81
+ for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
82
+ var kv_src : array<f32, 64>;
83
+ for (var i = 0u; i < params.d_model; i = i + 1u) {
84
+ kv_src[i] = kv_in[k_idx * params.d_model + i];
85
+ }
86
+ var score = 0.0;
87
+ for (var kd = 0u; kd < params.head_dim; kd = kd + 1u) {
88
+ let k_col = params.d_model + h_idx * params.head_dim + kd;
89
+ score = score + q_proj[kd] * project(&kv_src, k_idx, k_col);
90
+ }
91
+ score = score * params.scale_factor;
92
+ let attn = exp(score - max_logit) / sum_exp;
93
+ let v_col = params.d_model * 2u + h_idx * params.head_dim + d;
94
+ out_val = out_val + attn * project(&kv_src, k_idx, v_col);
95
+ }
96
+ output[q_idx * params.d_model + h_idx * params.head_dim + d] = out_val;
97
+ }
98
+ }