oidn-web 0.4.0 → 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 (106) hide show
  1. package/CHANGELOG.md +45 -0
  2. package/README.md +83 -19
  3. package/dist/oidn.js +3239 -2642
  4. package/dist/oidn.umd.cjs +512 -296
  5. package/lib/UNet.d.ts +54 -20
  6. package/lib/UNet.js +194 -118
  7. package/lib/UNet.js.map +1 -1
  8. package/lib/backend.d.ts +1 -8
  9. package/lib/backend.js +1 -9
  10. package/lib/backend.js.map +1 -1
  11. package/lib/finalRgbShader.d.ts +13 -0
  12. package/lib/finalRgbShader.js +160 -0
  13. package/lib/finalRgbShader.js.map +1 -0
  14. package/lib/graphOptimizer.js +1 -2
  15. package/lib/graphOptimizer.js.map +1 -1
  16. package/lib/hdrTransfer.d.ts +14 -0
  17. package/lib/hdrTransfer.js +61 -0
  18. package/lib/hdrTransfer.js.map +1 -0
  19. package/lib/main.d.ts +16 -7
  20. package/lib/main.js +4 -5
  21. package/lib/main.js.map +1 -1
  22. package/lib/nativeUNet.d.ts +39 -3
  23. package/lib/nativeUNet.js +449 -120
  24. package/lib/nativeUNet.js.map +1 -1
  25. package/lib/process.d.ts +5 -11
  26. package/lib/process.js +35 -49
  27. package/lib/process.js.map +1 -1
  28. package/lib/tileScheduler.d.ts +32 -4
  29. package/lib/tileScheduler.js +133 -20
  30. package/lib/tileScheduler.js.map +1 -1
  31. package/package.json +9 -2
  32. package/src/UNet.ts +287 -158
  33. package/src/backend.ts +1 -14
  34. package/src/finalRgbShader.ts +186 -0
  35. package/src/graphOptimizer.ts +1 -2
  36. package/src/hdrTransfer.ts +88 -0
  37. package/src/main.ts +28 -13
  38. package/src/nativeUNet.ts +515 -116
  39. package/src/process.ts +43 -70
  40. package/src/tileScheduler.ts +216 -24
  41. package/benchmarks/compare.mjs +0 -651
  42. package/benchmarks/leak.mjs +0 -255
  43. package/benchmarks/results/before-spatial.json +0 -391
  44. package/benchmarks/results/before-spatial.md +0 -47
  45. package/benchmarks/results/int8-scan.json +0 -2007
  46. package/benchmarks/results/int8-scan.md +0 -160
  47. package/benchmarks/results/int8-w8a8-scan.json +0 -2007
  48. package/benchmarks/results/int8-w8a8-scan.md +0 -160
  49. package/benchmarks/results/int8-weight-channel.json +0 -1413
  50. package/benchmarks/results/int8-weight-channel.md +0 -118
  51. package/benchmarks/results/int8-weight-only.json +0 -1437
  52. package/benchmarks/results/int8-weight-only.md +0 -118
  53. package/benchmarks/results/kernel-webnn-final.json +0 -1115
  54. package/benchmarks/results/kernel-webnn-final.md +0 -104
  55. package/benchmarks/results/latest-optimized.json +0 -375
  56. package/benchmarks/results/latest-optimized.md +0 -47
  57. package/benchmarks/results/latest.json +0 -391
  58. package/benchmarks/results/latest.md +0 -47
  59. package/benchmarks/results/profile-baseline.json +0 -331
  60. package/benchmarks/results/profile-baseline.md +0 -12
  61. package/benchmarks/results/profile-conv2x.json +0 -331
  62. package/benchmarks/results/profile-conv2x.md +0 -12
  63. package/benchmarks/results/profile-fast-init.json +0 -385
  64. package/benchmarks/results/profile-fast-init.md +0 -47
  65. package/benchmarks/results/profile-fp16-fma.json +0 -369
  66. package/benchmarks/results/profile-fp16-fma.md +0 -47
  67. package/benchmarks/results/profile-fp16-tiled-decoder.json +0 -369
  68. package/benchmarks/results/profile-fp16-tiled-decoder.md +0 -47
  69. package/benchmarks/results/profile-fp16-tiled-encoder.json +0 -369
  70. package/benchmarks/results/profile-fp16-tiled-encoder.md +0 -47
  71. package/benchmarks/results/profile-fp16-unfused-pool.json +0 -385
  72. package/benchmarks/results/profile-fp16-unfused-pool.md +0 -47
  73. package/benchmarks/results/profile-input-major.json +0 -347
  74. package/benchmarks/results/profile-input-major.md +0 -12
  75. package/benchmarks/results/profile-k16.json +0 -347
  76. package/benchmarks/results/profile-k16.md +0 -12
  77. package/benchmarks/results/profile-k4.json +0 -347
  78. package/benchmarks/results/profile-k4.md +0 -12
  79. package/benchmarks/results/profile-pool-reuse.json +0 -331
  80. package/benchmarks/results/profile-pool-reuse.md +0 -12
  81. package/benchmarks/results/profile-precompiled.json +0 -385
  82. package/benchmarks/results/profile-precompiled.md +0 -47
  83. package/benchmarks/results/profile-static-channels.json +0 -385
  84. package/benchmarks/results/profile-static-channels.md +0 -47
  85. package/benchmarks/results/profile-static-io.json +0 -385
  86. package/benchmarks/results/profile-static-io.md +0 -47
  87. package/benchmarks/results/profile-tiled-conv.json +0 -331
  88. package/benchmarks/results/profile-tiled-conv.md +0 -12
  89. package/benchmarks/results/profile-tiled-decoder.json +0 -347
  90. package/benchmarks/results/profile-tiled-decoder.md +0 -12
  91. package/benchmarks/results/profile-tiled-matmul.json +0 -331
  92. package/benchmarks/results/profile-tiled-matmul.md +0 -12
  93. package/benchmarks/results/profile-unfused-decoder.json +0 -379
  94. package/benchmarks/results/profile-unfused-decoder.md +0 -12
  95. package/benchmarks/results/profile-unfused-pool.json +0 -347
  96. package/benchmarks/results/profile-unfused-pool.md +0 -12
  97. package/benchmarks/results/spatial-auto.json +0 -575
  98. package/benchmarks/results/spatial-auto.md +0 -61
  99. package/benchmarks/results/subgroup-smoke.json +0 -1094
  100. package/benchmarks/results/subgroup-smoke.md +0 -104
  101. package/benchmarks/results/webnn-smoke.json +0 -739
  102. package/benchmarks/results/webnn-smoke.md +0 -76
  103. package/scripts/inspect-model.mjs +0 -64
  104. package/tests/modelSpec.test.mjs +0 -128
  105. package/tests/resourceLifecycle.test.mjs +0 -383
  106. package/tests/tileScheduler.test.mjs +0 -90
package/src/UNet.ts CHANGED
@@ -4,14 +4,15 @@ import {
4
4
  Tile,
5
5
  avgLogLum,
6
6
  hdrTransferFuncCPU,
7
- hdrTransferFuncInverseCPU
7
+ hdrTransferFuncInverseCPU,
8
+ type HDRTransfer
8
9
  } from './process';
9
10
  import {
10
11
  DynamicTileController,
11
12
  type DynamicTileSetting,
12
- fitTileDimension,
13
- OIDN_TILE_ALIGNMENT,
14
- waitForSubmittedGPUWork
13
+ planTileGrid,
14
+ type PlannedTile,
15
+ OIDN_TILE_ALIGNMENT
15
16
  } from './tileScheduler';
16
17
  import {
17
18
  detectUNetModelSpec,
@@ -20,6 +21,7 @@ import {
20
21
  } from './modelSpec';
21
22
  import {
22
23
  NativeUNetExecutor,
24
+ type NativeUNetGemmOptions,
23
25
  type NativeUNetKernelSetting,
24
26
  type NativeUNetPrecisionSetting
25
27
  } from './nativeUNet';
@@ -27,6 +29,9 @@ import { WebNNUNetExecutor } from './webnnUNet';
27
29
 
28
30
  export type UNetEngineSetting = 'auto' | 'wgsl' | 'webnn';
29
31
 
32
+ /** Upper bound on waiting for a display frame between tiles. */
33
+ const ANIMATION_FRAME_FALLBACK_MS = 100;
34
+
30
35
  interface HDRImageData {
31
36
  data: Float32Array;
32
37
  width: number;
@@ -48,9 +53,12 @@ interface GPUImageDataOutput {
48
53
  export interface UNetExecutionStats {
49
54
  width: number;
50
55
  height: number;
51
- tileWidth: number;
52
- tileHeight: number;
53
56
  tileCount: number;
57
+ tileColumns: number;
58
+ tileRows: number;
59
+ tileOverlap: number;
60
+ inputPixelCount: number;
61
+ inputShapeCount: number;
54
62
  durationMs: number;
55
63
  tileTimeMs: {
56
64
  min: number;
@@ -71,18 +79,11 @@ function isGPUImageData(
71
79
  }
72
80
 
73
81
  class UNet {
74
- private _device: GPUDevice | undefined;
75
-
76
- // TODO calculate the tile size from memory size
77
- // https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/core/unet_filter.cpp#L287
78
- private _tileWidth = 0;
79
- private _tileHeight = 0;
80
-
81
- private _tileOverlapX = 0;
82
- private _tileOverlapY = 0;
82
+ private _device: GPUDevice;
83
83
 
84
84
  private _aux;
85
85
  private _hdr;
86
+ private _hdrTransfer: HDRTransfer;
86
87
 
87
88
  private _dataProcessGPU?: GPUDataProcess;
88
89
  private _nativeExecutor?: NativeUNetExecutor;
@@ -93,10 +94,14 @@ class UNet {
93
94
 
94
95
  private _dynamicTileController: DynamicTileController;
95
96
  private _lastExecution?: UNetExecutionStats;
97
+ private _activeExecutionFailures = new Set<(reason: unknown) => void>();
98
+ private _deviceLostObserved = false;
99
+ private _deviceLostSettled = false;
100
+ private _deviceLostReason: unknown;
96
101
 
97
102
  constructor(
98
103
  hostTensors: Map<string, HostTensor>,
99
- backend: { device: GPUDevice; adapterInfo: GPUAdapterInfo },
104
+ device: GPUDevice,
100
105
  opts: {
101
106
  /**
102
107
  * If use auxiliary data.
@@ -106,20 +111,24 @@ class UNet {
106
111
  * If input is HDR image.
107
112
  */
108
113
  hdr?: boolean;
114
+ /** HDR transfer function expected by the trained model. */
115
+ hdrTransfer?: HDRTransfer;
109
116
  maxTileSize?: number;
110
117
  dynamicTile?: DynamicTileSetting;
111
- /** Reserved for explicit native WGSL selection. */
118
+ /** Native WGSL or the experimental WebNN backend. */
112
119
  engine?: UNetEngineSetting;
113
- /** Arithmetic/storage precision used by the native WGSL engine. */
120
+ /** Arithmetic/storage precision used by the native WGSL executor. */
114
121
  precision?: NativeUNetPrecisionSetting;
115
122
  /** Model-independent convolution kernel selection. */
116
123
  kernel?: NativeUNetKernelSetting;
124
+ gemm?: NativeUNetGemmOptions;
117
125
  /** Explicit descriptor for a new OIDN topology not in the built-in registry. */
118
126
  modelSpec?: UNetModelSpec;
119
127
  } = {}
120
128
  ) {
121
129
  this._aux = opts.aux || false;
122
130
  this._hdr = opts.hdr || false;
131
+ this._hdrTransfer = opts.hdrTransfer ?? 'pu';
123
132
  this._engine = opts.engine ?? 'auto';
124
133
  const modelSpec = opts.modelSpec ?? detectUNetModelSpec(hostTensors);
125
134
  const validatedModel = validateUNetModel(hostTensors, modelSpec);
@@ -139,7 +148,8 @@ class UNet {
139
148
  opts.dynamicTile
140
149
  );
141
150
 
142
- this._device = backend.device;
151
+ this._device = device;
152
+ this._observeDeviceLoss();
143
153
  if (this._engine === 'webnn') {
144
154
  this._webNNExecutor = new WebNNUNetExecutor(
145
155
  this._device,
@@ -150,7 +160,7 @@ class UNet {
150
160
  this._nativeExecutor = new NativeUNetExecutor(
151
161
  this._device,
152
162
  validatedModel,
153
- { precision: opts.precision, kernel: opts.kernel }
163
+ { precision: opts.precision, kernel: opts.kernel, gemm: opts.gemm }
154
164
  );
155
165
  }
156
166
  }
@@ -182,6 +192,40 @@ class UNet {
182
192
  await this._nativeExecutor!.prepare();
183
193
  }
184
194
 
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 } = {}
203
+ ) {
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);
227
+ }
228
+
185
229
  getRuntimeInfo() {
186
230
  return {
187
231
  configuredEngine: this._engine,
@@ -190,6 +234,7 @@ class UNet {
190
234
  kernel: this._nativeExecutor
191
235
  ? {
192
236
  configured: this._nativeExecutor.kernelSetting,
237
+ gemm: this._nativeExecutor.gemm,
193
238
  maxSpatialInputBlocks: this._nativeExecutor.maxSpatialInputBlocks,
194
239
  subgroupsAvailable: this._nativeExecutor.subgroupsAvailable
195
240
  }
@@ -201,6 +246,7 @@ class UNet {
201
246
  model: this._modelSpec.id,
202
247
  modelFamily: this._modelSpec.family,
203
248
  inputChannels: this._inputChannels,
249
+ hdrTransfer: this._hdrTransfer,
204
250
  dynamicTile: {
205
251
  enabled: this._dynamicTileController.enabled,
206
252
  currentTileSize: this._dynamicTileController.tileSize,
@@ -208,10 +254,44 @@ class UNet {
208
254
  maxTileSize: this._dynamicTileController.maxTileSize,
209
255
  targetTileTimeMs: this._dynamicTileController.targetTileTimeMs
210
256
  },
211
- lastExecution: this._lastExecution
257
+ lastExecution: this._lastExecution,
258
+ activeExecutionCount: this._activeExecutionFailures.size
212
259
  };
213
260
  }
214
261
 
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
+ );
283
+ }
284
+
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));
291
+ }
292
+ return () => this._activeExecutionFailures.delete(fail);
293
+ }
294
+
215
295
  /** Captures per-node GPU timestamps for the next native tile execution. */
216
296
  profileNextExecution() {
217
297
  return this._nativeExecutor?.profileNextExecution() ?? false;
@@ -221,54 +301,6 @@ class UNet {
221
301
  return this._nativeExecutor?.getLastExecutionProfile();
222
302
  }
223
303
 
224
- private _updateModel(width: number, height: number) {
225
- const maxTileSize = this._dynamicTileController.tileSize;
226
-
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;
235
-
236
- if (width <= maxTileSize) {
237
- tileOverlapX = 0;
238
- }
239
- if (height <= maxTileSize) {
240
- tileOverlapY = 0;
241
- }
242
-
243
- // Force width and height has same size. reduce the cache in memory
244
- const tileSize = Math.max(tileWidth, tileHeight);
245
- const tileOverlap = Math.max(tileOverlapX, tileOverlapY);
246
- tileWidth = tileSize;
247
- tileHeight = tileSize;
248
- tileOverlapX = tileOverlap;
249
- tileOverlapY = tileOverlap;
250
-
251
- if (
252
- tileWidth !== this._tileWidth ||
253
- tileHeight !== this._tileHeight ||
254
- tileOverlapX !== this._tileOverlapX ||
255
- tileOverlapY !== this._tileOverlapY
256
- ) {
257
- // console.log(tileWidth, tileHeight, tileOverlapX, tileOverlapY);
258
- this._tileWidth = tileWidth;
259
- this._tileHeight = tileHeight;
260
- this._tileOverlapX = tileOverlapX;
261
- this._tileOverlapY = tileOverlapY;
262
- }
263
- }
264
-
265
- private _getTileSizeWithOverlap() {
266
- return {
267
- width: this._tileWidth + 2 * this._tileOverlapX,
268
- height: this._tileHeight + 2 * this._tileOverlapY
269
- };
270
- }
271
-
272
304
  private _processImageData(
273
305
  color: ImageData | HDRImageData,
274
306
  albedo: ImageData | undefined,
@@ -326,9 +358,12 @@ class UNet {
326
358
  const tileData = new Float32Array(
327
359
  srcTile.width * srcTile.height * channels
328
360
  );
361
+ const height = data.length / (width * channels);
329
362
  for (let y = 0; y < srcTile.height; y++) {
330
363
  for (let x = 0; x < srcTile.width; x++) {
331
- 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;
332
367
  const i1 = (y * srcTile.width + x) * channels;
333
368
 
334
369
  for (let c = 0; c < channels; c++) {
@@ -381,36 +416,33 @@ class UNet {
381
416
  },
382
417
  outputTileData: ImageData | HDRImageData | undefined,
383
418
  outputImageData: ImageData | HDRImageData | undefined,
384
- i: number,
385
- j: number,
419
+ tile: PlannedTile,
420
+ isFirstTile: boolean,
386
421
  width: number,
387
422
  height: number,
388
423
  isHDR: boolean,
389
424
  denoiseAlpha?: boolean
390
425
  ) {
391
426
  const channels = this._aux ? 9 : 3;
392
- const tileOverlapX = this._tileOverlapX;
393
- const tileOverlapY = this._tileOverlapY;
394
- let srcTileSize = this._getTileSizeWithOverlap();
395
- let dstTileSize = { width: this._tileWidth, height: this._tileHeight };
396
-
397
- let srcX0 = i > 0 ? i * dstTileSize.width - tileOverlapX : 0;
398
- let srcX1 = Math.min(srcX0 + srcTileSize.width, width);
399
- srcX0 = Math.max(srcX1 - srcTileSize.width, 0);
400
-
401
- let srcY0 = j > 0 ? j * dstTileSize.height - tileOverlapY : 0;
402
- let srcY1 = Math.min(srcY0 + srcTileSize.height, height);
403
- srcY0 = Math.max(srcY1 - srcTileSize.height, 0);
404
-
405
- const srcTileWidth = srcTileSize.width;
406
- const srcTileHeight = srcTileSize.height;
407
-
408
- 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;
409
441
 
410
442
  let nativeOutputBuffer: GPUBuffer | undefined;
411
443
  let denoisedData: Float32Array | undefined;
412
444
  let inputScale = 1;
413
- const device = this._device!;
445
+ const device = this._device;
414
446
  let dataProcessGPU = this._dataProcessGPU;
415
447
 
416
448
  if (inputData instanceof Float32Array) {
@@ -423,7 +455,8 @@ class UNet {
423
455
  tileData = hdrTransferFuncCPU({
424
456
  data: tileData,
425
457
  channels,
426
- inputScale
458
+ inputScale,
459
+ transfer: this._hdrTransfer
427
460
  });
428
461
  }
429
462
  denoisedData = await (this._webNNExecutor ?? this._nativeExecutor!).executeCPU(
@@ -435,13 +468,14 @@ class UNet {
435
468
  if (!dataProcessGPU) {
436
469
  dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(
437
470
  device,
438
- isHDR
471
+ isHDR,
472
+ this._hdrTransfer
439
473
  );
440
474
  }
441
475
  dataProcessGPU.setImageSize(width, height);
442
476
  dataProcessGPU.setInputTile(srcTile);
443
477
  // Display the noisy input instead of prev denoised result
444
- if (i === 0 && j === 0) {
478
+ if (isFirstTile) {
445
479
  dataProcessGPU.copyInputDataToOutput(inputData.color);
446
480
  }
447
481
  const { color, albedo, normal } = dataProcessGPU.forward(
@@ -460,18 +494,13 @@ class UNet {
460
494
 
461
495
  let outBuffer: GPUBuffer;
462
496
 
463
- const dstWidth = Math.min(dstTileSize.width, width);
464
- const dstHeight = Math.min(dstTileSize.height, height);
465
- const dstTile = new Tile(i * dstWidth, j * dstHeight, dstWidth, dstHeight);
466
- dstTile.width = Math.min(dstTile.width, width - dstTile.x);
467
- dstTile.height = Math.min(dstTile.height, height - dstTile.y);
468
-
469
497
  if (inputData instanceof Float32Array) {
470
498
  if (isHDR) {
471
499
  denoisedData = hdrTransferFuncInverseCPU({
472
500
  data: denoisedData!,
473
501
  channels: 3,
474
- inputScale
502
+ inputScale,
503
+ transfer: this._hdrTransfer
475
504
  });
476
505
  }
477
506
 
@@ -480,13 +509,13 @@ class UNet {
480
509
  srcTile,
481
510
  dstTile,
482
511
  denoisedData!,
483
- srcTileSize.width,
512
+ srcTile.width,
484
513
  isHDR
485
514
  );
486
515
 
487
- for (let y = 0; y < dstHeight; y++) {
488
- for (let x = 0; x < dstWidth; x++) {
489
- 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;
490
519
  const i2 = ((y + dstTile.y) * width + (x + dstTile.x)) * 4;
491
520
  for (let c = 0; c < 4; c++) {
492
521
  outputTileData!.data[i1 + c] = outputImageData!.data[i2 + c];
@@ -509,7 +538,11 @@ class UNet {
509
538
  normal,
510
539
  done,
511
540
  progress,
512
- denoiseAlpha
541
+ denoiseAlpha,
542
+ tileOverlap,
543
+ wholeImage,
544
+ scheduling = 'event-loop',
545
+ error
513
546
  }: {
514
547
  color: T;
515
548
  albedo?: ImageData | GPUImageData;
@@ -518,14 +551,35 @@ class UNet {
518
551
  * If denoise alpha channel. Otherwise denoise RGB channels.
519
552
  */
520
553
  denoiseAlpha?: boolean;
521
- 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>;
522
576
  progress?: (
523
577
  outputData: T extends GPUImageData ? GPUImageDataOutput : T,
524
578
  tileData: (T extends GPUImageData ? GPUImageDataOutput : T) | undefined,
525
579
  tile: Tile,
526
580
  currentIdx: number,
527
581
  totalIdx: number
528
- ) => void;
582
+ ) => void | Promise<void>;
529
583
  }): () => void {
530
584
  if (this._aux && (!albedo || !normal)) {
531
585
  throw new Error('Normal map and albedo map are both required');
@@ -540,9 +594,24 @@ class UNet {
540
594
  const width = color.width;
541
595
  const height = color.height;
542
596
  const adaptiveTileSize = this._dynamicTileController.tileSize;
543
- const shouldAdaptTileSize =
544
- width > adaptiveTileSize || height > adaptiveTileSize;
545
- this._updateModel(width, height);
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;
546
615
 
547
616
  // TODO should fixed to be hdr when UNet is created.
548
617
  // weights of hdr and ldr is different
@@ -557,11 +626,6 @@ class UNet {
557
626
  hdr
558
627
  );
559
628
  }
560
- const tileWidth = this._tileWidth;
561
- const tileHeight = this._tileHeight;
562
- const tileCountH = Math.ceil(height / tileHeight);
563
- const tileCountW = Math.ceil(width / tileWidth);
564
-
565
629
  function makeImageData(width: number, height: number) {
566
630
  return hdr
567
631
  ? {
@@ -575,28 +639,80 @@ class UNet {
575
639
  const outputImageData = isGPUImageData(color)
576
640
  ? undefined
577
641
  : makeImageData(width, height);
578
- const outputTileData = isGPUImageData(color)
579
- ? undefined
580
- : makeImageData(Math.min(tileWidth, width), Math.min(tileHeight, height));
581
642
 
582
- 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;
583
648
 
584
649
  const now = () =>
585
650
  typeof performance === 'undefined' ? Date.now() : performance.now();
586
651
  const executionStartTime = now();
587
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
+ };
588
686
  const scheduleNextTile = (callback: () => void) => {
589
- if (typeof requestAnimationFrame === 'undefined') {
590
- setTimeout(callback, 0);
687
+ if (
688
+ scheduling === 'event-loop' ||
689
+ typeof requestAnimationFrame === 'undefined'
690
+ ) {
691
+ scheduledTimer = setTimeout(() => {
692
+ scheduledTimer = undefined;
693
+ callback();
694
+ }, 0);
591
695
  } else {
592
- requestAnimationFrame(callback);
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);
593
705
  }
594
706
  };
595
707
 
596
- const executeTile = async (i: number, j: number) => {
597
- if (aborted) {
708
+ const executeTile = async (tileIndex: number) => {
709
+ if (state !== 'active') {
598
710
  return;
599
711
  }
712
+ const tile = plan.tiles[tileIndex];
713
+ const outputTileData = isGPUImageData(color)
714
+ ? undefined
715
+ : makeImageData(tile.output.width, tile.output.height);
600
716
  const tileStartTime = now();
601
717
  const resGPUBuffer = await this._executeTile(
602
718
  isGPUImageData(color)
@@ -608,41 +724,47 @@ class UNet {
608
724
  : rawData,
609
725
  outputTileData,
610
726
  outputImageData,
611
- i,
612
- j,
727
+ tile,
728
+ tileIndex === 0,
613
729
  width,
614
730
  height,
615
731
  hdr,
616
732
  denoiseAlpha
617
733
  );
618
- if (aborted) return;
734
+ if (state !== 'active') return;
619
735
  const output = outputImageData || {
620
736
  data: resGPUBuffer,
621
737
  width,
622
738
  height
623
739
  };
624
- progress?.(
625
- output as any,
626
- // Is undefined if using webgpu buffer
627
- outputTileData as any,
628
- new Tile(i * tileWidth, j * tileHeight, tileWidth, tileHeight),
629
- i + j * tileCountW,
630
- tileCountW * tileCountH
631
- );
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
+ );
754
+ }
755
+ if (state !== 'active') return;
632
756
 
633
- const hasNextTile = i + 1 < tileCountW || j + 1 < tileCountH;
634
- const continueAfterGPUWork = () => {
757
+ const hasNextTile = tileIndex + 1 < plan.tiles.length;
758
+ await this._device.queue.onSubmittedWorkDone();
759
+ if (state !== 'active') return;
760
+ const continueAfterGPUWork = async () => {
635
761
  tileTimesMs.push(now() - tileStartTime);
636
- if (aborted) return;
762
+ if (state !== 'active') return;
637
763
 
638
764
  if (hasNextTile) {
639
765
  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
- }
766
+ if (state !== 'active') return;
767
+ void executeTile(tileIndex + 1).catch(settleError);
646
768
  });
647
769
  } else {
648
770
  const sortedTileTimes = [...tileTimesMs].sort((a, b) => a - b);
@@ -653,9 +775,12 @@ class UNet {
653
775
  this._lastExecution = {
654
776
  width,
655
777
  height,
656
- tileWidth,
657
- tileHeight,
658
- tileCount: tileCountW * tileCountH,
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,
659
784
  durationMs: now() - executionStartTime,
660
785
  tileTimeMs: {
661
786
  min: sortedTileTimes[0],
@@ -671,23 +796,27 @@ class UNet {
671
796
  if (shouldAdaptTileSize) {
672
797
  this._dynamicTileController.observe(tileTimesMs);
673
798
  }
674
- // console.log(memory());
675
- done(output as any);
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
+ }
676
807
  }
677
808
  };
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
- );
809
+ await continueAfterGPUWork();
685
810
  };
686
811
 
687
- executeTile(0, 0);
812
+ unregisterDeviceLoss = this._registerExecutionFailure(settleError);
813
+ void executeTile(0).catch(settleError);
688
814
 
689
815
  return () => {
690
- aborted = true;
816
+ if (state !== 'active') return;
817
+ state = 'aborted';
818
+ cancelScheduledTile();
819
+ unregisterDeviceLoss();
691
820
  };
692
821
  }
693
822