oidn-web 0.3.4 → 0.4.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (122) hide show
  1. package/CHANGELOG.md +49 -0
  2. package/README.md +140 -4
  3. package/benchmarks/compare.mjs +651 -0
  4. package/benchmarks/leak.mjs +255 -0
  5. package/benchmarks/results/before-spatial.json +391 -0
  6. package/benchmarks/results/before-spatial.md +47 -0
  7. package/benchmarks/results/int8-scan.json +2007 -0
  8. package/benchmarks/results/int8-scan.md +160 -0
  9. package/benchmarks/results/int8-w8a8-scan.json +2007 -0
  10. package/benchmarks/results/int8-w8a8-scan.md +160 -0
  11. package/benchmarks/results/int8-weight-channel.json +1413 -0
  12. package/benchmarks/results/int8-weight-channel.md +118 -0
  13. package/benchmarks/results/int8-weight-only.json +1437 -0
  14. package/benchmarks/results/int8-weight-only.md +118 -0
  15. package/benchmarks/results/kernel-webnn-final.json +1115 -0
  16. package/benchmarks/results/kernel-webnn-final.md +104 -0
  17. package/benchmarks/results/latest-optimized.json +375 -0
  18. package/benchmarks/results/latest-optimized.md +47 -0
  19. package/benchmarks/results/latest.json +391 -0
  20. package/benchmarks/results/latest.md +47 -0
  21. package/benchmarks/results/profile-baseline.json +331 -0
  22. package/benchmarks/results/profile-baseline.md +12 -0
  23. package/benchmarks/results/profile-conv2x.json +331 -0
  24. package/benchmarks/results/profile-conv2x.md +12 -0
  25. package/benchmarks/results/profile-fast-init.json +385 -0
  26. package/benchmarks/results/profile-fast-init.md +47 -0
  27. package/benchmarks/results/profile-fp16-fma.json +369 -0
  28. package/benchmarks/results/profile-fp16-fma.md +47 -0
  29. package/benchmarks/results/profile-fp16-tiled-decoder.json +369 -0
  30. package/benchmarks/results/profile-fp16-tiled-decoder.md +47 -0
  31. package/benchmarks/results/profile-fp16-tiled-encoder.json +369 -0
  32. package/benchmarks/results/profile-fp16-tiled-encoder.md +47 -0
  33. package/benchmarks/results/profile-fp16-unfused-pool.json +385 -0
  34. package/benchmarks/results/profile-fp16-unfused-pool.md +47 -0
  35. package/benchmarks/results/profile-input-major.json +347 -0
  36. package/benchmarks/results/profile-input-major.md +12 -0
  37. package/benchmarks/results/profile-k16.json +347 -0
  38. package/benchmarks/results/profile-k16.md +12 -0
  39. package/benchmarks/results/profile-k4.json +347 -0
  40. package/benchmarks/results/profile-k4.md +12 -0
  41. package/benchmarks/results/profile-pool-reuse.json +331 -0
  42. package/benchmarks/results/profile-pool-reuse.md +12 -0
  43. package/benchmarks/results/profile-precompiled.json +385 -0
  44. package/benchmarks/results/profile-precompiled.md +47 -0
  45. package/benchmarks/results/profile-static-channels.json +385 -0
  46. package/benchmarks/results/profile-static-channels.md +47 -0
  47. package/benchmarks/results/profile-static-io.json +385 -0
  48. package/benchmarks/results/profile-static-io.md +47 -0
  49. package/benchmarks/results/profile-tiled-conv.json +331 -0
  50. package/benchmarks/results/profile-tiled-conv.md +12 -0
  51. package/benchmarks/results/profile-tiled-decoder.json +347 -0
  52. package/benchmarks/results/profile-tiled-decoder.md +12 -0
  53. package/benchmarks/results/profile-tiled-matmul.json +331 -0
  54. package/benchmarks/results/profile-tiled-matmul.md +12 -0
  55. package/benchmarks/results/profile-unfused-decoder.json +379 -0
  56. package/benchmarks/results/profile-unfused-decoder.md +12 -0
  57. package/benchmarks/results/profile-unfused-pool.json +347 -0
  58. package/benchmarks/results/profile-unfused-pool.md +12 -0
  59. package/benchmarks/results/spatial-auto.json +575 -0
  60. package/benchmarks/results/spatial-auto.md +61 -0
  61. package/benchmarks/results/subgroup-smoke.json +1094 -0
  62. package/benchmarks/results/subgroup-smoke.md +104 -0
  63. package/benchmarks/results/webnn-smoke.json +739 -0
  64. package/benchmarks/results/webnn-smoke.md +76 -0
  65. package/dist/oidn.js +4189 -22603
  66. package/dist/oidn.umd.cjs +784 -5807
  67. package/lib/UNet.d.ts +66 -15
  68. package/lib/UNet.js +162 -257
  69. package/lib/UNet.js.map +1 -1
  70. package/lib/WGPUComputePass.d.ts +1 -1
  71. package/lib/WGPUComputePass.js +6 -4
  72. package/lib/WGPUComputePass.js.map +1 -1
  73. package/lib/backend.d.ts +8 -4
  74. package/lib/backend.js +36 -44
  75. package/lib/backend.js.map +1 -1
  76. package/lib/graphOptimizer.d.ts +54 -0
  77. package/lib/graphOptimizer.js +216 -0
  78. package/lib/graphOptimizer.js.map +1 -0
  79. package/lib/main.d.ts +33 -10
  80. package/lib/main.js +5 -0
  81. package/lib/main.js.map +1 -1
  82. package/lib/modelSpec.d.ts +80 -0
  83. package/lib/modelSpec.js +270 -0
  84. package/lib/modelSpec.js.map +1 -0
  85. package/lib/nativeUNet.d.ts +67 -0
  86. package/lib/nativeUNet.js +1735 -0
  87. package/lib/nativeUNet.js.map +1 -0
  88. package/lib/process.js +38 -35
  89. package/lib/process.js.map +1 -1
  90. package/lib/resourceTracker.d.ts +26 -0
  91. package/lib/resourceTracker.js +65 -0
  92. package/lib/resourceTracker.js.map +1 -0
  93. package/lib/tileScheduler.d.ts +33 -0
  94. package/lib/tileScheduler.js +86 -0
  95. package/lib/tileScheduler.js.map +1 -0
  96. package/lib/webnnUNet.d.ts +52 -0
  97. package/lib/webnnUNet.js +535 -0
  98. package/lib/webnnUNet.js.map +1 -0
  99. package/package.json +9 -5
  100. package/scripts/inspect-model.mjs +64 -0
  101. package/src/UNet.ts +236 -339
  102. package/src/WGPUComputePass.ts +6 -4
  103. package/src/backend.ts +42 -55
  104. package/src/graphOptimizer.ts +301 -0
  105. package/src/main.ts +71 -11
  106. package/src/modelSpec.ts +414 -0
  107. package/src/nativeUNet.ts +2256 -0
  108. package/src/process.ts +38 -36
  109. package/src/resourceTracker.ts +94 -0
  110. package/src/tileScheduler.ts +138 -0
  111. package/src/webnnUNet.ts +812 -0
  112. package/tests/modelSpec.test.mjs +128 -0
  113. package/tests/resourceLifecycle.test.mjs +383 -0
  114. package/tests/tileScheduler.test.mjs +90 -0
  115. package/lib/helper.d.ts +0 -4
  116. package/lib/helper.js +0 -33
  117. package/lib/helper.js.map +0 -1
  118. package/lib/kernels.d.ts +0 -1
  119. package/lib/kernels.js +0 -26
  120. package/lib/kernels.js.map +0 -1
  121. package/src/helper.ts +0 -43
  122. package/src/kernels.ts +0 -31
package/lib/UNet.d.ts CHANGED
@@ -1,6 +1,9 @@
1
1
  import { HostTensor } from './tza';
2
2
  import { Tile } from './process';
3
- import type { WebGPUBackend } from '@tensorflow/tfjs-backend-webgpu';
3
+ import { type DynamicTileSetting } from './tileScheduler';
4
+ import { type UNetModelSpec } from './modelSpec';
5
+ import { type NativeUNetKernelSetting, type NativeUNetPrecisionSetting } from './nativeUNet';
6
+ export type UNetEngineSetting = 'auto' | 'wgsl' | 'webnn';
4
7
  interface HDRImageData {
5
8
  data: Float32Array;
6
9
  width: number;
@@ -16,10 +19,21 @@ interface GPUImageDataOutput {
16
19
  width: number;
17
20
  height: number;
18
21
  }
22
+ export interface UNetExecutionStats {
23
+ width: number;
24
+ height: number;
25
+ tileWidth: number;
26
+ tileHeight: number;
27
+ tileCount: number;
28
+ durationMs: number;
29
+ tileTimeMs: {
30
+ min: number;
31
+ median: number;
32
+ mean: number;
33
+ max: number;
34
+ };
35
+ }
19
36
  declare class UNet {
20
- private _hostTensors;
21
- private _backend;
22
- private _tfModel;
23
37
  private _device;
24
38
  private _tileWidth;
25
39
  private _tileHeight;
@@ -28,10 +42,17 @@ declare class UNet {
28
42
  private _aux;
29
43
  private _hdr;
30
44
  private _dataProcessGPU?;
31
- private _maxTileSize;
32
- private _tensors;
33
- private _modelsCache;
34
- constructor(_hostTensors: Map<string, HostTensor>, _backend: WebGPUBackend, opts?: {
45
+ private _nativeExecutor?;
46
+ private _webNNExecutor?;
47
+ private _modelSpec;
48
+ private _inputChannels;
49
+ private _engine;
50
+ private _dynamicTileController;
51
+ private _lastExecution?;
52
+ constructor(hostTensors: Map<string, HostTensor>, backend: {
53
+ device: GPUDevice;
54
+ adapterInfo: GPUAdapterInfo;
55
+ }, opts?: {
35
56
  /**
36
57
  * If use auxiliary data.
37
58
  */
@@ -41,15 +62,45 @@ declare class UNet {
41
62
  */
42
63
  hdr?: boolean;
43
64
  maxTileSize?: number;
65
+ dynamicTile?: DynamicTileSetting;
66
+ /** Reserved for explicit native WGSL selection. */
67
+ engine?: UNetEngineSetting;
68
+ /** Arithmetic/storage precision used by the native WGSL engine. */
69
+ precision?: NativeUNetPrecisionSetting;
70
+ /** Model-independent convolution kernel selection. */
71
+ kernel?: NativeUNetKernelSetting;
72
+ /** Explicit descriptor for a new OIDN topology not in the built-in registry. */
73
+ modelSpec?: UNetModelSpec;
44
74
  });
45
75
  getDevice(): GPUDevice | undefined;
46
- private _buildModel;
47
- private _createConv;
48
- private _createConcatConv;
49
- private _createPooling;
50
- private _addUpsamplingLayer;
51
- private _addNet;
52
- private _addNetLarge;
76
+ /** Completes backend compilation before first interactive use. */
77
+ prepare(): Promise<void>;
78
+ getRuntimeInfo(): {
79
+ configuredEngine: UNetEngineSetting;
80
+ gpuEngine: "wgsl" | "webnn";
81
+ precision: import("./nativeUNet").NativeUNetPrecision;
82
+ kernel: {
83
+ configured: NativeUNetKernelSetting;
84
+ maxSpatialInputBlocks: number;
85
+ subgroupsAvailable: boolean;
86
+ } | undefined;
87
+ webnn: import("./webnnUNet").WebNNRuntimeSupport | undefined;
88
+ resources: import("./resourceTracker").OIDNResourceSnapshot;
89
+ model: string;
90
+ modelFamily: (string & {}) | "oidn-unet-small" | "oidn-unet-large";
91
+ inputChannels: number;
92
+ dynamicTile: {
93
+ enabled: boolean;
94
+ currentTileSize: number;
95
+ minTileSize: number;
96
+ maxTileSize: number;
97
+ targetTileTimeMs: number;
98
+ };
99
+ lastExecution: UNetExecutionStats | undefined;
100
+ };
101
+ /** Captures per-node GPU timestamps for the next native tile execution. */
102
+ profileNextExecution(): boolean;
103
+ getLastExecutionProfile(): Promise<import("./nativeUNet").NativeUNetExecutionProfile> | undefined;
53
104
  private _updateModel;
54
105
  private _getTileSizeWithOverlap;
55
106
  private _processImageData;
package/lib/UNet.js CHANGED
@@ -1,67 +1,15 @@
1
- import { tensor } from '@tensorflow/tfjs-core/dist/ops/tensor';
2
- import { tensor1d } from '@tensorflow/tfjs-core/dist/ops/tensor1d';
3
- import { pad4d } from '@tensorflow/tfjs-core/dist/ops/pad4d';
4
- import { slice4d } from '@tensorflow/tfjs-core/dist/ops/slice4d';
5
- import { concat4d } from '@tensorflow/tfjs-core/dist/ops/concat_4d';
6
- import { ENGINE } from '@tensorflow/tfjs-core/dist/engine';
7
- import { Conv2D, UpSampling2D } from '@tensorflow/tfjs-layers/dist/layers/convolutional';
8
- import { MaxPooling2D } from '@tensorflow/tfjs-layers/dist/layers/pooling';
9
- import { Concatenate } from '@tensorflow/tfjs-layers/dist/layers/merge';
10
- import { LayersModel } from '@tensorflow/tfjs-layers/dist/engine/training';
11
- import { Input as TFInput } from '@tensorflow/tfjs-layers/dist/engine/input_layer';
12
- import { Float16Array } from '@petamoriken/float16';
13
1
  import { GPUDataProcess, Tile, avgLogLum, hdrTransferFuncCPU, hdrTransferFuncInverseCPU } from './process';
14
- // import { profileAndLogKernelCode, memory } from './helper';
15
- function getTensorData(ubytes, type) {
16
- const buffer = ubytes.buffer;
17
- if (type === 'Float32') {
18
- return new Float32Array(ubytes.buffer);
19
- }
20
- const float16Data = new Float16Array(buffer);
21
- const float32Data = new Float32Array(float16Data.length);
22
- for (let i = 0; i < float32Data.length; ++i) {
23
- float32Data[i] = float16Data[i];
24
- }
25
- return float32Data;
26
- }
27
- function changeWeightShapes(weightData, dims) {
28
- const [O, C, H, W] = dims;
29
- const reorderedWeightData = new Float32Array(weightData.length);
30
- for (let o = 0; o < O; ++o) {
31
- for (let c = 0; c < C; ++c) {
32
- for (let h = 0; h < H; ++h) {
33
- for (let w = 0; w < W; ++w) {
34
- // Change OCHW to HWCO
35
- const idx = o * C * H * W + c * H * W + h * W + w;
36
- const idx2 = h * W * C * O + w * C * O + c * O + o;
37
- reorderedWeightData[idx2] = weightData[idx];
38
- }
39
- }
40
- }
41
- }
42
- return reorderedWeightData;
43
- }
2
+ import { DynamicTileController, fitTileDimension, OIDN_TILE_ALIGNMENT, waitForSubmittedGPUWork } from './tileScheduler';
3
+ import { detectUNetModelSpec, validateUNetModel } from './modelSpec';
4
+ import { NativeUNetExecutor } from './nativeUNet';
5
+ import { WebNNUNetExecutor } from './webnnUNet';
44
6
  function roundUp(a, b) {
45
7
  return Math.ceil(a / b) * b;
46
8
  }
47
- // Returns the smallest integer larger than or equal to a which has remainder c when divided by b
48
- function roundUp2(a, b, c) {
49
- return Math.ceil((a - c) / b) * b + c;
50
- }
51
9
  function isGPUImageData(data) {
52
10
  return data.data instanceof GPUBuffer || data.data instanceof GPUTexture;
53
11
  }
54
- const receptiveField = 174; // receptive field in pixels
55
- const receptiveFieldLarge = 202;
56
- // TODO metal is 32?
57
- const minTileAlignment = 1;
58
- const tileAlignment = 16; // required spatial alignment in pixels (padding may be necessary)
59
- const defaultTileOverlap = roundUp(receptiveField / 2, tileAlignment);
60
- const defaultTileOverlapLarge = roundUp(receptiveFieldLarge / 2, tileAlignment);
61
12
  class UNet {
62
- _hostTensors;
63
- _backend;
64
- _tfModel;
65
13
  _device;
66
14
  // TODO calculate the tile size from memory size
67
15
  // https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/core/unet_filter.cpp#L287
@@ -72,164 +20,101 @@ class UNet {
72
20
  _aux;
73
21
  _hdr;
74
22
  _dataProcessGPU;
75
- _maxTileSize;
76
- _tensors = new Map();
77
- _modelsCache = new Map();
78
- constructor(_hostTensors, _backend, opts = {}) {
79
- this._hostTensors = _hostTensors;
80
- this._backend = _backend;
23
+ _nativeExecutor;
24
+ _webNNExecutor;
25
+ _modelSpec;
26
+ _inputChannels;
27
+ _engine;
28
+ _dynamicTileController;
29
+ _lastExecution;
30
+ constructor(hostTensors, backend, opts = {}) {
81
31
  this._aux = opts.aux || false;
82
32
  this._hdr = opts.hdr || false;
83
- this._maxTileSize = roundUp(opts.maxTileSize ?? 512, 2);
84
- this._device = this._backend.device;
33
+ this._engine = opts.engine ?? 'auto';
34
+ const modelSpec = opts.modelSpec ?? detectUNetModelSpec(hostTensors);
35
+ const validatedModel = validateUNetModel(hostTensors, modelSpec);
36
+ this._modelSpec = validatedModel.spec;
37
+ this._inputChannels = validatedModel.inputChannels;
38
+ const expectedInputChannels = this._aux ? 9 : 3;
39
+ if (validatedModel.inputChannels !== expectedInputChannels) {
40
+ throw new Error(`OIDN model expects ${validatedModel.inputChannels} input channels, ` +
41
+ `but aux=${this._aux} provides ${expectedInputChannels}`);
42
+ }
43
+ this._dynamicTileController = new DynamicTileController(opts.maxTileSize ?? 512, opts.dynamicTile);
44
+ this._device = backend.device;
45
+ if (this._engine === 'webnn') {
46
+ this._webNNExecutor = new WebNNUNetExecutor(this._device, validatedModel, { precision: opts.precision });
47
+ }
48
+ else {
49
+ this._nativeExecutor = new NativeUNetExecutor(this._device, validatedModel, { precision: opts.precision, kernel: opts.kernel });
50
+ }
85
51
  }
86
52
  getDevice() {
87
53
  return this._device;
88
54
  }
89
- _buildModel(isLarge) {
90
- const aux = this._aux;
91
- const channels = 3 + (aux ? 6 : 0);
92
- const tileSize = this._getTileSizeWithOverlap();
93
- const cache = this._modelsCache;
94
- const key = [tileSize.width, tileSize.height].join(',');
95
- // We cache the model instead of disposing and recreate.
96
- // Because seems tfjs will also cache the layer and gpubuffers.
97
- // Recreating the model will cause memory leak.
98
- // Width and height can only be 256, 512, 768. So the cache won't be too large
99
- if (cache.has(key)) {
100
- this._tfModel = cache.get(key);
55
+ /** Completes backend compilation before first interactive use. */
56
+ async prepare() {
57
+ if (this._webNNExecutor) {
58
+ await this._webNNExecutor.prepare();
59
+ const overlap = roundUp(this._modelSpec.receptiveField / 2, OIDN_TILE_ALIGNMENT);
60
+ const outputTileEdges = [
61
+ this._dynamicTileController.tileSize,
62
+ this._dynamicTileController.minTileSize
63
+ ];
64
+ await this._webNNExecutor.prewarm([...new Set(outputTileEdges)].map((edge) => ({
65
+ width: edge + 2 * overlap,
66
+ height: edge + 2 * overlap
67
+ })));
101
68
  return;
102
69
  }
103
- const input = TFInput({
104
- name: 'input',
105
- shape: [tileSize.height, tileSize.width, channels],
106
- dtype: 'float32'
107
- });
108
- this._tfModel = new LayersModel({
109
- inputs: [input],
110
- outputs: isLarge ? this._addNetLarge(input) : this._addNet(input)
111
- });
112
- cache.set(key, this._tfModel);
113
- }
114
- _createConv(name, source, activation) {
115
- const weightTensorName = name + '.weight';
116
- const biasTensorName = name + '.bias';
117
- const tensors = this._tensors;
118
- let weightTensor = tensors.get(weightTensorName);
119
- let biasTensor = tensors.get(biasTensorName);
120
- const unetWeightTensor = this._hostTensors.get(weightTensorName);
121
- if (!weightTensor) {
122
- const weightDims = unetWeightTensor.desc.dims;
123
- weightTensor = tensor(changeWeightShapes(getTensorData(unetWeightTensor.data, unetWeightTensor.desc.dataType), weightDims), [weightDims[2], weightDims[3], weightDims[1], weightDims[0]], 'float32');
124
- tensors.set(weightTensorName, weightTensor);
125
- }
126
- if (!biasTensor) {
127
- const unetBiasTensor = this._hostTensors.get(name + '.bias');
128
- biasTensor = tensor1d(getTensorData(unetBiasTensor.data, unetBiasTensor.desc.dataType), 'float32');
129
- tensors.set(biasTensorName, biasTensor);
130
- }
131
- // TODO whats the purpose of padded dims ?
132
- const convLayer = new Conv2D({
133
- name,
134
- filters: unetWeightTensor.desc.dims[0],
135
- kernelSize: unetWeightTensor.desc.dims.slice(2, 4),
136
- useBias: true,
137
- activation,
138
- padding: 'same',
139
- weights: [weightTensor, biasTensor],
140
- trainable: false
141
- });
142
- return convLayer.apply(source);
70
+ await this._nativeExecutor.prepare();
143
71
  }
144
- _createConcatConv(name, source1, source2) {
145
- const concatLayer = new Concatenate({
146
- name: name + '/concat',
147
- trainable: false,
148
- axis: 3
149
- });
150
- //https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L40
151
- return this._createConv(name,
152
- // Concat on the channel
153
- concatLayer.apply([source1, source2]), 'relu');
154
- }
155
- _createPooling(source) {
156
- const poolingLayer = new MaxPooling2D({
157
- name: source.name + '/pooling',
158
- poolSize: [2, 2],
159
- strides: [2, 2],
160
- padding: 'same',
161
- trainable: false
162
- });
163
- // https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L33
164
- return poolingLayer.apply(source);
165
- }
166
- _addUpsamplingLayer(source) {
167
- const upsamplingLayer = new UpSampling2D({
168
- name: source.name + '/upsampling',
169
- size: [2, 2],
170
- trainable: false
171
- });
172
- return upsamplingLayer.apply(source);
72
+ getRuntimeInfo() {
73
+ return {
74
+ configuredEngine: this._engine,
75
+ gpuEngine: this._webNNExecutor ? 'webnn' : 'wgsl',
76
+ precision: (this._webNNExecutor ?? this._nativeExecutor).precision,
77
+ kernel: this._nativeExecutor
78
+ ? {
79
+ configured: this._nativeExecutor.kernelSetting,
80
+ maxSpatialInputBlocks: this._nativeExecutor.maxSpatialInputBlocks,
81
+ subgroupsAvailable: this._nativeExecutor.subgroupsAvailable
82
+ }
83
+ : undefined,
84
+ webnn: this._webNNExecutor?.support,
85
+ resources: (this._webNNExecutor ?? this._nativeExecutor).getResourceInfo(),
86
+ model: this._modelSpec.id,
87
+ modelFamily: this._modelSpec.family,
88
+ inputChannels: this._inputChannels,
89
+ dynamicTile: {
90
+ enabled: this._dynamicTileController.enabled,
91
+ currentTileSize: this._dynamicTileController.tileSize,
92
+ minTileSize: this._dynamicTileController.minTileSize,
93
+ maxTileSize: this._dynamicTileController.maxTileSize,
94
+ targetTileTimeMs: this._dynamicTileController.targetTileTimeMs
95
+ },
96
+ lastExecution: this._lastExecution
97
+ };
173
98
  }
174
- _addNet(input) {
175
- let x = this._createConv('enc_conv0', input, 'relu');
176
- const pool1 = (x = this._createPooling(this._createConv('enc_conv1', x, 'relu')));
177
- const pool2 = (x = this._createPooling(this._createConv('enc_conv2', x, 'relu')));
178
- const pool3 = (x = this._createPooling(this._createConv('enc_conv3', x, 'relu')));
179
- const pool4 = (x = this._createPooling(this._createConv('enc_conv4', x, 'relu')));
180
- x = this._createConv('enc_conv5a', pool4, 'relu');
181
- x = this._addUpsamplingLayer(this._createConv('enc_conv5b', x, 'relu'));
182
- x = this._createConcatConv('dec_conv4a', x, pool3);
183
- x = this._addUpsamplingLayer(this._createConv('dec_conv4b', x, 'relu'));
184
- x = this._createConcatConv('dec_conv3a', x, pool2);
185
- x = this._addUpsamplingLayer(this._createConv('dec_conv3b', x, 'relu'));
186
- x = this._createConcatConv('dec_conv2a', x, pool1);
187
- x = this._addUpsamplingLayer(this._createConv('dec_conv2b', x, 'relu'));
188
- x = this._createConcatConv('dec_conv1a', x, input);
189
- x = this._createConv('dec_conv1b', x, 'relu');
190
- x = this._createConv('dec_conv0', x, 'relu');
191
- return x;
99
+ /** Captures per-node GPU timestamps for the next native tile execution. */
100
+ profileNextExecution() {
101
+ return this._nativeExecutor?.profileNextExecution() ?? false;
192
102
  }
193
- _addNetLarge(input) {
194
- let x = this._createConv('enc_conv1a', input, 'relu');
195
- const pool1 = (x = this._createPooling(this._createConv('enc_conv1b', x, 'relu')));
196
- x = this._createConv('enc_conv2a', x, 'relu');
197
- const pool2 = (x = this._createPooling(this._createConv('enc_conv2b', x, 'relu')));
198
- x = this._createConv('enc_conv3a', x, 'relu');
199
- const pool3 = (x = this._createPooling(this._createConv('enc_conv3b', x, 'relu')));
200
- x = this._createConv('enc_conv4a', x, 'relu');
201
- const pool4 = (x = this._createPooling(this._createConv('enc_conv4b', x, 'relu')));
202
- x = this._createConv('enc_conv5a', pool4, 'relu');
203
- x = this._addUpsamplingLayer(this._createConv('enc_conv5b', x, 'relu'));
204
- x = this._createConcatConv('dec_conv4a', x, pool3);
205
- x = this._addUpsamplingLayer(this._createConv('dec_conv4b', x, 'relu'));
206
- x = this._createConcatConv('dec_conv3a', x, pool2);
207
- x = this._addUpsamplingLayer(this._createConv('dec_conv3b', x, 'relu'));
208
- x = this._createConcatConv('dec_conv2a', x, pool1);
209
- x = this._addUpsamplingLayer(this._createConv('dec_conv2b', x, 'relu'));
210
- x = this._createConcatConv('dec_conv1a', x, input);
211
- x = this._createConv('dec_conv1b', x, 'relu');
212
- x = this._createConv('dec_conv1c', x, 'relu');
213
- return x;
103
+ getLastExecutionProfile() {
104
+ return this._nativeExecutor?.getLastExecutionProfile();
214
105
  }
215
106
  _updateModel(width, height) {
216
- const isLarge = this._hostTensors.has('enc_conv1b.weight');
217
- const maxTileSize = this._maxTileSize;
218
- let tileWidth = maxTileSize;
219
- let tileHeight = maxTileSize;
220
- let tileOverlapX = isLarge ? defaultTileOverlapLarge : defaultTileOverlap;
221
- let tileOverlapY = isLarge ? defaultTileOverlapLarge : defaultTileOverlap;
222
- if (width < maxTileSize + defaultTileOverlap * 2) {
223
- tileWidth = roundUp(width, maxTileSize / 2);
224
- if (width <= maxTileSize) {
225
- tileOverlapX = 0;
226
- }
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;
227
115
  }
228
- if (height < maxTileSize + defaultTileOverlap * 2) {
229
- tileHeight = roundUp(height, maxTileSize / 2);
230
- if (height <= maxTileSize) {
231
- tileOverlapY = 0;
232
- }
116
+ if (height <= maxTileSize) {
117
+ tileOverlapY = 0;
233
118
  }
234
119
  // Force width and height has same size. reduce the cache in memory
235
120
  const tileSize = Math.max(tileWidth, tileHeight);
@@ -241,14 +126,12 @@ class UNet {
241
126
  if (tileWidth !== this._tileWidth ||
242
127
  tileHeight !== this._tileHeight ||
243
128
  tileOverlapX !== this._tileOverlapX ||
244
- tileOverlapY !== this._tileOverlapY ||
245
- !this._tfModel) {
129
+ tileOverlapY !== this._tileOverlapY) {
246
130
  // console.log(tileWidth, tileHeight, tileOverlapX, tileOverlapY);
247
131
  this._tileWidth = tileWidth;
248
132
  this._tileHeight = tileHeight;
249
133
  this._tileOverlapX = tileOverlapX;
250
134
  this._tileOverlapY = tileOverlapY;
251
- this._buildModel(isLarge);
252
135
  }
253
136
  }
254
137
  _getTileSizeWithOverlap() {
@@ -327,7 +210,7 @@ class UNet {
327
210
  }
328
211
  }
329
212
  }
330
- _executeTile(inputData, outputTileData, outputImageData, i, j, width, height, isHDR, denoiseAlpha) {
213
+ async _executeTile(inputData, outputTileData, outputImageData, i, j, width, height, isHDR, denoiseAlpha) {
331
214
  const channels = this._aux ? 9 : 3;
332
215
  const tileOverlapX = this._tileOverlapX;
333
216
  const tileOverlapY = this._tileOverlapY;
@@ -342,7 +225,8 @@ class UNet {
342
225
  const srcTileWidth = srcTileSize.width;
343
226
  const srcTileHeight = srcTileSize.height;
344
227
  const srcTile = new Tile(srcX0, srcY0, srcTileWidth, srcTileHeight);
345
- let tileTensor;
228
+ let nativeOutputBuffer;
229
+ let denoisedData;
346
230
  let inputScale = 1;
347
231
  const device = this._device;
348
232
  let dataProcessGPU = this._dataProcessGPU;
@@ -359,7 +243,7 @@ class UNet {
359
243
  inputScale
360
244
  });
361
245
  }
362
- tileTensor = tensor(tileData, [1, srcTileHeight, srcTileWidth, channels], 'float32');
246
+ denoisedData = await (this._webNNExecutor ?? this._nativeExecutor).executeCPU(tileData, srcTileWidth, srcTileHeight);
363
247
  }
364
248
  else {
365
249
  if (!dataProcessGPU) {
@@ -372,33 +256,15 @@ class UNet {
372
256
  dataProcessGPU.copyInputDataToOutput(inputData.color);
373
257
  }
374
258
  const { color, albedo, normal } = dataProcessGPU.forward(inputData.color, this._aux ? inputData.albedo : undefined, this._aux ? inputData.normal : undefined, denoiseAlpha);
375
- const createTensor = (buffer) => {
376
- const tmp = tensor({ buffer, zeroCopy: true }, [
377
- 1,
378
- srcTileHeight,
379
- srcTileWidth,
380
- 4
381
- ]);
382
- const ret = slice4d(tmp, [0, 0, 0, 0], [1, srcTileHeight, srcTileWidth, 3]);
383
- return ret;
384
- };
385
- if (this._aux) {
386
- const tensors = [color, albedo, normal].map((buffer) => createTensor(buffer));
387
- tileTensor = concat4d(tensors, 3);
388
- }
389
- else {
390
- tileTensor = createTensor(color);
391
- }
259
+ nativeOutputBuffer = await (this._webNNExecutor ?? this._nativeExecutor).execute(this._aux ? [color, albedo, normal] : [color], srcTileWidth, srcTileHeight);
392
260
  }
393
261
  let outBuffer;
394
- const outputTensor = this._tfModel.predict(tileTensor);
395
262
  const dstWidth = Math.min(dstTileSize.width, width);
396
263
  const dstHeight = Math.min(dstTileSize.height, height);
397
264
  const dstTile = new Tile(i * dstWidth, j * dstHeight, dstWidth, dstHeight);
398
265
  dstTile.width = Math.min(dstTile.width, width - dstTile.x);
399
266
  dstTile.height = Math.min(dstTile.height, height - dstTile.y);
400
267
  if (inputData instanceof Float32Array) {
401
- let denoisedData = outputTensor.dataSync();
402
268
  if (isHDR) {
403
269
  denoisedData = hdrTransferFuncInverseCPU({
404
270
  data: denoisedData,
@@ -419,16 +285,7 @@ class UNet {
419
285
  }
420
286
  else {
421
287
  dataProcessGPU.setOutputTile(dstTile, srcTile);
422
- // IMPORTANT
423
- // storage buffer has alignment. that 3 channels still needs 16 bytes data.
424
- // So we need to pad it to 4 channels.
425
- const outputTensor4Channnels = pad4d(outputTensor, [
426
- [0, 0],
427
- [0, 0],
428
- [0, 0],
429
- [0, 1]
430
- ]);
431
- outBuffer = dataProcessGPU.inverse(outputTensor4Channnels.dataToGPU().buffer, inputData.color);
288
+ outBuffer = dataProcessGPU.inverse(nativeOutputBuffer, inputData.color);
432
289
  }
433
290
  return outBuffer;
434
291
  }
@@ -443,6 +300,8 @@ class UNet {
443
300
  }
444
301
  const width = color.width;
445
302
  const height = color.height;
303
+ const adaptiveTileSize = this._dynamicTileController.tileSize;
304
+ const shouldAdaptTileSize = width > adaptiveTileSize || height > adaptiveTileSize;
446
305
  this._updateModel(width, height);
447
306
  // TODO should fixed to be hdr when UNet is created.
448
307
  // weights of hdr and ldr is different
@@ -471,22 +330,31 @@ class UNet {
471
330
  ? undefined
472
331
  : makeImageData(Math.min(tileWidth, width), Math.min(tileHeight, height));
473
332
  let aborted = false;
474
- const executeTile = (i, j) => {
333
+ const now = () => typeof performance === 'undefined' ? Date.now() : performance.now();
334
+ const executionStartTime = now();
335
+ const tileTimesMs = [];
336
+ const scheduleNextTile = (callback) => {
337
+ if (typeof requestAnimationFrame === 'undefined') {
338
+ setTimeout(callback, 0);
339
+ }
340
+ else {
341
+ requestAnimationFrame(callback);
342
+ }
343
+ };
344
+ const executeTile = async (i, j) => {
475
345
  if (aborted) {
476
346
  return;
477
347
  }
478
- let resGPUBuffer;
479
- // profileAndLogKernelCode(() => {
480
- ENGINE.startScope();
481
- resGPUBuffer = this._executeTile(isGPUImageData(color)
348
+ const tileStartTime = now();
349
+ const resGPUBuffer = await this._executeTile(isGPUImageData(color)
482
350
  ? {
483
351
  color: color.data,
484
352
  albedo: albedo?.data,
485
353
  normal: normal?.data
486
354
  }
487
355
  : rawData, outputTileData, outputImageData, i, j, width, height, hdr, denoiseAlpha);
488
- ENGINE.endScope();
489
- // }, true);
356
+ if (aborted)
357
+ return;
490
358
  const output = outputImageData || {
491
359
  data: resGPUBuffer,
492
360
  width,
@@ -495,20 +363,57 @@ class UNet {
495
363
  progress?.(output,
496
364
  // Is undefined if using webgpu buffer
497
365
  outputTileData, new Tile(i * tileWidth, j * tileHeight, tileWidth, tileHeight), i + j * tileCountW, tileCountW * tileCountH);
498
- if (i + 1 < tileCountW || j + 1 < tileCountH) {
499
- requestAnimationFrame(() => {
500
- if (i + 1 < tileCountW) {
501
- executeTile(i + 1, j);
502
- }
503
- else if (j + 1 < tileCountH) {
504
- executeTile(0, j + 1);
366
+ const hasNextTile = i + 1 < tileCountW || j + 1 < tileCountH;
367
+ const continueAfterGPUWork = () => {
368
+ tileTimesMs.push(now() - tileStartTime);
369
+ if (aborted)
370
+ return;
371
+ if (hasNextTile) {
372
+ scheduleNextTile(() => {
373
+ if (aborted)
374
+ return;
375
+ if (i + 1 < tileCountW) {
376
+ executeTile(i + 1, j);
377
+ }
378
+ else if (j + 1 < tileCountH) {
379
+ executeTile(0, j + 1);
380
+ }
381
+ });
382
+ }
383
+ else {
384
+ const sortedTileTimes = [...tileTimesMs].sort((a, b) => a - b);
385
+ const middle = Math.floor(sortedTileTimes.length / 2);
386
+ const medianTileTime = sortedTileTimes.length % 2
387
+ ? sortedTileTimes[middle]
388
+ : (sortedTileTimes[middle - 1] + sortedTileTimes[middle]) / 2;
389
+ this._lastExecution = {
390
+ width,
391
+ height,
392
+ tileWidth,
393
+ tileHeight,
394
+ tileCount: tileCountW * tileCountH,
395
+ durationMs: now() - executionStartTime,
396
+ tileTimeMs: {
397
+ min: sortedTileTimes[0],
398
+ median: medianTileTime,
399
+ mean: sortedTileTimes.reduce((sum, value) => sum + value, 0) /
400
+ sortedTileTimes.length,
401
+ max: sortedTileTimes[sortedTileTimes.length - 1]
402
+ }
403
+ };
404
+ // Adapt only from complete executions. Cancelled work is commonly
405
+ // contending with interactive rendering and is not representative.
406
+ if (shouldAdaptTileSize) {
407
+ this._dynamicTileController.observe(tileTimesMs);
505
408
  }
506
- });
507
- }
508
- else {
509
- // console.log(memory());
510
- done(output);
511
- }
409
+ // console.log(memory());
410
+ done(output);
411
+ }
412
+ };
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);
512
417
  };
513
418
  executeTile(0, 0);
514
419
  return () => {
@@ -516,9 +421,9 @@ class UNet {
516
421
  };
517
422
  }
518
423
  dispose() {
519
- this._tfModel?.dispose();
520
424
  this._dataProcessGPU?.dispose();
521
- this._tensors.forEach((tensor) => tensor.dispose());
425
+ this._nativeExecutor?.dispose();
426
+ this._webNNExecutor?.dispose();
522
427
  }
523
428
  }
524
429
  export default UNet;