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.
Files changed (65) hide show
  1. package/CHANGELOG.md +94 -0
  2. package/README.md +208 -8
  3. package/dist/oidn.js +4699 -22516
  4. package/dist/oidn.umd.cjs +989 -5796
  5. package/lib/UNet.d.ts +111 -26
  6. package/lib/UNet.js +310 -329
  7. package/lib/UNet.js.map +1 -1
  8. package/lib/WGPUComputePass.d.ts +1 -1
  9. package/lib/WGPUComputePass.js +6 -4
  10. package/lib/WGPUComputePass.js.map +1 -1
  11. package/lib/backend.d.ts +1 -4
  12. package/lib/backend.js +28 -44
  13. package/lib/backend.js.map +1 -1
  14. package/lib/finalRgbShader.d.ts +13 -0
  15. package/lib/finalRgbShader.js +160 -0
  16. package/lib/finalRgbShader.js.map +1 -0
  17. package/lib/graphOptimizer.d.ts +54 -0
  18. package/lib/graphOptimizer.js +215 -0
  19. package/lib/graphOptimizer.js.map +1 -0
  20. package/lib/hdrTransfer.d.ts +14 -0
  21. package/lib/hdrTransfer.js +61 -0
  22. package/lib/hdrTransfer.js.map +1 -0
  23. package/lib/main.d.ts +43 -11
  24. package/lib/main.js +9 -5
  25. package/lib/main.js.map +1 -1
  26. package/lib/modelSpec.d.ts +80 -0
  27. package/lib/modelSpec.js +270 -0
  28. package/lib/modelSpec.js.map +1 -0
  29. package/lib/nativeUNet.d.ts +103 -0
  30. package/lib/nativeUNet.js +2064 -0
  31. package/lib/nativeUNet.js.map +1 -0
  32. package/lib/process.d.ts +5 -11
  33. package/lib/process.js +38 -49
  34. package/lib/process.js.map +1 -1
  35. package/lib/resourceTracker.d.ts +26 -0
  36. package/lib/resourceTracker.js +65 -0
  37. package/lib/resourceTracker.js.map +1 -0
  38. package/lib/tileScheduler.d.ts +61 -0
  39. package/lib/tileScheduler.js +199 -0
  40. package/lib/tileScheduler.js.map +1 -0
  41. package/lib/webnnUNet.d.ts +52 -0
  42. package/lib/webnnUNet.js +535 -0
  43. package/lib/webnnUNet.js.map +1 -0
  44. package/package.json +16 -5
  45. package/src/UNet.ts +463 -437
  46. package/src/WGPUComputePass.ts +6 -4
  47. package/src/backend.ts +33 -59
  48. package/src/finalRgbShader.ts +186 -0
  49. package/src/graphOptimizer.ts +300 -0
  50. package/src/hdrTransfer.ts +88 -0
  51. package/src/main.ts +95 -20
  52. package/src/modelSpec.ts +414 -0
  53. package/src/nativeUNet.ts +2655 -0
  54. package/src/process.ts +46 -71
  55. package/src/resourceTracker.ts +94 -0
  56. package/src/tileScheduler.ts +330 -0
  57. package/src/webnnUNet.ts +812 -0
  58. package/lib/helper.d.ts +0 -4
  59. package/lib/helper.js +0 -33
  60. package/lib/helper.js.map +0 -1
  61. package/lib/kernels.d.ts +0 -1
  62. package/lib/kernels.js +0 -26
  63. package/lib/kernels.js.map +0 -1
  64. package/src/helper.ts +0 -43
  65. package/src/kernels.ts +0 -31
package/lib/UNet.js CHANGED
@@ -1,261 +1,162 @@
1
- import { tensor } from '@tensorflow/tfjs-core/dist/ops/tensor';
2
- import { tensor1d } from '@tensorflow/tfjs-core/dist/ops/tensor1d';
3
- import { pad4d } from '@tensorflow/tfjs-core/dist/ops/pad4d';
4
- import { slice4d } from '@tensorflow/tfjs-core/dist/ops/slice4d';
5
- import { concat4d } from '@tensorflow/tfjs-core/dist/ops/concat_4d';
6
- import { ENGINE } from '@tensorflow/tfjs-core/dist/engine';
7
- import { Conv2D, UpSampling2D } from '@tensorflow/tfjs-layers/dist/layers/convolutional';
8
- import { MaxPooling2D } from '@tensorflow/tfjs-layers/dist/layers/pooling';
9
- import { Concatenate } from '@tensorflow/tfjs-layers/dist/layers/merge';
10
- import { LayersModel } from '@tensorflow/tfjs-layers/dist/engine/training';
11
- import { Input as TFInput } from '@tensorflow/tfjs-layers/dist/engine/input_layer';
12
- import { Float16Array } from '@petamoriken/float16';
13
1
  import { GPUDataProcess, Tile, avgLogLum, hdrTransferFuncCPU, hdrTransferFuncInverseCPU } from './process';
14
- // import { profileAndLogKernelCode, memory } from './helper';
15
- function getTensorData(ubytes, type) {
16
- const buffer = ubytes.buffer;
17
- if (type === 'Float32') {
18
- return new Float32Array(ubytes.buffer);
19
- }
20
- const float16Data = new Float16Array(buffer);
21
- const float32Data = new Float32Array(float16Data.length);
22
- for (let i = 0; i < float32Data.length; ++i) {
23
- float32Data[i] = float16Data[i];
24
- }
25
- return float32Data;
26
- }
27
- function changeWeightShapes(weightData, dims) {
28
- const [O, C, H, W] = dims;
29
- const reorderedWeightData = new Float32Array(weightData.length);
30
- for (let o = 0; o < O; ++o) {
31
- for (let c = 0; c < C; ++c) {
32
- for (let h = 0; h < H; ++h) {
33
- for (let w = 0; w < W; ++w) {
34
- // Change OCHW to HWCO
35
- const idx = o * C * H * W + c * H * W + h * W + w;
36
- const idx2 = h * W * C * O + w * C * O + c * O + o;
37
- reorderedWeightData[idx2] = weightData[idx];
38
- }
39
- }
40
- }
41
- }
42
- return reorderedWeightData;
43
- }
2
+ import { DynamicTileController, planTileGrid, OIDN_TILE_ALIGNMENT } from './tileScheduler';
3
+ import { detectUNetModelSpec, validateUNetModel } from './modelSpec';
4
+ import { NativeUNetExecutor } from './nativeUNet';
5
+ import { WebNNUNetExecutor } from './webnnUNet';
6
+ /** Upper bound on waiting for a display frame between tiles. */
7
+ const ANIMATION_FRAME_FALLBACK_MS = 100;
44
8
  function roundUp(a, b) {
45
9
  return Math.ceil(a / b) * b;
46
10
  }
47
- // Returns the smallest integer larger than or equal to a which has remainder c when divided by b
48
- function roundUp2(a, b, c) {
49
- return Math.ceil((a - c) / b) * b + c;
50
- }
51
11
  function isGPUImageData(data) {
52
12
  return data.data instanceof GPUBuffer || data.data instanceof GPUTexture;
53
13
  }
54
- const receptiveField = 174; // receptive field in pixels
55
- const receptiveFieldLarge = 202;
56
- // TODO metal is 32?
57
- const minTileAlignment = 1;
58
- const tileAlignment = 16; // required spatial alignment in pixels (padding may be necessary)
59
- const defaultTileOverlap = roundUp(receptiveField / 2, tileAlignment);
60
- const defaultTileOverlapLarge = roundUp(receptiveFieldLarge / 2, tileAlignment);
61
14
  class UNet {
62
- _hostTensors;
63
- _backend;
64
- _tfModel;
65
15
  _device;
66
- // TODO calculate the tile size from memory size
67
- // https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/core/unet_filter.cpp#L287
68
- _tileWidth = 0;
69
- _tileHeight = 0;
70
- _tileOverlapX = 0;
71
- _tileOverlapY = 0;
72
16
  _aux;
73
17
  _hdr;
18
+ _hdrTransfer;
74
19
  _dataProcessGPU;
75
- _maxTileSize;
76
- _tensors = new Map();
77
- _modelsCache = new Map();
78
- constructor(_hostTensors, _backend, opts = {}) {
79
- this._hostTensors = _hostTensors;
80
- this._backend = _backend;
20
+ _nativeExecutor;
21
+ _webNNExecutor;
22
+ _modelSpec;
23
+ _inputChannels;
24
+ _engine;
25
+ _dynamicTileController;
26
+ _lastExecution;
27
+ _activeExecutionFailures = new Set();
28
+ _deviceLostObserved = false;
29
+ _deviceLostSettled = false;
30
+ _deviceLostReason;
31
+ constructor(hostTensors, device, opts = {}) {
81
32
  this._aux = opts.aux || false;
82
33
  this._hdr = opts.hdr || false;
83
- this._maxTileSize = roundUp(opts.maxTileSize ?? 512, 2);
84
- this._device = this._backend.device;
34
+ this._hdrTransfer = opts.hdrTransfer ?? 'pu';
35
+ this._engine = opts.engine ?? 'auto';
36
+ const modelSpec = opts.modelSpec ?? detectUNetModelSpec(hostTensors);
37
+ const validatedModel = validateUNetModel(hostTensors, modelSpec);
38
+ this._modelSpec = validatedModel.spec;
39
+ this._inputChannels = validatedModel.inputChannels;
40
+ const expectedInputChannels = this._aux ? 9 : 3;
41
+ if (validatedModel.inputChannels !== expectedInputChannels) {
42
+ throw new Error(`OIDN model expects ${validatedModel.inputChannels} input channels, ` +
43
+ `but aux=${this._aux} provides ${expectedInputChannels}`);
44
+ }
45
+ this._dynamicTileController = new DynamicTileController(opts.maxTileSize ?? 512, opts.dynamicTile);
46
+ this._device = device;
47
+ this._observeDeviceLoss();
48
+ if (this._engine === 'webnn') {
49
+ this._webNNExecutor = new WebNNUNetExecutor(this._device, validatedModel, { precision: opts.precision });
50
+ }
51
+ else {
52
+ this._nativeExecutor = new NativeUNetExecutor(this._device, validatedModel, { precision: opts.precision, kernel: opts.kernel, gemm: opts.gemm });
53
+ }
85
54
  }
86
55
  getDevice() {
87
56
  return this._device;
88
57
  }
89
- _buildModel(isLarge) {
90
- const aux = this._aux;
91
- const channels = 3 + (aux ? 6 : 0);
92
- const tileSize = this._getTileSizeWithOverlap();
93
- const cache = this._modelsCache;
94
- const key = [tileSize.width, tileSize.height].join(',');
95
- // We cache the model instead of disposing and recreate.
96
- // Because seems tfjs will also cache the layer and gpubuffers.
97
- // Recreating the model will cause memory leak.
98
- // Width and height can only be 256, 512, 768. So the cache won't be too large
99
- if (cache.has(key)) {
100
- this._tfModel = cache.get(key);
58
+ /** Completes backend compilation before first interactive use. */
59
+ async prepare() {
60
+ if (this._webNNExecutor) {
61
+ await this._webNNExecutor.prepare();
62
+ const overlap = roundUp(this._modelSpec.receptiveField / 2, OIDN_TILE_ALIGNMENT);
63
+ const outputTileEdges = [
64
+ this._dynamicTileController.tileSize,
65
+ this._dynamicTileController.minTileSize
66
+ ];
67
+ await this._webNNExecutor.prewarm([...new Set(outputTileEdges)].map((edge) => ({
68
+ width: edge + 2 * overlap,
69
+ height: edge + 2 * overlap
70
+ })));
101
71
  return;
102
72
  }
103
- const input = TFInput({
104
- name: 'input',
105
- shape: [tileSize.height, tileSize.width, channels],
106
- dtype: 'float32'
107
- });
108
- this._tfModel = new LayersModel({
109
- inputs: [input],
110
- outputs: isLarge ? this._addNetLarge(input) : this._addNet(input)
111
- });
112
- cache.set(key, this._tfModel);
113
- }
114
- _createConv(name, source, activation) {
115
- const weightTensorName = name + '.weight';
116
- const biasTensorName = name + '.bias';
117
- const tensors = this._tensors;
118
- let weightTensor = tensors.get(weightTensorName);
119
- let biasTensor = tensors.get(biasTensorName);
120
- const unetWeightTensor = this._hostTensors.get(weightTensorName);
121
- if (!weightTensor) {
122
- const weightDims = unetWeightTensor.desc.dims;
123
- weightTensor = tensor(changeWeightShapes(getTensorData(unetWeightTensor.data, unetWeightTensor.desc.dataType), weightDims), [weightDims[2], weightDims[3], weightDims[1], weightDims[0]], 'float32');
124
- tensors.set(weightTensorName, weightTensor);
125
- }
126
- if (!biasTensor) {
127
- const unetBiasTensor = this._hostTensors.get(name + '.bias');
128
- biasTensor = tensor1d(getTensorData(unetBiasTensor.data, unetBiasTensor.desc.dataType), 'float32');
129
- tensors.set(biasTensorName, biasTensor);
130
- }
131
- // TODO whats the purpose of padded dims ?
132
- const convLayer = new Conv2D({
133
- name,
134
- filters: unetWeightTensor.desc.dims[0],
135
- kernelSize: unetWeightTensor.desc.dims.slice(2, 4),
136
- useBias: true,
137
- activation,
138
- padding: 'same',
139
- weights: [weightTensor, biasTensor],
140
- trainable: false
141
- });
142
- return convLayer.apply(source);
143
- }
144
- _createConcatConv(name, source1, source2) {
145
- const concatLayer = new Concatenate({
146
- name: name + '/concat',
147
- trainable: false,
148
- axis: 3
149
- });
150
- //https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L40
151
- return this._createConv(name,
152
- // Concat on the channel
153
- concatLayer.apply([source1, source2]), 'relu');
73
+ await this._nativeExecutor.prepare();
154
74
  }
155
- _createPooling(source) {
156
- const poolingLayer = new MaxPooling2D({
157
- name: source.name + '/pooling',
158
- poolSize: [2, 2],
159
- strides: [2, 2],
160
- padding: 'same',
161
- trainable: false
162
- });
163
- // https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L33
164
- return poolingLayer.apply(source);
75
+ /**
76
+ * Prepares the input shapes selected for an image before its first denoise.
77
+ * Hosts can call this while they still display their model-loading state.
78
+ */
79
+ async prepareForImage(width, height, options = {}) {
80
+ const defaultTileOverlap = roundUp(this._modelSpec.receptiveField / 2, OIDN_TILE_ALIGNMENT);
81
+ const resolvedTileOverlap = options.tileOverlap === undefined
82
+ ? defaultTileOverlap
83
+ : roundUp(Math.max(0, options.tileOverlap), OIDN_TILE_ALIGNMENT);
84
+ const plan = planTileGrid(width, height, options.wholeImage
85
+ ? roundUp(Math.max(width, height), OIDN_TILE_ALIGNMENT)
86
+ : this._dynamicTileController.tileSize, resolvedTileOverlap);
87
+ const shapes = [...new Map(plan.tiles.map(({ input }) => [
88
+ `${input.width}x${input.height}`,
89
+ { width: input.width, height: input.height }
90
+ ])).values()];
91
+ await this._webNNExecutor?.prewarm(shapes);
92
+ this._nativeExecutor?.prewarm(shapes);
165
93
  }
166
- _addUpsamplingLayer(source) {
167
- const upsamplingLayer = new UpSampling2D({
168
- name: source.name + '/upsampling',
169
- size: [2, 2],
170
- trainable: false
171
- });
172
- return upsamplingLayer.apply(source);
173
- }
174
- _addNet(input) {
175
- let x = this._createConv('enc_conv0', input, 'relu');
176
- const pool1 = (x = this._createPooling(this._createConv('enc_conv1', x, 'relu')));
177
- const pool2 = (x = this._createPooling(this._createConv('enc_conv2', x, 'relu')));
178
- const pool3 = (x = this._createPooling(this._createConv('enc_conv3', x, 'relu')));
179
- const pool4 = (x = this._createPooling(this._createConv('enc_conv4', x, 'relu')));
180
- x = this._createConv('enc_conv5a', pool4, 'relu');
181
- x = this._addUpsamplingLayer(this._createConv('enc_conv5b', x, 'relu'));
182
- x = this._createConcatConv('dec_conv4a', x, pool3);
183
- x = this._addUpsamplingLayer(this._createConv('dec_conv4b', x, 'relu'));
184
- x = this._createConcatConv('dec_conv3a', x, pool2);
185
- x = this._addUpsamplingLayer(this._createConv('dec_conv3b', x, 'relu'));
186
- x = this._createConcatConv('dec_conv2a', x, pool1);
187
- x = this._addUpsamplingLayer(this._createConv('dec_conv2b', x, 'relu'));
188
- x = this._createConcatConv('dec_conv1a', x, input);
189
- x = this._createConv('dec_conv1b', x, 'relu');
190
- x = this._createConv('dec_conv0', x, 'relu');
191
- return x;
94
+ getRuntimeInfo() {
95
+ return {
96
+ configuredEngine: this._engine,
97
+ gpuEngine: this._webNNExecutor ? 'webnn' : 'wgsl',
98
+ precision: (this._webNNExecutor ?? this._nativeExecutor).precision,
99
+ kernel: this._nativeExecutor
100
+ ? {
101
+ configured: this._nativeExecutor.kernelSetting,
102
+ gemm: this._nativeExecutor.gemm,
103
+ maxSpatialInputBlocks: this._nativeExecutor.maxSpatialInputBlocks,
104
+ subgroupsAvailable: this._nativeExecutor.subgroupsAvailable
105
+ }
106
+ : undefined,
107
+ webnn: this._webNNExecutor?.support,
108
+ resources: (this._webNNExecutor ?? this._nativeExecutor).getResourceInfo(),
109
+ model: this._modelSpec.id,
110
+ modelFamily: this._modelSpec.family,
111
+ inputChannels: this._inputChannels,
112
+ hdrTransfer: this._hdrTransfer,
113
+ dynamicTile: {
114
+ enabled: this._dynamicTileController.enabled,
115
+ currentTileSize: this._dynamicTileController.tileSize,
116
+ minTileSize: this._dynamicTileController.minTileSize,
117
+ maxTileSize: this._dynamicTileController.maxTileSize,
118
+ targetTileTimeMs: this._dynamicTileController.targetTileTimeMs
119
+ },
120
+ lastExecution: this._lastExecution,
121
+ activeExecutionCount: this._activeExecutionFailures.size
122
+ };
192
123
  }
193
- _addNetLarge(input) {
194
- let x = this._createConv('enc_conv1a', input, 'relu');
195
- const pool1 = (x = this._createPooling(this._createConv('enc_conv1b', x, 'relu')));
196
- x = this._createConv('enc_conv2a', x, 'relu');
197
- const pool2 = (x = this._createPooling(this._createConv('enc_conv2b', x, 'relu')));
198
- x = this._createConv('enc_conv3a', x, 'relu');
199
- const pool3 = (x = this._createPooling(this._createConv('enc_conv3b', x, 'relu')));
200
- x = this._createConv('enc_conv4a', x, 'relu');
201
- const pool4 = (x = this._createPooling(this._createConv('enc_conv4b', x, 'relu')));
202
- x = this._createConv('enc_conv5a', pool4, 'relu');
203
- x = this._addUpsamplingLayer(this._createConv('enc_conv5b', x, 'relu'));
204
- x = this._createConcatConv('dec_conv4a', x, pool3);
205
- x = this._addUpsamplingLayer(this._createConv('dec_conv4b', x, 'relu'));
206
- x = this._createConcatConv('dec_conv3a', x, pool2);
207
- x = this._addUpsamplingLayer(this._createConv('dec_conv3b', x, 'relu'));
208
- x = this._createConcatConv('dec_conv2a', x, pool1);
209
- x = this._addUpsamplingLayer(this._createConv('dec_conv2b', x, 'relu'));
210
- x = this._createConcatConv('dec_conv1a', x, input);
211
- x = this._createConv('dec_conv1b', x, 'relu');
212
- x = this._createConv('dec_conv1c', x, 'relu');
213
- return x;
124
+ _observeDeviceLoss() {
125
+ // Object.create-based embedders/tests can bypass field initializers.
126
+ this._activeExecutionFailures ??= new Set();
127
+ if (this._deviceLostObserved)
128
+ return;
129
+ this._deviceLostObserved = true;
130
+ const deviceLost = this._device.lost;
131
+ if (!deviceLost)
132
+ return;
133
+ const failAll = (reason) => {
134
+ if (this._deviceLostSettled)
135
+ return;
136
+ this._deviceLostSettled = true;
137
+ this._deviceLostReason = reason;
138
+ const active = [...this._activeExecutionFailures];
139
+ this._activeExecutionFailures.clear();
140
+ for (const fail of active)
141
+ fail(reason);
142
+ };
143
+ void deviceLost.then((info) => failAll(new Error(`WebGPU device lost: ${info.message}`)), failAll);
214
144
  }
215
- _updateModel(width, height) {
216
- const isLarge = this._hostTensors.has('enc_conv1b.weight');
217
- const maxTileSize = this._maxTileSize;
218
- let tileWidth = maxTileSize;
219
- let tileHeight = maxTileSize;
220
- let tileOverlapX = isLarge ? defaultTileOverlapLarge : defaultTileOverlap;
221
- let tileOverlapY = isLarge ? defaultTileOverlapLarge : defaultTileOverlap;
222
- if (width < maxTileSize + defaultTileOverlap * 2) {
223
- tileWidth = roundUp(width, maxTileSize / 2);
224
- if (width <= maxTileSize) {
225
- tileOverlapX = 0;
226
- }
227
- }
228
- if (height < maxTileSize + defaultTileOverlap * 2) {
229
- tileHeight = roundUp(height, maxTileSize / 2);
230
- if (height <= maxTileSize) {
231
- tileOverlapY = 0;
232
- }
233
- }
234
- // Force width and height has same size. reduce the cache in memory
235
- const tileSize = Math.max(tileWidth, tileHeight);
236
- const tileOverlap = Math.max(tileOverlapX, tileOverlapY);
237
- tileWidth = tileSize;
238
- tileHeight = tileSize;
239
- tileOverlapX = tileOverlap;
240
- tileOverlapY = tileOverlap;
241
- if (tileWidth !== this._tileWidth ||
242
- tileHeight !== this._tileHeight ||
243
- tileOverlapX !== this._tileOverlapX ||
244
- tileOverlapY !== this._tileOverlapY ||
245
- !this._tfModel) {
246
- // console.log(tileWidth, tileHeight, tileOverlapX, tileOverlapY);
247
- this._tileWidth = tileWidth;
248
- this._tileHeight = tileHeight;
249
- this._tileOverlapX = tileOverlapX;
250
- this._tileOverlapY = tileOverlapY;
251
- this._buildModel(isLarge);
145
+ _registerExecutionFailure(fail) {
146
+ this._observeDeviceLoss();
147
+ this._activeExecutionFailures.add(fail);
148
+ if (this._deviceLostSettled) {
149
+ this._activeExecutionFailures.delete(fail);
150
+ queueMicrotask(() => fail(this._deviceLostReason));
252
151
  }
152
+ return () => this._activeExecutionFailures.delete(fail);
253
153
  }
254
- _getTileSizeWithOverlap() {
255
- return {
256
- width: this._tileWidth + 2 * this._tileOverlapX,
257
- height: this._tileHeight + 2 * this._tileOverlapY
258
- };
154
+ /** Captures per-node GPU timestamps for the next native tile execution. */
155
+ profileNextExecution() {
156
+ return this._nativeExecutor?.profileNextExecution() ?? false;
157
+ }
158
+ getLastExecutionProfile() {
159
+ return this._nativeExecutor?.getLastExecutionProfile();
259
160
  }
260
161
  _processImageData(color, albedo, normal, isHDR) {
261
162
  const rawData = color.data;
@@ -296,9 +197,12 @@ class UNet {
296
197
  }
297
198
  _readTile(data, channels, srcTile, width) {
298
199
  const tileData = new Float32Array(srcTile.width * srcTile.height * channels);
200
+ const height = data.length / (width * channels);
299
201
  for (let y = 0; y < srcTile.height; y++) {
300
202
  for (let x = 0; x < srcTile.width; x++) {
301
- const i2 = ((y + srcTile.y) * width + (x + srcTile.x)) * channels;
203
+ const sourceX = Math.min(width - 1, x + srcTile.x);
204
+ const sourceY = Math.min(height - 1, y + srcTile.y);
205
+ const i2 = (sourceY * width + sourceX) * channels;
302
206
  const i1 = (y * srcTile.width + x) * channels;
303
207
  for (let c = 0; c < channels; c++) {
304
208
  tileData[i1 + c] = data[i2 + c];
@@ -327,22 +231,14 @@ class UNet {
327
231
  }
328
232
  }
329
233
  }
330
- _executeTile(inputData, outputTileData, outputImageData, i, j, width, height, isHDR, denoiseAlpha) {
234
+ async _executeTile(inputData, outputTileData, outputImageData, tile, isFirstTile, width, height, isHDR, denoiseAlpha) {
331
235
  const channels = this._aux ? 9 : 3;
332
- const tileOverlapX = this._tileOverlapX;
333
- const tileOverlapY = this._tileOverlapY;
334
- let srcTileSize = this._getTileSizeWithOverlap();
335
- let dstTileSize = { width: this._tileWidth, height: this._tileHeight };
336
- let srcX0 = i > 0 ? i * dstTileSize.width - tileOverlapX : 0;
337
- let srcX1 = Math.min(srcX0 + srcTileSize.width, width);
338
- srcX0 = Math.max(srcX1 - srcTileSize.width, 0);
339
- let srcY0 = j > 0 ? j * dstTileSize.height - tileOverlapY : 0;
340
- let srcY1 = Math.min(srcY0 + srcTileSize.height, height);
341
- srcY0 = Math.max(srcY1 - srcTileSize.height, 0);
342
- const srcTileWidth = srcTileSize.width;
343
- const srcTileHeight = srcTileSize.height;
344
- const srcTile = new Tile(srcX0, srcY0, srcTileWidth, srcTileHeight);
345
- let tileTensor;
236
+ const srcTile = new Tile(tile.input.x, tile.input.y, tile.input.width, tile.input.height);
237
+ const dstTile = new Tile(tile.output.x, tile.output.y, tile.output.width, tile.output.height);
238
+ const srcTileWidth = srcTile.width;
239
+ const srcTileHeight = srcTile.height;
240
+ let nativeOutputBuffer;
241
+ let denoisedData;
346
242
  let inputScale = 1;
347
243
  const device = this._device;
348
244
  let dataProcessGPU = this._dataProcessGPU;
@@ -356,60 +252,39 @@ class UNet {
356
252
  tileData = hdrTransferFuncCPU({
357
253
  data: tileData,
358
254
  channels,
359
- inputScale
255
+ inputScale,
256
+ transfer: this._hdrTransfer
360
257
  });
361
258
  }
362
- tileTensor = tensor(tileData, [1, srcTileHeight, srcTileWidth, channels], 'float32');
259
+ denoisedData = await (this._webNNExecutor ?? this._nativeExecutor).executeCPU(tileData, srcTileWidth, srcTileHeight);
363
260
  }
364
261
  else {
365
262
  if (!dataProcessGPU) {
366
- dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(device, isHDR);
263
+ dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(device, isHDR, this._hdrTransfer);
367
264
  }
368
265
  dataProcessGPU.setImageSize(width, height);
369
266
  dataProcessGPU.setInputTile(srcTile);
370
267
  // Display the noisy input instead of prev denoised result
371
- if (i === 0 && j === 0) {
268
+ if (isFirstTile) {
372
269
  dataProcessGPU.copyInputDataToOutput(inputData.color);
373
270
  }
374
271
  const { color, albedo, normal } = dataProcessGPU.forward(inputData.color, this._aux ? inputData.albedo : undefined, this._aux ? inputData.normal : undefined, denoiseAlpha);
375
- const createTensor = (buffer) => {
376
- const tmp = tensor({ buffer, zeroCopy: true }, [
377
- 1,
378
- srcTileHeight,
379
- srcTileWidth,
380
- 4
381
- ]);
382
- const ret = slice4d(tmp, [0, 0, 0, 0], [1, srcTileHeight, srcTileWidth, 3]);
383
- return ret;
384
- };
385
- if (this._aux) {
386
- const tensors = [color, albedo, normal].map((buffer) => createTensor(buffer));
387
- tileTensor = concat4d(tensors, 3);
388
- }
389
- else {
390
- tileTensor = createTensor(color);
391
- }
272
+ nativeOutputBuffer = await (this._webNNExecutor ?? this._nativeExecutor).execute(this._aux ? [color, albedo, normal] : [color], srcTileWidth, srcTileHeight);
392
273
  }
393
274
  let outBuffer;
394
- const outputTensor = this._tfModel.predict(tileTensor);
395
- const dstWidth = Math.min(dstTileSize.width, width);
396
- const dstHeight = Math.min(dstTileSize.height, height);
397
- const dstTile = new Tile(i * dstWidth, j * dstHeight, dstWidth, dstHeight);
398
- dstTile.width = Math.min(dstTile.width, width - dstTile.x);
399
- dstTile.height = Math.min(dstTile.height, height - dstTile.y);
400
275
  if (inputData instanceof Float32Array) {
401
- let denoisedData = outputTensor.dataSync();
402
276
  if (isHDR) {
403
277
  denoisedData = hdrTransferFuncInverseCPU({
404
278
  data: denoisedData,
405
279
  channels: 3,
406
- inputScale
280
+ inputScale,
281
+ transfer: this._hdrTransfer
407
282
  });
408
283
  }
409
- this._writeTile(outputImageData, srcTile, dstTile, denoisedData, srcTileSize.width, isHDR);
410
- for (let y = 0; y < dstHeight; y++) {
411
- for (let x = 0; x < dstWidth; x++) {
412
- const i1 = (y * dstWidth + x) * 4;
284
+ this._writeTile(outputImageData, srcTile, dstTile, denoisedData, srcTile.width, isHDR);
285
+ for (let y = 0; y < dstTile.height; y++) {
286
+ for (let x = 0; x < dstTile.width; x++) {
287
+ const i1 = (y * dstTile.width + x) * 4;
413
288
  const i2 = ((y + dstTile.y) * width + (x + dstTile.x)) * 4;
414
289
  for (let c = 0; c < 4; c++) {
415
290
  outputTileData.data[i1 + c] = outputImageData.data[i2 + c];
@@ -419,20 +294,11 @@ class UNet {
419
294
  }
420
295
  else {
421
296
  dataProcessGPU.setOutputTile(dstTile, srcTile);
422
- // IMPORTANT
423
- // storage buffer has alignment. that 3 channels still needs 16 bytes data.
424
- // So we need to pad it to 4 channels.
425
- const outputTensor4Channnels = pad4d(outputTensor, [
426
- [0, 0],
427
- [0, 0],
428
- [0, 0],
429
- [0, 1]
430
- ]);
431
- outBuffer = dataProcessGPU.inverse(outputTensor4Channnels.dataToGPU().buffer, inputData.color);
297
+ outBuffer = dataProcessGPU.inverse(nativeOutputBuffer, inputData.color);
432
298
  }
433
299
  return outBuffer;
434
300
  }
435
- tileExecute({ color, albedo, normal, done, progress, denoiseAlpha }) {
301
+ tileExecute({ color, albedo, normal, done, progress, denoiseAlpha, tileOverlap, wholeImage, scheduling = 'event-loop', error }) {
436
302
  if (this._aux && (!albedo || !normal)) {
437
303
  throw new Error('Normal map and albedo map are both required');
438
304
  }
@@ -443,7 +309,17 @@ class UNet {
443
309
  }
444
310
  const width = color.width;
445
311
  const height = color.height;
446
- this._updateModel(width, height);
312
+ const adaptiveTileSize = this._dynamicTileController.tileSize;
313
+ // The planner aligns the maximum down, so round up to keep one tile.
314
+ const requestedTileSize = wholeImage
315
+ ? roundUp(Math.max(width, height), OIDN_TILE_ALIGNMENT)
316
+ : adaptiveTileSize;
317
+ const defaultTileOverlap = roundUp(this._modelSpec.receptiveField / 2, OIDN_TILE_ALIGNMENT);
318
+ const resolvedTileOverlap = tileOverlap === undefined
319
+ ? defaultTileOverlap
320
+ : roundUp(Math.max(0, tileOverlap), OIDN_TILE_ALIGNMENT);
321
+ const plan = planTileGrid(width, height, requestedTileSize, resolvedTileOverlap);
322
+ const shouldAdaptTileSize = plan.tiles.length > 1;
447
323
  // TODO should fixed to be hdr when UNet is created.
448
324
  // weights of hdr and ldr is different
449
325
  const hdr = this._hdr || false;
@@ -451,10 +327,6 @@ class UNet {
451
327
  if (!isGPUImageData(color)) {
452
328
  rawData = this._processImageData(color, albedo, normal, hdr);
453
329
  }
454
- const tileWidth = this._tileWidth;
455
- const tileHeight = this._tileHeight;
456
- const tileCountH = Math.ceil(height / tileHeight);
457
- const tileCountW = Math.ceil(width / tileWidth);
458
330
  function makeImageData(width, height) {
459
331
  return hdr
460
332
  ? {
@@ -467,58 +339,167 @@ class UNet {
467
339
  const outputImageData = isGPUImageData(color)
468
340
  ? undefined
469
341
  : makeImageData(width, height);
470
- const outputTileData = isGPUImageData(color)
471
- ? undefined
472
- : makeImageData(Math.min(tileWidth, width), Math.min(tileHeight, height));
473
- let aborted = false;
474
- const executeTile = (i, j) => {
475
- if (aborted) {
342
+ let state = 'active';
343
+ let scheduledTimer;
344
+ let scheduledAnimationFrame;
345
+ let unregisterDeviceLoss = () => false;
346
+ const now = () => typeof performance === 'undefined' ? Date.now() : performance.now();
347
+ const executionStartTime = now();
348
+ const tileTimesMs = [];
349
+ const cancelScheduledTile = () => {
350
+ if (scheduledTimer !== undefined) {
351
+ clearTimeout(scheduledTimer);
352
+ scheduledTimer = undefined;
353
+ }
354
+ if (scheduledAnimationFrame !== undefined &&
355
+ typeof cancelAnimationFrame !== 'undefined') {
356
+ cancelAnimationFrame(scheduledAnimationFrame);
357
+ scheduledAnimationFrame = undefined;
358
+ }
359
+ };
360
+ const reportCallbackFailure = (reason) => {
361
+ // An error callback is the terminal observer and cannot report its own
362
+ // failure through the same channel. Keep that failure handled.
363
+ console.error('OIDN error callback failed', reason);
364
+ };
365
+ const settleError = (reason) => {
366
+ if (state !== 'active')
476
367
  return;
368
+ state = 'settled';
369
+ cancelScheduledTile();
370
+ unregisterDeviceLoss();
371
+ if (error) {
372
+ try {
373
+ void Promise.resolve(error(reason)).catch(reportCallbackFailure);
374
+ }
375
+ catch (callbackReason) {
376
+ reportCallbackFailure(callbackReason);
377
+ }
378
+ }
379
+ else {
380
+ console.error('OIDN execution failed', reason);
381
+ }
382
+ };
383
+ const scheduleNextTile = (callback) => {
384
+ if (scheduling === 'event-loop' ||
385
+ typeof requestAnimationFrame === 'undefined') {
386
+ scheduledTimer = setTimeout(() => {
387
+ scheduledTimer = undefined;
388
+ callback();
389
+ }, 0);
390
+ }
391
+ else {
392
+ // Hidden documents pause requestAnimationFrame. Race it against a
393
+ // timer so animation-frame scheduling still completes in background
394
+ // tabs, minimized windows, and offscreen iframes.
395
+ const run = () => {
396
+ cancelScheduledTile();
397
+ callback();
398
+ };
399
+ scheduledAnimationFrame = requestAnimationFrame(run);
400
+ scheduledTimer = setTimeout(run, ANIMATION_FRAME_FALLBACK_MS);
477
401
  }
478
- let resGPUBuffer;
479
- // profileAndLogKernelCode(() => {
480
- ENGINE.startScope();
481
- resGPUBuffer = this._executeTile(isGPUImageData(color)
402
+ };
403
+ const executeTile = async (tileIndex) => {
404
+ if (state !== 'active') {
405
+ return;
406
+ }
407
+ const tile = plan.tiles[tileIndex];
408
+ const outputTileData = isGPUImageData(color)
409
+ ? undefined
410
+ : makeImageData(tile.output.width, tile.output.height);
411
+ const tileStartTime = now();
412
+ const resGPUBuffer = await this._executeTile(isGPUImageData(color)
482
413
  ? {
483
414
  color: color.data,
484
415
  albedo: albedo?.data,
485
416
  normal: normal?.data
486
417
  }
487
- : rawData, outputTileData, outputImageData, i, j, width, height, hdr, denoiseAlpha);
488
- ENGINE.endScope();
489
- // }, true);
418
+ : rawData, outputTileData, outputImageData, tile, tileIndex === 0, width, height, hdr, denoiseAlpha);
419
+ if (state !== 'active')
420
+ return;
490
421
  const output = outputImageData || {
491
422
  data: resGPUBuffer,
492
423
  width,
493
424
  height
494
425
  };
495
- progress?.(output,
496
- // Is undefined if using webgpu buffer
497
- outputTileData, new Tile(i * tileWidth, j * tileHeight, tileWidth, tileHeight), i + j * tileCountW, tileCountW * tileCountH);
498
- if (i + 1 < tileCountW || j + 1 < tileCountH) {
499
- requestAnimationFrame(() => {
500
- if (i + 1 < tileCountW) {
501
- executeTile(i + 1, j);
426
+ if (progress) {
427
+ await progress(output,
428
+ // Is undefined if using webgpu buffer
429
+ outputTileData, new Tile(tile.output.x, tile.output.y, tile.output.width, tile.output.height), tileIndex, plan.tiles.length);
430
+ }
431
+ if (state !== 'active')
432
+ return;
433
+ const hasNextTile = tileIndex + 1 < plan.tiles.length;
434
+ await this._device.queue.onSubmittedWorkDone();
435
+ if (state !== 'active')
436
+ return;
437
+ const continueAfterGPUWork = async () => {
438
+ tileTimesMs.push(now() - tileStartTime);
439
+ if (state !== 'active')
440
+ return;
441
+ if (hasNextTile) {
442
+ scheduleNextTile(() => {
443
+ if (state !== 'active')
444
+ return;
445
+ void executeTile(tileIndex + 1).catch(settleError);
446
+ });
447
+ }
448
+ else {
449
+ const sortedTileTimes = [...tileTimesMs].sort((a, b) => a - b);
450
+ const middle = Math.floor(sortedTileTimes.length / 2);
451
+ const medianTileTime = sortedTileTimes.length % 2
452
+ ? sortedTileTimes[middle]
453
+ : (sortedTileTimes[middle - 1] + sortedTileTimes[middle]) / 2;
454
+ this._lastExecution = {
455
+ width,
456
+ height,
457
+ tileCount: plan.tiles.length,
458
+ tileColumns: plan.columns,
459
+ tileRows: plan.rows,
460
+ tileOverlap: plan.overlap,
461
+ inputPixelCount: plan.inputPixelCount,
462
+ inputShapeCount: plan.inputShapeCount,
463
+ durationMs: now() - executionStartTime,
464
+ tileTimeMs: {
465
+ min: sortedTileTimes[0],
466
+ median: medianTileTime,
467
+ mean: sortedTileTimes.reduce((sum, value) => sum + value, 0) /
468
+ sortedTileTimes.length,
469
+ max: sortedTileTimes[sortedTileTimes.length - 1]
470
+ }
471
+ };
472
+ // Adapt only from complete executions. Cancelled work is commonly
473
+ // contending with interactive rendering and is not representative.
474
+ if (shouldAdaptTileSize) {
475
+ this._dynamicTileController.observe(tileTimesMs);
502
476
  }
503
- else if (j + 1 < tileCountH) {
504
- executeTile(0, j + 1);
477
+ // GPU inference is complete; device loss can no longer affect this
478
+ // execution. Deregister before the user callback resolves so a
479
+ // completed execution never remains retained by the device watcher.
480
+ unregisterDeviceLoss();
481
+ await done(output);
482
+ if (state === 'active') {
483
+ state = 'settled';
505
484
  }
506
- });
507
- }
508
- else {
509
- // console.log(memory());
510
- done(output);
511
- }
485
+ }
486
+ };
487
+ await continueAfterGPUWork();
512
488
  };
513
- executeTile(0, 0);
489
+ unregisterDeviceLoss = this._registerExecutionFailure(settleError);
490
+ void executeTile(0).catch(settleError);
514
491
  return () => {
515
- aborted = true;
492
+ if (state !== 'active')
493
+ return;
494
+ state = 'aborted';
495
+ cancelScheduledTile();
496
+ unregisterDeviceLoss();
516
497
  };
517
498
  }
518
499
  dispose() {
519
- this._tfModel?.dispose();
520
500
  this._dataProcessGPU?.dispose();
521
- this._tensors.forEach((tensor) => tensor.dispose());
501
+ this._nativeExecutor?.dispose();
502
+ this._webNNExecutor?.dispose();
522
503
  }
523
504
  }
524
505
  export default UNet;