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