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.
Files changed (122) hide show
  1. package/CHANGELOG.md +49 -0
  2. package/README.md +140 -4
  3. package/benchmarks/compare.mjs +651 -0
  4. package/benchmarks/leak.mjs +255 -0
  5. package/benchmarks/results/before-spatial.json +391 -0
  6. package/benchmarks/results/before-spatial.md +47 -0
  7. package/benchmarks/results/int8-scan.json +2007 -0
  8. package/benchmarks/results/int8-scan.md +160 -0
  9. package/benchmarks/results/int8-w8a8-scan.json +2007 -0
  10. package/benchmarks/results/int8-w8a8-scan.md +160 -0
  11. package/benchmarks/results/int8-weight-channel.json +1413 -0
  12. package/benchmarks/results/int8-weight-channel.md +118 -0
  13. package/benchmarks/results/int8-weight-only.json +1437 -0
  14. package/benchmarks/results/int8-weight-only.md +118 -0
  15. package/benchmarks/results/kernel-webnn-final.json +1115 -0
  16. package/benchmarks/results/kernel-webnn-final.md +104 -0
  17. package/benchmarks/results/latest-optimized.json +375 -0
  18. package/benchmarks/results/latest-optimized.md +47 -0
  19. package/benchmarks/results/latest.json +391 -0
  20. package/benchmarks/results/latest.md +47 -0
  21. package/benchmarks/results/profile-baseline.json +331 -0
  22. package/benchmarks/results/profile-baseline.md +12 -0
  23. package/benchmarks/results/profile-conv2x.json +331 -0
  24. package/benchmarks/results/profile-conv2x.md +12 -0
  25. package/benchmarks/results/profile-fast-init.json +385 -0
  26. package/benchmarks/results/profile-fast-init.md +47 -0
  27. package/benchmarks/results/profile-fp16-fma.json +369 -0
  28. package/benchmarks/results/profile-fp16-fma.md +47 -0
  29. package/benchmarks/results/profile-fp16-tiled-decoder.json +369 -0
  30. package/benchmarks/results/profile-fp16-tiled-decoder.md +47 -0
  31. package/benchmarks/results/profile-fp16-tiled-encoder.json +369 -0
  32. package/benchmarks/results/profile-fp16-tiled-encoder.md +47 -0
  33. package/benchmarks/results/profile-fp16-unfused-pool.json +385 -0
  34. package/benchmarks/results/profile-fp16-unfused-pool.md +47 -0
  35. package/benchmarks/results/profile-input-major.json +347 -0
  36. package/benchmarks/results/profile-input-major.md +12 -0
  37. package/benchmarks/results/profile-k16.json +347 -0
  38. package/benchmarks/results/profile-k16.md +12 -0
  39. package/benchmarks/results/profile-k4.json +347 -0
  40. package/benchmarks/results/profile-k4.md +12 -0
  41. package/benchmarks/results/profile-pool-reuse.json +331 -0
  42. package/benchmarks/results/profile-pool-reuse.md +12 -0
  43. package/benchmarks/results/profile-precompiled.json +385 -0
  44. package/benchmarks/results/profile-precompiled.md +47 -0
  45. package/benchmarks/results/profile-static-channels.json +385 -0
  46. package/benchmarks/results/profile-static-channels.md +47 -0
  47. package/benchmarks/results/profile-static-io.json +385 -0
  48. package/benchmarks/results/profile-static-io.md +47 -0
  49. package/benchmarks/results/profile-tiled-conv.json +331 -0
  50. package/benchmarks/results/profile-tiled-conv.md +12 -0
  51. package/benchmarks/results/profile-tiled-decoder.json +347 -0
  52. package/benchmarks/results/profile-tiled-decoder.md +12 -0
  53. package/benchmarks/results/profile-tiled-matmul.json +331 -0
  54. package/benchmarks/results/profile-tiled-matmul.md +12 -0
  55. package/benchmarks/results/profile-unfused-decoder.json +379 -0
  56. package/benchmarks/results/profile-unfused-decoder.md +12 -0
  57. package/benchmarks/results/profile-unfused-pool.json +347 -0
  58. package/benchmarks/results/profile-unfused-pool.md +12 -0
  59. package/benchmarks/results/spatial-auto.json +575 -0
  60. package/benchmarks/results/spatial-auto.md +61 -0
  61. package/benchmarks/results/subgroup-smoke.json +1094 -0
  62. package/benchmarks/results/subgroup-smoke.md +104 -0
  63. package/benchmarks/results/webnn-smoke.json +739 -0
  64. package/benchmarks/results/webnn-smoke.md +76 -0
  65. package/dist/oidn.js +4189 -22603
  66. package/dist/oidn.umd.cjs +784 -5807
  67. package/lib/UNet.d.ts +66 -15
  68. package/lib/UNet.js +162 -257
  69. package/lib/UNet.js.map +1 -1
  70. package/lib/WGPUComputePass.d.ts +1 -1
  71. package/lib/WGPUComputePass.js +6 -4
  72. package/lib/WGPUComputePass.js.map +1 -1
  73. package/lib/backend.d.ts +8 -4
  74. package/lib/backend.js +36 -44
  75. package/lib/backend.js.map +1 -1
  76. package/lib/graphOptimizer.d.ts +54 -0
  77. package/lib/graphOptimizer.js +216 -0
  78. package/lib/graphOptimizer.js.map +1 -0
  79. package/lib/main.d.ts +33 -10
  80. package/lib/main.js +5 -0
  81. package/lib/main.js.map +1 -1
  82. package/lib/modelSpec.d.ts +80 -0
  83. package/lib/modelSpec.js +270 -0
  84. package/lib/modelSpec.js.map +1 -0
  85. package/lib/nativeUNet.d.ts +67 -0
  86. package/lib/nativeUNet.js +1735 -0
  87. package/lib/nativeUNet.js.map +1 -0
  88. package/lib/process.js +38 -35
  89. package/lib/process.js.map +1 -1
  90. package/lib/resourceTracker.d.ts +26 -0
  91. package/lib/resourceTracker.js +65 -0
  92. package/lib/resourceTracker.js.map +1 -0
  93. package/lib/tileScheduler.d.ts +33 -0
  94. package/lib/tileScheduler.js +86 -0
  95. package/lib/tileScheduler.js.map +1 -0
  96. package/lib/webnnUNet.d.ts +52 -0
  97. package/lib/webnnUNet.js +535 -0
  98. package/lib/webnnUNet.js.map +1 -0
  99. package/package.json +9 -5
  100. package/scripts/inspect-model.mjs +64 -0
  101. package/src/UNet.ts +236 -339
  102. package/src/WGPUComputePass.ts +6 -4
  103. package/src/backend.ts +42 -55
  104. package/src/graphOptimizer.ts +301 -0
  105. package/src/main.ts +71 -11
  106. package/src/modelSpec.ts +414 -0
  107. package/src/nativeUNet.ts +2256 -0
  108. package/src/process.ts +38 -36
  109. package/src/resourceTracker.ts +94 -0
  110. package/src/tileScheduler.ts +138 -0
  111. package/src/webnnUNet.ts +812 -0
  112. package/tests/modelSpec.test.mjs +128 -0
  113. package/tests/resourceLifecycle.test.mjs +383 -0
  114. package/tests/tileScheduler.test.mjs +90 -0
  115. package/lib/helper.d.ts +0 -4
  116. package/lib/helper.js +0 -33
  117. package/lib/helper.js.map +0 -1
  118. package/lib/kernels.d.ts +0 -1
  119. package/lib/kernels.js +0 -26
  120. package/lib/kernels.js.map +0 -1
  121. package/src/helper.ts +0 -43
  122. package/src/kernels.ts +0 -31
@@ -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
- Object.keys(this._outputBuffers).forEach((key) => {
160
- (this._outputBuffers as any)[key].buffer.destroy();
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
- try {
8
- const gpuDescriptor: GPURequestAdapterOptions = {
9
- powerPreference: 'high-performance'
10
- };
11
-
12
- const adapter = (await navigator.gpu.requestAdapter(gpuDescriptor))!;
13
- const deviceDescriptor: GPUDeviceDescriptor = {};
14
-
15
- const requiredFeatures = [];
16
- if (adapter.features.has('timestamp-query')) {
17
- requiredFeatures.push('timestamp-query');
18
- }
19
- if (adapter.features.has('bgra8unorm-storage')) {
20
- requiredFeatures.push(['bgra8unorm-storage']);
21
- }
22
- deviceDescriptor.requiredFeatures =
23
- requiredFeatures as Iterable<GPUFeatureName>;
24
-
25
- const adapterLimits = adapter.limits;
26
- deviceDescriptor.requiredLimits = {
27
- maxComputeWorkgroupStorageSize:
28
- adapterLimits.maxComputeWorkgroupStorageSize,
29
- maxComputeWorkgroupsPerDimension:
30
- adapterLimits.maxComputeWorkgroupsPerDimension,
31
- maxStorageBufferBindingSize: adapterLimits.maxStorageBufferBindingSize,
32
- maxBufferSize: adapterLimits.maxBufferSize,
33
- maxComputeWorkgroupSizeX: adapterLimits.maxComputeWorkgroupSizeX,
34
- maxComputeInvocationsPerWorkgroup:
35
- adapterLimits.maxComputeInvocationsPerWorkgroup
36
- };
37
- const device = await adapter.requestDevice(deviceDescriptor);
38
- const adapterInfo =
39
- // requestAdapterInfo is deprecated
40
- // @ts-ignore
41
- adapter!.info ?? (await adapter!.requestAdapterInfo?.());
42
-
43
- return initWebGPUBackendWithDevice(device, adapterInfo);
44
- } catch (e) {}
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
- adapter: GPUAdapterInfo
46
+ adapterInfo: GPUAdapterInfo
50
47
  ) {
51
- // TODO multiple device and adapter in one backend
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!, opts);
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())