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/WGPUComputePass.ts
CHANGED
|
@@ -152,13 +152,15 @@ export class WGPUComputePass<I extends string, O extends string> {
|
|
|
152
152
|
return this._outputBuffers[name].buffer;
|
|
153
153
|
}
|
|
154
154
|
|
|
155
|
-
dispose() {
|
|
155
|
+
dispose(destroyOutputBuffers = true) {
|
|
156
156
|
Object.keys(this._uniformBuffers).forEach((key) => {
|
|
157
157
|
(this._uniformBuffers as any)[key].destroy();
|
|
158
158
|
});
|
|
159
|
-
|
|
160
|
-
(this._outputBuffers
|
|
161
|
-
|
|
159
|
+
if (destroyOutputBuffers) {
|
|
160
|
+
Object.keys(this._outputBuffers).forEach((key) => {
|
|
161
|
+
(this._outputBuffers as any)[key].buffer.destroy();
|
|
162
|
+
});
|
|
163
|
+
}
|
|
162
164
|
}
|
|
163
165
|
|
|
164
166
|
private _createBuffer(params: WGPUComputePassOutput) {
|
package/src/backend.ts
CHANGED
|
@@ -1,62 +1,49 @@
|
|
|
1
|
-
import { WebGPUBackend } from '@tensorflow/tfjs-backend-webgpu/dist/base';
|
|
2
|
-
import { ENGINE } from '@tensorflow/tfjs-core/dist/engine';
|
|
3
|
-
|
|
4
|
-
import './kernels';
|
|
5
|
-
|
|
6
1
|
export async function initWebGPUBackend() {
|
|
7
|
-
|
|
8
|
-
|
|
9
|
-
|
|
10
|
-
|
|
11
|
-
|
|
12
|
-
|
|
13
|
-
|
|
14
|
-
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
|
|
18
|
-
|
|
19
|
-
|
|
20
|
-
|
|
21
|
-
|
|
22
|
-
|
|
23
|
-
|
|
24
|
-
|
|
25
|
-
|
|
26
|
-
|
|
27
|
-
|
|
28
|
-
|
|
29
|
-
|
|
30
|
-
|
|
31
|
-
|
|
32
|
-
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
|
|
38
|
-
|
|
39
|
-
|
|
40
|
-
|
|
41
|
-
|
|
42
|
-
|
|
43
|
-
|
|
44
|
-
|
|
2
|
+
if (!navigator.gpu) throw new Error('WebGPU is not available');
|
|
3
|
+
const gpuDescriptor: GPURequestAdapterOptions = {
|
|
4
|
+
powerPreference: 'high-performance'
|
|
5
|
+
};
|
|
6
|
+
|
|
7
|
+
const adapter = await navigator.gpu.requestAdapter(gpuDescriptor);
|
|
8
|
+
if (!adapter) throw new Error('No WebGPU adapter is available');
|
|
9
|
+
const deviceDescriptor: GPUDeviceDescriptor = {};
|
|
10
|
+
|
|
11
|
+
const requiredFeatures: GPUFeatureName[] = [];
|
|
12
|
+
if (adapter.features.has('timestamp-query')) {
|
|
13
|
+
requiredFeatures.push('timestamp-query');
|
|
14
|
+
}
|
|
15
|
+
if (adapter.features.has('bgra8unorm-storage')) {
|
|
16
|
+
requiredFeatures.push('bgra8unorm-storage');
|
|
17
|
+
}
|
|
18
|
+
if (adapter.features.has('shader-f16')) {
|
|
19
|
+
requiredFeatures.push('shader-f16');
|
|
20
|
+
}
|
|
21
|
+
deviceDescriptor.requiredFeatures = requiredFeatures;
|
|
22
|
+
|
|
23
|
+
const adapterLimits = adapter.limits;
|
|
24
|
+
deviceDescriptor.requiredLimits = {
|
|
25
|
+
maxComputeWorkgroupStorageSize:
|
|
26
|
+
adapterLimits.maxComputeWorkgroupStorageSize,
|
|
27
|
+
maxComputeWorkgroupsPerDimension:
|
|
28
|
+
adapterLimits.maxComputeWorkgroupsPerDimension,
|
|
29
|
+
maxStorageBufferBindingSize: adapterLimits.maxStorageBufferBindingSize,
|
|
30
|
+
maxBufferSize: adapterLimits.maxBufferSize,
|
|
31
|
+
maxComputeWorkgroupSizeX: adapterLimits.maxComputeWorkgroupSizeX,
|
|
32
|
+
maxComputeInvocationsPerWorkgroup:
|
|
33
|
+
adapterLimits.maxComputeInvocationsPerWorkgroup
|
|
34
|
+
};
|
|
35
|
+
const device = await adapter.requestDevice(deviceDescriptor);
|
|
36
|
+
const adapterInfo =
|
|
37
|
+
// requestAdapterInfo is deprecated
|
|
38
|
+
// @ts-ignore
|
|
39
|
+
adapter.info ?? (await adapter.requestAdapterInfo?.());
|
|
40
|
+
|
|
41
|
+
return initWebGPUBackendWithDevice(device, adapterInfo);
|
|
45
42
|
}
|
|
46
43
|
|
|
47
44
|
export async function initWebGPUBackendWithDevice(
|
|
48
45
|
device: GPUDevice,
|
|
49
|
-
|
|
46
|
+
adapterInfo: GPUAdapterInfo
|
|
50
47
|
) {
|
|
51
|
-
|
|
52
|
-
let backend = ENGINE.findBackend('webgpu-oidn');
|
|
53
|
-
if (backend != null) {
|
|
54
|
-
return backend as WebGPUBackend;
|
|
55
|
-
}
|
|
56
|
-
|
|
57
|
-
backend = new WebGPUBackend(device, adapter);
|
|
58
|
-
ENGINE.registerBackend('webgpu-oidn', () => backend);
|
|
59
|
-
await ENGINE.setBackend('webgpu-oidn');
|
|
60
|
-
|
|
61
|
-
return backend as WebGPUBackend;
|
|
48
|
+
return { device, adapterInfo };
|
|
62
49
|
}
|
|
@@ -0,0 +1,301 @@
|
|
|
1
|
+
import type {
|
|
2
|
+
Conv2DNodeSpec,
|
|
3
|
+
MaxPool2DNodeSpec,
|
|
4
|
+
ModelNodeSpec,
|
|
5
|
+
UNetModelSpec,
|
|
6
|
+
Upsample2DNodeSpec,
|
|
7
|
+
UNetModelGraph
|
|
8
|
+
} from './modelSpec';
|
|
9
|
+
|
|
10
|
+
export interface FusedConvPoolNodeSpec {
|
|
11
|
+
op: 'fusedConvReluMaxPool2d';
|
|
12
|
+
id: string;
|
|
13
|
+
input: string;
|
|
14
|
+
conv: Conv2DNodeSpec;
|
|
15
|
+
pool: MaxPool2DNodeSpec;
|
|
16
|
+
}
|
|
17
|
+
|
|
18
|
+
export interface FusedUpsampleConcatConvNodeSpec {
|
|
19
|
+
op: 'fusedUpsampleConcatConv2d';
|
|
20
|
+
id: string;
|
|
21
|
+
/** Inputs stay in concat order because that order selects weight channels. */
|
|
22
|
+
inputs: readonly {
|
|
23
|
+
value: string;
|
|
24
|
+
upsample?: Upsample2DNodeSpec;
|
|
25
|
+
}[];
|
|
26
|
+
conv: Conv2DNodeSpec;
|
|
27
|
+
}
|
|
28
|
+
|
|
29
|
+
export type ExecutableModelNode =
|
|
30
|
+
| ModelNodeSpec
|
|
31
|
+
| FusedConvPoolNodeSpec
|
|
32
|
+
| FusedUpsampleConcatConvNodeSpec;
|
|
33
|
+
|
|
34
|
+
export interface GraphOptimizationOptions {
|
|
35
|
+
fuseConvPool?: boolean;
|
|
36
|
+
fuseUpsampleConcatConv?: boolean;
|
|
37
|
+
}
|
|
38
|
+
|
|
39
|
+
export interface OptimizedModelGraph {
|
|
40
|
+
spec: UNetModelSpec;
|
|
41
|
+
nodes: readonly ExecutableModelNode[];
|
|
42
|
+
fusions: {
|
|
43
|
+
convPool: number;
|
|
44
|
+
upsampleConcatConv: number;
|
|
45
|
+
};
|
|
46
|
+
}
|
|
47
|
+
|
|
48
|
+
function nodeInputs(node: ModelNodeSpec): readonly string[] {
|
|
49
|
+
return node.op === 'concat' ? node.inputs : [node.input];
|
|
50
|
+
}
|
|
51
|
+
|
|
52
|
+
function buildConsumers(spec: UNetModelSpec) {
|
|
53
|
+
const consumers = new Map<string, ModelNodeSpec[]>();
|
|
54
|
+
for (const node of spec.nodes) {
|
|
55
|
+
for (const input of nodeInputs(node)) {
|
|
56
|
+
const list = consumers.get(input) ?? [];
|
|
57
|
+
list.push(node);
|
|
58
|
+
consumers.set(input, list);
|
|
59
|
+
}
|
|
60
|
+
}
|
|
61
|
+
return consumers;
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
function onlyConsumer<T extends ModelNodeSpec['op']>(
|
|
65
|
+
consumers: ReadonlyMap<string, ModelNodeSpec[]>,
|
|
66
|
+
value: string,
|
|
67
|
+
op: T
|
|
68
|
+
): Extract<ModelNodeSpec, { op: T }> | undefined {
|
|
69
|
+
const list = consumers.get(value);
|
|
70
|
+
if (list?.length !== 1 || list[0].op !== op) return undefined;
|
|
71
|
+
return list[0] as Extract<ModelNodeSpec, { op: T }>;
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
/**
|
|
75
|
+
* Applies topology-only fusions. It never depends on a particular OIDN model
|
|
76
|
+
* name, so new descriptors automatically benefit from known graph patterns.
|
|
77
|
+
*/
|
|
78
|
+
export function optimizeModelGraph(
|
|
79
|
+
validated: UNetModelGraph,
|
|
80
|
+
options: GraphOptimizationOptions = {}
|
|
81
|
+
): OptimizedModelGraph {
|
|
82
|
+
const spec = validated.spec;
|
|
83
|
+
const consumers = buildConsumers(spec);
|
|
84
|
+
const nodesById = new Map(spec.nodes.map((node) => [node.id, node]));
|
|
85
|
+
const eliminated = new Set<string>();
|
|
86
|
+
const fusedAt = new Map<string, ExecutableModelNode>();
|
|
87
|
+
let convPool = 0;
|
|
88
|
+
let upsampleConcatConv = 0;
|
|
89
|
+
|
|
90
|
+
if (options.fuseConvPool !== false) {
|
|
91
|
+
for (const node of spec.nodes) {
|
|
92
|
+
if (node.op !== 'conv2d' || node.activation !== 'relu') continue;
|
|
93
|
+
const pool = onlyConsumer(consumers, node.id, 'maxPool2d');
|
|
94
|
+
if (!pool || pool.size !== 2 || pool.stride !== 2) continue;
|
|
95
|
+
|
|
96
|
+
eliminated.add(node.id);
|
|
97
|
+
fusedAt.set(pool.id, {
|
|
98
|
+
op: 'fusedConvReluMaxPool2d',
|
|
99
|
+
id: pool.id,
|
|
100
|
+
input: node.input,
|
|
101
|
+
conv: node,
|
|
102
|
+
pool
|
|
103
|
+
});
|
|
104
|
+
convPool++;
|
|
105
|
+
}
|
|
106
|
+
}
|
|
107
|
+
|
|
108
|
+
if (options.fuseUpsampleConcatConv !== false) {
|
|
109
|
+
for (const node of spec.nodes) {
|
|
110
|
+
if (node.op !== 'conv2d') continue;
|
|
111
|
+
const concat = nodesById.get(node.input);
|
|
112
|
+
if (concat?.op !== 'concat' || concat.inputs.length !== 2) continue;
|
|
113
|
+
if (onlyConsumer(consumers, concat.id, 'conv2d') !== node) continue;
|
|
114
|
+
// The native blocked layout can remove concat only when the source
|
|
115
|
+
// boundary is also a vec4 boundary. Other graphs keep the generic ops
|
|
116
|
+
// and can use the compatibility engine until a scalar-tail kernel exists.
|
|
117
|
+
const firstInputChannels = validated.channelsByValue.get(concat.inputs[0]);
|
|
118
|
+
if (firstInputChannels === undefined || firstInputChannels % 4 !== 0) {
|
|
119
|
+
continue;
|
|
120
|
+
}
|
|
121
|
+
|
|
122
|
+
const inputs = concat.inputs.map((value) => {
|
|
123
|
+
const candidate = nodesById.get(value);
|
|
124
|
+
if (
|
|
125
|
+
candidate?.op === 'upsample2d' &&
|
|
126
|
+
candidate.scale === 2 &&
|
|
127
|
+
candidate.mode === 'nearest' &&
|
|
128
|
+
onlyConsumer(consumers, candidate.id, 'concat') === concat
|
|
129
|
+
) {
|
|
130
|
+
return { value: candidate.input, upsample: candidate };
|
|
131
|
+
}
|
|
132
|
+
return { value };
|
|
133
|
+
});
|
|
134
|
+
const upsampleCount = inputs.filter((input) => input.upsample).length;
|
|
135
|
+
if (upsampleCount !== 1) continue;
|
|
136
|
+
|
|
137
|
+
eliminated.add(concat.id);
|
|
138
|
+
for (const value of concat.inputs) {
|
|
139
|
+
const candidate = nodesById.get(value);
|
|
140
|
+
if (candidate?.op === 'upsample2d') eliminated.add(candidate.id);
|
|
141
|
+
}
|
|
142
|
+
fusedAt.set(node.id, {
|
|
143
|
+
op: 'fusedUpsampleConcatConv2d',
|
|
144
|
+
id: node.id,
|
|
145
|
+
inputs,
|
|
146
|
+
conv: node
|
|
147
|
+
});
|
|
148
|
+
upsampleConcatConv++;
|
|
149
|
+
}
|
|
150
|
+
}
|
|
151
|
+
|
|
152
|
+
const nodes: ExecutableModelNode[] = [];
|
|
153
|
+
for (const node of spec.nodes) {
|
|
154
|
+
const fused = fusedAt.get(node.id);
|
|
155
|
+
if (fused) {
|
|
156
|
+
nodes.push(fused);
|
|
157
|
+
} else if (!eliminated.has(node.id)) {
|
|
158
|
+
nodes.push(node);
|
|
159
|
+
}
|
|
160
|
+
}
|
|
161
|
+
|
|
162
|
+
return {
|
|
163
|
+
spec,
|
|
164
|
+
nodes,
|
|
165
|
+
fusions: { convPool, upsampleConcatConv }
|
|
166
|
+
};
|
|
167
|
+
}
|
|
168
|
+
|
|
169
|
+
export interface ModelValueShape {
|
|
170
|
+
width: number;
|
|
171
|
+
height: number;
|
|
172
|
+
channels: number;
|
|
173
|
+
}
|
|
174
|
+
|
|
175
|
+
export interface PlannedModelNode {
|
|
176
|
+
node: ExecutableModelNode;
|
|
177
|
+
outputShape: ModelValueShape;
|
|
178
|
+
/** Last planned node that reads the output; output itself uses nodes.length. */
|
|
179
|
+
lastUse: number;
|
|
180
|
+
}
|
|
181
|
+
|
|
182
|
+
export interface ModelExecutionPlan extends OptimizedModelGraph {
|
|
183
|
+
inputShape: ModelValueShape;
|
|
184
|
+
valueShapes: ReadonlyMap<string, ModelValueShape>;
|
|
185
|
+
plannedNodes: readonly PlannedModelNode[];
|
|
186
|
+
}
|
|
187
|
+
|
|
188
|
+
function executableInputs(node: ExecutableModelNode): readonly string[] {
|
|
189
|
+
if (node.op === 'concat') return node.inputs;
|
|
190
|
+
if (node.op === 'fusedUpsampleConcatConv2d') {
|
|
191
|
+
return node.inputs.map((input) => input.value);
|
|
192
|
+
}
|
|
193
|
+
return [node.input];
|
|
194
|
+
}
|
|
195
|
+
|
|
196
|
+
function sameSpatialShape(
|
|
197
|
+
left: ModelValueShape,
|
|
198
|
+
right: ModelValueShape
|
|
199
|
+
): boolean {
|
|
200
|
+
return left.width === right.width && left.height === right.height;
|
|
201
|
+
}
|
|
202
|
+
|
|
203
|
+
/** Resolve all runtime shapes and value lifetimes before allocating GPU data. */
|
|
204
|
+
export function planModelExecution(
|
|
205
|
+
validated: UNetModelGraph,
|
|
206
|
+
width: number,
|
|
207
|
+
height: number,
|
|
208
|
+
options?: GraphOptimizationOptions
|
|
209
|
+
): ModelExecutionPlan {
|
|
210
|
+
if (!Number.isInteger(width) || width <= 0 || !Number.isInteger(height) || height <= 0) {
|
|
211
|
+
throw new Error(`Invalid model input size ${width}x${height}`);
|
|
212
|
+
}
|
|
213
|
+
|
|
214
|
+
const graph = optimizeModelGraph(validated, options);
|
|
215
|
+
const inputShape = { width, height, channels: validated.inputChannels };
|
|
216
|
+
const valueShapes = new Map<string, ModelValueShape>([
|
|
217
|
+
[validated.spec.input, inputShape]
|
|
218
|
+
]);
|
|
219
|
+
const outputShapes: ModelValueShape[] = [];
|
|
220
|
+
|
|
221
|
+
const shapeOf = (value: string, nodeId: string) => {
|
|
222
|
+
const shape = valueShapes.get(value);
|
|
223
|
+
if (!shape) throw new Error(`Planned node ${nodeId} reads missing value ${value}`);
|
|
224
|
+
return shape;
|
|
225
|
+
};
|
|
226
|
+
|
|
227
|
+
for (const node of graph.nodes) {
|
|
228
|
+
let outputShape: ModelValueShape;
|
|
229
|
+
if (node.op === 'conv2d') {
|
|
230
|
+
const input = shapeOf(node.input, node.id);
|
|
231
|
+
outputShape = {
|
|
232
|
+
width: input.width,
|
|
233
|
+
height: input.height,
|
|
234
|
+
channels: validated.convChannels.get(node.id)!.outputChannels
|
|
235
|
+
};
|
|
236
|
+
} else if (node.op === 'maxPool2d') {
|
|
237
|
+
const input = shapeOf(node.input, node.id);
|
|
238
|
+
outputShape = {
|
|
239
|
+
width: Math.ceil(input.width / 2),
|
|
240
|
+
height: Math.ceil(input.height / 2),
|
|
241
|
+
channels: input.channels
|
|
242
|
+
};
|
|
243
|
+
} else if (node.op === 'upsample2d') {
|
|
244
|
+
const input = shapeOf(node.input, node.id);
|
|
245
|
+
outputShape = {
|
|
246
|
+
width: input.width * 2,
|
|
247
|
+
height: input.height * 2,
|
|
248
|
+
channels: input.channels
|
|
249
|
+
};
|
|
250
|
+
} else if (node.op === 'concat') {
|
|
251
|
+
const inputs = node.inputs.map((value) => shapeOf(value, node.id));
|
|
252
|
+
if (inputs.some((shape) => !sameSpatialShape(shape, inputs[0]))) {
|
|
253
|
+
throw new Error(`Concat ${node.id} has mismatched spatial shapes`);
|
|
254
|
+
}
|
|
255
|
+
outputShape = {
|
|
256
|
+
width: inputs[0].width,
|
|
257
|
+
height: inputs[0].height,
|
|
258
|
+
channels: inputs.reduce((sum, shape) => sum + shape.channels, 0)
|
|
259
|
+
};
|
|
260
|
+
} else if (node.op === 'fusedConvReluMaxPool2d') {
|
|
261
|
+
const input = shapeOf(node.input, node.id);
|
|
262
|
+
outputShape = {
|
|
263
|
+
width: Math.ceil(input.width / 2),
|
|
264
|
+
height: Math.ceil(input.height / 2),
|
|
265
|
+
channels: validated.convChannels.get(node.conv.id)!.outputChannels
|
|
266
|
+
};
|
|
267
|
+
} else {
|
|
268
|
+
const inputs = node.inputs.map((input) => {
|
|
269
|
+
const source = shapeOf(input.value, node.id);
|
|
270
|
+
return input.upsample
|
|
271
|
+
? { ...source, width: source.width * 2, height: source.height * 2 }
|
|
272
|
+
: source;
|
|
273
|
+
});
|
|
274
|
+
if (inputs.some((shape) => !sameSpatialShape(shape, inputs[0]))) {
|
|
275
|
+
throw new Error(`Fused decoder ${node.id} has mismatched spatial shapes`);
|
|
276
|
+
}
|
|
277
|
+
outputShape = {
|
|
278
|
+
width: inputs[0].width,
|
|
279
|
+
height: inputs[0].height,
|
|
280
|
+
channels: validated.convChannels.get(node.conv.id)!.outputChannels
|
|
281
|
+
};
|
|
282
|
+
}
|
|
283
|
+
|
|
284
|
+
valueShapes.set(node.id, outputShape);
|
|
285
|
+
outputShapes.push(outputShape);
|
|
286
|
+
}
|
|
287
|
+
|
|
288
|
+
const lastUses = new Map<string, number>();
|
|
289
|
+
graph.nodes.forEach((node, index) => {
|
|
290
|
+
for (const input of executableInputs(node)) lastUses.set(input, index);
|
|
291
|
+
});
|
|
292
|
+
lastUses.set(validated.spec.output, graph.nodes.length);
|
|
293
|
+
|
|
294
|
+
const plannedNodes = graph.nodes.map((node, index) => ({
|
|
295
|
+
node,
|
|
296
|
+
outputShape: outputShapes[index],
|
|
297
|
+
lastUse: lastUses.get(node.id) ?? index
|
|
298
|
+
}));
|
|
299
|
+
|
|
300
|
+
return { ...graph, inputShape, valueShapes, plannedNodes };
|
|
301
|
+
}
|
package/src/main.ts
CHANGED
|
@@ -1,17 +1,80 @@
|
|
|
1
1
|
import { parseTZA } from './tza';
|
|
2
2
|
import UNet from './UNet';
|
|
3
|
+
import type { UNetEngineSetting, UNetExecutionStats } from './UNet';
|
|
3
4
|
import { initWebGPUBackend, initWebGPUBackendWithDevice } from './backend';
|
|
5
|
+
import type { DynamicTileSetting } from './tileScheduler';
|
|
6
|
+
import type { UNetModelSpec } from './modelSpec';
|
|
7
|
+
import type {
|
|
8
|
+
NativeUNetKernelSetting,
|
|
9
|
+
NativeUNetPrecisionSetting
|
|
10
|
+
} from './nativeUNet';
|
|
4
11
|
|
|
5
12
|
export { parseTZA, UNet };
|
|
13
|
+
export type { DynamicTileOptions, DynamicTileSetting } from './tileScheduler';
|
|
14
|
+
export {
|
|
15
|
+
detectUNetModelSpec,
|
|
16
|
+
OIDN_UNET_LARGE_SPEC,
|
|
17
|
+
OIDN_UNET_SMALL_SPEC,
|
|
18
|
+
validateUNetModel
|
|
19
|
+
} from './modelSpec';
|
|
20
|
+
export type {
|
|
21
|
+
ModelNodeSpec,
|
|
22
|
+
UNetModelGraph,
|
|
23
|
+
UNetModelSpec,
|
|
24
|
+
ValidatedUNetModel
|
|
25
|
+
} from './modelSpec';
|
|
26
|
+
export { optimizeModelGraph, planModelExecution } from './graphOptimizer';
|
|
27
|
+
export type {
|
|
28
|
+
ExecutableModelNode,
|
|
29
|
+
GraphOptimizationOptions,
|
|
30
|
+
ModelExecutionPlan,
|
|
31
|
+
ModelValueShape,
|
|
32
|
+
OptimizedModelGraph
|
|
33
|
+
} from './graphOptimizer';
|
|
34
|
+
export {
|
|
35
|
+
NativeUNetExecutor,
|
|
36
|
+
resolveNativeUNetPrecision
|
|
37
|
+
} from './nativeUNet';
|
|
38
|
+
export type {
|
|
39
|
+
NativeUNetExecutionProfile,
|
|
40
|
+
NativeUNetKernel,
|
|
41
|
+
NativeUNetKernelSetting,
|
|
42
|
+
NativeUNetLayerTiming,
|
|
43
|
+
NativeUNetOptions,
|
|
44
|
+
NativeUNetPrecision,
|
|
45
|
+
NativeUNetPrecisionSetting
|
|
46
|
+
} from './nativeUNet';
|
|
47
|
+
|
|
48
|
+
export interface UNetOptions {
|
|
49
|
+
aux?: boolean;
|
|
50
|
+
hdr?: boolean;
|
|
51
|
+
/** Hard upper bound for an output tile edge. Defaults to 512. */
|
|
52
|
+
maxTileSize?: number;
|
|
53
|
+
/** Adaptive GPU-time-based tile sizing. Enabled by default. */
|
|
54
|
+
dynamicTile?: DynamicTileSetting;
|
|
55
|
+
/** `auto` uses stable WGSL; `webnn` opts into the experimental WebNN backend. */
|
|
56
|
+
engine?: UNetEngineSetting;
|
|
57
|
+
/** `auto` selects FP16 when shader-f16 was enabled on the GPUDevice. */
|
|
58
|
+
precision?: NativeUNetPrecisionSetting;
|
|
59
|
+
/** `auto` selects kernels from precision, operation shape, and GPU limits. */
|
|
60
|
+
kernel?: NativeUNetKernelSetting;
|
|
61
|
+
/** Versioned topology descriptor for future/custom OIDN TZA models. */
|
|
62
|
+
modelSpec?: UNetModelSpec;
|
|
63
|
+
}
|
|
64
|
+
|
|
65
|
+
export type { UNetEngineSetting, UNetExecutionStats } from './UNet';
|
|
66
|
+
export { WebNNUNetExecutor } from './webnnUNet';
|
|
67
|
+
export type { WebNNRuntimeSupport, WebNNUNetOptions } from './webnnUNet';
|
|
68
|
+
export type {
|
|
69
|
+
OIDNResourceKind,
|
|
70
|
+
OIDNResourceSnapshot,
|
|
71
|
+
OIDNResourceStats
|
|
72
|
+
} from './resourceTracker';
|
|
6
73
|
|
|
7
74
|
export async function initUNetFromBuffer(
|
|
8
75
|
tzaBuffer: ArrayBuffer,
|
|
9
76
|
backendParams?: { device: GPUDevice; adapterInfo: GPUAdapterInfo },
|
|
10
|
-
opts?:
|
|
11
|
-
aux?: boolean;
|
|
12
|
-
hdr?: boolean;
|
|
13
|
-
maxTileSize?: number;
|
|
14
|
-
}
|
|
77
|
+
opts?: UNetOptions
|
|
15
78
|
) {
|
|
16
79
|
const backend = await (backendParams
|
|
17
80
|
? initWebGPUBackendWithDevice(
|
|
@@ -20,18 +83,15 @@ export async function initUNetFromBuffer(
|
|
|
20
83
|
)
|
|
21
84
|
: initWebGPUBackend());
|
|
22
85
|
const tensors = parseTZA(tzaBuffer);
|
|
23
|
-
const unet = new UNet(tensors, backend
|
|
86
|
+
const unet = new UNet(tensors, backend, opts);
|
|
87
|
+
await unet.prepare();
|
|
24
88
|
return unet;
|
|
25
89
|
}
|
|
26
90
|
|
|
27
91
|
export async function initUNetFromURL(
|
|
28
92
|
modelPath: string,
|
|
29
93
|
backendParams?: { device: GPUDevice; adapterInfo: GPUAdapterInfo },
|
|
30
|
-
opts?:
|
|
31
|
-
aux?: boolean;
|
|
32
|
-
hdr?: boolean;
|
|
33
|
-
maxTileSize?: number;
|
|
34
|
-
}
|
|
94
|
+
opts?: UNetOptions
|
|
35
95
|
) {
|
|
36
96
|
return fetch(modelPath)
|
|
37
97
|
.then((res) => res.arrayBuffer())
|