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,378 @@
1
+ import { _pair } from './Tensor.js';
2
+ import { Graph } from './Graph.js';
3
+ import { CPUEngine } from './CPUEngine.js';
4
+
5
+
6
+ export class WasmEngine extends CPUEngine {
7
+ constructor(wasmModule) {
8
+ super();
9
+ this.wasmModule = wasmModule;
10
+ this.api = wasmModule.instance.exports;
11
+ this.mem = this.api.memory;
12
+ this.pointers = new Map();
13
+ console.log("[VolvoxAI] WASM Engine ready (Tier 2 Fallback).");
14
+ }
15
+ static async init(wasmUrl) {
16
+ try {
17
+ let buffer;
18
+ if (typeof process !== "undefined" && process.versions && process.versions.node) {
19
+ const fs = await import('fs');
20
+ let url = wasmUrl;
21
+ if (!url.startsWith('/')) url = './' + url;
22
+ buffer = fs.readFileSync(url);
23
+ } else {
24
+ const response = await fetch(wasmUrl);
25
+ if (!response.ok) throw new Error("WASM file not found.");
26
+ buffer = await response.arrayBuffer();
27
+ }
28
+ const env = { expf: Math.exp, logf: Math.log, powf: Math.pow };
29
+ const module = await WebAssembly.instantiate(buffer, { env, math: env });
30
+ return new WasmEngine(module);
31
+ } catch (e) {
32
+ console.warn(`[VolvoxAI] Failed to load WASM from ${wasmUrl}:`, e);
33
+ return null;
34
+ }
35
+ }
36
+ createGraph() {
37
+ return new Graph();
38
+ }
39
+ _alloc(tensor) {
40
+ if (!this.pointers.has(tensor.name)) {
41
+ const ptr = this.api.alloc_bytes(tensor.sizeBytes);
42
+ // Grow WASM memory if the allocation exceeds current buffer
43
+ const needed = ptr + tensor.sizeBytes;
44
+ const currentBytes = this.mem.buffer.byteLength;
45
+ if (needed > currentBytes) {
46
+ const pagesNeeded = Math.ceil((needed - currentBytes) / 65536);
47
+ this.mem.grow(pagesNeeded);
48
+ }
49
+ this.pointers.set(tensor.name, ptr);
50
+ // Create a Float32Array view directly into WASM memory for CPU fallbacks
51
+ const wasmView = new Float32Array(this.mem.buffer, ptr, tensor.sizeBytes / 4);
52
+ if (tensor.isWeight && tensor.buffer) {
53
+ let view = tensor.buffer;
54
+ // The C matmul_f32 reads weights as [d_in, d_out]. PyTorch-style [d_out, d_in]
55
+ // weights (node.wLayout==='dout') are transposed to [d_in, d_out] on the heap.
56
+ const dw = this._doutWeights && this._doutWeights.get(tensor.name);
57
+ if (dw) {
58
+ const { din, dout } = dw;
59
+ const t = new Float32Array(din * dout);
60
+ for (let j = 0; j < dout; j++) for (let k = 0; k < din; k++) t[k * dout + j] = view[j * din + k];
61
+ view = t;
62
+ }
63
+ const src = new Uint8Array(view.buffer, view.byteOffset, view.byteLength);
64
+ const bytesToCopy = Math.min(view.byteLength, tensor.sizeBytes);
65
+ new Uint8Array(this.mem.buffer, ptr, bytesToCopy).set(src.subarray(0, bytesToCopy));
66
+ }
67
+ tensor.buffer = wasmView;
68
+ }
69
+ return this.pointers.get(tensor.name);
70
+ }
71
+ allocateGraph(graph) {
72
+ return this.compile(graph);
73
+ }
74
+ compile(graph) {
75
+ console.log("[VolvoxAI WASM] Allocating graph tensors on WASM heap...");
76
+ if (this.api.reset_heap) this.api.reset_heap();
77
+ this.pointers.clear();
78
+ // Weights consumed by a [d_out, d_in]-layout MatMul are transposed to the C
79
+ // kernel's [d_in, d_out] layout during _alloc (keyed by weight name).
80
+ this._doutWeights = new Map();
81
+ for (const n of graph.nodes) {
82
+ if ((n.opType === "MatMul" || n.opType === "Linear" || n.opType === "Gemm") && n.wLayout === "dout" && n.inputs.weight && !n.inputs.scale) {
83
+ const din = n.inputs.input.shape[n.inputs.input.shape.length - 1];
84
+ const dout = n.outputs.out.shape[n.outputs.out.shape.length - 1];
85
+ this._doutWeights.set(n.inputs.weight.name, { din, dout });
86
+ }
87
+ }
88
+ for (const [name, tensor] of graph.tensors.entries()) {
89
+ this._alloc(tensor);
90
+ }
91
+ // Re-create views because mem.grow detaches old views
92
+ for (const [name, tensor] of graph.tensors.entries()) {
93
+ const ptr = this.pointers.get(name);
94
+ tensor.buffer = new Float32Array(this.mem.buffer, ptr, tensor.sizeBytes / 4);
95
+ }
96
+ this.graph = graph;
97
+ return this;
98
+ }
99
+ async execute(inputs) {
100
+ for (const [name, data] of Object.entries(inputs)) {
101
+ const tensor = this.graph.tensors.get(name);
102
+ if (!tensor) continue;
103
+ const ptr = this.pointers.get(name);
104
+ new Float32Array(this.mem.buffer, ptr, data.length).set(data);
105
+ }
106
+
107
+ for (const node of this.graph.nodes) {
108
+ try {
109
+ const inPtr = this.pointers.get(node.inputs.input?.name);
110
+ const outPtr = this.pointers.get(node.outputs.out?.name);
111
+ const wPtr = this.pointers.get(node.inputs.weight?.name);
112
+ const bPtr = this.pointers.get(node.inputs.bias?.name);
113
+
114
+ if (node.opType === "Conv2D") {
115
+ const wPtr = this.pointers.get(node.inputs.weight.name);
116
+ const bPtr = node.inputs.bias ? this.pointers.get(node.inputs.bias.name) : 0;
117
+ const inShape = node.inputs.input.shape;
118
+ const outShape = node.outputs.out.shape;
119
+ const wShape = node.inputs.weight.shape;
120
+ const [sy, sx] = _pair(node.params.stride, 1);
121
+ const [dy, dx] = _pair(node.params.dilation, 1);
122
+ const padPair = _pair(node.params.padding, 0);
123
+ const pads = node.params.pads || [padPair[0], padPair[1], padPair[0], padPair[1]];
124
+ const groups = node.params.groups || 1;
125
+ const relu = node.params.relu ? 1 : 0;
126
+ this.api.conv2d_f32(
127
+ inPtr, outPtr, wPtr, bPtr,
128
+ inShape[0], inShape[1], inShape[2], inShape[3],
129
+ wShape[0], wShape[1], wShape[2], wShape[3],
130
+ outShape[1], outShape[2], sy, sx, pads[0], pads[1],
131
+ groups, relu, dy, dx
132
+ );
133
+ } else if (node.opType === "ConvTranspose2D") {
134
+ const wPtr = this.pointers.get(node.inputs.weight.name);
135
+ const bPtr = node.inputs.bias ? this.pointers.get(node.inputs.bias.name) : 0;
136
+ const [b, in_h, in_w, in_c] = node.inputs.input.shape;
137
+ const [out_b, out_h, out_w, out_c] = node.outputs.out.shape;
138
+ const kh = node.params.kernel[0], kw = node.params.kernel[1];
139
+ const sh = node.params.stride ? node.params.stride[0] : 1;
140
+ const sw = node.params.stride ? node.params.stride[1] : 1;
141
+ const ph = node.params.padding ? node.params.padding[0] : 0;
142
+ const pw = node.params.padding ? node.params.padding[1] : 0;
143
+ this.api.conv_transpose2d_f32(
144
+ inPtr, wPtr, bPtr, outPtr,
145
+ b, in_h, in_w, in_c, out_h, out_w, out_c,
146
+ kh, kw, sh, sw, ph, pw
147
+ );
148
+ } else if (node.opType === "ReduceSum") {
149
+ const inShape = node.inputs.input ? node.inputs.input.shape : node.inputs.data.shape;
150
+ const shape = inShape.length === 2 ? inShape : [1, node.inputs.input ? node.inputs.input.buffer.length : node.inputs.data.buffer.length];
151
+ this.api.reduce_sum_f32(inPtr, outPtr, shape[0], shape[1]);
152
+ } else if (node.opType === "ReduceMean") {
153
+ const inShape = node.inputs.input ? node.inputs.input.shape : node.inputs.data.shape;
154
+ const shape = inShape.length === 2 ? inShape : [1, node.inputs.input ? node.inputs.input.buffer.length : node.inputs.data.buffer.length];
155
+ this.api.reduce_mean_f32(inPtr, outPtr, shape[0], shape[1]);
156
+ } else if (node.opType === "MatMul") {
157
+ const wPtr = this.pointers.get(node.inputs.weight.name);
158
+ const bPtr = node.inputs.bias ? this.pointers.get(node.inputs.bias.name) : 0;
159
+ const d_in = node.inputs.input.shape[node.inputs.input.shape.length-1];
160
+ const d_out = node.outputs.out.shape[node.outputs.out.shape.length-1];
161
+ const flatSeq = node.inputs.input.shape.slice(0, -1).reduce((a,b)=>a*b,1);
162
+ if (node.inputs.scale) {
163
+ const sPtr = this.pointers.get(node.inputs.scale.name);
164
+ this.api.matmul_int8_f32(inPtr, wPtr, sPtr, bPtr, outPtr, flatSeq, d_in, d_out);
165
+ } else {
166
+ this.api.matmul_f32(inPtr, wPtr, bPtr, outPtr, flatSeq, d_in, d_out);
167
+ }
168
+ } else if (node.opType === "LayerNorm") {
169
+ const flatSeq = node.inputs.input.shape.slice(0, -1).reduce((a,b)=>a*b,1);
170
+ const d_model = node.params.d_model;
171
+ this.api.layernorm_f32(inPtr, wPtr, bPtr, outPtr, flatSeq, d_model, 1e-5);
172
+ } else if (node.opType === "SDPA") {
173
+ const qkvPtr = this.pointers.get(node.inputs.qkv.name);
174
+ const seqLen = node.inputs.qkv.shape[1]; // [batch, seq, d]
175
+ const d_model = node.outputs.out.shape[node.outputs.out.shape.length - 1];
176
+ const heads = node.params.heads || node.params.num_heads || 8;
177
+ const head_dim = node.params.head_dim || d_model / heads;
178
+ const scale = node.params.scale !== undefined ? node.params.scale : 1.0 / Math.sqrt(head_dim);
179
+ this.api.sdpa_f32(qkvPtr, outPtr, seqLen, d_model, heads, head_dim, scale);
180
+ } else if (node.opType === "CrossSDPA") {
181
+ const qPtr = this.pointers.get(node.inputs.q.name);
182
+ const kPtr = this.pointers.get(node.inputs.k.name);
183
+ const vPtr = this.pointers.get(node.inputs.v.name);
184
+ const seqQ = node.inputs.q.shape[1];
185
+ const seqKV = node.inputs.k.shape[1];
186
+ const d_model = node.outputs.out.shape[node.outputs.out.shape.length - 1];
187
+ const heads = node.params.heads || node.params.num_heads || 8;
188
+ const head_dim = node.params.head_dim || d_model / heads;
189
+ this.api.cross_sdpa_f32(qPtr, kPtr, vPtr, outPtr, seqQ, seqKV, d_model, heads, head_dim, 1.0 / Math.sqrt(head_dim));
190
+ } else if (node.opType === "Embedding") {
191
+ const seqLen = node.inputs.input.shape.reduce((a, b) => a * b, 1);
192
+ const d_model = node.outputs.out.shape[node.outputs.out.shape.length - 1];
193
+ this.api.embedding_f32(inPtr, wPtr, outPtr, seqLen, d_model);
194
+ } else if (node.opType === "ReLU") {
195
+ const inShape = node.inputs.input.shape;
196
+ const elements = inShape.reduce((a, b) => a * b, 1);
197
+ this.api.relu_f32(inPtr, outPtr, elements);
198
+ } else if (node.opType === "GELU") {
199
+ const elements = node.inputs.input.sizeBytes / 4;
200
+ this.api.gelu_f32(inPtr, outPtr, elements);
201
+ } else if (node.opType === "Add") {
202
+ this._cpuAdd(node);
203
+ } else if (node.opType === "Mul") {
204
+ this._cpuMul(node);
205
+ } else if (node.opType === "Conv1D") {
206
+ const inShape = node.inputs.input.shape;
207
+ const wShape = node.inputs.weight.shape;
208
+ const [st] = _pair(node.params.stride, 1);
209
+ const [pd] = _pair(node.params.padding, 0);
210
+ const groups = node.params.groups || 1;
211
+ const relu = node.params.relu ? 1 : 0;
212
+ this.api.conv1d_f32(
213
+ inPtr, outPtr, wPtr, bPtr,
214
+ inShape[1], inShape[2], wShape[0], wShape[1], wShape[2], st, pd, groups, relu
215
+ );
216
+ } else if (node.opType === "UpsampleNearest2D") {
217
+ const inShape = node.inputs.input.shape;
218
+ this.api.upsample_nearest2x_f32(inPtr, outPtr, inShape[3], inShape[1], inShape[2]);
219
+ } else if (node.opType === "Concat" || node.opType === "Concat2") {
220
+ // N-way channel concat via the heap-view buffers (like Add/Mul above);
221
+ // the C concat2_f32 only handled two inputs.
222
+ this._cpuConcat2(node);
223
+ } else if (node.opType === "ProfileY") {
224
+ const inShape = node.inputs.input.shape;
225
+ this.api.profile_y_f32(inPtr, outPtr, inShape[3], inShape[1], inShape[2]);
226
+ } else if (node.opType === "ProfileX") {
227
+ const inShape = node.inputs.input.shape;
228
+ this.api.profile_x_f32(inPtr, outPtr, inShape[3], inShape[1], inShape[2]);
229
+ } else if (node.opType === "InterpLinear1D") {
230
+ const inShape = node.inputs.input.shape; // [1, C, L]
231
+ const outL = node.params.size;
232
+ this.api.interp1d_f32(inPtr, outPtr, inShape[1], inShape[2], outL);
233
+ } else if (node.opType === "SpatialSoftargmaxY") {
234
+ const inShape = node.inputs.input.shape;
235
+ this.api.spatial_softargmax_y_f32(inPtr, outPtr, inShape[3], inShape[1], inShape[2]);
236
+ } else if (node.opType === "Sigmoid") {
237
+ const inS = node.inputs.input.shape;
238
+ const elements = inS.reduce((a,b)=>a*b, 1);
239
+ this.api.sigmoid_f32(inPtr, outPtr, elements);
240
+ } else if (node.opType === "Clip") {
241
+ let minVal = node.params.min !== undefined ? node.params.min : -1e9;
242
+ let maxVal = node.params.max !== undefined ? node.params.max : 1e9;
243
+ if (node.inputs.min) {
244
+ const p = this.pointers.get(node.inputs.min.name);
245
+ minVal = new Float32Array(this.mem.buffer, p, 1)[0];
246
+ }
247
+ if (node.inputs.max) {
248
+ const p = this.pointers.get(node.inputs.max.name);
249
+ maxVal = new Float32Array(this.mem.buffer, p, 1)[0];
250
+ }
251
+ const inS = node.inputs.input.shape;
252
+ const elements = inS.reduce((a,b)=>a*b, 1);
253
+ this.api.clip_f32(inPtr, outPtr, elements, minVal, maxVal);
254
+ } else if (node.opType === "HardSwish") {
255
+ const inS = node.inputs.input.shape;
256
+ const elements = inS.reduce((a,b)=>a*b, 1);
257
+ this.api.hardswish_f32(inPtr, outPtr, elements);
258
+ } else if (node.opType === "LeakyReLU") {
259
+ const inS = node.inputs.input.shape;
260
+ const elements = inS.reduce((a,b)=>a*b, 1);
261
+ const alpha = node.params.alpha || 0.01;
262
+ if (this.api.leakyrelu_f32) this.api.leakyrelu_f32(inPtr, outPtr, elements, alpha);
263
+ else this._cpuLeakyReLU(node);
264
+ } else if (node.opType === "PReLU") {
265
+ if (this.api.prelu_f32) {
266
+ const inS = node.inputs.input.shape;
267
+ const wPtr = this.pointers.get(node.inputs.weight.name);
268
+ this.api.prelu_f32(inPtr, wPtr, outPtr, inS[0], inS[1], inS[2], inS[3]);
269
+ } else {
270
+ this._cpuPReLU(node);
271
+ }
272
+ } else if (node.opType === "HardSigmoid") {
273
+ const inS = node.inputs.input.shape;
274
+ const elements = inS.reduce((a,b)=>a*b, 1);
275
+ this.api.hardsigmoid_f32(inPtr, outPtr, elements);
276
+ } else if (node.opType === "Reshape") {
277
+ const inS = node.inputs.input.shape;
278
+ const elements = inS.reduce((a,b)=>a*b, 1);
279
+ this.api.copy_f32(inPtr, outPtr, elements);
280
+ } else if (node.opType === "Transpose") {
281
+ this._cpuTranspose(node);
282
+ } else if (node.opType === "GlobalAveragePool") {
283
+ const inShape = node.inputs.input.shape;
284
+ this.api.global_average_pool_f32(inPtr, outPtr, inShape[0], inShape[1], inShape[2], inShape[3]);
285
+ } else if (node.opType === "BatchNorm2D") {
286
+ const wPtr = this.pointers.get(node.inputs.weight.name);
287
+ const bPtr = this.pointers.get(node.inputs.bias.name);
288
+ const rmPtr = this.pointers.get(node.inputs.running_mean.name);
289
+ const rvPtr = this.pointers.get(node.inputs.running_var.name);
290
+ const [b, h, w, c] = node.inputs.input.shape;
291
+ const eps = node.params.eps || 1e-5;
292
+ this.api.batch_norm2d_f32(inPtr, wPtr, bPtr, rmPtr, rvPtr, outPtr, b, h, w, c, eps);
293
+ } else if (node.opType === "ResizeNearest2D") {
294
+ this._cpuResize(node);
295
+ } else if (node.opType === "Resize") {
296
+ const [b, in_h, in_w, c] = node.inputs.input.shape;
297
+ const [ob, out_h, out_w] = node.outputs.out.shape;
298
+ this.api.resize_bilinear_f32(inPtr, outPtr, b, in_h, in_w, c, out_h, out_w);
299
+ } else if (node.opType === "Cast") {
300
+ this._cpuCast(node); // Keep cast in JS due to Float32Array aliasing
301
+ } else if (node.opType === "Slice") {
302
+ const in_s = [1,1,1,1].slice(0, 4-node.inputs.input.shape.length).concat(node.inputs.input.shape);
303
+ const out_s = [1,1,1,1].slice(0, 4-node.outputs.out.shape.length).concat(node.outputs.out.shape);
304
+ const starts = node.params.starts || [0,0,0,0]; const steps = node.params.steps || [1,1,1,1]; const axes = node.params.axes || [0,1,2,3];
305
+ const st = [0,0,0,0]; const sp = [1,1,1,1];
306
+ for(let i=0; i<axes.length; i++) {
307
+ let ax = axes[i]; if(ax < 0) ax += node.inputs.input.shape.length; ax += (4 - node.inputs.input.shape.length);
308
+ st[ax] = starts[i] < 0 ? starts[i] + in_s[ax] : starts[i]; sp[ax] = steps[i];
309
+ }
310
+ this.api.slice_4d_f32(inPtr, outPtr, in_s[0], in_s[1], in_s[2], in_s[3], out_s[0], out_s[1], out_s[2], out_s[3], st[0], st[1], st[2], st[3], sp[0], sp[1], sp[2], sp[3]);
311
+ } else if (node.opType === "Split") {
312
+ this._cpuSplit(node); // Delegate multi-output loop to CPU
313
+ } else if (node.opType === "Gather") {
314
+ const in_s = [1,1,1,1].slice(0, 4-node.inputs.input.shape.length).concat(node.inputs.input.shape);
315
+ let axis = node.params.axis || 0; if (axis < 0) axis += node.inputs.input.shape.length; axis += (4 - node.inputs.input.shape.length);
316
+ const idxPtr = this.pointers.get(node.inputs.indices.name);
317
+ this.api.gather_4d_f32(inPtr, idxPtr, outPtr, in_s[0], in_s[1], in_s[2], in_s[3], node.inputs.indices.sizeBytes/4, axis);
318
+ } else if (node.opType === "GatherElements") {
319
+ this._cpuGatherElements(node); // Complex enough to fallback to CPU JS loop
320
+ } else if (node.opType === "NonMaxSuppression") {
321
+ this._cpuNonMaxSuppression(node);
322
+ } else if (node.opType === "Where" || node.opType === "Mask") {
323
+ const condPtr = this.pointers.get(node.inputs.cond ? node.inputs.cond.name : node.inputs.condition.name);
324
+ const aPtr = this.pointers.get(node.inputs.x ? node.inputs.x.name : node.inputs.a.name);
325
+ const bPtr = this.pointers.get(node.inputs.y ? node.inputs.y.name : node.inputs.b.name);
326
+ this.api.where_f32(condPtr, aPtr, bPtr, outPtr, node.outputs.out.sizeBytes / 4);
327
+ } else if (node.opType === "Pad") {
328
+ const pads = node.params.pads;
329
+ const pt = pads.length === 8 ? pads[1] : pads[0];
330
+ const pl = pads.length === 8 ? pads[2] : pads[1];
331
+ const pb = pads.length === 8 ? pads[5] : pads[2];
332
+ const pr = pads.length === 8 ? pads[6] : pads[3];
333
+ const val = node.params.value || 0.0;
334
+ const inShape = node.inputs.input ? node.inputs.input.shape : node.inputs.data.shape;
335
+ const s = inShape.length === 4 ? inShape : [1, inShape[0] || 1, inShape[1] || 1, 1];
336
+ this.api.pad_2d_f32(inPtr, outPtr, val, s[0], s[1], s[2], s[3], pt, pb, pl, pr);
337
+ } else if (node.opType === "AveragePool2D" || node.opType === "AveragePool") {
338
+ const [b, in_h, in_w, c] = (node.inputs.input || node.inputs.x).shape;
339
+ const [ob, out_h, out_w] = node.outputs.out.shape;
340
+ const ky = node.params.kernel[0], kx = node.params.kernel[1];
341
+ const sy = node.params.stride ? node.params.stride[0] : 1;
342
+ const sx = node.params.stride ? node.params.stride[1] : 1;
343
+ const py = node.params.padding ? node.params.padding[0] : 0;
344
+ const px = node.params.padding ? node.params.padding[1] : 0;
345
+ this.api.averagepool2d_f32(inPtr, outPtr, b, in_h, in_w, c, ky, kx, sy, sx, py, px, out_h, out_w);
346
+ } else if (node.opType === "Div") {
347
+ this._cpuDiv(node);
348
+ } else if (node.opType === "MaxPool2D") {
349
+ const inShape = node.inputs.input.shape;
350
+ const outShape = node.outputs.out.shape;
351
+ const ky = node.params.kernel[0], kx = node.params.kernel[1];
352
+ const sy = node.params.stride[0], sx = node.params.stride[1];
353
+ const py = node.params.padding ? node.params.padding[0] : 0;
354
+ const px = node.params.padding ? node.params.padding[1] : 0;
355
+ this.api.maxpool2d_f32(inPtr, outPtr, inShape[1], inShape[2], inShape[3], outShape[1], outShape[2], ky, kx, sy, sx, py, px);
356
+ } else if (node.opType === "MeanHeight") {
357
+ const inShape = node.inputs.input.shape;
358
+ this.api.mean_height_f32(inPtr, outPtr, inShape[3], inShape[1], inShape[2]);
359
+ } else if (node.opType === "Flatten" || node.opType === "Squeeze" || node.opType === "Unsqueeze" || node.opType === "Dropout" || node.opType === "Reshape" || node.opType === "Identity") {
360
+ this._cpuReshape(node);
361
+ } else {
362
+ super._runNode(node);
363
+ }
364
+ } catch (e) {
365
+ console.error("[WasmEngine] Execution failed at node:", node, e);
366
+ throw e;
367
+ }
368
+ }
369
+
370
+ const results = {};
371
+ for (const name of this.graph.outputNames) {
372
+ const tensor = this.graph.tensors.get(name);
373
+ const ptr = this.pointers.get(name);
374
+ results[name] = new Float32Array(this.mem.buffer, ptr, tensor.sizeBytes / 4).slice();
375
+ }
376
+ return results;
377
+ }
378
+ };
@@ -0,0 +1,169 @@
1
+ // Experimental WebNN (navigator.ml) backend. Maps the Volvox graph onto MLGraphBuilder
2
+ // ops and runs on the platform NPU/GPU. It covers feed-forward/vision graphs; ops it
3
+ // can't express throw during build so VolvoxAI cleanly falls back to a lower
4
+ // tier. Note: graph.tensors is a Map and node.inputs/outputs hold Tensor OBJECTS.
5
+ export class WebNNEngine {
6
+ constructor(context) {
7
+ this.context = context;
8
+ this.graph = null;
9
+ this.compiledGraph = null;
10
+ this.operands = {};
11
+ this.inputs = [];
12
+ this.outputs = [];
13
+ }
14
+
15
+ async allocateGraph(graph) {
16
+ this.graph = graph;
17
+ this.operands = {}; this.inputs = []; this.outputs = [];
18
+ const builder = new MLGraphBuilder(this.context);
19
+ // MLOperandDescriptor changed across WebNN versions (dimensions→shape, added dataType).
20
+ // Provide all spellings so constants/inputs build on old and new implementations.
21
+ const desc = (shape) => { const d = shape.length ? shape : [1]; return { dataType: 'float32', type: 'float32', shape: d, dimensions: d }; };
22
+
23
+ // Weights become constants; anything else not produced by a node is a graph input.
24
+ const generated = new Set();
25
+ for (const node of graph.nodes)
26
+ for (const t of Object.values(node.outputs)) if (t && t.name) generated.add(t.name);
27
+
28
+ for (const t of graph.tensors.values()) {
29
+ if (t.isWeight && t.buffer) {
30
+ this.operands[t.name] = builder.constant(desc(t.shape), t.buffer);
31
+ generated.add(t.name);
32
+ } else if (!generated.has(t.name)) {
33
+ this.inputs.push(t.name);
34
+ this.operands[t.name] = builder.input(t.name, desc(t.shape));
35
+ generated.add(t.name);
36
+ }
37
+ }
38
+
39
+ const getOp = (name) => {
40
+ if (!name || !this.operands[name]) {
41
+ console.warn(`[WebNN] missing operand: ${name}`);
42
+ this.operands[name] = builder.constant(desc([1]), new Float32Array([0]));
43
+ }
44
+ return this.operands[name];
45
+ };
46
+ const nin = (node, key) => (node.inputs[key] ? node.inputs[key].name : null);
47
+
48
+ for (const node of graph.nodes) {
49
+ const op = node.opType;
50
+ const outName = Object.values(node.outputs)[0].name;
51
+ try {
52
+ if (op === "MatMul" || op === "Linear" || op === "Gemm") {
53
+ const a = getOp(nin(node, "input") || nin(node, "a"));
54
+ const w = getOp(nin(node, "weight") || nin(node, "b"));
55
+ // node.wLayout: 'dout' = [d_out, d_in] (needs Wᵀ), 'din' = [d_in, d_out] (x·W).
56
+ let res = node.wLayout === "dout" ? builder.gemm(a, w, { bTranspose: true }) : builder.matmul(a, w);
57
+ if (nin(node, "bias")) res = builder.add(res, getOp(nin(node, "bias")));
58
+ this.operands[outName] = res;
59
+ } else if (op === "Add") {
60
+ this.operands[outName] = builder.add(getOp(nin(node, "a") || nin(node, "input")), getOp(nin(node, "b")));
61
+ } else if (op === "Mul") {
62
+ this.operands[outName] = builder.mul(getOp(nin(node, "a") || nin(node, "input")), getOp(nin(node, "b")));
63
+ } else if (op === "ReLU") {
64
+ this.operands[outName] = builder.relu(getOp(nin(node, "input")));
65
+ } else if (op === "GELU") {
66
+ this.operands[outName] = builder.gelu(getOp(nin(node, "input")));
67
+ } else if (op === "SiLU" || op === "Swish") {
68
+ const x = getOp(nin(node, "input"));
69
+ this.operands[outName] = builder.mul(x, builder.sigmoid(x));
70
+ } else if (op === "Sigmoid") {
71
+ this.operands[outName] = builder.sigmoid(getOp(nin(node, "input")));
72
+ } else if (op === "Softmax") {
73
+ this.operands[outName] = builder.softmax(getOp(nin(node, "input")));
74
+ } else if (op === "Reshape" || op === "Flatten") {
75
+ const shape = Object.values(node.outputs)[0].shape;
76
+ this.operands[outName] = builder.reshape(getOp(nin(node, "input")), shape.length ? shape : [1]);
77
+ } else if (op === "LayerNorm") {
78
+ const input = getOp(nin(node, "input"));
79
+ const scale = nin(node, "weight") ? getOp(nin(node, "weight")) : undefined;
80
+ const bias = nin(node, "bias") ? getOp(nin(node, "bias")) : undefined;
81
+ const inShape = node.inputs.input.shape;
82
+ this.operands[outName] = builder.layerNormalization(input, { axes: [inShape.length - 1], scale, bias });
83
+ } else if (op === "Conv2D") {
84
+ this.operands[outName] = builder.conv2d(getOp(nin(node, "input")), getOp(nin(node, "weight")),
85
+ { bias: nin(node, "bias") ? getOp(nin(node, "bias")) : undefined,
86
+ strides: node.params.stride, padding: node.params.padding, groups: node.params.groups || 1 });
87
+ } else if (op === "Embedding") {
88
+ // Row gather: table[vocab, d_model] indexed by token ids. WebNN gather needs
89
+ // integer indices; the ids arrive as f32, so cast first.
90
+ const idx = builder.cast(getOp(nin(node, "input")), "int32");
91
+ this.operands[outName] = builder.gather(getOp(nin(node, "weight")), idx, { axis: 0 });
92
+ } else if (op === "SDPA") {
93
+ // SDPA is not a hardware/API primitive — decompose it into ops WebNN has:
94
+ // softmax( (Q·Kᵀ)·scale + causal_mask ) · V, batched over heads.
95
+ const qkvT = node.inputs.qkv;
96
+ const seq = qkvT.shape[1];
97
+ const d = Math.floor(qkvT.shape[2] / 3); // d_model
98
+ const heads = node.params.heads || 8;
99
+ const hd = Math.floor(d / heads); // head_dim
100
+ const scale = node.params.scale !== undefined ? node.params.scale : 1 / Math.sqrt(hd);
101
+ const qkv = builder.reshape(getOp(nin(node, "qkv")), [seq, 3 * d]);
102
+ const Q = builder.slice(qkv, [0, 0], [seq, d]);
103
+ const K = builder.slice(qkv, [0, d], [seq, d]);
104
+ const V = builder.slice(qkv, [0, 2 * d], [seq, d]);
105
+ const toHeads = (x) => builder.transpose(builder.reshape(x, [seq, heads, hd]), { permutation: [1, 0, 2] }); // [heads, seq, hd]
106
+ const Qh = toHeads(Q), Vh = toHeads(V);
107
+ const Kh = builder.transpose(builder.reshape(K, [seq, heads, hd]), { permutation: [1, 2, 0] }); // [heads, hd, seq]
108
+ let scores = builder.matmul(Qh, Kh); // [heads, seq, seq]
109
+ scores = builder.mul(scores, builder.constant(desc([1]), new Float32Array([scale])));
110
+ const mask = new Float32Array(seq * seq); // causal: 0 for j<=i, -inf above
111
+ for (let i = 0; i < seq; i++) for (let j = 0; j < seq; j++) mask[i * seq + j] = j <= i ? 0 : -1e9;
112
+ scores = builder.add(scores, builder.constant(desc([seq, seq]), mask)); // broadcast over heads
113
+ const attn = builder.softmax(scores, 2); // over the last axis (keys)
114
+ let out = builder.matmul(attn, Vh); // [heads, seq, hd]
115
+ out = builder.transpose(out, { permutation: [1, 0, 2] }); // [seq, heads, hd]
116
+ this.operands[outName] = builder.reshape(out, [1, seq, d]);
117
+ } else {
118
+ // Ops WebNN can't express here or that have not been wired yet.
119
+ // Throw so VolvoxAI falls back to a tier that implements them.
120
+ throw new Error(`Unsupported op in WebNNEngine: ${op}`);
121
+ }
122
+ } catch (e) {
123
+ console.warn(`[WebNN] cannot map op ${op}: ${e.message}`);
124
+ throw e; // bubble up → VolvoxAI tries the next tier
125
+ }
126
+ }
127
+
128
+ const outputOperands = {};
129
+ for (const name of graph.outputNames) {
130
+ if (this.operands[name]) { outputOperands[name] = this.operands[name]; this.outputs.push(name); }
131
+ }
132
+ console.log(`[WebNN] building graph, outputs:`, this.outputs);
133
+ this.compiledGraph = await builder.build(outputOperands);
134
+ console.log(`[WebNN] graph compiled.`);
135
+ return this;
136
+ }
137
+
138
+ async execute(inputsMap) {
139
+ if (!this.compiledGraph) throw new Error("WebNN graph not compiled.");
140
+ const ctx = this.context;
141
+ const numel = (name) => { const t = this.graph.tensors.get(name); return t ? t.shape.reduce((a, b) => a * b, 1) : 1; };
142
+ const shapeOf = (name) => { const t = this.graph.tensors.get(name); return t && t.shape.length ? t.shape : [1]; };
143
+
144
+ // Legacy API: MLContext.compute(graph, inputs, outputs) with plain ArrayBufferViews.
145
+ if (typeof ctx.compute === "function") {
146
+ const ins = {}, outs = {};
147
+ for (const name of this.inputs) ins[name] = inputsMap[name] || new Float32Array(numel(name));
148
+ for (const name of this.outputs) outs[name] = new Float32Array(numel(name));
149
+ return (await ctx.compute(this.compiledGraph, ins, outs)).outputs;
150
+ }
151
+
152
+ // Current API: MLTensor + dispatch + readTensor.
153
+ const inT = {}, outT = {};
154
+ for (const name of this.inputs) {
155
+ const t = await ctx.createTensor({ dataType: "float32", shape: shapeOf(name), dimensions: shapeOf(name), writable: true });
156
+ ctx.writeTensor(t, inputsMap[name] || new Float32Array(numel(name)));
157
+ inT[name] = t;
158
+ }
159
+ for (const name of this.outputs)
160
+ outT[name] = await ctx.createTensor({ dataType: "float32", shape: shapeOf(name), dimensions: shapeOf(name), readable: true });
161
+
162
+ ctx.dispatch(this.compiledGraph, inT, outT);
163
+
164
+ const results = {};
165
+ for (const name of this.outputs) { const ab = await ctx.readTensor(outT[name]); results[name] = new Float32Array(ab); outT[name].destroy?.(); }
166
+ for (const name of this.inputs) inT[name].destroy?.();
167
+ return results;
168
+ }
169
+ }
package/js/index.js ADDED
@@ -0,0 +1,11 @@
1
+ export * from './Tensor.js';
2
+ export * from './Graph.js';
3
+ export * from './CPUEngine.js';
4
+ export * from './WasmEngine.js';
5
+ // ShaderLibrary is intentionally NOT re-exported here: it statically imports
6
+ // every .wgsl file (resolved only by the bundler's text loader), which would
7
+ // break plain Node imports of the engine. GraphExecutor loads it lazily.
8
+ export * from './GraphExecutor.js';
9
+ export * from './GraphLoader.js';
10
+ export * from './VolvoxAI.js';
11
+ export * from './Tokenizer.js';
package/js/ops/add.js ADDED
@@ -0,0 +1,31 @@
1
+ export function _cpuAdd(node) {
2
+
3
+ const a = node.inputs.a;
4
+ const b = node.inputs.b;
5
+ const outBuf = node.outputs.out.buffer;
6
+ let aBuf = a.buffer;
7
+ let bBuf = b.buffer;
8
+ let aShape = a.shape;
9
+ let bShape = b.shape;
10
+
11
+ // Swap to ensure a is the larger tensor
12
+ if (aBuf.length < bBuf.length) {
13
+ let temp = aBuf; aBuf = bBuf; bBuf = temp;
14
+ let tempS = aShape; aShape = bShape; bShape = tempS;
15
+ }
16
+
17
+ const elements = aBuf.length;
18
+ if (bBuf.length === 1) {
19
+ for (let i = 0; i < elements; i++) outBuf[i] = aBuf[i] + bBuf[0];
20
+ } else if (bBuf.length === elements) {
21
+ for (let i = 0; i < elements; i++) outBuf[i] = aBuf[i] + bBuf[i];
22
+ } else if (aShape.length === 4 && bBuf.length === aShape[3]) {
23
+ const c = aShape[3];
24
+ for (let i = 0; i < elements; i++) outBuf[i] = aBuf[i] + bBuf[i % c];
25
+ } else if (aShape.length === 4 && bShape.length === 4 && bShape[0] === 1 && bShape[1] === 1 && bShape[2] === 1 && bShape[3] === aShape[3]) {
26
+ const c = aShape[3];
27
+ for (let i = 0; i < elements; i++) outBuf[i] = aBuf[i] + bBuf[i % c];
28
+ } else {
29
+ console.warn("[VolvoxAI CPU] Executing Add is not fully implemented yet for shapes", a.shape, b.shape);
30
+ }
31
+ }
@@ -0,0 +1,33 @@
1
+ export function _cpuArgMax(node) {
2
+
3
+ // ONNX ArgMax over a single axis. Output holds the winning indices (stored as
4
+ // float32 here); works for any rank and either keepdims setting because the
5
+ // output element count is outerBlock * innerBlock regardless.
6
+ const input = node.inputs.input || node.inputs.data;
7
+ const out = node.outputs.out;
8
+ const inBuf = input.buffer;
9
+ const outBuf = out.buffer;
10
+ const shape = input.shape;
11
+ let axis = node.params.axis !== undefined ? node.params.axis : 0;
12
+ if (axis < 0) axis += shape.length;
13
+ const axisSize = shape[axis];
14
+
15
+ let innerBlock = 1;
16
+ for (let i = axis + 1; i < shape.length; i++) innerBlock *= shape[i];
17
+ let outerBlock = 1;
18
+ for (let i = 0; i < axis; i++) outerBlock *= shape[i];
19
+
20
+ let o = 0;
21
+ for (let ob = 0; ob < outerBlock; ob++) {
22
+ for (let ib = 0; ib < innerBlock; ib++) {
23
+ const base = ob * axisSize * innerBlock + ib;
24
+ let best = inBuf[base];
25
+ let bestIdx = 0;
26
+ for (let a = 1; a < axisSize; a++) {
27
+ const v = inBuf[base + a * innerBlock];
28
+ if (v > best) { best = v; bestIdx = a; }
29
+ }
30
+ outBuf[o++] = bestIdx;
31
+ }
32
+ }
33
+ }
@@ -0,0 +1,38 @@
1
+ export function _cpuAveragePool2D(node) {
2
+ const input = node.inputs.input || node.inputs.x;
3
+ const inBuf = input.buffer;
4
+ const outBuf = node.outputs.out.buffer;
5
+
6
+ const [b, in_h, in_w, c] = input.shape;
7
+ const [out_b, out_h, out_w, out_c] = node.outputs.out.shape;
8
+
9
+ const kh = node.params.kernel[0], kw = node.params.kernel[1];
10
+ const sh = node.params.stride ? node.params.stride[0] : 1;
11
+ const sw = node.params.stride ? node.params.stride[1] : 1;
12
+ const ph = node.params.padding ? node.params.padding[0] : 0;
13
+ const pw = node.params.padding ? node.params.padding[1] : 0;
14
+
15
+ for (let batch = 0; batch < b; batch++) {
16
+ for (let y = 0; y < out_h; y++) {
17
+ for (let x = 0; x < out_w; x++) {
18
+ for (let chan = 0; chan < c; chan++) {
19
+ let sum = 0.0;
20
+ let count = 0;
21
+ for (let ky = 0; ky < kh; ky++) {
22
+ for (let kx = 0; kx < kw; kx++) {
23
+ const in_y = y * sh - ph + ky;
24
+ const in_x = x * sw - pw + kx;
25
+ if (in_y >= 0 && in_y < in_h && in_x >= 0 && in_x < in_w) {
26
+ const inIdx = ((batch * in_h + in_y) * in_w + in_x) * c + chan;
27
+ sum += inBuf[inIdx];
28
+ count++;
29
+ }
30
+ }
31
+ }
32
+ const outIdx = ((batch * out_h + y) * out_w + x) * out_c + chan;
33
+ outBuf[outIdx] = count > 0 ? sum / count : 0.0;
34
+ }
35
+ }
36
+ }
37
+ }
38
+ }