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,15 @@
1
+ export function _cpuWhere(node) {
2
+ const cond = node.inputs.cond || node.inputs.condition;
3
+ const a = node.inputs.x || node.inputs.a;
4
+ const b = node.inputs.y || node.inputs.b;
5
+ const outBuf = node.outputs.out.buffer;
6
+
7
+ const condBuf = cond.buffer;
8
+ const aBuf = a.buffer;
9
+ const bBuf = b.buffer;
10
+
11
+ const elements = outBuf.length;
12
+ for (let i = 0; i < elements; i++) {
13
+ outBuf[i] = condBuf[i] !== 0 ? aBuf[i] : bBuf[i];
14
+ }
15
+ }
package/package.json ADDED
@@ -0,0 +1,33 @@
1
+ {
2
+ "name": "volvoxai",
3
+ "version": "0.1.0",
4
+ "description": "A Zero-Dependency, Bare-Metal Deep Learning Engine for the Browser and Node.js",
5
+ "main": "js/index.js",
6
+ "type": "module",
7
+ "bin": {
8
+ "volvox": "./bin/volvox.js"
9
+ },
10
+ "files": [
11
+ "bin/",
12
+ "dist/",
13
+ "docs/",
14
+ "js/",
15
+ "shaders/",
16
+ "volvoxai.wasm"
17
+ ],
18
+ "scripts": {
19
+ "build": "esbuild js/index.js --bundle --format=esm --loader:.wgsl=text --outfile=dist/volvoxai.js",
20
+ "build:min": "esbuild js/index.js --bundle --minify --format=esm --loader:.wgsl=text --outfile=dist/volvoxai.min.js",
21
+ "build:all": "npm run build && npm run build:min",
22
+ "test": "echo 'Run tests manually for now'",
23
+ "start": "node ./bin/volvox.js"
24
+ },
25
+ "dependencies": {
26
+ "commander": "^11.1.0"
27
+ },
28
+ "devDependencies": {
29
+ "esbuild": "^0.28.1"
30
+ },
31
+ "author": "",
32
+ "license": "MIT"
33
+ }
@@ -0,0 +1,13 @@
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 }
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
+ output[idx] = a[idx] + b[idx];
13
+ }
@@ -0,0 +1,23 @@
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> c : array<f32>;
4
+ @group(0) @binding(3) var<storage, read_write> output : array<f32>;
5
+
6
+ struct Params {
7
+ size : u32,
8
+ relu : u32,
9
+ }
10
+ @group(0) @binding(4) var<uniform> params : Params;
11
+
12
+ @compute @workgroup_size(64)
13
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
14
+ let i = gid.x;
15
+ if (i >= params.size) { return; }
16
+ var v = a[i] + b[i] + c[i];
17
+ if (params.relu == 1u) {
18
+ v = max(v, 0.0);
19
+ } else if (params.relu >= 2u) {
20
+ v = min(max(v, 0.0), 6.0);
21
+ }
22
+ output[i] = v;
23
+ }
@@ -0,0 +1,22 @@
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 {
6
+ size : u32,
7
+ relu : u32,
8
+ }
9
+ @group(0) @binding(3) var<uniform> params : Params;
10
+
11
+ @compute @workgroup_size(64)
12
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
13
+ let i = gid.x;
14
+ if (i >= params.size) { return; }
15
+ var v = a[i] + b[i];
16
+ if (params.relu == 1u) {
17
+ v = max(v, 0.0);
18
+ } else if (params.relu >= 2u) {
19
+ v = min(max(v, 0.0), 6.0);
20
+ }
21
+ output[i] = v;
22
+ }
@@ -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
+ struct Params { b: u32, in_h: u32, in_w: u32, c: u32, out_h: u32, out_w: u32, kh: u32, kw: u32, sh: u32, sw: u32, ph: u32, pw: u32 }
4
+ @group(0) @binding(2) var<uniform> params : Params;
5
+ @compute @workgroup_size(8, 8, 1)
6
+ fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
7
+ let x = global_id.x; let y = global_id.y; let z = global_id.z;
8
+ let b = z / params.c;
9
+ let c = z - b * params.c;
10
+ if (x >= params.out_w || y >= params.out_h || b >= params.b) { return; }
11
+ var sum = 0.0; var count = 0u;
12
+ for (var ky = 0u; ky < params.kh; ky = ky + 1u) {
13
+ for (var kx = 0u; kx < params.kw; kx = kx + 1u) {
14
+ let in_y = i32(y * params.sh) - i32(params.ph) + i32(ky);
15
+ let in_x = i32(x * params.sw) - i32(params.pw) + i32(kx);
16
+ if (in_y >= 0 && in_y < i32(params.in_h) && in_x >= 0 && in_x < i32(params.in_w)) {
17
+ sum = sum + input[((b * params.in_h + u32(in_y)) * params.in_w + u32(in_x)) * params.c + c];
18
+ count = count + 1u;
19
+ }
20
+ }
21
+ }
22
+ if (count == 0u) { count = 1u; }
23
+ output[((b * params.out_h + y) * params.out_w + x) * params.c + c] = sum / f32(count);
24
+ }
@@ -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> bias : array<f32>;
4
+ @group(0) @binding(3) var<storage, read> running_mean : array<f32>;
5
+ @group(0) @binding(4) var<storage, read> running_var : array<f32>;
6
+ @group(0) @binding(5) var<storage, read_write> output : array<f32>;
7
+ struct Params { b : u32, c : u32, h : u32, w : u32, eps : f32 }
8
+ @group(0) @binding(6) var<uniform> params : Params;
9
+ @compute @workgroup_size(64)
10
+ fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
11
+ let idx = global_id.x;
12
+ let total = params.b * params.c * params.h * params.w;
13
+ if (idx >= total) { return; }
14
+ let c = idx % params.c;
15
+ let mean = running_mean[c];
16
+ let var_val = running_var[c];
17
+ let gamma = weight[c];
18
+ let beta = bias[c];
19
+ let inv_std = 1.0 / sqrt(var_val + params.eps);
20
+ output[idx] = (input[idx] - mean) * inv_std * gamma + beta;
21
+ }
@@ -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
+ undefined
33
+ output[idx] = out_val;
34
+ }
@@ -0,0 +1,26 @@
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
+ @group(0) @binding(3) var<storage, read> md : array<u32>;
5
+ @compute @workgroup_size(64)
6
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
7
+ let idx = gid.x;
8
+ let total = md[0];
9
+ if (idx >= total) { return; }
10
+ let rank = md[1];
11
+ var rem = idx;
12
+ var a_idx = 0u;
13
+ var b_idx = 0u;
14
+ for (var d = 0u; d < rank; d = d + 1u) {
15
+ let os = md[2u + d];
16
+ let coord = rem / os;
17
+ rem = rem - coord * os;
18
+ a_idx = a_idx + coord * md[2u + rank + d];
19
+ b_idx = b_idx + coord * md[2u + 2u * rank + d];
20
+ }
21
+ let av = a[a_idx];
22
+ let bv = b[b_idx];
23
+ var out_val = 0.0;
24
+ //__BINOP__
25
+ output[idx] = out_val;
26
+ }
@@ -0,0 +1,10 @@
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, min_v : f32, max_v : 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
+ output[idx] = clamp(input[idx], params.min_v, params.max_v);
10
+ }
@@ -0,0 +1,16 @@
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 { a_size : u32, b_size : 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.a_size) {
12
+ output[idx] = a[idx];
13
+ } else if (idx < params.a_size + params.b_size) {
14
+ output[idx] = b[idx - params.a_size];
15
+ }
16
+ }
@@ -0,0 +1,10 @@
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, 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.size) { return; }
9
+ output[params.offset + i] = input[i];
10
+ }
@@ -0,0 +1,16 @@
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
+ size : u32,
6
+ offset : u32,
7
+ }
8
+ @group(0) @binding(2) var<uniform> params : Params;
9
+
10
+ @compute @workgroup_size(64)
11
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
12
+ let i = gid.x;
13
+ if (i >= params.size) { return; }
14
+ let v = input[i];
15
+ output[params.offset + i] = 1.0 / (1.0 + exp(-v));
16
+ }
@@ -0,0 +1,37 @@
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
+ in_c : u32, in_l : u32,
8
+ out_c : u32, k : u32,
9
+ stride : u32, pad : u32, relu : u32
10
+ }
11
+ @group(0) @binding(4) var<uniform> params : Params;
12
+
13
+ @compute @workgroup_size(64)
14
+ fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
15
+ let x = global_id.x;
16
+ let oc = global_id.y;
17
+ let out_l = (params.in_l + 2u * params.pad - params.k) / params.stride + 1u;
18
+
19
+ if (x >= out_l || oc >= params.out_c) { return; }
20
+
21
+ var sum = bias[oc];
22
+ let w_base = oc * params.in_c * params.k;
23
+
24
+ for (var ic = 0u; ic < params.in_c; ic = ic + 1u) {
25
+ let in_base = ic * params.in_l;
26
+ let w_ic_base = w_base + ic * params.k;
27
+ for (var k = 0u; k < params.k; k = k + 1u) {
28
+ let in_x = i32(x * params.stride + k) - i32(params.pad);
29
+ if (in_x >= 0 && in_x < i32(params.in_l)) {
30
+ sum = sum + input[in_base + u32(in_x)] * weight[w_ic_base + k];
31
+ }
32
+ }
33
+ }
34
+
35
+ if (params.relu == 1u && sum < 0.0) { sum = 0.0; }
36
+ output[oc * out_l + x] = sum;
37
+ }
@@ -0,0 +1,80 @@
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
+ @compute @workgroup_size(8, 8, 1)
28
+ fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
29
+ let ox = gid.x;
30
+ let oy = gid.y;
31
+ let z = gid.z;
32
+ let nb = z / params.out_c;
33
+ let oc = z - nb * params.out_c;
34
+ if (nb >= params.n || ox >= params.out_w || oy >= params.out_h) { return; }
35
+
36
+ var sum = 0.0;
37
+ if (params.groups == params.in_c) {
38
+ let mult = params.out_c / params.in_c;
39
+ let ic = oc / mult;
40
+ let m = oc - ic * mult;
41
+ for (var yy = 0u; yy < params.kh; yy = yy + 1u) {
42
+ let iy = i32(oy * params.sy + yy * params.dy) - i32(params.pt);
43
+ if (iy < 0 || iy >= i32(params.in_h)) { continue; }
44
+ for (var xx = 0u; xx < params.kw; xx = xx + 1u) {
45
+ let ix = i32(ox * params.sx + xx * params.dx) - i32(params.pl);
46
+ if (ix < 0 || ix >= i32(params.in_w)) { continue; }
47
+ let ii = ((nb * params.in_h + u32(iy)) * params.in_w + u32(ix)) * params.in_c + ic;
48
+ let wi = (((yy * params.kw + xx) * params.in_c + ic) * mult) + m;
49
+ sum = sum + input[ii] * weight[wi];
50
+ }
51
+ }
52
+ } else {
53
+ let out_per_g = params.out_c / params.groups;
54
+ let in_per_g = params.in_c / params.groups;
55
+ let g = oc / out_per_g;
56
+ let ic0 = g * in_per_g;
57
+ for (var icl = 0u; icl < in_per_g; icl = icl + 1u) {
58
+ let ic = ic0 + icl;
59
+ for (var yy = 0u; yy < params.kh; yy = yy + 1u) {
60
+ let iy = i32(oy * params.sy + yy * params.dy) - i32(params.pt);
61
+ if (iy < 0 || iy >= i32(params.in_h)) { continue; }
62
+ for (var xx = 0u; xx < params.kw; xx = xx + 1u) {
63
+ let ix = i32(ox * params.sx + xx * params.dx) - i32(params.pl);
64
+ if (ix < 0 || ix >= i32(params.in_w)) { continue; }
65
+ let ii = ((nb * params.in_h + u32(iy)) * params.in_w + u32(ix)) * params.in_c + ic;
66
+ let wi = (((yy * params.kw + xx) * params.in_c + ic) * params.out_c) + oc;
67
+ sum = sum + input[ii] * weight[wi];
68
+ }
69
+ }
70
+ }
71
+ }
72
+
73
+ var v = sum + bias[oc];
74
+ if (params.relu == 1u) {
75
+ v = max(v, 0.0);
76
+ } else if (params.relu >= 2u) {
77
+ v = min(max(v, 0.0), 6.0);
78
+ }
79
+ output[((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc] = v;
80
+ }
@@ -0,0 +1,74 @@
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 c4_count = (params.out_c + 3u) / 4u;
42
+ let nb = gid.z / c4_count;
43
+ let ch0 = (gid.z - nb * c4_count) * 4u;
44
+ if (nb >= params.n || ox >= params.out_w || oy >= params.out_h || ch0 >= params.out_c) { return; }
45
+
46
+ let has1 = ch0 + 1u < params.out_c;
47
+ let has2 = ch0 + 2u < params.out_c;
48
+ let has3 = ch0 + 3u < params.out_c;
49
+ var sum0 = 0.0;
50
+ var sum1 = 0.0;
51
+ var sum2 = 0.0;
52
+ var sum3 = 0.0;
53
+
54
+ for (var yy = 0u; yy < params.kh; yy = yy + 1u) {
55
+ let iy = i32(oy * params.sy + yy * params.dy) - i32(params.pt);
56
+ if (iy < 0 || iy >= i32(params.in_h)) { continue; }
57
+ for (var xx = 0u; xx < params.kw; xx = xx + 1u) {
58
+ let ix = i32(ox * params.sx + xx * params.dx) - i32(params.pl);
59
+ if (ix < 0 || ix >= i32(params.in_w)) { continue; }
60
+ let input_base = ((nb * params.in_h + u32(iy)) * params.in_w + u32(ix)) * params.in_c + ch0;
61
+ let weight_base = ((yy * params.kw + xx) * params.in_c) + ch0;
62
+ sum0 = sum0 + input[input_base] * weight[weight_base];
63
+ if (has1) { sum1 = sum1 + input[input_base + 1u] * weight[weight_base + 1u]; }
64
+ if (has2) { sum2 = sum2 + input[input_base + 2u] * weight[weight_base + 2u]; }
65
+ if (has3) { sum3 = sum3 + input[input_base + 3u] * weight[weight_base + 3u]; }
66
+ }
67
+ }
68
+
69
+ let output_base = ((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + ch0;
70
+ output[output_base] = apply_relu(sum0 + bias[ch0]);
71
+ if (has1) { output[output_base + 1u] = apply_relu(sum1 + bias[ch0 + 1u]); }
72
+ if (has2) { output[output_base + 2u] = apply_relu(sum2 + bias[ch0 + 2u]); }
73
+ if (has3) { output[output_base + 3u] = apply_relu(sum3 + bias[ch0 + 3u]); }
74
+ }
@@ -0,0 +1,66 @@
1
+ @group(0) @binding(0) var<storage, read> input : array<vec4<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 c8_count = (params.out_c + 7u) / 8u;
42
+ let nb = gid.z / c8_count;
43
+ let ch0 = (gid.z - nb * c8_count) * 8u;
44
+ if (nb >= params.n || ox >= params.out_w || oy >= params.out_h || ch0 >= params.out_c) { return; }
45
+
46
+ var sum0 = vec4<f32>(0.0);
47
+ var sum1 = vec4<f32>(0.0);
48
+
49
+ for (var yy = 0u; yy < params.kh; yy = yy + 1u) {
50
+ let iy = i32(oy * params.sy + yy * params.dy) - i32(params.pt);
51
+ if (iy < 0 || iy >= i32(params.in_h)) { continue; }
52
+ for (var xx = 0u; xx < params.kw; xx = xx + 1u) {
53
+ let ix = i32(ox * params.sx + xx * params.dx) - i32(params.pl);
54
+ if (ix < 0 || ix >= i32(params.in_w)) { continue; }
55
+ let input_base = (((nb * params.in_h + u32(iy)) * params.in_w + u32(ix)) * params.in_c + ch0) / 4u;
56
+ let weight_base = (((yy * params.kw + xx) * params.in_c) + ch0) / 4u;
57
+ sum0 = sum0 + input[input_base] * weight[weight_base];
58
+ sum1 = sum1 + input[input_base + 1u] * weight[weight_base + 1u];
59
+ }
60
+ }
61
+
62
+ let output_base = (((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + ch0) / 4u;
63
+ let bias_base = ch0 / 4u;
64
+ output[output_base] = apply_relu4(sum0 + bias[bias_base]);
65
+ output[output_base + 1u] = apply_relu4(sum1 + bias[bias_base + 1u]);
66
+ }
@@ -0,0 +1,67 @@
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 + 15u) / 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 || oc0 >= params.out_c) { 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
+
51
+ let input_base = ((nb * params.in_h + oy) * params.in_w + ox) * params.in_c;
52
+ for (var ic = 0u; ic < params.in_c; ic = ic + 1u) {
53
+ let x = input[input_base + ic];
54
+ let wbase = (ic * params.out_c + oc0) / 4u;
55
+ sum0 = sum0 + x * weight[wbase];
56
+ sum1 = sum1 + x * weight[wbase + 1u];
57
+ sum2 = sum2 + x * weight[wbase + 2u];
58
+ sum3 = sum3 + x * weight[wbase + 3u];
59
+ }
60
+
61
+ let output_base = (((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc0) / 4u;
62
+ let bias_base = oc0 / 4u;
63
+ output[output_base] = apply_relu4(sum0 + bias[bias_base]);
64
+ output[output_base + 1u] = apply_relu4(sum1 + bias[bias_base + 1u]);
65
+ output[output_base + 2u] = apply_relu4(sum2 + bias[bias_base + 2u]);
66
+ output[output_base + 3u] = apply_relu4(sum3 + bias[bias_base + 3u]);
67
+ }