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/lib/UNet.d.ts CHANGED
@@ -1,8 +1,8 @@
1
1
  import { HostTensor } from './tza';
2
- import { Tile } from './process';
2
+ import { Tile, type HDRTransfer } from './process';
3
3
  import { type DynamicTileSetting } from './tileScheduler';
4
4
  import { type UNetModelSpec } from './modelSpec';
5
- import { type NativeUNetKernelSetting, type NativeUNetPrecisionSetting } from './nativeUNet';
5
+ import { type NativeUNetGemmOptions, type NativeUNetKernelSetting, type NativeUNetPrecisionSetting } from './nativeUNet';
6
6
  export type UNetEngineSetting = 'auto' | 'wgsl' | 'webnn';
7
7
  interface HDRImageData {
8
8
  data: Float32Array;
@@ -22,9 +22,12 @@ interface GPUImageDataOutput {
22
22
  export interface UNetExecutionStats {
23
23
  width: number;
24
24
  height: number;
25
- tileWidth: number;
26
- tileHeight: number;
27
25
  tileCount: number;
26
+ tileColumns: number;
27
+ tileRows: number;
28
+ tileOverlap: number;
29
+ inputPixelCount: number;
30
+ inputShapeCount: number;
28
31
  durationMs: number;
29
32
  tileTimeMs: {
30
33
  min: number;
@@ -35,12 +38,9 @@ export interface UNetExecutionStats {
35
38
  }
36
39
  declare class UNet {
37
40
  private _device;
38
- private _tileWidth;
39
- private _tileHeight;
40
- private _tileOverlapX;
41
- private _tileOverlapY;
42
41
  private _aux;
43
42
  private _hdr;
43
+ private _hdrTransfer;
44
44
  private _dataProcessGPU?;
45
45
  private _nativeExecutor?;
46
46
  private _webNNExecutor?;
@@ -49,10 +49,11 @@ declare class UNet {
49
49
  private _engine;
50
50
  private _dynamicTileController;
51
51
  private _lastExecution?;
52
- constructor(hostTensors: Map<string, HostTensor>, backend: {
53
- device: GPUDevice;
54
- adapterInfo: GPUAdapterInfo;
55
- }, opts?: {
52
+ private _activeExecutionFailures;
53
+ private _deviceLostObserved;
54
+ private _deviceLostSettled;
55
+ private _deviceLostReason;
56
+ constructor(hostTensors: Map<string, HostTensor>, device: GPUDevice, opts?: {
56
57
  /**
57
58
  * If use auxiliary data.
58
59
  */
@@ -61,26 +62,38 @@ declare class UNet {
61
62
  * If input is HDR image.
62
63
  */
63
64
  hdr?: boolean;
65
+ /** HDR transfer function expected by the trained model. */
66
+ hdrTransfer?: HDRTransfer;
64
67
  maxTileSize?: number;
65
68
  dynamicTile?: DynamicTileSetting;
66
- /** Reserved for explicit native WGSL selection. */
69
+ /** Native WGSL or the experimental WebNN backend. */
67
70
  engine?: UNetEngineSetting;
68
- /** Arithmetic/storage precision used by the native WGSL engine. */
71
+ /** Arithmetic/storage precision used by the native WGSL executor. */
69
72
  precision?: NativeUNetPrecisionSetting;
70
73
  /** Model-independent convolution kernel selection. */
71
74
  kernel?: NativeUNetKernelSetting;
75
+ gemm?: NativeUNetGemmOptions;
72
76
  /** Explicit descriptor for a new OIDN topology not in the built-in registry. */
73
77
  modelSpec?: UNetModelSpec;
74
78
  });
75
- getDevice(): GPUDevice | undefined;
79
+ getDevice(): GPUDevice;
76
80
  /** Completes backend compilation before first interactive use. */
77
81
  prepare(): Promise<void>;
82
+ /**
83
+ * Prepares the input shapes selected for an image before its first denoise.
84
+ * Hosts can call this while they still display their model-loading state.
85
+ */
86
+ prepareForImage(width: number, height: number, options?: {
87
+ tileOverlap?: number;
88
+ wholeImage?: boolean;
89
+ }): Promise<void>;
78
90
  getRuntimeInfo(): {
79
91
  configuredEngine: UNetEngineSetting;
80
92
  gpuEngine: "wgsl" | "webnn";
81
93
  precision: import("./nativeUNet").NativeUNetPrecision;
82
94
  kernel: {
83
95
  configured: NativeUNetKernelSetting;
96
+ gemm: Readonly<Required<NativeUNetGemmOptions>>;
84
97
  maxSpatialInputBlocks: number;
85
98
  subgroupsAvailable: boolean;
86
99
  } | undefined;
@@ -89,6 +102,7 @@ declare class UNet {
89
102
  model: string;
90
103
  modelFamily: (string & {}) | "oidn-unet-small" | "oidn-unet-large";
91
104
  inputChannels: number;
105
+ hdrTransfer: HDRTransfer;
92
106
  dynamicTile: {
93
107
  enabled: boolean;
94
108
  currentTileSize: number;
@@ -97,17 +111,18 @@ declare class UNet {
97
111
  targetTileTimeMs: number;
98
112
  };
99
113
  lastExecution: UNetExecutionStats | undefined;
114
+ activeExecutionCount: number;
100
115
  };
116
+ private _observeDeviceLoss;
117
+ private _registerExecutionFailure;
101
118
  /** Captures per-node GPU timestamps for the next native tile execution. */
102
119
  profileNextExecution(): boolean;
103
120
  getLastExecutionProfile(): Promise<import("./nativeUNet").NativeUNetExecutionProfile> | undefined;
104
- private _updateModel;
105
- private _getTileSizeWithOverlap;
106
121
  private _processImageData;
107
122
  private _readTile;
108
123
  private _writeTile;
109
124
  private _executeTile;
110
- tileExecute<T extends ImageData | HDRImageData | GPUImageData>({ color, albedo, normal, done, progress, denoiseAlpha }: {
125
+ tileExecute<T extends ImageData | HDRImageData | GPUImageData>({ color, albedo, normal, done, progress, denoiseAlpha, tileOverlap, wholeImage, scheduling, error }: {
111
126
  color: T;
112
127
  albedo?: ImageData | GPUImageData;
113
128
  normal?: ImageData | GPUImageData;
@@ -115,8 +130,27 @@ declare class UNet {
115
130
  * If denoise alpha channel. Otherwise denoise RGB channels.
116
131
  */
117
132
  denoiseAlpha?: boolean;
118
- done: (outputData: T extends GPUImageData ? GPUImageDataOutput : T) => void;
119
- progress?: (outputData: T extends GPUImageData ? GPUImageDataOutput : T, tileData: (T extends GPUImageData ? GPUImageDataOutput : T) | undefined, tile: Tile, currentIdx: number, totalIdx: number) => void;
133
+ /**
134
+ * Execute the complete input image as one tile, ignoring `maxTileSize`.
135
+ * The image must fit the device's buffer and dispatch limits.
136
+ */
137
+ wholeImage?: boolean;
138
+ /**
139
+ * Per-side context for boundaries shared with another tile. Defaults to
140
+ * half of the model receptive field rounded up to 16 pixels.
141
+ */
142
+ tileOverlap?: number;
143
+ /**
144
+ * How JavaScript yields between completed GPU tiles. `event-loop`
145
+ * (default) continues on the next macrotask. `animation-frame` waits for
146
+ * the next display frame, bounded by a short timer so hidden pages still
147
+ * complete.
148
+ */
149
+ scheduling?: 'animation-frame' | 'event-loop';
150
+ done: (outputData: T extends GPUImageData ? GPUImageDataOutput : T) => void | Promise<void>;
151
+ /** Receives asynchronous execution, queue, and callback failures. */
152
+ error?: (reason: unknown) => void | Promise<void>;
153
+ progress?: (outputData: T extends GPUImageData ? GPUImageDataOutput : T, tileData: (T extends GPUImageData ? GPUImageDataOutput : T) | undefined, tile: Tile, currentIdx: number, totalIdx: number) => void | Promise<void>;
120
154
  }): () => void;
121
155
  dispose(): void;
122
156
  }
package/lib/UNet.js CHANGED
@@ -1,8 +1,10 @@
1
1
  import { GPUDataProcess, Tile, avgLogLum, hdrTransferFuncCPU, hdrTransferFuncInverseCPU } from './process';
2
- import { DynamicTileController, fitTileDimension, OIDN_TILE_ALIGNMENT, waitForSubmittedGPUWork } from './tileScheduler';
2
+ import { DynamicTileController, planTileGrid, OIDN_TILE_ALIGNMENT } from './tileScheduler';
3
3
  import { detectUNetModelSpec, validateUNetModel } from './modelSpec';
4
4
  import { NativeUNetExecutor } from './nativeUNet';
5
5
  import { WebNNUNetExecutor } from './webnnUNet';
6
+ /** Upper bound on waiting for a display frame between tiles. */
7
+ const ANIMATION_FRAME_FALLBACK_MS = 100;
6
8
  function roundUp(a, b) {
7
9
  return Math.ceil(a / b) * b;
8
10
  }
@@ -11,14 +13,9 @@ function isGPUImageData(data) {
11
13
  }
12
14
  class UNet {
13
15
  _device;
14
- // TODO calculate the tile size from memory size
15
- // https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/core/unet_filter.cpp#L287
16
- _tileWidth = 0;
17
- _tileHeight = 0;
18
- _tileOverlapX = 0;
19
- _tileOverlapY = 0;
20
16
  _aux;
21
17
  _hdr;
18
+ _hdrTransfer;
22
19
  _dataProcessGPU;
23
20
  _nativeExecutor;
24
21
  _webNNExecutor;
@@ -27,9 +24,14 @@ class UNet {
27
24
  _engine;
28
25
  _dynamicTileController;
29
26
  _lastExecution;
30
- constructor(hostTensors, backend, opts = {}) {
27
+ _activeExecutionFailures = new Set();
28
+ _deviceLostObserved = false;
29
+ _deviceLostSettled = false;
30
+ _deviceLostReason;
31
+ constructor(hostTensors, device, opts = {}) {
31
32
  this._aux = opts.aux || false;
32
33
  this._hdr = opts.hdr || false;
34
+ this._hdrTransfer = opts.hdrTransfer ?? 'pu';
33
35
  this._engine = opts.engine ?? 'auto';
34
36
  const modelSpec = opts.modelSpec ?? detectUNetModelSpec(hostTensors);
35
37
  const validatedModel = validateUNetModel(hostTensors, modelSpec);
@@ -41,12 +43,13 @@ class UNet {
41
43
  `but aux=${this._aux} provides ${expectedInputChannels}`);
42
44
  }
43
45
  this._dynamicTileController = new DynamicTileController(opts.maxTileSize ?? 512, opts.dynamicTile);
44
- this._device = backend.device;
46
+ this._device = device;
47
+ this._observeDeviceLoss();
45
48
  if (this._engine === 'webnn') {
46
49
  this._webNNExecutor = new WebNNUNetExecutor(this._device, validatedModel, { precision: opts.precision });
47
50
  }
48
51
  else {
49
- this._nativeExecutor = new NativeUNetExecutor(this._device, validatedModel, { precision: opts.precision, kernel: opts.kernel });
52
+ this._nativeExecutor = new NativeUNetExecutor(this._device, validatedModel, { precision: opts.precision, kernel: opts.kernel, gemm: opts.gemm });
50
53
  }
51
54
  }
52
55
  getDevice() {
@@ -69,6 +72,25 @@ class UNet {
69
72
  }
70
73
  await this._nativeExecutor.prepare();
71
74
  }
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);
93
+ }
72
94
  getRuntimeInfo() {
73
95
  return {
74
96
  configuredEngine: this._engine,
@@ -77,6 +99,7 @@ class UNet {
77
99
  kernel: this._nativeExecutor
78
100
  ? {
79
101
  configured: this._nativeExecutor.kernelSetting,
102
+ gemm: this._nativeExecutor.gemm,
80
103
  maxSpatialInputBlocks: this._nativeExecutor.maxSpatialInputBlocks,
81
104
  subgroupsAvailable: this._nativeExecutor.subgroupsAvailable
82
105
  }
@@ -86,6 +109,7 @@ class UNet {
86
109
  model: this._modelSpec.id,
87
110
  modelFamily: this._modelSpec.family,
88
111
  inputChannels: this._inputChannels,
112
+ hdrTransfer: this._hdrTransfer,
89
113
  dynamicTile: {
90
114
  enabled: this._dynamicTileController.enabled,
91
115
  currentTileSize: this._dynamicTileController.tileSize,
@@ -93,8 +117,39 @@ class UNet {
93
117
  maxTileSize: this._dynamicTileController.maxTileSize,
94
118
  targetTileTimeMs: this._dynamicTileController.targetTileTimeMs
95
119
  },
96
- lastExecution: this._lastExecution
120
+ lastExecution: this._lastExecution,
121
+ activeExecutionCount: this._activeExecutionFailures.size
122
+ };
123
+ }
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);
97
142
  };
143
+ void deviceLost.then((info) => failAll(new Error(`WebGPU device lost: ${info.message}`)), failAll);
144
+ }
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));
151
+ }
152
+ return () => this._activeExecutionFailures.delete(fail);
98
153
  }
99
154
  /** Captures per-node GPU timestamps for the next native tile execution. */
100
155
  profileNextExecution() {
@@ -103,43 +158,6 @@ class UNet {
103
158
  getLastExecutionProfile() {
104
159
  return this._nativeExecutor?.getLastExecutionProfile();
105
160
  }
106
- _updateModel(width, height) {
107
- const maxTileSize = this._dynamicTileController.tileSize;
108
- let tileWidth = fitTileDimension(width, maxTileSize);
109
- let tileHeight = fitTileDimension(height, maxTileSize);
110
- const defaultTileOverlap = roundUp(this._modelSpec.receptiveField / 2, OIDN_TILE_ALIGNMENT);
111
- let tileOverlapX = defaultTileOverlap;
112
- let tileOverlapY = defaultTileOverlap;
113
- if (width <= maxTileSize) {
114
- tileOverlapX = 0;
115
- }
116
- if (height <= maxTileSize) {
117
- tileOverlapY = 0;
118
- }
119
- // Force width and height has same size. reduce the cache in memory
120
- const tileSize = Math.max(tileWidth, tileHeight);
121
- const tileOverlap = Math.max(tileOverlapX, tileOverlapY);
122
- tileWidth = tileSize;
123
- tileHeight = tileSize;
124
- tileOverlapX = tileOverlap;
125
- tileOverlapY = tileOverlap;
126
- if (tileWidth !== this._tileWidth ||
127
- tileHeight !== this._tileHeight ||
128
- tileOverlapX !== this._tileOverlapX ||
129
- tileOverlapY !== this._tileOverlapY) {
130
- // console.log(tileWidth, tileHeight, tileOverlapX, tileOverlapY);
131
- this._tileWidth = tileWidth;
132
- this._tileHeight = tileHeight;
133
- this._tileOverlapX = tileOverlapX;
134
- this._tileOverlapY = tileOverlapY;
135
- }
136
- }
137
- _getTileSizeWithOverlap() {
138
- return {
139
- width: this._tileWidth + 2 * this._tileOverlapX,
140
- height: this._tileHeight + 2 * this._tileOverlapY
141
- };
142
- }
143
161
  _processImageData(color, albedo, normal, isHDR) {
144
162
  const rawData = color.data;
145
163
  const pixelsCount = rawData.length / 4;
@@ -179,9 +197,12 @@ class UNet {
179
197
  }
180
198
  _readTile(data, channels, srcTile, width) {
181
199
  const tileData = new Float32Array(srcTile.width * srcTile.height * channels);
200
+ const height = data.length / (width * channels);
182
201
  for (let y = 0; y < srcTile.height; y++) {
183
202
  for (let x = 0; x < srcTile.width; x++) {
184
- 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;
185
206
  const i1 = (y * srcTile.width + x) * channels;
186
207
  for (let c = 0; c < channels; c++) {
187
208
  tileData[i1 + c] = data[i2 + c];
@@ -210,21 +231,12 @@ class UNet {
210
231
  }
211
232
  }
212
233
  }
213
- async _executeTile(inputData, outputTileData, outputImageData, i, j, width, height, isHDR, denoiseAlpha) {
234
+ async _executeTile(inputData, outputTileData, outputImageData, tile, isFirstTile, width, height, isHDR, denoiseAlpha) {
214
235
  const channels = this._aux ? 9 : 3;
215
- const tileOverlapX = this._tileOverlapX;
216
- const tileOverlapY = this._tileOverlapY;
217
- let srcTileSize = this._getTileSizeWithOverlap();
218
- let dstTileSize = { width: this._tileWidth, height: this._tileHeight };
219
- let srcX0 = i > 0 ? i * dstTileSize.width - tileOverlapX : 0;
220
- let srcX1 = Math.min(srcX0 + srcTileSize.width, width);
221
- srcX0 = Math.max(srcX1 - srcTileSize.width, 0);
222
- let srcY0 = j > 0 ? j * dstTileSize.height - tileOverlapY : 0;
223
- let srcY1 = Math.min(srcY0 + srcTileSize.height, height);
224
- srcY0 = Math.max(srcY1 - srcTileSize.height, 0);
225
- const srcTileWidth = srcTileSize.width;
226
- const srcTileHeight = srcTileSize.height;
227
- const srcTile = new Tile(srcX0, srcY0, srcTileWidth, srcTileHeight);
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;
228
240
  let nativeOutputBuffer;
229
241
  let denoisedData;
230
242
  let inputScale = 1;
@@ -240,42 +252,39 @@ class UNet {
240
252
  tileData = hdrTransferFuncCPU({
241
253
  data: tileData,
242
254
  channels,
243
- inputScale
255
+ inputScale,
256
+ transfer: this._hdrTransfer
244
257
  });
245
258
  }
246
259
  denoisedData = await (this._webNNExecutor ?? this._nativeExecutor).executeCPU(tileData, srcTileWidth, srcTileHeight);
247
260
  }
248
261
  else {
249
262
  if (!dataProcessGPU) {
250
- dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(device, isHDR);
263
+ dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(device, isHDR, this._hdrTransfer);
251
264
  }
252
265
  dataProcessGPU.setImageSize(width, height);
253
266
  dataProcessGPU.setInputTile(srcTile);
254
267
  // Display the noisy input instead of prev denoised result
255
- if (i === 0 && j === 0) {
268
+ if (isFirstTile) {
256
269
  dataProcessGPU.copyInputDataToOutput(inputData.color);
257
270
  }
258
271
  const { color, albedo, normal } = dataProcessGPU.forward(inputData.color, this._aux ? inputData.albedo : undefined, this._aux ? inputData.normal : undefined, denoiseAlpha);
259
272
  nativeOutputBuffer = await (this._webNNExecutor ?? this._nativeExecutor).execute(this._aux ? [color, albedo, normal] : [color], srcTileWidth, srcTileHeight);
260
273
  }
261
274
  let outBuffer;
262
- const dstWidth = Math.min(dstTileSize.width, width);
263
- const dstHeight = Math.min(dstTileSize.height, height);
264
- const dstTile = new Tile(i * dstWidth, j * dstHeight, dstWidth, dstHeight);
265
- dstTile.width = Math.min(dstTile.width, width - dstTile.x);
266
- dstTile.height = Math.min(dstTile.height, height - dstTile.y);
267
275
  if (inputData instanceof Float32Array) {
268
276
  if (isHDR) {
269
277
  denoisedData = hdrTransferFuncInverseCPU({
270
278
  data: denoisedData,
271
279
  channels: 3,
272
- inputScale
280
+ inputScale,
281
+ transfer: this._hdrTransfer
273
282
  });
274
283
  }
275
- this._writeTile(outputImageData, srcTile, dstTile, denoisedData, srcTileSize.width, isHDR);
276
- for (let y = 0; y < dstHeight; y++) {
277
- for (let x = 0; x < dstWidth; x++) {
278
- 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;
279
288
  const i2 = ((y + dstTile.y) * width + (x + dstTile.x)) * 4;
280
289
  for (let c = 0; c < 4; c++) {
281
290
  outputTileData.data[i1 + c] = outputImageData.data[i2 + c];
@@ -289,7 +298,7 @@ class UNet {
289
298
  }
290
299
  return outBuffer;
291
300
  }
292
- tileExecute({ color, albedo, normal, done, progress, denoiseAlpha }) {
301
+ tileExecute({ color, albedo, normal, done, progress, denoiseAlpha, tileOverlap, wholeImage, scheduling = 'event-loop', error }) {
293
302
  if (this._aux && (!albedo || !normal)) {
294
303
  throw new Error('Normal map and albedo map are both required');
295
304
  }
@@ -301,8 +310,16 @@ class UNet {
301
310
  const width = color.width;
302
311
  const height = color.height;
303
312
  const adaptiveTileSize = this._dynamicTileController.tileSize;
304
- const shouldAdaptTileSize = width > adaptiveTileSize || height > adaptiveTileSize;
305
- this._updateModel(width, height);
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;
306
323
  // TODO should fixed to be hdr when UNet is created.
307
324
  // weights of hdr and ldr is different
308
325
  const hdr = this._hdr || false;
@@ -310,10 +327,6 @@ class UNet {
310
327
  if (!isGPUImageData(color)) {
311
328
  rawData = this._processImageData(color, albedo, normal, hdr);
312
329
  }
313
- const tileWidth = this._tileWidth;
314
- const tileHeight = this._tileHeight;
315
- const tileCountH = Math.ceil(height / tileHeight);
316
- const tileCountW = Math.ceil(width / tileWidth);
317
330
  function makeImageData(width, height) {
318
331
  return hdr
319
332
  ? {
@@ -326,25 +339,75 @@ class UNet {
326
339
  const outputImageData = isGPUImageData(color)
327
340
  ? undefined
328
341
  : makeImageData(width, height);
329
- const outputTileData = isGPUImageData(color)
330
- ? undefined
331
- : makeImageData(Math.min(tileWidth, width), Math.min(tileHeight, height));
332
- let aborted = false;
342
+ let state = 'active';
343
+ let scheduledTimer;
344
+ let scheduledAnimationFrame;
345
+ let unregisterDeviceLoss = () => false;
333
346
  const now = () => typeof performance === 'undefined' ? Date.now() : performance.now();
334
347
  const executionStartTime = now();
335
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')
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
+ };
336
383
  const scheduleNextTile = (callback) => {
337
- if (typeof requestAnimationFrame === 'undefined') {
338
- setTimeout(callback, 0);
384
+ if (scheduling === 'event-loop' ||
385
+ typeof requestAnimationFrame === 'undefined') {
386
+ scheduledTimer = setTimeout(() => {
387
+ scheduledTimer = undefined;
388
+ callback();
389
+ }, 0);
339
390
  }
340
391
  else {
341
- requestAnimationFrame(callback);
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);
342
401
  }
343
402
  };
344
- const executeTile = async (i, j) => {
345
- if (aborted) {
403
+ const executeTile = async (tileIndex) => {
404
+ if (state !== 'active') {
346
405
  return;
347
406
  }
407
+ const tile = plan.tiles[tileIndex];
408
+ const outputTileData = isGPUImageData(color)
409
+ ? undefined
410
+ : makeImageData(tile.output.width, tile.output.height);
348
411
  const tileStartTime = now();
349
412
  const resGPUBuffer = await this._executeTile(isGPUImageData(color)
350
413
  ? {
@@ -352,32 +415,34 @@ class UNet {
352
415
  albedo: albedo?.data,
353
416
  normal: normal?.data
354
417
  }
355
- : rawData, outputTileData, outputImageData, i, j, width, height, hdr, denoiseAlpha);
356
- if (aborted)
418
+ : rawData, outputTileData, outputImageData, tile, tileIndex === 0, width, height, hdr, denoiseAlpha);
419
+ if (state !== 'active')
357
420
  return;
358
421
  const output = outputImageData || {
359
422
  data: resGPUBuffer,
360
423
  width,
361
424
  height
362
425
  };
363
- progress?.(output,
364
- // Is undefined if using webgpu buffer
365
- outputTileData, new Tile(i * tileWidth, j * tileHeight, tileWidth, tileHeight), i + j * tileCountW, tileCountW * tileCountH);
366
- const hasNextTile = i + 1 < tileCountW || j + 1 < tileCountH;
367
- const continueAfterGPUWork = () => {
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 () => {
368
438
  tileTimesMs.push(now() - tileStartTime);
369
- if (aborted)
439
+ if (state !== 'active')
370
440
  return;
371
441
  if (hasNextTile) {
372
442
  scheduleNextTile(() => {
373
- if (aborted)
443
+ if (state !== 'active')
374
444
  return;
375
- if (i + 1 < tileCountW) {
376
- executeTile(i + 1, j);
377
- }
378
- else if (j + 1 < tileCountH) {
379
- executeTile(0, j + 1);
380
- }
445
+ void executeTile(tileIndex + 1).catch(settleError);
381
446
  });
382
447
  }
383
448
  else {
@@ -389,9 +454,12 @@ class UNet {
389
454
  this._lastExecution = {
390
455
  width,
391
456
  height,
392
- tileWidth,
393
- tileHeight,
394
- tileCount: tileCountW * tileCountH,
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,
395
463
  durationMs: now() - executionStartTime,
396
464
  tileTimeMs: {
397
465
  min: sortedTileTimes[0],
@@ -406,18 +474,26 @@ class UNet {
406
474
  if (shouldAdaptTileSize) {
407
475
  this._dynamicTileController.observe(tileTimesMs);
408
476
  }
409
- // console.log(memory());
410
- done(output);
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';
484
+ }
411
485
  }
412
486
  };
413
- // requestAnimationFrame only throttles JavaScript submission. Waiting
414
- // for the queue here keeps at most one OIDN tile in flight, so aborting
415
- // cannot leave a long tail of already-submitted GPU work.
416
- void waitForSubmittedGPUWork(this._device.queue).then(continueAfterGPUWork);
487
+ await continueAfterGPUWork();
417
488
  };
418
- executeTile(0, 0);
489
+ unregisterDeviceLoss = this._registerExecutionFailure(settleError);
490
+ void executeTile(0).catch(settleError);
419
491
  return () => {
420
- aborted = true;
492
+ if (state !== 'active')
493
+ return;
494
+ state = 'aborted';
495
+ cancelScheduledTile();
496
+ unregisterDeviceLoss();
421
497
  };
422
498
  }
423
499
  dispose() {