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
package/src/UNet.ts CHANGED
@@ -1,23 +1,4 @@
1
- // import * as tfjs from '@tensorflow/tfjs-core';
2
- import { Tensor, Tensor1D, Tensor4D } from '@tensorflow/tfjs-core';
3
- import type { SymbolicTensor } from '@tensorflow/tfjs-layers';
4
- import { tensor } from '@tensorflow/tfjs-core/dist/ops/tensor';
5
- import { tensor1d } from '@tensorflow/tfjs-core/dist/ops/tensor1d';
6
- import { mirrorPad } from '@tensorflow/tfjs-core/dist/ops/mirror_pad';
7
- import { pad4d } from '@tensorflow/tfjs-core/dist/ops/pad4d';
8
- import { slice4d } from '@tensorflow/tfjs-core/dist/ops/slice4d';
9
- import { concat4d } from '@tensorflow/tfjs-core/dist/ops/concat_4d';
10
- import { ENGINE } from '@tensorflow/tfjs-core/dist/engine';
11
- import {
12
- Conv2D,
13
- UpSampling2D
14
- } from '@tensorflow/tfjs-layers/dist/layers/convolutional';
15
- import { MaxPooling2D } from '@tensorflow/tfjs-layers/dist/layers/pooling';
16
- import { Concatenate } from '@tensorflow/tfjs-layers/dist/layers/merge';
17
- import { LayersModel } from '@tensorflow/tfjs-layers/dist/engine/training';
18
- import { Input as TFInput } from '@tensorflow/tfjs-layers/dist/engine/input_layer';
19
1
  import { HostTensor } from './tza';
20
- import { Float16Array } from '@petamoriken/float16';
21
2
  import {
22
3
  GPUDataProcess,
23
4
  Tile,
@@ -25,43 +6,26 @@ import {
25
6
  hdrTransferFuncCPU,
26
7
  hdrTransferFuncInverseCPU
27
8
  } from './process';
28
- import type { WebGPUBackend } from '@tensorflow/tfjs-backend-webgpu';
29
-
30
- // import { profileAndLogKernelCode, memory } from './helper';
31
-
32
- function getTensorData(
33
- ubytes: Uint8Array,
34
- type: HostTensor['desc']['dataType']
35
- ) {
36
- const buffer = ubytes.buffer;
37
- if (type === 'Float32') {
38
- return new Float32Array(ubytes.buffer);
39
- }
40
- const float16Data = new Float16Array(buffer);
41
- const float32Data = new Float32Array(float16Data.length);
42
- for (let i = 0; i < float32Data.length; ++i) {
43
- float32Data[i] = float16Data[i];
44
- }
45
- return float32Data;
46
- }
9
+ import {
10
+ DynamicTileController,
11
+ type DynamicTileSetting,
12
+ fitTileDimension,
13
+ OIDN_TILE_ALIGNMENT,
14
+ waitForSubmittedGPUWork
15
+ } from './tileScheduler';
16
+ import {
17
+ detectUNetModelSpec,
18
+ validateUNetModel,
19
+ type UNetModelSpec
20
+ } from './modelSpec';
21
+ import {
22
+ NativeUNetExecutor,
23
+ type NativeUNetKernelSetting,
24
+ type NativeUNetPrecisionSetting
25
+ } from './nativeUNet';
26
+ import { WebNNUNetExecutor } from './webnnUNet';
47
27
 
48
- function changeWeightShapes(weightData: Float32Array, dims: number[]) {
49
- const [O, C, H, W] = dims;
50
- const reorderedWeightData = new Float32Array(weightData.length);
51
- for (let o = 0; o < O; ++o) {
52
- for (let c = 0; c < C; ++c) {
53
- for (let h = 0; h < H; ++h) {
54
- for (let w = 0; w < W; ++w) {
55
- // Change OCHW to HWCO
56
- const idx = o * C * H * W + c * H * W + h * W + w;
57
- const idx2 = h * W * C * O + w * C * O + c * O + o;
58
- reorderedWeightData[idx2] = weightData[idx];
59
- }
60
- }
61
- }
62
- }
63
- return reorderedWeightData;
64
- }
28
+ export type UNetEngineSetting = 'auto' | 'wgsl' | 'webnn';
65
29
 
66
30
  interface HDRImageData {
67
31
  data: Float32Array;
@@ -81,13 +45,24 @@ interface GPUImageDataOutput {
81
45
  height: number;
82
46
  }
83
47
 
48
+ export interface UNetExecutionStats {
49
+ width: number;
50
+ height: number;
51
+ tileWidth: number;
52
+ tileHeight: number;
53
+ tileCount: number;
54
+ durationMs: number;
55
+ tileTimeMs: {
56
+ min: number;
57
+ median: number;
58
+ mean: number;
59
+ max: number;
60
+ };
61
+ }
62
+
84
63
  function roundUp(a: number, b: number) {
85
64
  return Math.ceil(a / b) * b;
86
65
  }
87
- // Returns the smallest integer larger than or equal to a which has remainder c when divided by b
88
- function roundUp2(a: number, b: number, c: number) {
89
- return Math.ceil((a - c) / b) * b + c;
90
- }
91
66
 
92
67
  function isGPUImageData(
93
68
  data: ImageData | GPUImageData | HDRImageData
@@ -95,18 +70,7 @@ function isGPUImageData(
95
70
  return data.data instanceof GPUBuffer || data.data instanceof GPUTexture;
96
71
  }
97
72
 
98
- const receptiveField = 174; // receptive field in pixels
99
- const receptiveFieldLarge = 202;
100
- // TODO metal is 32?
101
- const minTileAlignment = 1;
102
-
103
- const tileAlignment = 16; // required spatial alignment in pixels (padding may be necessary)
104
-
105
- const defaultTileOverlap = roundUp(receptiveField / 2, tileAlignment);
106
- const defaultTileOverlapLarge = roundUp(receptiveFieldLarge / 2, tileAlignment);
107
-
108
73
  class UNet {
109
- private _tfModel: LayersModel | undefined;
110
74
  private _device: GPUDevice | undefined;
111
75
 
112
76
  // TODO calculate the tile size from memory size
@@ -121,15 +85,18 @@ class UNet {
121
85
  private _hdr;
122
86
 
123
87
  private _dataProcessGPU?: GPUDataProcess;
88
+ private _nativeExecutor?: NativeUNetExecutor;
89
+ private _webNNExecutor?: WebNNUNetExecutor;
90
+ private _modelSpec: UNetModelSpec;
91
+ private _inputChannels: number;
92
+ private _engine: UNetEngineSetting;
124
93
 
125
- private _maxTileSize;
126
-
127
- private _tensors = new Map<string, Tensor>();
128
- private _modelsCache = new Map<string, LayersModel>();
94
+ private _dynamicTileController: DynamicTileController;
95
+ private _lastExecution?: UNetExecutionStats;
129
96
 
130
97
  constructor(
131
- private _hostTensors: Map<string, HostTensor>,
132
- private _backend: WebGPUBackend,
98
+ hostTensors: Map<string, HostTensor>,
99
+ backend: { device: GPUDevice; adapterInfo: GPUAdapterInfo },
133
100
  opts: {
134
101
  /**
135
102
  * If use auxiliary data.
@@ -140,228 +107,137 @@ class UNet {
140
107
  */
141
108
  hdr?: boolean;
142
109
  maxTileSize?: number;
110
+ dynamicTile?: DynamicTileSetting;
111
+ /** Reserved for explicit native WGSL selection. */
112
+ engine?: UNetEngineSetting;
113
+ /** Arithmetic/storage precision used by the native WGSL engine. */
114
+ precision?: NativeUNetPrecisionSetting;
115
+ /** Model-independent convolution kernel selection. */
116
+ kernel?: NativeUNetKernelSetting;
117
+ /** Explicit descriptor for a new OIDN topology not in the built-in registry. */
118
+ modelSpec?: UNetModelSpec;
143
119
  } = {}
144
120
  ) {
145
121
  this._aux = opts.aux || false;
146
122
  this._hdr = opts.hdr || false;
123
+ this._engine = opts.engine ?? 'auto';
124
+ const modelSpec = opts.modelSpec ?? detectUNetModelSpec(hostTensors);
125
+ const validatedModel = validateUNetModel(hostTensors, modelSpec);
126
+ this._modelSpec = validatedModel.spec;
127
+ this._inputChannels = validatedModel.inputChannels;
128
+
129
+ const expectedInputChannels = this._aux ? 9 : 3;
130
+ if (validatedModel.inputChannels !== expectedInputChannels) {
131
+ throw new Error(
132
+ `OIDN model expects ${validatedModel.inputChannels} input channels, ` +
133
+ `but aux=${this._aux} provides ${expectedInputChannels}`
134
+ );
135
+ }
147
136
 
148
- this._maxTileSize = roundUp(opts.maxTileSize ?? 512, 2);
137
+ this._dynamicTileController = new DynamicTileController(
138
+ opts.maxTileSize ?? 512,
139
+ opts.dynamicTile
140
+ );
149
141
 
150
- this._device = this._backend.device;
142
+ this._device = backend.device;
143
+ if (this._engine === 'webnn') {
144
+ this._webNNExecutor = new WebNNUNetExecutor(
145
+ this._device,
146
+ validatedModel,
147
+ { precision: opts.precision }
148
+ );
149
+ } else {
150
+ this._nativeExecutor = new NativeUNetExecutor(
151
+ this._device,
152
+ validatedModel,
153
+ { precision: opts.precision, kernel: opts.kernel }
154
+ );
155
+ }
151
156
  }
152
157
 
153
158
  getDevice() {
154
159
  return this._device;
155
160
  }
156
161
 
157
- private _buildModel(isLarge: boolean) {
158
- const aux = this._aux;
159
- const channels = 3 + (aux ? 6 : 0);
160
- const tileSize = this._getTileSizeWithOverlap();
161
- const cache = this._modelsCache;
162
- const key = [tileSize.width, tileSize.height].join(',');
163
-
164
- // We cache the model instead of disposing and recreate.
165
- // Because seems tfjs will also cache the layer and gpubuffers.
166
- // Recreating the model will cause memory leak.
167
-
168
- // Width and height can only be 256, 512, 768. So the cache won't be too large
169
- if (cache.has(key)) {
170
- this._tfModel = cache.get(key);
171
- return;
172
- }
173
-
174
- const input = TFInput({
175
- name: 'input',
176
- shape: [tileSize.height, tileSize.width, channels],
177
- dtype: 'float32'
178
- });
179
-
180
- this._tfModel = new LayersModel({
181
- inputs: [input],
182
- outputs: isLarge ? this._addNetLarge(input) : this._addNet(input)
183
- });
184
- cache.set(key, this._tfModel);
185
- }
186
-
187
- private _createConv(
188
- name: string,
189
- source: SymbolicTensor,
190
- activation?: 'relu'
191
- ) {
192
- const weightTensorName = name + '.weight';
193
- const biasTensorName = name + '.bias';
194
- const tensors = this._tensors;
195
- let weightTensor = tensors.get(weightTensorName);
196
- let biasTensor = tensors.get(biasTensorName);
197
- const unetWeightTensor = this._hostTensors.get(weightTensorName)!;
198
-
199
- if (!weightTensor) {
200
- const weightDims = unetWeightTensor.desc.dims;
201
- weightTensor = tensor(
202
- changeWeightShapes(
203
- getTensorData(unetWeightTensor.data, unetWeightTensor.desc.dataType),
204
- weightDims
205
- ),
206
- [weightDims[2], weightDims[3], weightDims[1], weightDims[0]],
207
- 'float32'
162
+ /** Completes backend compilation before first interactive use. */
163
+ async prepare() {
164
+ if (this._webNNExecutor) {
165
+ await this._webNNExecutor.prepare();
166
+ const overlap = roundUp(
167
+ this._modelSpec.receptiveField / 2,
168
+ OIDN_TILE_ALIGNMENT
208
169
  );
209
-
210
- tensors.set(weightTensorName, weightTensor);
211
- }
212
- if (!biasTensor) {
213
- const unetBiasTensor = this._hostTensors.get(name + '.bias')!;
214
- biasTensor = tensor1d(
215
- getTensorData(unetBiasTensor.data, unetBiasTensor.desc.dataType),
216
- 'float32'
170
+ const outputTileEdges = [
171
+ this._dynamicTileController.tileSize,
172
+ this._dynamicTileController.minTileSize
173
+ ];
174
+ await this._webNNExecutor.prewarm(
175
+ [...new Set(outputTileEdges)].map((edge) => ({
176
+ width: edge + 2 * overlap,
177
+ height: edge + 2 * overlap
178
+ }))
217
179
  );
218
- tensors.set(biasTensorName, biasTensor);
180
+ return;
219
181
  }
220
- // TODO whats the purpose of padded dims ?
221
- const convLayer = new Conv2D({
222
- name,
223
- filters: unetWeightTensor.desc.dims[0],
224
- kernelSize: unetWeightTensor.desc.dims.slice(2, 4) as [number, number],
225
- useBias: true,
226
- activation,
227
- padding: 'same',
228
- weights: [weightTensor, biasTensor],
229
- trainable: false
230
- });
231
-
232
- return convLayer.apply(source) as SymbolicTensor;
182
+ await this._nativeExecutor!.prepare();
233
183
  }
234
184
 
235
- private _createConcatConv(
236
- name: string,
237
- source1: SymbolicTensor,
238
- source2: SymbolicTensor
239
- ) {
240
- const concatLayer = new Concatenate({
241
- name: name + '/concat',
242
- trainable: false,
243
- axis: 3
244
- });
245
- //https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L40
246
- return this._createConv(
247
- name,
248
- // Concat on the channel
249
- concatLayer.apply([source1, source2]) as SymbolicTensor,
250
- 'relu'
251
- ) as SymbolicTensor;
185
+ getRuntimeInfo() {
186
+ return {
187
+ configuredEngine: this._engine,
188
+ gpuEngine: this._webNNExecutor ? 'webnn' as const : 'wgsl' as const,
189
+ precision: (this._webNNExecutor ?? this._nativeExecutor!).precision,
190
+ kernel: this._nativeExecutor
191
+ ? {
192
+ configured: this._nativeExecutor.kernelSetting,
193
+ maxSpatialInputBlocks: this._nativeExecutor.maxSpatialInputBlocks,
194
+ subgroupsAvailable: this._nativeExecutor.subgroupsAvailable
195
+ }
196
+ : undefined,
197
+ webnn: this._webNNExecutor?.support,
198
+ resources: (
199
+ this._webNNExecutor ?? this._nativeExecutor!
200
+ ).getResourceInfo(),
201
+ model: this._modelSpec.id,
202
+ modelFamily: this._modelSpec.family,
203
+ inputChannels: this._inputChannels,
204
+ dynamicTile: {
205
+ enabled: this._dynamicTileController.enabled,
206
+ currentTileSize: this._dynamicTileController.tileSize,
207
+ minTileSize: this._dynamicTileController.minTileSize,
208
+ maxTileSize: this._dynamicTileController.maxTileSize,
209
+ targetTileTimeMs: this._dynamicTileController.targetTileTimeMs
210
+ },
211
+ lastExecution: this._lastExecution
212
+ };
252
213
  }
253
214
 
254
- private _createPooling(source: SymbolicTensor) {
255
- const poolingLayer = new MaxPooling2D({
256
- name: source.name + '/pooling',
257
- poolSize: [2, 2],
258
- strides: [2, 2],
259
- padding: 'same',
260
- trainable: false
261
- });
262
- // https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L33
263
- return poolingLayer.apply(source) as SymbolicTensor;
215
+ /** Captures per-node GPU timestamps for the next native tile execution. */
216
+ profileNextExecution() {
217
+ return this._nativeExecutor?.profileNextExecution() ?? false;
264
218
  }
265
219
 
266
- private _addUpsamplingLayer(source: SymbolicTensor) {
267
- const upsamplingLayer = new UpSampling2D({
268
- name: source.name + '/upsampling',
269
- size: [2, 2],
270
- trainable: false
271
- });
272
- return upsamplingLayer.apply(source) as SymbolicTensor;
220
+ getLastExecutionProfile() {
221
+ return this._nativeExecutor?.getLastExecutionProfile();
273
222
  }
274
223
 
275
- private _addNet(input: SymbolicTensor) {
276
- let x = this._createConv('enc_conv0', input, 'relu');
277
- const pool1 = (x = this._createPooling(
278
- this._createConv('enc_conv1', x, 'relu')
279
- ));
280
- const pool2 = (x = this._createPooling(
281
- this._createConv('enc_conv2', x, 'relu')
282
- ));
283
- const pool3 = (x = this._createPooling(
284
- this._createConv('enc_conv3', x, 'relu')
285
- ));
286
- const pool4 = (x = this._createPooling(
287
- this._createConv('enc_conv4', x, 'relu')
288
- ));
289
- x = this._createConv('enc_conv5a', pool4, 'relu');
290
- x = this._addUpsamplingLayer(this._createConv('enc_conv5b', x, 'relu'));
291
-
292
- x = this._createConcatConv('dec_conv4a', x, pool3);
293
- x = this._addUpsamplingLayer(this._createConv('dec_conv4b', x, 'relu'));
294
-
295
- x = this._createConcatConv('dec_conv3a', x, pool2);
296
- x = this._addUpsamplingLayer(this._createConv('dec_conv3b', x, 'relu'));
297
-
298
- x = this._createConcatConv('dec_conv2a', x, pool1);
299
- x = this._addUpsamplingLayer(this._createConv('dec_conv2b', x, 'relu'));
300
-
301
- x = this._createConcatConv('dec_conv1a', x, input);
302
- x = this._createConv('dec_conv1b', x, 'relu');
303
- x = this._createConv('dec_conv0', x, 'relu');
304
-
305
- return x;
306
- }
224
+ private _updateModel(width: number, height: number) {
225
+ const maxTileSize = this._dynamicTileController.tileSize;
307
226
 
308
- private _addNetLarge(input: SymbolicTensor) {
309
- let x = this._createConv('enc_conv1a', input, 'relu');
310
- const pool1 = (x = this._createPooling(
311
- this._createConv('enc_conv1b', x, 'relu')
312
- ));
313
- x = this._createConv('enc_conv2a', x, 'relu');
314
- const pool2 = (x = this._createPooling(
315
- this._createConv('enc_conv2b', x, 'relu')
316
- ));
317
- x = this._createConv('enc_conv3a', x, 'relu');
318
- const pool3 = (x = this._createPooling(
319
- this._createConv('enc_conv3b', x, 'relu')
320
- ));
321
- x = this._createConv('enc_conv4a', x, 'relu');
322
- const pool4 = (x = this._createPooling(
323
- this._createConv('enc_conv4b', x, 'relu')
324
- ));
325
-
326
- x = this._createConv('enc_conv5a', pool4, 'relu');
327
- x = this._addUpsamplingLayer(this._createConv('enc_conv5b', x, 'relu'));
328
-
329
- x = this._createConcatConv('dec_conv4a', x, pool3);
330
- x = this._addUpsamplingLayer(this._createConv('dec_conv4b', x, 'relu'));
331
-
332
- x = this._createConcatConv('dec_conv3a', x, pool2);
333
- x = this._addUpsamplingLayer(this._createConv('dec_conv3b', x, 'relu'));
334
-
335
- x = this._createConcatConv('dec_conv2a', x, pool1);
336
- x = this._addUpsamplingLayer(this._createConv('dec_conv2b', x, 'relu'));
337
-
338
- x = this._createConcatConv('dec_conv1a', x, input);
339
- x = this._createConv('dec_conv1b', x, 'relu');
340
- x = this._createConv('dec_conv1c', x, 'relu');
341
-
342
- return x;
343
- }
227
+ let tileWidth = fitTileDimension(width, maxTileSize);
228
+ let tileHeight = fitTileDimension(height, maxTileSize);
229
+ const defaultTileOverlap = roundUp(
230
+ this._modelSpec.receptiveField / 2,
231
+ OIDN_TILE_ALIGNMENT
232
+ );
233
+ let tileOverlapX = defaultTileOverlap;
234
+ let tileOverlapY = defaultTileOverlap;
344
235
 
345
- private _updateModel(width: number, height: number) {
346
- const isLarge = this._hostTensors.has('enc_conv1b.weight');
347
- const maxTileSize = this._maxTileSize;
348
-
349
- let tileWidth = maxTileSize;
350
- let tileHeight = maxTileSize;
351
- let tileOverlapX = isLarge ? defaultTileOverlapLarge : defaultTileOverlap;
352
- let tileOverlapY = isLarge ? defaultTileOverlapLarge : defaultTileOverlap;
353
-
354
- if (width < maxTileSize + defaultTileOverlap * 2) {
355
- tileWidth = roundUp(width, maxTileSize / 2);
356
- if (width <= maxTileSize) {
357
- tileOverlapX = 0;
358
- }
236
+ if (width <= maxTileSize) {
237
+ tileOverlapX = 0;
359
238
  }
360
- if (height < maxTileSize + defaultTileOverlap * 2) {
361
- tileHeight = roundUp(height, maxTileSize / 2);
362
- if (height <= maxTileSize) {
363
- tileOverlapY = 0;
364
- }
239
+ if (height <= maxTileSize) {
240
+ tileOverlapY = 0;
365
241
  }
366
242
 
367
243
  // Force width and height has same size. reduce the cache in memory
@@ -376,16 +252,13 @@ class UNet {
376
252
  tileWidth !== this._tileWidth ||
377
253
  tileHeight !== this._tileHeight ||
378
254
  tileOverlapX !== this._tileOverlapX ||
379
- tileOverlapY !== this._tileOverlapY ||
380
- !this._tfModel
255
+ tileOverlapY !== this._tileOverlapY
381
256
  ) {
382
257
  // console.log(tileWidth, tileHeight, tileOverlapX, tileOverlapY);
383
258
  this._tileWidth = tileWidth;
384
259
  this._tileHeight = tileHeight;
385
260
  this._tileOverlapX = tileOverlapX;
386
261
  this._tileOverlapY = tileOverlapY;
387
-
388
- this._buildModel(isLarge);
389
262
  }
390
263
  }
391
264
 
@@ -497,7 +370,7 @@ class UNet {
497
370
  }
498
371
  }
499
372
 
500
- private _executeTile(
373
+ private async _executeTile(
501
374
  inputData:
502
375
  | Float32Array
503
376
  | {
@@ -534,7 +407,8 @@ class UNet {
534
407
 
535
408
  const srcTile = new Tile(srcX0, srcY0, srcTileWidth, srcTileHeight);
536
409
 
537
- let tileTensor!: Tensor;
410
+ let nativeOutputBuffer: GPUBuffer | undefined;
411
+ let denoisedData: Float32Array | undefined;
538
412
  let inputScale = 1;
539
413
  const device = this._device!;
540
414
  let dataProcessGPU = this._dataProcessGPU;
@@ -552,11 +426,11 @@ class UNet {
552
426
  inputScale
553
427
  });
554
428
  }
555
- tileTensor = tensor(
429
+ denoisedData = await (this._webNNExecutor ?? this._nativeExecutor!).executeCPU(
556
430
  tileData,
557
- [1, srcTileHeight, srcTileWidth, channels],
558
- 'float32'
559
- ) as Tensor4D;
431
+ srcTileWidth,
432
+ srcTileHeight
433
+ );
560
434
  } else {
561
435
  if (!dataProcessGPU) {
562
436
  dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(
@@ -577,33 +451,14 @@ class UNet {
577
451
  denoiseAlpha
578
452
  );
579
453
 
580
- const createTensor = (buffer: GPUBuffer) => {
581
- const tmp = tensor({ buffer, zeroCopy: true }, [
582
- 1,
583
- srcTileHeight,
584
- srcTileWidth,
585
- 4
586
- ]) as Tensor4D;
587
- const ret = slice4d(
588
- tmp,
589
- [0, 0, 0, 0],
590
- [1, srcTileHeight, srcTileWidth, 3]
591
- );
592
- return ret;
593
- };
594
-
595
- if (this._aux) {
596
- const tensors = [color, albedo, normal].map((buffer) =>
597
- createTensor(buffer!)
598
- );
599
- tileTensor = concat4d(tensors, 3);
600
- } else {
601
- tileTensor = createTensor(color);
602
- }
454
+ nativeOutputBuffer = await (this._webNNExecutor ?? this._nativeExecutor!).execute(
455
+ this._aux ? [color, albedo!, normal!] : [color],
456
+ srcTileWidth,
457
+ srcTileHeight
458
+ );
603
459
  }
604
460
 
605
461
  let outBuffer: GPUBuffer;
606
- const outputTensor = this._tfModel!.predict(tileTensor) as Tensor;
607
462
 
608
463
  const dstWidth = Math.min(dstTileSize.width, width);
609
464
  const dstHeight = Math.min(dstTileSize.height, height);
@@ -612,10 +467,9 @@ class UNet {
612
467
  dstTile.height = Math.min(dstTile.height, height - dstTile.y);
613
468
 
614
469
  if (inputData instanceof Float32Array) {
615
- let denoisedData = outputTensor.dataSync();
616
470
  if (isHDR) {
617
471
  denoisedData = hdrTransferFuncInverseCPU({
618
- data: denoisedData as Float32Array,
472
+ data: denoisedData!,
619
473
  channels: 3,
620
474
  inputScale
621
475
  });
@@ -625,7 +479,7 @@ class UNet {
625
479
  outputImageData!,
626
480
  srcTile,
627
481
  dstTile,
628
- denoisedData as Float32Array,
482
+ denoisedData!,
629
483
  srcTileSize.width,
630
484
  isHDR
631
485
  );
@@ -641,17 +495,8 @@ class UNet {
641
495
  }
642
496
  } else {
643
497
  dataProcessGPU!.setOutputTile(dstTile, srcTile);
644
- // IMPORTANT
645
- // storage buffer has alignment. that 3 channels still needs 16 bytes data.
646
- // So we need to pad it to 4 channels.
647
- const outputTensor4Channnels = pad4d(outputTensor as Tensor4D, [
648
- [0, 0],
649
- [0, 0],
650
- [0, 0],
651
- [0, 1]
652
- ]);
653
498
  outBuffer = dataProcessGPU!.inverse(
654
- outputTensor4Channnels.dataToGPU().buffer!,
499
+ nativeOutputBuffer!,
655
500
  inputData.color
656
501
  );
657
502
  }
@@ -694,6 +539,9 @@ class UNet {
694
539
 
695
540
  const width = color.width;
696
541
  const height = color.height;
542
+ const adaptiveTileSize = this._dynamicTileController.tileSize;
543
+ const shouldAdaptTileSize =
544
+ width > adaptiveTileSize || height > adaptiveTileSize;
697
545
  this._updateModel(width, height);
698
546
 
699
547
  // TODO should fixed to be hdr when UNet is created.
@@ -733,14 +581,24 @@ class UNet {
733
581
 
734
582
  let aborted = false;
735
583
 
736
- const executeTile = (i: number, j: number) => {
584
+ const now = () =>
585
+ typeof performance === 'undefined' ? Date.now() : performance.now();
586
+ const executionStartTime = now();
587
+ const tileTimesMs: number[] = [];
588
+ const scheduleNextTile = (callback: () => void) => {
589
+ if (typeof requestAnimationFrame === 'undefined') {
590
+ setTimeout(callback, 0);
591
+ } else {
592
+ requestAnimationFrame(callback);
593
+ }
594
+ };
595
+
596
+ const executeTile = async (i: number, j: number) => {
737
597
  if (aborted) {
738
598
  return;
739
599
  }
740
- let resGPUBuffer;
741
- // profileAndLogKernelCode(() => {
742
- ENGINE.startScope();
743
- resGPUBuffer = this._executeTile(
600
+ const tileStartTime = now();
601
+ const resGPUBuffer = await this._executeTile(
744
602
  isGPUImageData(color)
745
603
  ? {
746
604
  color: color.data,
@@ -757,8 +615,7 @@ class UNet {
757
615
  hdr,
758
616
  denoiseAlpha
759
617
  );
760
- ENGINE.endScope();
761
- // }, true);
618
+ if (aborted) return;
762
619
  const output = outputImageData || {
763
620
  data: resGPUBuffer,
764
621
  width,
@@ -773,18 +630,58 @@ class UNet {
773
630
  tileCountW * tileCountH
774
631
  );
775
632
 
776
- if (i + 1 < tileCountW || j + 1 < tileCountH) {
777
- requestAnimationFrame(() => {
778
- if (i + 1 < tileCountW) {
779
- executeTile(i + 1, j);
780
- } else if (j + 1 < tileCountH) {
781
- executeTile(0, j + 1);
633
+ const hasNextTile = i + 1 < tileCountW || j + 1 < tileCountH;
634
+ const continueAfterGPUWork = () => {
635
+ tileTimesMs.push(now() - tileStartTime);
636
+ if (aborted) return;
637
+
638
+ if (hasNextTile) {
639
+ scheduleNextTile(() => {
640
+ if (aborted) return;
641
+ if (i + 1 < tileCountW) {
642
+ executeTile(i + 1, j);
643
+ } else if (j + 1 < tileCountH) {
644
+ executeTile(0, j + 1);
645
+ }
646
+ });
647
+ } else {
648
+ const sortedTileTimes = [...tileTimesMs].sort((a, b) => a - b);
649
+ const middle = Math.floor(sortedTileTimes.length / 2);
650
+ const medianTileTime = sortedTileTimes.length % 2
651
+ ? sortedTileTimes[middle]
652
+ : (sortedTileTimes[middle - 1] + sortedTileTimes[middle]) / 2;
653
+ this._lastExecution = {
654
+ width,
655
+ height,
656
+ tileWidth,
657
+ tileHeight,
658
+ tileCount: tileCountW * tileCountH,
659
+ durationMs: now() - executionStartTime,
660
+ tileTimeMs: {
661
+ min: sortedTileTimes[0],
662
+ median: medianTileTime,
663
+ mean:
664
+ sortedTileTimes.reduce((sum, value) => sum + value, 0) /
665
+ sortedTileTimes.length,
666
+ max: sortedTileTimes[sortedTileTimes.length - 1]
667
+ }
668
+ };
669
+ // Adapt only from complete executions. Cancelled work is commonly
670
+ // contending with interactive rendering and is not representative.
671
+ if (shouldAdaptTileSize) {
672
+ this._dynamicTileController.observe(tileTimesMs);
782
673
  }
783
- });
784
- } else {
785
- // console.log(memory());
786
- done(output as any);
787
- }
674
+ // console.log(memory());
675
+ done(output as any);
676
+ }
677
+ };
678
+
679
+ // requestAnimationFrame only throttles JavaScript submission. Waiting
680
+ // for the queue here keeps at most one OIDN tile in flight, so aborting
681
+ // cannot leave a long tail of already-submitted GPU work.
682
+ void waitForSubmittedGPUWork(this._device!.queue).then(
683
+ continueAfterGPUWork
684
+ );
788
685
  };
789
686
 
790
687
  executeTile(0, 0);
@@ -795,9 +692,9 @@ class UNet {
795
692
  }
796
693
 
797
694
  dispose() {
798
- this._tfModel?.dispose();
799
695
  this._dataProcessGPU?.dispose();
800
- this._tensors.forEach((tensor) => tensor.dispose());
696
+ this._nativeExecutor?.dispose();
697
+ this._webNNExecutor?.dispose();
801
698
  }
802
699
  }
803
700