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/src/UNet.ts CHANGED
@@ -1,67 +1,36 @@
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,
24
5
  avgLogLum,
25
6
  hdrTransferFuncCPU,
26
- hdrTransferFuncInverseCPU
7
+ hdrTransferFuncInverseCPU,
8
+ type HDRTransfer
27
9
  } from './process';
28
- import type { WebGPUBackend } from '@tensorflow/tfjs-backend-webgpu';
10
+ import {
11
+ DynamicTileController,
12
+ type DynamicTileSetting,
13
+ planTileGrid,
14
+ type PlannedTile,
15
+ OIDN_TILE_ALIGNMENT
16
+ } from './tileScheduler';
17
+ import {
18
+ detectUNetModelSpec,
19
+ validateUNetModel,
20
+ type UNetModelSpec
21
+ } from './modelSpec';
22
+ import {
23
+ NativeUNetExecutor,
24
+ type NativeUNetGemmOptions,
25
+ type NativeUNetKernelSetting,
26
+ type NativeUNetPrecisionSetting
27
+ } from './nativeUNet';
28
+ import { WebNNUNetExecutor } from './webnnUNet';
29
29
 
30
- // import { profileAndLogKernelCode, memory } from './helper';
30
+ export type UNetEngineSetting = 'auto' | 'wgsl' | 'webnn';
31
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
- }
47
-
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
- }
32
+ /** Upper bound on waiting for a display frame between tiles. */
33
+ const ANIMATION_FRAME_FALLBACK_MS = 100;
65
34
 
66
35
  interface HDRImageData {
67
36
  data: Float32Array;
@@ -81,13 +50,27 @@ interface GPUImageDataOutput {
81
50
  height: number;
82
51
  }
83
52
 
53
+ export interface UNetExecutionStats {
54
+ width: number;
55
+ height: number;
56
+ tileCount: number;
57
+ tileColumns: number;
58
+ tileRows: number;
59
+ tileOverlap: number;
60
+ inputPixelCount: number;
61
+ inputShapeCount: number;
62
+ durationMs: number;
63
+ tileTimeMs: {
64
+ min: number;
65
+ median: number;
66
+ mean: number;
67
+ max: number;
68
+ };
69
+ }
70
+
84
71
  function roundUp(a: number, b: number) {
85
72
  return Math.ceil(a / b) * b;
86
73
  }
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
74
 
92
75
  function isGPUImageData(
93
76
  data: ImageData | GPUImageData | HDRImageData
@@ -95,41 +78,30 @@ function isGPUImageData(
95
78
  return data.data instanceof GPUBuffer || data.data instanceof GPUTexture;
96
79
  }
97
80
 
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
81
  class UNet {
109
- private _tfModel: LayersModel | undefined;
110
- private _device: GPUDevice | undefined;
111
-
112
- // TODO calculate the tile size from memory size
113
- // https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/core/unet_filter.cpp#L287
114
- private _tileWidth = 0;
115
- private _tileHeight = 0;
116
-
117
- private _tileOverlapX = 0;
118
- private _tileOverlapY = 0;
82
+ private _device: GPUDevice;
119
83
 
120
84
  private _aux;
121
85
  private _hdr;
86
+ private _hdrTransfer: HDRTransfer;
122
87
 
123
88
  private _dataProcessGPU?: GPUDataProcess;
124
-
125
- private _maxTileSize;
126
-
127
- private _tensors = new Map<string, Tensor>();
128
- private _modelsCache = new Map<string, LayersModel>();
89
+ private _nativeExecutor?: NativeUNetExecutor;
90
+ private _webNNExecutor?: WebNNUNetExecutor;
91
+ private _modelSpec: UNetModelSpec;
92
+ private _inputChannels: number;
93
+ private _engine: UNetEngineSetting;
94
+
95
+ private _dynamicTileController: DynamicTileController;
96
+ private _lastExecution?: UNetExecutionStats;
97
+ private _activeExecutionFailures = new Set<(reason: unknown) => void>();
98
+ private _deviceLostObserved = false;
99
+ private _deviceLostSettled = false;
100
+ private _deviceLostReason: unknown;
129
101
 
130
102
  constructor(
131
- private _hostTensors: Map<string, HostTensor>,
132
- private _backend: WebGPUBackend,
103
+ hostTensors: Map<string, HostTensor>,
104
+ device: GPUDevice,
133
105
  opts: {
134
106
  /**
135
107
  * If use auxiliary data.
@@ -139,261 +111,194 @@ class UNet {
139
111
  * If input is HDR image.
140
112
  */
141
113
  hdr?: boolean;
114
+ /** HDR transfer function expected by the trained model. */
115
+ hdrTransfer?: HDRTransfer;
142
116
  maxTileSize?: number;
117
+ dynamicTile?: DynamicTileSetting;
118
+ /** Native WGSL or the experimental WebNN backend. */
119
+ engine?: UNetEngineSetting;
120
+ /** Arithmetic/storage precision used by the native WGSL executor. */
121
+ precision?: NativeUNetPrecisionSetting;
122
+ /** Model-independent convolution kernel selection. */
123
+ kernel?: NativeUNetKernelSetting;
124
+ gemm?: NativeUNetGemmOptions;
125
+ /** Explicit descriptor for a new OIDN topology not in the built-in registry. */
126
+ modelSpec?: UNetModelSpec;
143
127
  } = {}
144
128
  ) {
145
129
  this._aux = opts.aux || false;
146
130
  this._hdr = opts.hdr || false;
131
+ this._hdrTransfer = opts.hdrTransfer ?? 'pu';
132
+ this._engine = opts.engine ?? 'auto';
133
+ const modelSpec = opts.modelSpec ?? detectUNetModelSpec(hostTensors);
134
+ const validatedModel = validateUNetModel(hostTensors, modelSpec);
135
+ this._modelSpec = validatedModel.spec;
136
+ this._inputChannels = validatedModel.inputChannels;
137
+
138
+ const expectedInputChannels = this._aux ? 9 : 3;
139
+ if (validatedModel.inputChannels !== expectedInputChannels) {
140
+ throw new Error(
141
+ `OIDN model expects ${validatedModel.inputChannels} input channels, ` +
142
+ `but aux=${this._aux} provides ${expectedInputChannels}`
143
+ );
144
+ }
147
145
 
148
- this._maxTileSize = roundUp(opts.maxTileSize ?? 512, 2);
146
+ this._dynamicTileController = new DynamicTileController(
147
+ opts.maxTileSize ?? 512,
148
+ opts.dynamicTile
149
+ );
149
150
 
150
- this._device = this._backend.device;
151
+ this._device = device;
152
+ this._observeDeviceLoss();
153
+ if (this._engine === 'webnn') {
154
+ this._webNNExecutor = new WebNNUNetExecutor(
155
+ this._device,
156
+ validatedModel,
157
+ { precision: opts.precision }
158
+ );
159
+ } else {
160
+ this._nativeExecutor = new NativeUNetExecutor(
161
+ this._device,
162
+ validatedModel,
163
+ { precision: opts.precision, kernel: opts.kernel, gemm: opts.gemm }
164
+ );
165
+ }
151
166
  }
152
167
 
153
168
  getDevice() {
154
169
  return this._device;
155
170
  }
156
171
 
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'
172
+ /** Completes backend compilation before first interactive use. */
173
+ async prepare() {
174
+ if (this._webNNExecutor) {
175
+ await this._webNNExecutor.prepare();
176
+ const overlap = roundUp(
177
+ this._modelSpec.receptiveField / 2,
178
+ OIDN_TILE_ALIGNMENT
208
179
  );
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'
180
+ const outputTileEdges = [
181
+ this._dynamicTileController.tileSize,
182
+ this._dynamicTileController.minTileSize
183
+ ];
184
+ await this._webNNExecutor.prewarm(
185
+ [...new Set(outputTileEdges)].map((edge) => ({
186
+ width: edge + 2 * overlap,
187
+ height: edge + 2 * overlap
188
+ }))
217
189
  );
218
- tensors.set(biasTensorName, biasTensor);
190
+ return;
219
191
  }
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;
192
+ await this._nativeExecutor!.prepare();
233
193
  }
234
194
 
235
- private _createConcatConv(
236
- name: string,
237
- source1: SymbolicTensor,
238
- source2: SymbolicTensor
195
+ /**
196
+ * Prepares the input shapes selected for an image before its first denoise.
197
+ * Hosts can call this while they still display their model-loading state.
198
+ */
199
+ async prepareForImage(
200
+ width: number,
201
+ height: number,
202
+ options: { tileOverlap?: number; wholeImage?: boolean } = {}
239
203
  ) {
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;
252
- }
253
-
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;
264
- }
265
-
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;
204
+ const defaultTileOverlap = roundUp(
205
+ this._modelSpec.receptiveField / 2,
206
+ OIDN_TILE_ALIGNMENT
207
+ );
208
+ const resolvedTileOverlap = options.tileOverlap === undefined
209
+ ? defaultTileOverlap
210
+ : roundUp(Math.max(0, options.tileOverlap), OIDN_TILE_ALIGNMENT);
211
+ const plan = planTileGrid(
212
+ width,
213
+ height,
214
+ options.wholeImage
215
+ ? roundUp(Math.max(width, height), OIDN_TILE_ALIGNMENT)
216
+ : this._dynamicTileController.tileSize,
217
+ resolvedTileOverlap
218
+ );
219
+ const shapes = [...new Map(
220
+ plan.tiles.map(({ input }) => [
221
+ `${input.width}x${input.height}`,
222
+ { width: input.width, height: input.height }
223
+ ])
224
+ ).values()];
225
+ await this._webNNExecutor?.prewarm(shapes);
226
+ this._nativeExecutor?.prewarm(shapes);
273
227
  }
274
228
 
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;
229
+ getRuntimeInfo() {
230
+ return {
231
+ configuredEngine: this._engine,
232
+ gpuEngine: this._webNNExecutor ? 'webnn' as const : 'wgsl' as const,
233
+ precision: (this._webNNExecutor ?? this._nativeExecutor!).precision,
234
+ kernel: this._nativeExecutor
235
+ ? {
236
+ configured: this._nativeExecutor.kernelSetting,
237
+ gemm: this._nativeExecutor.gemm,
238
+ maxSpatialInputBlocks: this._nativeExecutor.maxSpatialInputBlocks,
239
+ subgroupsAvailable: this._nativeExecutor.subgroupsAvailable
240
+ }
241
+ : undefined,
242
+ webnn: this._webNNExecutor?.support,
243
+ resources: (
244
+ this._webNNExecutor ?? this._nativeExecutor!
245
+ ).getResourceInfo(),
246
+ model: this._modelSpec.id,
247
+ modelFamily: this._modelSpec.family,
248
+ inputChannels: this._inputChannels,
249
+ hdrTransfer: this._hdrTransfer,
250
+ dynamicTile: {
251
+ enabled: this._dynamicTileController.enabled,
252
+ currentTileSize: this._dynamicTileController.tileSize,
253
+ minTileSize: this._dynamicTileController.minTileSize,
254
+ maxTileSize: this._dynamicTileController.maxTileSize,
255
+ targetTileTimeMs: this._dynamicTileController.targetTileTimeMs
256
+ },
257
+ lastExecution: this._lastExecution,
258
+ activeExecutionCount: this._activeExecutionFailures.size
259
+ };
306
260
  }
307
261
 
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;
262
+ private _observeDeviceLoss() {
263
+ // Object.create-based embedders/tests can bypass field initializers.
264
+ this._activeExecutionFailures ??= new Set();
265
+ if (this._deviceLostObserved) return;
266
+ this._deviceLostObserved = true;
267
+ const deviceLost = (this._device as GPUDevice & {
268
+ lost?: Promise<GPUDeviceLostInfo>;
269
+ }).lost;
270
+ if (!deviceLost) return;
271
+ const failAll = (reason: unknown) => {
272
+ if (this._deviceLostSettled) return;
273
+ this._deviceLostSettled = true;
274
+ this._deviceLostReason = reason;
275
+ const active = [...this._activeExecutionFailures];
276
+ this._activeExecutionFailures.clear();
277
+ for (const fail of active) fail(reason);
278
+ };
279
+ void deviceLost.then(
280
+ (info) => failAll(new Error(`WebGPU device lost: ${info.message}`)),
281
+ failAll
282
+ );
343
283
  }
344
284
 
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
- }
359
- }
360
- if (height < maxTileSize + defaultTileOverlap * 2) {
361
- tileHeight = roundUp(height, maxTileSize / 2);
362
- if (height <= maxTileSize) {
363
- tileOverlapY = 0;
364
- }
285
+ private _registerExecutionFailure(fail: (reason: unknown) => void) {
286
+ this._observeDeviceLoss();
287
+ this._activeExecutionFailures.add(fail);
288
+ if (this._deviceLostSettled) {
289
+ this._activeExecutionFailures.delete(fail);
290
+ queueMicrotask(() => fail(this._deviceLostReason));
365
291
  }
292
+ return () => this._activeExecutionFailures.delete(fail);
293
+ }
366
294
 
367
- // Force width and height has same size. reduce the cache in memory
368
- const tileSize = Math.max(tileWidth, tileHeight);
369
- const tileOverlap = Math.max(tileOverlapX, tileOverlapY);
370
- tileWidth = tileSize;
371
- tileHeight = tileSize;
372
- tileOverlapX = tileOverlap;
373
- tileOverlapY = tileOverlap;
374
-
375
- if (
376
- tileWidth !== this._tileWidth ||
377
- tileHeight !== this._tileHeight ||
378
- tileOverlapX !== this._tileOverlapX ||
379
- tileOverlapY !== this._tileOverlapY ||
380
- !this._tfModel
381
- ) {
382
- // console.log(tileWidth, tileHeight, tileOverlapX, tileOverlapY);
383
- this._tileWidth = tileWidth;
384
- this._tileHeight = tileHeight;
385
- this._tileOverlapX = tileOverlapX;
386
- this._tileOverlapY = tileOverlapY;
387
-
388
- this._buildModel(isLarge);
389
- }
295
+ /** Captures per-node GPU timestamps for the next native tile execution. */
296
+ profileNextExecution() {
297
+ return this._nativeExecutor?.profileNextExecution() ?? false;
390
298
  }
391
299
 
392
- private _getTileSizeWithOverlap() {
393
- return {
394
- width: this._tileWidth + 2 * this._tileOverlapX,
395
- height: this._tileHeight + 2 * this._tileOverlapY
396
- };
300
+ getLastExecutionProfile() {
301
+ return this._nativeExecutor?.getLastExecutionProfile();
397
302
  }
398
303
 
399
304
  private _processImageData(
@@ -453,9 +358,12 @@ class UNet {
453
358
  const tileData = new Float32Array(
454
359
  srcTile.width * srcTile.height * channels
455
360
  );
361
+ const height = data.length / (width * channels);
456
362
  for (let y = 0; y < srcTile.height; y++) {
457
363
  for (let x = 0; x < srcTile.width; x++) {
458
- const i2 = ((y + srcTile.y) * width + (x + srcTile.x)) * channels;
364
+ const sourceX = Math.min(width - 1, x + srcTile.x);
365
+ const sourceY = Math.min(height - 1, y + srcTile.y);
366
+ const i2 = (sourceY * width + sourceX) * channels;
459
367
  const i1 = (y * srcTile.width + x) * channels;
460
368
 
461
369
  for (let c = 0; c < channels; c++) {
@@ -497,7 +405,7 @@ class UNet {
497
405
  }
498
406
  }
499
407
 
500
- private _executeTile(
408
+ private async _executeTile(
501
409
  inputData:
502
410
  | Float32Array
503
411
  | {
@@ -508,35 +416,33 @@ class UNet {
508
416
  },
509
417
  outputTileData: ImageData | HDRImageData | undefined,
510
418
  outputImageData: ImageData | HDRImageData | undefined,
511
- i: number,
512
- j: number,
419
+ tile: PlannedTile,
420
+ isFirstTile: boolean,
513
421
  width: number,
514
422
  height: number,
515
423
  isHDR: boolean,
516
424
  denoiseAlpha?: boolean
517
425
  ) {
518
426
  const channels = this._aux ? 9 : 3;
519
- const tileOverlapX = this._tileOverlapX;
520
- const tileOverlapY = this._tileOverlapY;
521
- let srcTileSize = this._getTileSizeWithOverlap();
522
- let dstTileSize = { width: this._tileWidth, height: this._tileHeight };
523
-
524
- let srcX0 = i > 0 ? i * dstTileSize.width - tileOverlapX : 0;
525
- let srcX1 = Math.min(srcX0 + srcTileSize.width, width);
526
- srcX0 = Math.max(srcX1 - srcTileSize.width, 0);
527
-
528
- let srcY0 = j > 0 ? j * dstTileSize.height - tileOverlapY : 0;
529
- let srcY1 = Math.min(srcY0 + srcTileSize.height, height);
530
- srcY0 = Math.max(srcY1 - srcTileSize.height, 0);
531
-
532
- const srcTileWidth = srcTileSize.width;
533
- const srcTileHeight = srcTileSize.height;
534
-
535
- const srcTile = new Tile(srcX0, srcY0, srcTileWidth, srcTileHeight);
427
+ const srcTile = new Tile(
428
+ tile.input.x,
429
+ tile.input.y,
430
+ tile.input.width,
431
+ tile.input.height
432
+ );
433
+ const dstTile = new Tile(
434
+ tile.output.x,
435
+ tile.output.y,
436
+ tile.output.width,
437
+ tile.output.height
438
+ );
439
+ const srcTileWidth = srcTile.width;
440
+ const srcTileHeight = srcTile.height;
536
441
 
537
- let tileTensor!: Tensor;
442
+ let nativeOutputBuffer: GPUBuffer | undefined;
443
+ let denoisedData: Float32Array | undefined;
538
444
  let inputScale = 1;
539
- const device = this._device!;
445
+ const device = this._device;
540
446
  let dataProcessGPU = this._dataProcessGPU;
541
447
 
542
448
  if (inputData instanceof Float32Array) {
@@ -549,25 +455,27 @@ class UNet {
549
455
  tileData = hdrTransferFuncCPU({
550
456
  data: tileData,
551
457
  channels,
552
- inputScale
458
+ inputScale,
459
+ transfer: this._hdrTransfer
553
460
  });
554
461
  }
555
- tileTensor = tensor(
462
+ denoisedData = await (this._webNNExecutor ?? this._nativeExecutor!).executeCPU(
556
463
  tileData,
557
- [1, srcTileHeight, srcTileWidth, channels],
558
- 'float32'
559
- ) as Tensor4D;
464
+ srcTileWidth,
465
+ srcTileHeight
466
+ );
560
467
  } else {
561
468
  if (!dataProcessGPU) {
562
469
  dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(
563
470
  device,
564
- isHDR
471
+ isHDR,
472
+ this._hdrTransfer
565
473
  );
566
474
  }
567
475
  dataProcessGPU.setImageSize(width, height);
568
476
  dataProcessGPU.setInputTile(srcTile);
569
477
  // Display the noisy input instead of prev denoised result
570
- if (i === 0 && j === 0) {
478
+ if (isFirstTile) {
571
479
  dataProcessGPU.copyInputDataToOutput(inputData.color);
572
480
  }
573
481
  const { color, albedo, normal } = dataProcessGPU.forward(
@@ -577,47 +485,22 @@ class UNet {
577
485
  denoiseAlpha
578
486
  );
579
487
 
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
- }
488
+ nativeOutputBuffer = await (this._webNNExecutor ?? this._nativeExecutor!).execute(
489
+ this._aux ? [color, albedo!, normal!] : [color],
490
+ srcTileWidth,
491
+ srcTileHeight
492
+ );
603
493
  }
604
494
 
605
495
  let outBuffer: GPUBuffer;
606
- const outputTensor = this._tfModel!.predict(tileTensor) as Tensor;
607
-
608
- const dstWidth = Math.min(dstTileSize.width, width);
609
- const dstHeight = Math.min(dstTileSize.height, height);
610
- const dstTile = new Tile(i * dstWidth, j * dstHeight, dstWidth, dstHeight);
611
- dstTile.width = Math.min(dstTile.width, width - dstTile.x);
612
- dstTile.height = Math.min(dstTile.height, height - dstTile.y);
613
496
 
614
497
  if (inputData instanceof Float32Array) {
615
- let denoisedData = outputTensor.dataSync();
616
498
  if (isHDR) {
617
499
  denoisedData = hdrTransferFuncInverseCPU({
618
- data: denoisedData as Float32Array,
500
+ data: denoisedData!,
619
501
  channels: 3,
620
- inputScale
502
+ inputScale,
503
+ transfer: this._hdrTransfer
621
504
  });
622
505
  }
623
506
 
@@ -625,14 +508,14 @@ class UNet {
625
508
  outputImageData!,
626
509
  srcTile,
627
510
  dstTile,
628
- denoisedData as Float32Array,
629
- srcTileSize.width,
511
+ denoisedData!,
512
+ srcTile.width,
630
513
  isHDR
631
514
  );
632
515
 
633
- for (let y = 0; y < dstHeight; y++) {
634
- for (let x = 0; x < dstWidth; x++) {
635
- const i1 = (y * dstWidth + x) * 4;
516
+ for (let y = 0; y < dstTile.height; y++) {
517
+ for (let x = 0; x < dstTile.width; x++) {
518
+ const i1 = (y * dstTile.width + x) * 4;
636
519
  const i2 = ((y + dstTile.y) * width + (x + dstTile.x)) * 4;
637
520
  for (let c = 0; c < 4; c++) {
638
521
  outputTileData!.data[i1 + c] = outputImageData!.data[i2 + c];
@@ -641,17 +524,8 @@ class UNet {
641
524
  }
642
525
  } else {
643
526
  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
527
  outBuffer = dataProcessGPU!.inverse(
654
- outputTensor4Channnels.dataToGPU().buffer!,
528
+ nativeOutputBuffer!,
655
529
  inputData.color
656
530
  );
657
531
  }
@@ -664,7 +538,11 @@ class UNet {
664
538
  normal,
665
539
  done,
666
540
  progress,
667
- denoiseAlpha
541
+ denoiseAlpha,
542
+ tileOverlap,
543
+ wholeImage,
544
+ scheduling = 'event-loop',
545
+ error
668
546
  }: {
669
547
  color: T;
670
548
  albedo?: ImageData | GPUImageData;
@@ -673,14 +551,35 @@ class UNet {
673
551
  * If denoise alpha channel. Otherwise denoise RGB channels.
674
552
  */
675
553
  denoiseAlpha?: boolean;
676
- done: (outputData: T extends GPUImageData ? GPUImageDataOutput : T) => void;
554
+ /**
555
+ * Execute the complete input image as one tile, ignoring `maxTileSize`.
556
+ * The image must fit the device's buffer and dispatch limits.
557
+ */
558
+ wholeImage?: boolean;
559
+ /**
560
+ * Per-side context for boundaries shared with another tile. Defaults to
561
+ * half of the model receptive field rounded up to 16 pixels.
562
+ */
563
+ tileOverlap?: number;
564
+ /**
565
+ * How JavaScript yields between completed GPU tiles. `event-loop`
566
+ * (default) continues on the next macrotask. `animation-frame` waits for
567
+ * the next display frame, bounded by a short timer so hidden pages still
568
+ * complete.
569
+ */
570
+ scheduling?: 'animation-frame' | 'event-loop';
571
+ done: (
572
+ outputData: T extends GPUImageData ? GPUImageDataOutput : T
573
+ ) => void | Promise<void>;
574
+ /** Receives asynchronous execution, queue, and callback failures. */
575
+ error?: (reason: unknown) => void | Promise<void>;
677
576
  progress?: (
678
577
  outputData: T extends GPUImageData ? GPUImageDataOutput : T,
679
578
  tileData: (T extends GPUImageData ? GPUImageDataOutput : T) | undefined,
680
579
  tile: Tile,
681
580
  currentIdx: number,
682
581
  totalIdx: number
683
- ) => void;
582
+ ) => void | Promise<void>;
684
583
  }): () => void {
685
584
  if (this._aux && (!albedo || !normal)) {
686
585
  throw new Error('Normal map and albedo map are both required');
@@ -694,7 +593,25 @@ class UNet {
694
593
 
695
594
  const width = color.width;
696
595
  const height = color.height;
697
- this._updateModel(width, height);
596
+ const adaptiveTileSize = this._dynamicTileController.tileSize;
597
+ // The planner aligns the maximum down, so round up to keep one tile.
598
+ const requestedTileSize = wholeImage
599
+ ? roundUp(Math.max(width, height), OIDN_TILE_ALIGNMENT)
600
+ : adaptiveTileSize;
601
+ const defaultTileOverlap = roundUp(
602
+ this._modelSpec.receptiveField / 2,
603
+ OIDN_TILE_ALIGNMENT
604
+ );
605
+ const resolvedTileOverlap = tileOverlap === undefined
606
+ ? defaultTileOverlap
607
+ : roundUp(Math.max(0, tileOverlap), OIDN_TILE_ALIGNMENT);
608
+ const plan = planTileGrid(
609
+ width,
610
+ height,
611
+ requestedTileSize,
612
+ resolvedTileOverlap
613
+ );
614
+ const shouldAdaptTileSize = plan.tiles.length > 1;
698
615
 
699
616
  // TODO should fixed to be hdr when UNet is created.
700
617
  // weights of hdr and ldr is different
@@ -709,11 +626,6 @@ class UNet {
709
626
  hdr
710
627
  );
711
628
  }
712
- const tileWidth = this._tileWidth;
713
- const tileHeight = this._tileHeight;
714
- const tileCountH = Math.ceil(height / tileHeight);
715
- const tileCountW = Math.ceil(width / tileWidth);
716
-
717
629
  function makeImageData(width: number, height: number) {
718
630
  return hdr
719
631
  ? {
@@ -727,20 +639,82 @@ class UNet {
727
639
  const outputImageData = isGPUImageData(color)
728
640
  ? undefined
729
641
  : makeImageData(width, height);
730
- const outputTileData = isGPUImageData(color)
731
- ? undefined
732
- : makeImageData(Math.min(tileWidth, width), Math.min(tileHeight, height));
733
642
 
734
- let aborted = false;
643
+ type ExecutionState = 'active' | 'aborted' | 'settled';
644
+ let state: ExecutionState = 'active';
645
+ let scheduledTimer: ReturnType<typeof setTimeout> | undefined;
646
+ let scheduledAnimationFrame: number | undefined;
647
+ let unregisterDeviceLoss = () => false;
648
+
649
+ const now = () =>
650
+ typeof performance === 'undefined' ? Date.now() : performance.now();
651
+ const executionStartTime = now();
652
+ const tileTimesMs: number[] = [];
653
+ const cancelScheduledTile = () => {
654
+ if (scheduledTimer !== undefined) {
655
+ clearTimeout(scheduledTimer);
656
+ scheduledTimer = undefined;
657
+ }
658
+ if (
659
+ scheduledAnimationFrame !== undefined &&
660
+ typeof cancelAnimationFrame !== 'undefined'
661
+ ) {
662
+ cancelAnimationFrame(scheduledAnimationFrame);
663
+ scheduledAnimationFrame = undefined;
664
+ }
665
+ };
666
+ const reportCallbackFailure = (reason: unknown) => {
667
+ // An error callback is the terminal observer and cannot report its own
668
+ // failure through the same channel. Keep that failure handled.
669
+ console.error('OIDN error callback failed', reason);
670
+ };
671
+ const settleError = (reason: unknown) => {
672
+ if (state !== 'active') return;
673
+ state = 'settled';
674
+ cancelScheduledTile();
675
+ unregisterDeviceLoss();
676
+ if (error) {
677
+ try {
678
+ void Promise.resolve(error(reason)).catch(reportCallbackFailure);
679
+ } catch (callbackReason) {
680
+ reportCallbackFailure(callbackReason);
681
+ }
682
+ } else {
683
+ console.error('OIDN execution failed', reason);
684
+ }
685
+ };
686
+ const scheduleNextTile = (callback: () => void) => {
687
+ if (
688
+ scheduling === 'event-loop' ||
689
+ typeof requestAnimationFrame === 'undefined'
690
+ ) {
691
+ scheduledTimer = setTimeout(() => {
692
+ scheduledTimer = undefined;
693
+ callback();
694
+ }, 0);
695
+ } else {
696
+ // Hidden documents pause requestAnimationFrame. Race it against a
697
+ // timer so animation-frame scheduling still completes in background
698
+ // tabs, minimized windows, and offscreen iframes.
699
+ const run = () => {
700
+ cancelScheduledTile();
701
+ callback();
702
+ };
703
+ scheduledAnimationFrame = requestAnimationFrame(run);
704
+ scheduledTimer = setTimeout(run, ANIMATION_FRAME_FALLBACK_MS);
705
+ }
706
+ };
735
707
 
736
- const executeTile = (i: number, j: number) => {
737
- if (aborted) {
708
+ const executeTile = async (tileIndex: number) => {
709
+ if (state !== 'active') {
738
710
  return;
739
711
  }
740
- let resGPUBuffer;
741
- // profileAndLogKernelCode(() => {
742
- ENGINE.startScope();
743
- resGPUBuffer = this._executeTile(
712
+ const tile = plan.tiles[tileIndex];
713
+ const outputTileData = isGPUImageData(color)
714
+ ? undefined
715
+ : makeImageData(tile.output.width, tile.output.height);
716
+ const tileStartTime = now();
717
+ const resGPUBuffer = await this._executeTile(
744
718
  isGPUImageData(color)
745
719
  ? {
746
720
  color: color.data,
@@ -750,54 +724,106 @@ class UNet {
750
724
  : rawData,
751
725
  outputTileData,
752
726
  outputImageData,
753
- i,
754
- j,
727
+ tile,
728
+ tileIndex === 0,
755
729
  width,
756
730
  height,
757
731
  hdr,
758
732
  denoiseAlpha
759
733
  );
760
- ENGINE.endScope();
761
- // }, true);
734
+ if (state !== 'active') return;
762
735
  const output = outputImageData || {
763
736
  data: resGPUBuffer,
764
737
  width,
765
738
  height
766
739
  };
767
- progress?.(
768
- output as any,
769
- // Is undefined if using webgpu buffer
770
- outputTileData as any,
771
- new Tile(i * tileWidth, j * tileHeight, tileWidth, tileHeight),
772
- i + j * tileCountW,
773
- tileCountW * tileCountH
774
- );
775
-
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);
782
- }
783
- });
784
- } else {
785
- // console.log(memory());
786
- done(output as any);
740
+ if (progress) {
741
+ await progress(
742
+ output as any,
743
+ // Is undefined if using webgpu buffer
744
+ outputTileData as any,
745
+ new Tile(
746
+ tile.output.x,
747
+ tile.output.y,
748
+ tile.output.width,
749
+ tile.output.height
750
+ ),
751
+ tileIndex,
752
+ plan.tiles.length
753
+ );
787
754
  }
755
+ if (state !== 'active') return;
756
+
757
+ const hasNextTile = tileIndex + 1 < plan.tiles.length;
758
+ await this._device.queue.onSubmittedWorkDone();
759
+ if (state !== 'active') return;
760
+ const continueAfterGPUWork = async () => {
761
+ tileTimesMs.push(now() - tileStartTime);
762
+ if (state !== 'active') return;
763
+
764
+ if (hasNextTile) {
765
+ scheduleNextTile(() => {
766
+ if (state !== 'active') return;
767
+ void executeTile(tileIndex + 1).catch(settleError);
768
+ });
769
+ } else {
770
+ const sortedTileTimes = [...tileTimesMs].sort((a, b) => a - b);
771
+ const middle = Math.floor(sortedTileTimes.length / 2);
772
+ const medianTileTime = sortedTileTimes.length % 2
773
+ ? sortedTileTimes[middle]
774
+ : (sortedTileTimes[middle - 1] + sortedTileTimes[middle]) / 2;
775
+ this._lastExecution = {
776
+ width,
777
+ height,
778
+ tileCount: plan.tiles.length,
779
+ tileColumns: plan.columns,
780
+ tileRows: plan.rows,
781
+ tileOverlap: plan.overlap,
782
+ inputPixelCount: plan.inputPixelCount,
783
+ inputShapeCount: plan.inputShapeCount,
784
+ durationMs: now() - executionStartTime,
785
+ tileTimeMs: {
786
+ min: sortedTileTimes[0],
787
+ median: medianTileTime,
788
+ mean:
789
+ sortedTileTimes.reduce((sum, value) => sum + value, 0) /
790
+ sortedTileTimes.length,
791
+ max: sortedTileTimes[sortedTileTimes.length - 1]
792
+ }
793
+ };
794
+ // Adapt only from complete executions. Cancelled work is commonly
795
+ // contending with interactive rendering and is not representative.
796
+ if (shouldAdaptTileSize) {
797
+ this._dynamicTileController.observe(tileTimesMs);
798
+ }
799
+ // GPU inference is complete; device loss can no longer affect this
800
+ // execution. Deregister before the user callback resolves so a
801
+ // completed execution never remains retained by the device watcher.
802
+ unregisterDeviceLoss();
803
+ await done(output as any);
804
+ if (state === 'active') {
805
+ state = 'settled';
806
+ }
807
+ }
808
+ };
809
+ await continueAfterGPUWork();
788
810
  };
789
811
 
790
- executeTile(0, 0);
812
+ unregisterDeviceLoss = this._registerExecutionFailure(settleError);
813
+ void executeTile(0).catch(settleError);
791
814
 
792
815
  return () => {
793
- aborted = true;
816
+ if (state !== 'active') return;
817
+ state = 'aborted';
818
+ cancelScheduledTile();
819
+ unregisterDeviceLoss();
794
820
  };
795
821
  }
796
822
 
797
823
  dispose() {
798
- this._tfModel?.dispose();
799
824
  this._dataProcessGPU?.dispose();
800
- this._tensors.forEach((tensor) => tensor.dispose());
825
+ this._nativeExecutor?.dispose();
826
+ this._webNNExecutor?.dispose();
801
827
  }
802
828
  }
803
829