oidn-web 0.3.5 → 0.5.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 +94 -0
- package/README.md +208 -8
- package/dist/oidn.js +4699 -22516
- package/dist/oidn.umd.cjs +989 -5796
- package/lib/UNet.d.ts +111 -26
- package/lib/UNet.js +310 -329
- 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 +1 -4
- package/lib/backend.js +28 -44
- package/lib/backend.js.map +1 -1
- package/lib/finalRgbShader.d.ts +13 -0
- package/lib/finalRgbShader.js +160 -0
- package/lib/finalRgbShader.js.map +1 -0
- package/lib/graphOptimizer.d.ts +54 -0
- package/lib/graphOptimizer.js +215 -0
- package/lib/graphOptimizer.js.map +1 -0
- package/lib/hdrTransfer.d.ts +14 -0
- package/lib/hdrTransfer.js +61 -0
- package/lib/hdrTransfer.js.map +1 -0
- package/lib/main.d.ts +43 -11
- package/lib/main.js +9 -5
- 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 +103 -0
- package/lib/nativeUNet.js +2064 -0
- package/lib/nativeUNet.js.map +1 -0
- package/lib/process.d.ts +5 -11
- package/lib/process.js +38 -49
- 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 +61 -0
- package/lib/tileScheduler.js +199 -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 +16 -5
- package/src/UNet.ts +463 -437
- package/src/WGPUComputePass.ts +6 -4
- package/src/backend.ts +33 -59
- package/src/finalRgbShader.ts +186 -0
- package/src/graphOptimizer.ts +300 -0
- package/src/hdrTransfer.ts +88 -0
- package/src/main.ts +95 -20
- package/src/modelSpec.ts +414 -0
- package/src/nativeUNet.ts +2655 -0
- package/src/process.ts +46 -71
- package/src/resourceTracker.ts +94 -0
- package/src/tileScheduler.ts +330 -0
- package/src/webnnUNet.ts +812 -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/main.ts
CHANGED
|
@@ -1,37 +1,112 @@
|
|
|
1
1
|
import { parseTZA } from './tza';
|
|
2
2
|
import UNet from './UNet';
|
|
3
|
-
import {
|
|
3
|
+
import type { UNetEngineSetting, UNetExecutionStats } from './UNet';
|
|
4
|
+
import { initWebGPUBackend } from './backend';
|
|
5
|
+
import type { DynamicTileSetting } from './tileScheduler';
|
|
6
|
+
import type { UNetModelSpec } from './modelSpec';
|
|
7
|
+
import type {
|
|
8
|
+
NativeUNetGemmOptions,
|
|
9
|
+
NativeUNetKernelSetting,
|
|
10
|
+
NativeUNetPrecisionSetting
|
|
11
|
+
} from './nativeUNet';
|
|
12
|
+
import type { HDRTransfer } from './process';
|
|
4
13
|
|
|
5
14
|
export { parseTZA, UNet };
|
|
15
|
+
export { planTileGrid } from './tileScheduler';
|
|
16
|
+
export type {
|
|
17
|
+
DynamicTileOptions,
|
|
18
|
+
DynamicTileSetting,
|
|
19
|
+
PlannedTile,
|
|
20
|
+
TilePlan,
|
|
21
|
+
TileRect
|
|
22
|
+
} from './tileScheduler';
|
|
23
|
+
export {
|
|
24
|
+
detectUNetModelSpec,
|
|
25
|
+
OIDN_UNET_LARGE_SPEC,
|
|
26
|
+
OIDN_UNET_SMALL_SPEC,
|
|
27
|
+
validateUNetModel
|
|
28
|
+
} from './modelSpec';
|
|
29
|
+
export type {
|
|
30
|
+
ModelNodeSpec,
|
|
31
|
+
UNetModelGraph,
|
|
32
|
+
UNetModelSpec,
|
|
33
|
+
ValidatedUNetModel
|
|
34
|
+
} from './modelSpec';
|
|
35
|
+
export { optimizeModelGraph, planModelExecution } from './graphOptimizer';
|
|
36
|
+
export type {
|
|
37
|
+
ExecutableModelNode,
|
|
38
|
+
GraphOptimizationOptions,
|
|
39
|
+
ModelExecutionPlan,
|
|
40
|
+
ModelValueShape,
|
|
41
|
+
OptimizedModelGraph
|
|
42
|
+
} from './graphOptimizer';
|
|
43
|
+
export {
|
|
44
|
+
NativeUNetExecutor,
|
|
45
|
+
resolveNativeUNetPrecision
|
|
46
|
+
} from './nativeUNet';
|
|
47
|
+
export type {
|
|
48
|
+
NativeUNetGemmWorkgroup,
|
|
49
|
+
NativeUNetGemmOptions,
|
|
50
|
+
NativeUNetExecutionProfile,
|
|
51
|
+
NativeUNetKernel,
|
|
52
|
+
NativeUNetKernelSetting,
|
|
53
|
+
NativeUNetLayerTiming,
|
|
54
|
+
NativeUNetOptions,
|
|
55
|
+
NativeUNetPrecision,
|
|
56
|
+
NativeUNetPrecisionSetting
|
|
57
|
+
} from './nativeUNet';
|
|
58
|
+
|
|
59
|
+
export interface UNetOptions {
|
|
60
|
+
aux?: boolean;
|
|
61
|
+
hdr?: boolean;
|
|
62
|
+
/** HDR transfer function expected by the trained model. Defaults to PU. */
|
|
63
|
+
hdrTransfer?: HDRTransfer;
|
|
64
|
+
/** Hard upper bound for an output tile edge. Defaults to 512. */
|
|
65
|
+
maxTileSize?: number;
|
|
66
|
+
/** Adaptive GPU-time-based tile sizing. Enabled by default. */
|
|
67
|
+
dynamicTile?: DynamicTileSetting;
|
|
68
|
+
/** `auto` uses native WGSL; `webnn` opts into the WebNN backend. */
|
|
69
|
+
engine?: UNetEngineSetting;
|
|
70
|
+
/** `auto` selects FP16 when shader-f16 was enabled on the GPUDevice. */
|
|
71
|
+
precision?: NativeUNetPrecisionSetting;
|
|
72
|
+
/** `auto` uses implicit GEMM for FP16/FP32 convolutions, except the direct output layer. */
|
|
73
|
+
kernel?: NativeUNetKernelSetting;
|
|
74
|
+
/**
|
|
75
|
+
* Experimental implicit-GEMM tuning for native WGSL execution. These knobs
|
|
76
|
+
* exist for benchmarking and may change or be removed in any release.
|
|
77
|
+
* @experimental
|
|
78
|
+
*/
|
|
79
|
+
gemm?: NativeUNetGemmOptions;
|
|
80
|
+
/** Versioned topology descriptor for future/custom OIDN TZA models. */
|
|
81
|
+
modelSpec?: UNetModelSpec;
|
|
82
|
+
}
|
|
83
|
+
|
|
84
|
+
export type { UNetEngineSetting, UNetExecutionStats } from './UNet';
|
|
85
|
+
export type { HDRTransfer } from './process';
|
|
86
|
+
export { WebNNUNetExecutor } from './webnnUNet';
|
|
87
|
+
export type { WebNNRuntimeSupport, WebNNUNetOptions } from './webnnUNet';
|
|
88
|
+
export type {
|
|
89
|
+
OIDNResourceKind,
|
|
90
|
+
OIDNResourceSnapshot,
|
|
91
|
+
OIDNResourceStats
|
|
92
|
+
} from './resourceTracker';
|
|
6
93
|
|
|
7
94
|
export async function initUNetFromBuffer(
|
|
8
95
|
tzaBuffer: ArrayBuffer,
|
|
9
|
-
backendParams?: { device: GPUDevice
|
|
10
|
-
opts?:
|
|
11
|
-
aux?: boolean;
|
|
12
|
-
hdr?: boolean;
|
|
13
|
-
maxTileSize?: number;
|
|
14
|
-
}
|
|
96
|
+
backendParams?: { device: GPUDevice },
|
|
97
|
+
opts?: UNetOptions
|
|
15
98
|
) {
|
|
16
|
-
const
|
|
17
|
-
? initWebGPUBackendWithDevice(
|
|
18
|
-
backendParams.device,
|
|
19
|
-
backendParams.adapterInfo
|
|
20
|
-
)
|
|
21
|
-
: initWebGPUBackend());
|
|
99
|
+
const device = backendParams?.device ?? await initWebGPUBackend();
|
|
22
100
|
const tensors = parseTZA(tzaBuffer);
|
|
23
|
-
const unet = new UNet(tensors,
|
|
101
|
+
const unet = new UNet(tensors, device, opts);
|
|
102
|
+
await unet.prepare();
|
|
24
103
|
return unet;
|
|
25
104
|
}
|
|
26
105
|
|
|
27
106
|
export async function initUNetFromURL(
|
|
28
107
|
modelPath: string,
|
|
29
|
-
backendParams?: { device: GPUDevice
|
|
30
|
-
opts?:
|
|
31
|
-
aux?: boolean;
|
|
32
|
-
hdr?: boolean;
|
|
33
|
-
maxTileSize?: number;
|
|
34
|
-
}
|
|
108
|
+
backendParams?: { device: GPUDevice },
|
|
109
|
+
opts?: UNetOptions
|
|
35
110
|
) {
|
|
36
111
|
return fetch(modelPath)
|
|
37
112
|
.then((res) => res.arrayBuffer())
|
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
|
+
}
|