oidn-web 0.3.4 → 0.4.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.
- package/CHANGELOG.md +49 -0
- package/README.md +140 -4
- package/benchmarks/compare.mjs +651 -0
- package/benchmarks/leak.mjs +255 -0
- package/benchmarks/results/before-spatial.json +391 -0
- package/benchmarks/results/before-spatial.md +47 -0
- package/benchmarks/results/int8-scan.json +2007 -0
- package/benchmarks/results/int8-scan.md +160 -0
- package/benchmarks/results/int8-w8a8-scan.json +2007 -0
- package/benchmarks/results/int8-w8a8-scan.md +160 -0
- package/benchmarks/results/int8-weight-channel.json +1413 -0
- package/benchmarks/results/int8-weight-channel.md +118 -0
- package/benchmarks/results/int8-weight-only.json +1437 -0
- package/benchmarks/results/int8-weight-only.md +118 -0
- package/benchmarks/results/kernel-webnn-final.json +1115 -0
- package/benchmarks/results/kernel-webnn-final.md +104 -0
- package/benchmarks/results/latest-optimized.json +375 -0
- package/benchmarks/results/latest-optimized.md +47 -0
- package/benchmarks/results/latest.json +391 -0
- package/benchmarks/results/latest.md +47 -0
- package/benchmarks/results/profile-baseline.json +331 -0
- package/benchmarks/results/profile-baseline.md +12 -0
- package/benchmarks/results/profile-conv2x.json +331 -0
- package/benchmarks/results/profile-conv2x.md +12 -0
- package/benchmarks/results/profile-fast-init.json +385 -0
- package/benchmarks/results/profile-fast-init.md +47 -0
- package/benchmarks/results/profile-fp16-fma.json +369 -0
- package/benchmarks/results/profile-fp16-fma.md +47 -0
- package/benchmarks/results/profile-fp16-tiled-decoder.json +369 -0
- package/benchmarks/results/profile-fp16-tiled-decoder.md +47 -0
- package/benchmarks/results/profile-fp16-tiled-encoder.json +369 -0
- package/benchmarks/results/profile-fp16-tiled-encoder.md +47 -0
- package/benchmarks/results/profile-fp16-unfused-pool.json +385 -0
- package/benchmarks/results/profile-fp16-unfused-pool.md +47 -0
- package/benchmarks/results/profile-input-major.json +347 -0
- package/benchmarks/results/profile-input-major.md +12 -0
- package/benchmarks/results/profile-k16.json +347 -0
- package/benchmarks/results/profile-k16.md +12 -0
- package/benchmarks/results/profile-k4.json +347 -0
- package/benchmarks/results/profile-k4.md +12 -0
- package/benchmarks/results/profile-pool-reuse.json +331 -0
- package/benchmarks/results/profile-pool-reuse.md +12 -0
- package/benchmarks/results/profile-precompiled.json +385 -0
- package/benchmarks/results/profile-precompiled.md +47 -0
- package/benchmarks/results/profile-static-channels.json +385 -0
- package/benchmarks/results/profile-static-channels.md +47 -0
- package/benchmarks/results/profile-static-io.json +385 -0
- package/benchmarks/results/profile-static-io.md +47 -0
- package/benchmarks/results/profile-tiled-conv.json +331 -0
- package/benchmarks/results/profile-tiled-conv.md +12 -0
- package/benchmarks/results/profile-tiled-decoder.json +347 -0
- package/benchmarks/results/profile-tiled-decoder.md +12 -0
- package/benchmarks/results/profile-tiled-matmul.json +331 -0
- package/benchmarks/results/profile-tiled-matmul.md +12 -0
- package/benchmarks/results/profile-unfused-decoder.json +379 -0
- package/benchmarks/results/profile-unfused-decoder.md +12 -0
- package/benchmarks/results/profile-unfused-pool.json +347 -0
- package/benchmarks/results/profile-unfused-pool.md +12 -0
- package/benchmarks/results/spatial-auto.json +575 -0
- package/benchmarks/results/spatial-auto.md +61 -0
- package/benchmarks/results/subgroup-smoke.json +1094 -0
- package/benchmarks/results/subgroup-smoke.md +104 -0
- package/benchmarks/results/webnn-smoke.json +739 -0
- package/benchmarks/results/webnn-smoke.md +76 -0
- package/dist/oidn.js +4189 -22603
- package/dist/oidn.umd.cjs +784 -5807
- package/lib/UNet.d.ts +66 -15
- package/lib/UNet.js +162 -257
- package/lib/UNet.js.map +1 -1
- package/lib/WGPUComputePass.d.ts +1 -1
- package/lib/WGPUComputePass.js +6 -4
- package/lib/WGPUComputePass.js.map +1 -1
- package/lib/backend.d.ts +8 -4
- package/lib/backend.js +36 -44
- package/lib/backend.js.map +1 -1
- package/lib/graphOptimizer.d.ts +54 -0
- package/lib/graphOptimizer.js +216 -0
- package/lib/graphOptimizer.js.map +1 -0
- package/lib/main.d.ts +33 -10
- package/lib/main.js +5 -0
- package/lib/main.js.map +1 -1
- package/lib/modelSpec.d.ts +80 -0
- package/lib/modelSpec.js +270 -0
- package/lib/modelSpec.js.map +1 -0
- package/lib/nativeUNet.d.ts +67 -0
- package/lib/nativeUNet.js +1735 -0
- package/lib/nativeUNet.js.map +1 -0
- package/lib/process.js +38 -35
- package/lib/process.js.map +1 -1
- package/lib/resourceTracker.d.ts +26 -0
- package/lib/resourceTracker.js +65 -0
- package/lib/resourceTracker.js.map +1 -0
- package/lib/tileScheduler.d.ts +33 -0
- package/lib/tileScheduler.js +86 -0
- package/lib/tileScheduler.js.map +1 -0
- package/lib/webnnUNet.d.ts +52 -0
- package/lib/webnnUNet.js +535 -0
- package/lib/webnnUNet.js.map +1 -0
- package/package.json +9 -5
- package/scripts/inspect-model.mjs +64 -0
- package/src/UNet.ts +236 -339
- package/src/WGPUComputePass.ts +6 -4
- package/src/backend.ts +42 -55
- package/src/graphOptimizer.ts +301 -0
- package/src/main.ts +71 -11
- package/src/modelSpec.ts +414 -0
- package/src/nativeUNet.ts +2256 -0
- package/src/process.ts +38 -36
- package/src/resourceTracker.ts +94 -0
- package/src/tileScheduler.ts +138 -0
- package/src/webnnUNet.ts +812 -0
- package/tests/modelSpec.test.mjs +128 -0
- package/tests/resourceLifecycle.test.mjs +383 -0
- package/tests/tileScheduler.test.mjs +90 -0
- package/lib/helper.d.ts +0 -4
- package/lib/helper.js +0 -33
- package/lib/helper.js.map +0 -1
- package/lib/kernels.d.ts +0 -1
- package/lib/kernels.js +0 -26
- package/lib/kernels.js.map +0 -1
- package/src/helper.ts +0 -43
- package/src/kernels.ts +0 -31
package/src/modelSpec.ts
ADDED
|
@@ -0,0 +1,414 @@
|
|
|
1
|
+
import type { HostTensor } from './tza';
|
|
2
|
+
|
|
3
|
+
/**
|
|
4
|
+
* Versioned, runtime-independent description of an OIDN network.
|
|
5
|
+
*
|
|
6
|
+
* TZA contains named tensors but no executable graph. Keeping the graph in a
|
|
7
|
+
* small declarative descriptor makes model upgrades independent from the WGSL
|
|
8
|
+
* kernels: a new OIDN topology only needs a new descriptor and validation
|
|
9
|
+
* fixture unless it introduces a genuinely new operation.
|
|
10
|
+
*/
|
|
11
|
+
export interface UNetModelSpec {
|
|
12
|
+
schemaVersion: 1;
|
|
13
|
+
id: string;
|
|
14
|
+
family: 'oidn-unet-small' | 'oidn-unet-large' | (string & {});
|
|
15
|
+
input: string;
|
|
16
|
+
output: string;
|
|
17
|
+
receptiveField: number;
|
|
18
|
+
nodes: readonly ModelNodeSpec[];
|
|
19
|
+
/** Reject unrecognised tensors so an upstream topology change is explicit. */
|
|
20
|
+
allowAdditionalTensors?: boolean;
|
|
21
|
+
}
|
|
22
|
+
|
|
23
|
+
export type ModelActivation = 'identity' | 'relu';
|
|
24
|
+
|
|
25
|
+
export interface Conv2DNodeSpec {
|
|
26
|
+
op: 'conv2d';
|
|
27
|
+
id: string;
|
|
28
|
+
input: string;
|
|
29
|
+
weight: string;
|
|
30
|
+
bias: string;
|
|
31
|
+
activation: ModelActivation;
|
|
32
|
+
padding: 'same';
|
|
33
|
+
}
|
|
34
|
+
|
|
35
|
+
export interface MaxPool2DNodeSpec {
|
|
36
|
+
op: 'maxPool2d';
|
|
37
|
+
id: string;
|
|
38
|
+
input: string;
|
|
39
|
+
size: 2;
|
|
40
|
+
stride: 2;
|
|
41
|
+
padding: 'same';
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
export interface Upsample2DNodeSpec {
|
|
45
|
+
op: 'upsample2d';
|
|
46
|
+
id: string;
|
|
47
|
+
input: string;
|
|
48
|
+
scale: 2;
|
|
49
|
+
mode: 'nearest';
|
|
50
|
+
}
|
|
51
|
+
|
|
52
|
+
export interface ConcatNodeSpec {
|
|
53
|
+
op: 'concat';
|
|
54
|
+
id: string;
|
|
55
|
+
inputs: readonly string[];
|
|
56
|
+
axis: 'channels';
|
|
57
|
+
}
|
|
58
|
+
|
|
59
|
+
export type ModelNodeSpec =
|
|
60
|
+
| Conv2DNodeSpec
|
|
61
|
+
| MaxPool2DNodeSpec
|
|
62
|
+
| Upsample2DNodeSpec
|
|
63
|
+
| ConcatNodeSpec;
|
|
64
|
+
|
|
65
|
+
export interface ValidatedConvTensor {
|
|
66
|
+
weight: HostTensor;
|
|
67
|
+
bias: HostTensor;
|
|
68
|
+
inputChannels: number;
|
|
69
|
+
outputChannels: number;
|
|
70
|
+
kernelHeight: number;
|
|
71
|
+
kernelWidth: number;
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
export interface ModelConvChannels {
|
|
75
|
+
inputChannels: number;
|
|
76
|
+
outputChannels: number;
|
|
77
|
+
}
|
|
78
|
+
|
|
79
|
+
export interface UNetModelGraph {
|
|
80
|
+
spec: UNetModelSpec;
|
|
81
|
+
inputChannels: number;
|
|
82
|
+
outputChannels: number;
|
|
83
|
+
channelsByValue: ReadonlyMap<string, number>;
|
|
84
|
+
convChannels: ReadonlyMap<string, ModelConvChannels>;
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
export interface ValidatedUNetModel extends UNetModelGraph {
|
|
88
|
+
tensorDataType: HostTensor['desc']['dataType'];
|
|
89
|
+
convTensors: ReadonlyMap<string, ValidatedConvTensor>;
|
|
90
|
+
}
|
|
91
|
+
|
|
92
|
+
function conv(id: string, input: string): Conv2DNodeSpec {
|
|
93
|
+
return {
|
|
94
|
+
op: 'conv2d',
|
|
95
|
+
id,
|
|
96
|
+
input,
|
|
97
|
+
weight: `${id}.weight`,
|
|
98
|
+
bias: `${id}.bias`,
|
|
99
|
+
activation: 'relu',
|
|
100
|
+
padding: 'same'
|
|
101
|
+
};
|
|
102
|
+
}
|
|
103
|
+
|
|
104
|
+
function pool(id: string, input: string): MaxPool2DNodeSpec {
|
|
105
|
+
return {
|
|
106
|
+
op: 'maxPool2d',
|
|
107
|
+
id,
|
|
108
|
+
input,
|
|
109
|
+
size: 2,
|
|
110
|
+
stride: 2,
|
|
111
|
+
padding: 'same'
|
|
112
|
+
};
|
|
113
|
+
}
|
|
114
|
+
|
|
115
|
+
function upsample(id: string, input: string): Upsample2DNodeSpec {
|
|
116
|
+
return {
|
|
117
|
+
op: 'upsample2d',
|
|
118
|
+
id,
|
|
119
|
+
input,
|
|
120
|
+
scale: 2,
|
|
121
|
+
mode: 'nearest'
|
|
122
|
+
};
|
|
123
|
+
}
|
|
124
|
+
|
|
125
|
+
function concat(
|
|
126
|
+
id: string,
|
|
127
|
+
first: string,
|
|
128
|
+
second: string
|
|
129
|
+
): ConcatNodeSpec {
|
|
130
|
+
return {
|
|
131
|
+
op: 'concat',
|
|
132
|
+
id,
|
|
133
|
+
inputs: [first, second],
|
|
134
|
+
axis: 'channels'
|
|
135
|
+
};
|
|
136
|
+
}
|
|
137
|
+
|
|
138
|
+
export const OIDN_UNET_SMALL_SPEC: UNetModelSpec = {
|
|
139
|
+
schemaVersion: 1,
|
|
140
|
+
id: 'oidn-unet-small-v1',
|
|
141
|
+
family: 'oidn-unet-small',
|
|
142
|
+
input: 'input',
|
|
143
|
+
output: 'dec_conv0',
|
|
144
|
+
receptiveField: 174,
|
|
145
|
+
nodes: [
|
|
146
|
+
conv('enc_conv0', 'input'),
|
|
147
|
+
conv('enc_conv1', 'enc_conv0'),
|
|
148
|
+
pool('pool1', 'enc_conv1'),
|
|
149
|
+
conv('enc_conv2', 'pool1'),
|
|
150
|
+
pool('pool2', 'enc_conv2'),
|
|
151
|
+
conv('enc_conv3', 'pool2'),
|
|
152
|
+
pool('pool3', 'enc_conv3'),
|
|
153
|
+
conv('enc_conv4', 'pool3'),
|
|
154
|
+
pool('pool4', 'enc_conv4'),
|
|
155
|
+
conv('enc_conv5a', 'pool4'),
|
|
156
|
+
conv('enc_conv5b', 'enc_conv5a'),
|
|
157
|
+
upsample('up4', 'enc_conv5b'),
|
|
158
|
+
concat('concat4', 'up4', 'pool3'),
|
|
159
|
+
conv('dec_conv4a', 'concat4'),
|
|
160
|
+
conv('dec_conv4b', 'dec_conv4a'),
|
|
161
|
+
upsample('up3', 'dec_conv4b'),
|
|
162
|
+
concat('concat3', 'up3', 'pool2'),
|
|
163
|
+
conv('dec_conv3a', 'concat3'),
|
|
164
|
+
conv('dec_conv3b', 'dec_conv3a'),
|
|
165
|
+
upsample('up2', 'dec_conv3b'),
|
|
166
|
+
concat('concat2', 'up2', 'pool1'),
|
|
167
|
+
conv('dec_conv2a', 'concat2'),
|
|
168
|
+
conv('dec_conv2b', 'dec_conv2a'),
|
|
169
|
+
upsample('up1', 'dec_conv2b'),
|
|
170
|
+
concat('concat1', 'up1', 'input'),
|
|
171
|
+
conv('dec_conv1a', 'concat1'),
|
|
172
|
+
conv('dec_conv1b', 'dec_conv1a'),
|
|
173
|
+
conv('dec_conv0', 'dec_conv1b')
|
|
174
|
+
]
|
|
175
|
+
};
|
|
176
|
+
|
|
177
|
+
export const OIDN_UNET_LARGE_SPEC: UNetModelSpec = {
|
|
178
|
+
schemaVersion: 1,
|
|
179
|
+
id: 'oidn-unet-large-v1',
|
|
180
|
+
family: 'oidn-unet-large',
|
|
181
|
+
input: 'input',
|
|
182
|
+
output: 'dec_conv1c',
|
|
183
|
+
receptiveField: 202,
|
|
184
|
+
nodes: [
|
|
185
|
+
conv('enc_conv1a', 'input'),
|
|
186
|
+
conv('enc_conv1b', 'enc_conv1a'),
|
|
187
|
+
pool('pool1', 'enc_conv1b'),
|
|
188
|
+
conv('enc_conv2a', 'pool1'),
|
|
189
|
+
conv('enc_conv2b', 'enc_conv2a'),
|
|
190
|
+
pool('pool2', 'enc_conv2b'),
|
|
191
|
+
conv('enc_conv3a', 'pool2'),
|
|
192
|
+
conv('enc_conv3b', 'enc_conv3a'),
|
|
193
|
+
pool('pool3', 'enc_conv3b'),
|
|
194
|
+
conv('enc_conv4a', 'pool3'),
|
|
195
|
+
conv('enc_conv4b', 'enc_conv4a'),
|
|
196
|
+
pool('pool4', 'enc_conv4b'),
|
|
197
|
+
conv('enc_conv5a', 'pool4'),
|
|
198
|
+
conv('enc_conv5b', 'enc_conv5a'),
|
|
199
|
+
upsample('up4', 'enc_conv5b'),
|
|
200
|
+
concat('concat4', 'up4', 'pool3'),
|
|
201
|
+
conv('dec_conv4a', 'concat4'),
|
|
202
|
+
conv('dec_conv4b', 'dec_conv4a'),
|
|
203
|
+
upsample('up3', 'dec_conv4b'),
|
|
204
|
+
concat('concat3', 'up3', 'pool2'),
|
|
205
|
+
conv('dec_conv3a', 'concat3'),
|
|
206
|
+
conv('dec_conv3b', 'dec_conv3a'),
|
|
207
|
+
upsample('up2', 'dec_conv3b'),
|
|
208
|
+
concat('concat2', 'up2', 'pool1'),
|
|
209
|
+
conv('dec_conv2a', 'concat2'),
|
|
210
|
+
conv('dec_conv2b', 'dec_conv2a'),
|
|
211
|
+
upsample('up1', 'dec_conv2b'),
|
|
212
|
+
concat('concat1', 'up1', 'input'),
|
|
213
|
+
conv('dec_conv1a', 'concat1'),
|
|
214
|
+
conv('dec_conv1b', 'dec_conv1a'),
|
|
215
|
+
conv('dec_conv1c', 'dec_conv1b')
|
|
216
|
+
]
|
|
217
|
+
};
|
|
218
|
+
|
|
219
|
+
const BUILTIN_MODEL_SPECS = [
|
|
220
|
+
OIDN_UNET_SMALL_SPEC,
|
|
221
|
+
OIDN_UNET_LARGE_SPEC
|
|
222
|
+
] as const;
|
|
223
|
+
|
|
224
|
+
function expectedTensorNames(spec: UNetModelSpec): Set<string> {
|
|
225
|
+
const names = new Set<string>();
|
|
226
|
+
for (const node of spec.nodes) {
|
|
227
|
+
if (node.op === 'conv2d') {
|
|
228
|
+
names.add(node.weight);
|
|
229
|
+
names.add(node.bias);
|
|
230
|
+
}
|
|
231
|
+
}
|
|
232
|
+
return names;
|
|
233
|
+
}
|
|
234
|
+
|
|
235
|
+
function tensorByteSize(tensor: HostTensor): number {
|
|
236
|
+
return tensor.desc.getByteSize();
|
|
237
|
+
}
|
|
238
|
+
|
|
239
|
+
function describeNames(names: Iterable<string>): string {
|
|
240
|
+
return [...names].sort().join(', ');
|
|
241
|
+
}
|
|
242
|
+
|
|
243
|
+
export function detectUNetModelSpec(
|
|
244
|
+
tensors: ReadonlyMap<string, HostTensor>,
|
|
245
|
+
specs: readonly UNetModelSpec[] = BUILTIN_MODEL_SPECS
|
|
246
|
+
): UNetModelSpec {
|
|
247
|
+
const matches = specs.filter((spec) => {
|
|
248
|
+
const expected = expectedTensorNames(spec);
|
|
249
|
+
if ([...expected].some((name) => !tensors.has(name))) return false;
|
|
250
|
+
return (
|
|
251
|
+
spec.allowAdditionalTensors === true ||
|
|
252
|
+
[...tensors.keys()].every((name) => expected.has(name))
|
|
253
|
+
);
|
|
254
|
+
});
|
|
255
|
+
|
|
256
|
+
if (matches.length === 1) return matches[0];
|
|
257
|
+
if (matches.length > 1) {
|
|
258
|
+
throw new Error(
|
|
259
|
+
`Ambiguous OIDN model topology: ${matches.map((spec) => spec.id).join(', ')}`
|
|
260
|
+
);
|
|
261
|
+
}
|
|
262
|
+
|
|
263
|
+
throw new Error(
|
|
264
|
+
`Unsupported OIDN model topology. TZA tensors: ${describeNames(tensors.keys())}`
|
|
265
|
+
);
|
|
266
|
+
}
|
|
267
|
+
|
|
268
|
+
function requireTensor(
|
|
269
|
+
tensors: ReadonlyMap<string, HostTensor>,
|
|
270
|
+
name: string,
|
|
271
|
+
modelId: string
|
|
272
|
+
): HostTensor {
|
|
273
|
+
const tensor = tensors.get(name);
|
|
274
|
+
if (!tensor) {
|
|
275
|
+
throw new Error(`Model ${modelId} is missing tensor ${name}`);
|
|
276
|
+
}
|
|
277
|
+
if (tensor.data.byteLength !== tensorByteSize(tensor)) {
|
|
278
|
+
throw new Error(
|
|
279
|
+
`Tensor ${name} has ${tensor.data.byteLength} bytes, expected ${tensorByteSize(tensor)}`
|
|
280
|
+
);
|
|
281
|
+
}
|
|
282
|
+
return tensor;
|
|
283
|
+
}
|
|
284
|
+
|
|
285
|
+
/** Validate tensor layout, shapes, dtypes, graph order, and channel flow. */
|
|
286
|
+
export function validateUNetModel(
|
|
287
|
+
tensors: ReadonlyMap<string, HostTensor>,
|
|
288
|
+
spec: UNetModelSpec = detectUNetModelSpec(tensors)
|
|
289
|
+
): ValidatedUNetModel {
|
|
290
|
+
if (spec.schemaVersion !== 1) {
|
|
291
|
+
throw new Error(`Unsupported model descriptor schema ${spec.schemaVersion}`);
|
|
292
|
+
}
|
|
293
|
+
|
|
294
|
+
const expected = expectedTensorNames(spec);
|
|
295
|
+
if (!spec.allowAdditionalTensors) {
|
|
296
|
+
const additional = [...tensors.keys()].filter((name) => !expected.has(name));
|
|
297
|
+
if (additional.length > 0) {
|
|
298
|
+
throw new Error(
|
|
299
|
+
`Model ${spec.id} has unexpected tensors: ${describeNames(additional)}`
|
|
300
|
+
);
|
|
301
|
+
}
|
|
302
|
+
}
|
|
303
|
+
|
|
304
|
+
const channelsByValue = new Map<string, number>();
|
|
305
|
+
const convTensors = new Map<string, ValidatedConvTensor>();
|
|
306
|
+
const convChannels = new Map<string, ModelConvChannels>();
|
|
307
|
+
const producedValues = new Set<string>([spec.input]);
|
|
308
|
+
let inputChannels: number | undefined;
|
|
309
|
+
let commonDataType: HostTensor['desc']['dataType'] | undefined;
|
|
310
|
+
|
|
311
|
+
const getChannels = (value: string, nodeId: string) => {
|
|
312
|
+
const channels = channelsByValue.get(value);
|
|
313
|
+
if (channels === undefined) {
|
|
314
|
+
throw new Error(
|
|
315
|
+
`Model ${spec.id} node ${nodeId} reads unknown or forward value ${value}`
|
|
316
|
+
);
|
|
317
|
+
}
|
|
318
|
+
return channels;
|
|
319
|
+
};
|
|
320
|
+
|
|
321
|
+
for (const node of spec.nodes) {
|
|
322
|
+
if (producedValues.has(node.id)) {
|
|
323
|
+
throw new Error(`Model ${spec.id} produces duplicate value ${node.id}`);
|
|
324
|
+
}
|
|
325
|
+
|
|
326
|
+
if (node.op === 'conv2d') {
|
|
327
|
+
const weight = requireTensor(tensors, node.weight, spec.id);
|
|
328
|
+
const bias = requireTensor(tensors, node.bias, spec.id);
|
|
329
|
+
const dims = weight.desc.dims;
|
|
330
|
+
|
|
331
|
+
if (weight.desc.layout !== 'oihw' || dims.length !== 4) {
|
|
332
|
+
throw new Error(`Tensor ${node.weight} must use OIHW layout`);
|
|
333
|
+
}
|
|
334
|
+
if (dims[2] !== 3 || dims[3] !== 3) {
|
|
335
|
+
throw new Error(`Tensor ${node.weight} must use a 3x3 kernel`);
|
|
336
|
+
}
|
|
337
|
+
if (bias.desc.layout !== 'x' || bias.desc.dims.length !== 1) {
|
|
338
|
+
throw new Error(`Tensor ${node.bias} must be a one-dimensional bias`);
|
|
339
|
+
}
|
|
340
|
+
if (bias.desc.dims[0] !== dims[0]) {
|
|
341
|
+
throw new Error(
|
|
342
|
+
`Tensor ${node.bias} has ${bias.desc.dims[0]} channels, expected ${dims[0]}`
|
|
343
|
+
);
|
|
344
|
+
}
|
|
345
|
+
if (weight.desc.dataType !== bias.desc.dataType) {
|
|
346
|
+
throw new Error(`Weight and bias dtype differ for ${node.id}`);
|
|
347
|
+
}
|
|
348
|
+
if (commonDataType && commonDataType !== weight.desc.dataType) {
|
|
349
|
+
throw new Error(`Mixed tensor dtypes are not supported by model ${spec.id}`);
|
|
350
|
+
}
|
|
351
|
+
commonDataType = weight.desc.dataType;
|
|
352
|
+
|
|
353
|
+
if (node.input === spec.input && inputChannels === undefined) {
|
|
354
|
+
inputChannels = dims[1];
|
|
355
|
+
channelsByValue.set(spec.input, inputChannels);
|
|
356
|
+
}
|
|
357
|
+
const actualInputChannels = getChannels(node.input, node.id);
|
|
358
|
+
if (actualInputChannels !== dims[1]) {
|
|
359
|
+
throw new Error(
|
|
360
|
+
`Tensor ${node.weight} expects ${dims[1]} input channels, ` +
|
|
361
|
+
`but ${node.input} provides ${actualInputChannels}`
|
|
362
|
+
);
|
|
363
|
+
}
|
|
364
|
+
|
|
365
|
+
channelsByValue.set(node.id, dims[0]);
|
|
366
|
+
convTensors.set(node.id, {
|
|
367
|
+
weight,
|
|
368
|
+
bias,
|
|
369
|
+
inputChannels: dims[1],
|
|
370
|
+
outputChannels: dims[0],
|
|
371
|
+
kernelHeight: dims[2],
|
|
372
|
+
kernelWidth: dims[3]
|
|
373
|
+
});
|
|
374
|
+
convChannels.set(node.id, {
|
|
375
|
+
inputChannels: dims[1],
|
|
376
|
+
outputChannels: dims[0]
|
|
377
|
+
});
|
|
378
|
+
} else if (node.op === 'concat') {
|
|
379
|
+
if (node.inputs.length < 2) {
|
|
380
|
+
throw new Error(`Concat ${node.id} requires at least two inputs`);
|
|
381
|
+
}
|
|
382
|
+
const channels = node.inputs.reduce(
|
|
383
|
+
(sum, value) => sum + getChannels(value, node.id),
|
|
384
|
+
0
|
|
385
|
+
);
|
|
386
|
+
channelsByValue.set(node.id, channels);
|
|
387
|
+
} else {
|
|
388
|
+
channelsByValue.set(node.id, getChannels(node.input, node.id));
|
|
389
|
+
}
|
|
390
|
+
|
|
391
|
+
producedValues.add(node.id);
|
|
392
|
+
}
|
|
393
|
+
|
|
394
|
+
if (inputChannels === undefined || commonDataType === undefined) {
|
|
395
|
+
throw new Error(`Model ${spec.id} has no convolution reading its input`);
|
|
396
|
+
}
|
|
397
|
+
const outputChannels = channelsByValue.get(spec.output);
|
|
398
|
+
if (outputChannels === undefined) {
|
|
399
|
+
throw new Error(`Model ${spec.id} output ${spec.output} is not produced`);
|
|
400
|
+
}
|
|
401
|
+
if (outputChannels !== 3) {
|
|
402
|
+
throw new Error(`Model ${spec.id} must produce 3 channels, got ${outputChannels}`);
|
|
403
|
+
}
|
|
404
|
+
|
|
405
|
+
return {
|
|
406
|
+
spec,
|
|
407
|
+
inputChannels,
|
|
408
|
+
outputChannels,
|
|
409
|
+
tensorDataType: commonDataType,
|
|
410
|
+
channelsByValue,
|
|
411
|
+
convChannels,
|
|
412
|
+
convTensors
|
|
413
|
+
};
|
|
414
|
+
}
|