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/lib/UNet.js
CHANGED
|
@@ -1,261 +1,162 @@
|
|
|
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
|
-
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
|
|
18
|
-
|
|
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, planTileGrid, OIDN_TILE_ALIGNMENT } from './tileScheduler';
|
|
3
|
+
import { detectUNetModelSpec, validateUNetModel } from './modelSpec';
|
|
4
|
+
import { NativeUNetExecutor } from './nativeUNet';
|
|
5
|
+
import { WebNNUNetExecutor } from './webnnUNet';
|
|
6
|
+
/** Upper bound on waiting for a display frame between tiles. */
|
|
7
|
+
const ANIMATION_FRAME_FALLBACK_MS = 100;
|
|
44
8
|
function roundUp(a, b) {
|
|
45
9
|
return Math.ceil(a / b) * b;
|
|
46
10
|
}
|
|
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
11
|
function isGPUImageData(data) {
|
|
52
12
|
return data.data instanceof GPUBuffer || data.data instanceof GPUTexture;
|
|
53
13
|
}
|
|
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
14
|
class UNet {
|
|
62
|
-
_hostTensors;
|
|
63
|
-
_backend;
|
|
64
|
-
_tfModel;
|
|
65
15
|
_device;
|
|
66
|
-
// TODO calculate the tile size from memory size
|
|
67
|
-
// https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/core/unet_filter.cpp#L287
|
|
68
|
-
_tileWidth = 0;
|
|
69
|
-
_tileHeight = 0;
|
|
70
|
-
_tileOverlapX = 0;
|
|
71
|
-
_tileOverlapY = 0;
|
|
72
16
|
_aux;
|
|
73
17
|
_hdr;
|
|
18
|
+
_hdrTransfer;
|
|
74
19
|
_dataProcessGPU;
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
|
|
78
|
-
|
|
79
|
-
|
|
80
|
-
|
|
20
|
+
_nativeExecutor;
|
|
21
|
+
_webNNExecutor;
|
|
22
|
+
_modelSpec;
|
|
23
|
+
_inputChannels;
|
|
24
|
+
_engine;
|
|
25
|
+
_dynamicTileController;
|
|
26
|
+
_lastExecution;
|
|
27
|
+
_activeExecutionFailures = new Set();
|
|
28
|
+
_deviceLostObserved = false;
|
|
29
|
+
_deviceLostSettled = false;
|
|
30
|
+
_deviceLostReason;
|
|
31
|
+
constructor(hostTensors, device, opts = {}) {
|
|
81
32
|
this._aux = opts.aux || false;
|
|
82
33
|
this._hdr = opts.hdr || false;
|
|
83
|
-
this.
|
|
84
|
-
this.
|
|
34
|
+
this._hdrTransfer = opts.hdrTransfer ?? 'pu';
|
|
35
|
+
this._engine = opts.engine ?? 'auto';
|
|
36
|
+
const modelSpec = opts.modelSpec ?? detectUNetModelSpec(hostTensors);
|
|
37
|
+
const validatedModel = validateUNetModel(hostTensors, modelSpec);
|
|
38
|
+
this._modelSpec = validatedModel.spec;
|
|
39
|
+
this._inputChannels = validatedModel.inputChannels;
|
|
40
|
+
const expectedInputChannels = this._aux ? 9 : 3;
|
|
41
|
+
if (validatedModel.inputChannels !== expectedInputChannels) {
|
|
42
|
+
throw new Error(`OIDN model expects ${validatedModel.inputChannels} input channels, ` +
|
|
43
|
+
`but aux=${this._aux} provides ${expectedInputChannels}`);
|
|
44
|
+
}
|
|
45
|
+
this._dynamicTileController = new DynamicTileController(opts.maxTileSize ?? 512, opts.dynamicTile);
|
|
46
|
+
this._device = device;
|
|
47
|
+
this._observeDeviceLoss();
|
|
48
|
+
if (this._engine === 'webnn') {
|
|
49
|
+
this._webNNExecutor = new WebNNUNetExecutor(this._device, validatedModel, { precision: opts.precision });
|
|
50
|
+
}
|
|
51
|
+
else {
|
|
52
|
+
this._nativeExecutor = new NativeUNetExecutor(this._device, validatedModel, { precision: opts.precision, kernel: opts.kernel, gemm: opts.gemm });
|
|
53
|
+
}
|
|
85
54
|
}
|
|
86
55
|
getDevice() {
|
|
87
56
|
return this._device;
|
|
88
57
|
}
|
|
89
|
-
|
|
90
|
-
|
|
91
|
-
|
|
92
|
-
|
|
93
|
-
|
|
94
|
-
|
|
95
|
-
|
|
96
|
-
|
|
97
|
-
|
|
98
|
-
|
|
99
|
-
|
|
100
|
-
|
|
58
|
+
/** Completes backend compilation before first interactive use. */
|
|
59
|
+
async prepare() {
|
|
60
|
+
if (this._webNNExecutor) {
|
|
61
|
+
await this._webNNExecutor.prepare();
|
|
62
|
+
const overlap = roundUp(this._modelSpec.receptiveField / 2, OIDN_TILE_ALIGNMENT);
|
|
63
|
+
const outputTileEdges = [
|
|
64
|
+
this._dynamicTileController.tileSize,
|
|
65
|
+
this._dynamicTileController.minTileSize
|
|
66
|
+
];
|
|
67
|
+
await this._webNNExecutor.prewarm([...new Set(outputTileEdges)].map((edge) => ({
|
|
68
|
+
width: edge + 2 * overlap,
|
|
69
|
+
height: edge + 2 * overlap
|
|
70
|
+
})));
|
|
101
71
|
return;
|
|
102
72
|
}
|
|
103
|
-
|
|
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);
|
|
143
|
-
}
|
|
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');
|
|
73
|
+
await this._nativeExecutor.prepare();
|
|
154
74
|
}
|
|
155
|
-
|
|
156
|
-
|
|
157
|
-
|
|
158
|
-
|
|
159
|
-
|
|
160
|
-
|
|
161
|
-
|
|
162
|
-
|
|
163
|
-
|
|
164
|
-
|
|
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);
|
|
165
93
|
}
|
|
166
|
-
|
|
167
|
-
|
|
168
|
-
|
|
169
|
-
|
|
170
|
-
|
|
171
|
-
|
|
172
|
-
|
|
173
|
-
|
|
174
|
-
|
|
175
|
-
|
|
176
|
-
|
|
177
|
-
|
|
178
|
-
|
|
179
|
-
|
|
180
|
-
|
|
181
|
-
|
|
182
|
-
|
|
183
|
-
|
|
184
|
-
|
|
185
|
-
|
|
186
|
-
|
|
187
|
-
|
|
188
|
-
|
|
189
|
-
|
|
190
|
-
|
|
191
|
-
|
|
94
|
+
getRuntimeInfo() {
|
|
95
|
+
return {
|
|
96
|
+
configuredEngine: this._engine,
|
|
97
|
+
gpuEngine: this._webNNExecutor ? 'webnn' : 'wgsl',
|
|
98
|
+
precision: (this._webNNExecutor ?? this._nativeExecutor).precision,
|
|
99
|
+
kernel: this._nativeExecutor
|
|
100
|
+
? {
|
|
101
|
+
configured: this._nativeExecutor.kernelSetting,
|
|
102
|
+
gemm: this._nativeExecutor.gemm,
|
|
103
|
+
maxSpatialInputBlocks: this._nativeExecutor.maxSpatialInputBlocks,
|
|
104
|
+
subgroupsAvailable: this._nativeExecutor.subgroupsAvailable
|
|
105
|
+
}
|
|
106
|
+
: undefined,
|
|
107
|
+
webnn: this._webNNExecutor?.support,
|
|
108
|
+
resources: (this._webNNExecutor ?? this._nativeExecutor).getResourceInfo(),
|
|
109
|
+
model: this._modelSpec.id,
|
|
110
|
+
modelFamily: this._modelSpec.family,
|
|
111
|
+
inputChannels: this._inputChannels,
|
|
112
|
+
hdrTransfer: this._hdrTransfer,
|
|
113
|
+
dynamicTile: {
|
|
114
|
+
enabled: this._dynamicTileController.enabled,
|
|
115
|
+
currentTileSize: this._dynamicTileController.tileSize,
|
|
116
|
+
minTileSize: this._dynamicTileController.minTileSize,
|
|
117
|
+
maxTileSize: this._dynamicTileController.maxTileSize,
|
|
118
|
+
targetTileTimeMs: this._dynamicTileController.targetTileTimeMs
|
|
119
|
+
},
|
|
120
|
+
lastExecution: this._lastExecution,
|
|
121
|
+
activeExecutionCount: this._activeExecutionFailures.size
|
|
122
|
+
};
|
|
192
123
|
}
|
|
193
|
-
|
|
194
|
-
|
|
195
|
-
|
|
196
|
-
|
|
197
|
-
|
|
198
|
-
|
|
199
|
-
const
|
|
200
|
-
|
|
201
|
-
|
|
202
|
-
|
|
203
|
-
|
|
204
|
-
|
|
205
|
-
|
|
206
|
-
|
|
207
|
-
|
|
208
|
-
|
|
209
|
-
|
|
210
|
-
|
|
211
|
-
|
|
212
|
-
|
|
213
|
-
return x;
|
|
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);
|
|
142
|
+
};
|
|
143
|
+
void deviceLost.then((info) => failAll(new Error(`WebGPU device lost: ${info.message}`)), failAll);
|
|
214
144
|
}
|
|
215
|
-
|
|
216
|
-
|
|
217
|
-
|
|
218
|
-
|
|
219
|
-
|
|
220
|
-
|
|
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
|
-
}
|
|
227
|
-
}
|
|
228
|
-
if (height < maxTileSize + defaultTileOverlap * 2) {
|
|
229
|
-
tileHeight = roundUp(height, maxTileSize / 2);
|
|
230
|
-
if (height <= maxTileSize) {
|
|
231
|
-
tileOverlapY = 0;
|
|
232
|
-
}
|
|
233
|
-
}
|
|
234
|
-
// Force width and height has same size. reduce the cache in memory
|
|
235
|
-
const tileSize = Math.max(tileWidth, tileHeight);
|
|
236
|
-
const tileOverlap = Math.max(tileOverlapX, tileOverlapY);
|
|
237
|
-
tileWidth = tileSize;
|
|
238
|
-
tileHeight = tileSize;
|
|
239
|
-
tileOverlapX = tileOverlap;
|
|
240
|
-
tileOverlapY = tileOverlap;
|
|
241
|
-
if (tileWidth !== this._tileWidth ||
|
|
242
|
-
tileHeight !== this._tileHeight ||
|
|
243
|
-
tileOverlapX !== this._tileOverlapX ||
|
|
244
|
-
tileOverlapY !== this._tileOverlapY ||
|
|
245
|
-
!this._tfModel) {
|
|
246
|
-
// console.log(tileWidth, tileHeight, tileOverlapX, tileOverlapY);
|
|
247
|
-
this._tileWidth = tileWidth;
|
|
248
|
-
this._tileHeight = tileHeight;
|
|
249
|
-
this._tileOverlapX = tileOverlapX;
|
|
250
|
-
this._tileOverlapY = tileOverlapY;
|
|
251
|
-
this._buildModel(isLarge);
|
|
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));
|
|
252
151
|
}
|
|
152
|
+
return () => this._activeExecutionFailures.delete(fail);
|
|
253
153
|
}
|
|
254
|
-
|
|
255
|
-
|
|
256
|
-
|
|
257
|
-
|
|
258
|
-
|
|
154
|
+
/** Captures per-node GPU timestamps for the next native tile execution. */
|
|
155
|
+
profileNextExecution() {
|
|
156
|
+
return this._nativeExecutor?.profileNextExecution() ?? false;
|
|
157
|
+
}
|
|
158
|
+
getLastExecutionProfile() {
|
|
159
|
+
return this._nativeExecutor?.getLastExecutionProfile();
|
|
259
160
|
}
|
|
260
161
|
_processImageData(color, albedo, normal, isHDR) {
|
|
261
162
|
const rawData = color.data;
|
|
@@ -296,9 +197,12 @@ class UNet {
|
|
|
296
197
|
}
|
|
297
198
|
_readTile(data, channels, srcTile, width) {
|
|
298
199
|
const tileData = new Float32Array(srcTile.width * srcTile.height * channels);
|
|
200
|
+
const height = data.length / (width * channels);
|
|
299
201
|
for (let y = 0; y < srcTile.height; y++) {
|
|
300
202
|
for (let x = 0; x < srcTile.width; x++) {
|
|
301
|
-
const
|
|
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;
|
|
302
206
|
const i1 = (y * srcTile.width + x) * channels;
|
|
303
207
|
for (let c = 0; c < channels; c++) {
|
|
304
208
|
tileData[i1 + c] = data[i2 + c];
|
|
@@ -327,22 +231,14 @@ class UNet {
|
|
|
327
231
|
}
|
|
328
232
|
}
|
|
329
233
|
}
|
|
330
|
-
_executeTile(inputData, outputTileData, outputImageData,
|
|
234
|
+
async _executeTile(inputData, outputTileData, outputImageData, tile, isFirstTile, width, height, isHDR, denoiseAlpha) {
|
|
331
235
|
const channels = this._aux ? 9 : 3;
|
|
332
|
-
const
|
|
333
|
-
const
|
|
334
|
-
|
|
335
|
-
|
|
336
|
-
let
|
|
337
|
-
let
|
|
338
|
-
srcX0 = Math.max(srcX1 - srcTileSize.width, 0);
|
|
339
|
-
let srcY0 = j > 0 ? j * dstTileSize.height - tileOverlapY : 0;
|
|
340
|
-
let srcY1 = Math.min(srcY0 + srcTileSize.height, height);
|
|
341
|
-
srcY0 = Math.max(srcY1 - srcTileSize.height, 0);
|
|
342
|
-
const srcTileWidth = srcTileSize.width;
|
|
343
|
-
const srcTileHeight = srcTileSize.height;
|
|
344
|
-
const srcTile = new Tile(srcX0, srcY0, srcTileWidth, srcTileHeight);
|
|
345
|
-
let tileTensor;
|
|
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;
|
|
240
|
+
let nativeOutputBuffer;
|
|
241
|
+
let denoisedData;
|
|
346
242
|
let inputScale = 1;
|
|
347
243
|
const device = this._device;
|
|
348
244
|
let dataProcessGPU = this._dataProcessGPU;
|
|
@@ -356,60 +252,39 @@ class UNet {
|
|
|
356
252
|
tileData = hdrTransferFuncCPU({
|
|
357
253
|
data: tileData,
|
|
358
254
|
channels,
|
|
359
|
-
inputScale
|
|
255
|
+
inputScale,
|
|
256
|
+
transfer: this._hdrTransfer
|
|
360
257
|
});
|
|
361
258
|
}
|
|
362
|
-
|
|
259
|
+
denoisedData = await (this._webNNExecutor ?? this._nativeExecutor).executeCPU(tileData, srcTileWidth, srcTileHeight);
|
|
363
260
|
}
|
|
364
261
|
else {
|
|
365
262
|
if (!dataProcessGPU) {
|
|
366
|
-
dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(device, isHDR);
|
|
263
|
+
dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(device, isHDR, this._hdrTransfer);
|
|
367
264
|
}
|
|
368
265
|
dataProcessGPU.setImageSize(width, height);
|
|
369
266
|
dataProcessGPU.setInputTile(srcTile);
|
|
370
267
|
// Display the noisy input instead of prev denoised result
|
|
371
|
-
if (
|
|
268
|
+
if (isFirstTile) {
|
|
372
269
|
dataProcessGPU.copyInputDataToOutput(inputData.color);
|
|
373
270
|
}
|
|
374
271
|
const { color, albedo, normal } = dataProcessGPU.forward(inputData.color, this._aux ? inputData.albedo : undefined, this._aux ? inputData.normal : undefined, denoiseAlpha);
|
|
375
|
-
|
|
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
|
-
}
|
|
272
|
+
nativeOutputBuffer = await (this._webNNExecutor ?? this._nativeExecutor).execute(this._aux ? [color, albedo, normal] : [color], srcTileWidth, srcTileHeight);
|
|
392
273
|
}
|
|
393
274
|
let outBuffer;
|
|
394
|
-
const outputTensor = this._tfModel.predict(tileTensor);
|
|
395
|
-
const dstWidth = Math.min(dstTileSize.width, width);
|
|
396
|
-
const dstHeight = Math.min(dstTileSize.height, height);
|
|
397
|
-
const dstTile = new Tile(i * dstWidth, j * dstHeight, dstWidth, dstHeight);
|
|
398
|
-
dstTile.width = Math.min(dstTile.width, width - dstTile.x);
|
|
399
|
-
dstTile.height = Math.min(dstTile.height, height - dstTile.y);
|
|
400
275
|
if (inputData instanceof Float32Array) {
|
|
401
|
-
let denoisedData = outputTensor.dataSync();
|
|
402
276
|
if (isHDR) {
|
|
403
277
|
denoisedData = hdrTransferFuncInverseCPU({
|
|
404
278
|
data: denoisedData,
|
|
405
279
|
channels: 3,
|
|
406
|
-
inputScale
|
|
280
|
+
inputScale,
|
|
281
|
+
transfer: this._hdrTransfer
|
|
407
282
|
});
|
|
408
283
|
}
|
|
409
|
-
this._writeTile(outputImageData, srcTile, dstTile, denoisedData,
|
|
410
|
-
for (let y = 0; y <
|
|
411
|
-
for (let x = 0; x <
|
|
412
|
-
const i1 = (y *
|
|
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;
|
|
413
288
|
const i2 = ((y + dstTile.y) * width + (x + dstTile.x)) * 4;
|
|
414
289
|
for (let c = 0; c < 4; c++) {
|
|
415
290
|
outputTileData.data[i1 + c] = outputImageData.data[i2 + c];
|
|
@@ -419,20 +294,11 @@ class UNet {
|
|
|
419
294
|
}
|
|
420
295
|
else {
|
|
421
296
|
dataProcessGPU.setOutputTile(dstTile, srcTile);
|
|
422
|
-
|
|
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);
|
|
297
|
+
outBuffer = dataProcessGPU.inverse(nativeOutputBuffer, inputData.color);
|
|
432
298
|
}
|
|
433
299
|
return outBuffer;
|
|
434
300
|
}
|
|
435
|
-
tileExecute({ color, albedo, normal, done, progress, denoiseAlpha }) {
|
|
301
|
+
tileExecute({ color, albedo, normal, done, progress, denoiseAlpha, tileOverlap, wholeImage, scheduling = 'event-loop', error }) {
|
|
436
302
|
if (this._aux && (!albedo || !normal)) {
|
|
437
303
|
throw new Error('Normal map and albedo map are both required');
|
|
438
304
|
}
|
|
@@ -443,7 +309,17 @@ class UNet {
|
|
|
443
309
|
}
|
|
444
310
|
const width = color.width;
|
|
445
311
|
const height = color.height;
|
|
446
|
-
this.
|
|
312
|
+
const adaptiveTileSize = this._dynamicTileController.tileSize;
|
|
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;
|
|
447
323
|
// TODO should fixed to be hdr when UNet is created.
|
|
448
324
|
// weights of hdr and ldr is different
|
|
449
325
|
const hdr = this._hdr || false;
|
|
@@ -451,10 +327,6 @@ class UNet {
|
|
|
451
327
|
if (!isGPUImageData(color)) {
|
|
452
328
|
rawData = this._processImageData(color, albedo, normal, hdr);
|
|
453
329
|
}
|
|
454
|
-
const tileWidth = this._tileWidth;
|
|
455
|
-
const tileHeight = this._tileHeight;
|
|
456
|
-
const tileCountH = Math.ceil(height / tileHeight);
|
|
457
|
-
const tileCountW = Math.ceil(width / tileWidth);
|
|
458
330
|
function makeImageData(width, height) {
|
|
459
331
|
return hdr
|
|
460
332
|
? {
|
|
@@ -467,58 +339,167 @@ class UNet {
|
|
|
467
339
|
const outputImageData = isGPUImageData(color)
|
|
468
340
|
? undefined
|
|
469
341
|
: makeImageData(width, height);
|
|
470
|
-
|
|
471
|
-
|
|
472
|
-
|
|
473
|
-
let
|
|
474
|
-
const
|
|
475
|
-
|
|
342
|
+
let state = 'active';
|
|
343
|
+
let scheduledTimer;
|
|
344
|
+
let scheduledAnimationFrame;
|
|
345
|
+
let unregisterDeviceLoss = () => false;
|
|
346
|
+
const now = () => typeof performance === 'undefined' ? Date.now() : performance.now();
|
|
347
|
+
const executionStartTime = now();
|
|
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')
|
|
476
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
|
+
};
|
|
383
|
+
const scheduleNextTile = (callback) => {
|
|
384
|
+
if (scheduling === 'event-loop' ||
|
|
385
|
+
typeof requestAnimationFrame === 'undefined') {
|
|
386
|
+
scheduledTimer = setTimeout(() => {
|
|
387
|
+
scheduledTimer = undefined;
|
|
388
|
+
callback();
|
|
389
|
+
}, 0);
|
|
390
|
+
}
|
|
391
|
+
else {
|
|
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);
|
|
477
401
|
}
|
|
478
|
-
|
|
479
|
-
|
|
480
|
-
|
|
481
|
-
|
|
402
|
+
};
|
|
403
|
+
const executeTile = async (tileIndex) => {
|
|
404
|
+
if (state !== 'active') {
|
|
405
|
+
return;
|
|
406
|
+
}
|
|
407
|
+
const tile = plan.tiles[tileIndex];
|
|
408
|
+
const outputTileData = isGPUImageData(color)
|
|
409
|
+
? undefined
|
|
410
|
+
: makeImageData(tile.output.width, tile.output.height);
|
|
411
|
+
const tileStartTime = now();
|
|
412
|
+
const resGPUBuffer = await this._executeTile(isGPUImageData(color)
|
|
482
413
|
? {
|
|
483
414
|
color: color.data,
|
|
484
415
|
albedo: albedo?.data,
|
|
485
416
|
normal: normal?.data
|
|
486
417
|
}
|
|
487
|
-
: rawData, outputTileData, outputImageData,
|
|
488
|
-
|
|
489
|
-
|
|
418
|
+
: rawData, outputTileData, outputImageData, tile, tileIndex === 0, width, height, hdr, denoiseAlpha);
|
|
419
|
+
if (state !== 'active')
|
|
420
|
+
return;
|
|
490
421
|
const output = outputImageData || {
|
|
491
422
|
data: resGPUBuffer,
|
|
492
423
|
width,
|
|
493
424
|
height
|
|
494
425
|
};
|
|
495
|
-
progress
|
|
496
|
-
|
|
497
|
-
|
|
498
|
-
|
|
499
|
-
|
|
500
|
-
|
|
501
|
-
|
|
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 () => {
|
|
438
|
+
tileTimesMs.push(now() - tileStartTime);
|
|
439
|
+
if (state !== 'active')
|
|
440
|
+
return;
|
|
441
|
+
if (hasNextTile) {
|
|
442
|
+
scheduleNextTile(() => {
|
|
443
|
+
if (state !== 'active')
|
|
444
|
+
return;
|
|
445
|
+
void executeTile(tileIndex + 1).catch(settleError);
|
|
446
|
+
});
|
|
447
|
+
}
|
|
448
|
+
else {
|
|
449
|
+
const sortedTileTimes = [...tileTimesMs].sort((a, b) => a - b);
|
|
450
|
+
const middle = Math.floor(sortedTileTimes.length / 2);
|
|
451
|
+
const medianTileTime = sortedTileTimes.length % 2
|
|
452
|
+
? sortedTileTimes[middle]
|
|
453
|
+
: (sortedTileTimes[middle - 1] + sortedTileTimes[middle]) / 2;
|
|
454
|
+
this._lastExecution = {
|
|
455
|
+
width,
|
|
456
|
+
height,
|
|
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,
|
|
463
|
+
durationMs: now() - executionStartTime,
|
|
464
|
+
tileTimeMs: {
|
|
465
|
+
min: sortedTileTimes[0],
|
|
466
|
+
median: medianTileTime,
|
|
467
|
+
mean: sortedTileTimes.reduce((sum, value) => sum + value, 0) /
|
|
468
|
+
sortedTileTimes.length,
|
|
469
|
+
max: sortedTileTimes[sortedTileTimes.length - 1]
|
|
470
|
+
}
|
|
471
|
+
};
|
|
472
|
+
// Adapt only from complete executions. Cancelled work is commonly
|
|
473
|
+
// contending with interactive rendering and is not representative.
|
|
474
|
+
if (shouldAdaptTileSize) {
|
|
475
|
+
this._dynamicTileController.observe(tileTimesMs);
|
|
502
476
|
}
|
|
503
|
-
|
|
504
|
-
|
|
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';
|
|
505
484
|
}
|
|
506
|
-
}
|
|
507
|
-
}
|
|
508
|
-
|
|
509
|
-
// console.log(memory());
|
|
510
|
-
done(output);
|
|
511
|
-
}
|
|
485
|
+
}
|
|
486
|
+
};
|
|
487
|
+
await continueAfterGPUWork();
|
|
512
488
|
};
|
|
513
|
-
|
|
489
|
+
unregisterDeviceLoss = this._registerExecutionFailure(settleError);
|
|
490
|
+
void executeTile(0).catch(settleError);
|
|
514
491
|
return () => {
|
|
515
|
-
|
|
492
|
+
if (state !== 'active')
|
|
493
|
+
return;
|
|
494
|
+
state = 'aborted';
|
|
495
|
+
cancelScheduledTile();
|
|
496
|
+
unregisterDeviceLoss();
|
|
516
497
|
};
|
|
517
498
|
}
|
|
518
499
|
dispose() {
|
|
519
|
-
this._tfModel?.dispose();
|
|
520
500
|
this._dataProcessGPU?.dispose();
|
|
521
|
-
this.
|
|
501
|
+
this._nativeExecutor?.dispose();
|
|
502
|
+
this._webNNExecutor?.dispose();
|
|
522
503
|
}
|
|
523
504
|
}
|
|
524
505
|
export default UNet;
|