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.
- package/CHANGELOG.md +49 -0
- package/README.md +140 -4
- package/benchmarks/compare.mjs +651 -0
- package/benchmarks/leak.mjs +255 -0
- package/benchmarks/results/before-spatial.json +391 -0
- package/benchmarks/results/before-spatial.md +47 -0
- package/benchmarks/results/int8-scan.json +2007 -0
- package/benchmarks/results/int8-scan.md +160 -0
- package/benchmarks/results/int8-w8a8-scan.json +2007 -0
- package/benchmarks/results/int8-w8a8-scan.md +160 -0
- package/benchmarks/results/int8-weight-channel.json +1413 -0
- package/benchmarks/results/int8-weight-channel.md +118 -0
- package/benchmarks/results/int8-weight-only.json +1437 -0
- package/benchmarks/results/int8-weight-only.md +118 -0
- package/benchmarks/results/kernel-webnn-final.json +1115 -0
- package/benchmarks/results/kernel-webnn-final.md +104 -0
- package/benchmarks/results/latest-optimized.json +375 -0
- package/benchmarks/results/latest-optimized.md +47 -0
- package/benchmarks/results/latest.json +391 -0
- package/benchmarks/results/latest.md +47 -0
- package/benchmarks/results/profile-baseline.json +331 -0
- package/benchmarks/results/profile-baseline.md +12 -0
- package/benchmarks/results/profile-conv2x.json +331 -0
- package/benchmarks/results/profile-conv2x.md +12 -0
- package/benchmarks/results/profile-fast-init.json +385 -0
- package/benchmarks/results/profile-fast-init.md +47 -0
- package/benchmarks/results/profile-fp16-fma.json +369 -0
- package/benchmarks/results/profile-fp16-fma.md +47 -0
- package/benchmarks/results/profile-fp16-tiled-decoder.json +369 -0
- package/benchmarks/results/profile-fp16-tiled-decoder.md +47 -0
- package/benchmarks/results/profile-fp16-tiled-encoder.json +369 -0
- package/benchmarks/results/profile-fp16-tiled-encoder.md +47 -0
- package/benchmarks/results/profile-fp16-unfused-pool.json +385 -0
- package/benchmarks/results/profile-fp16-unfused-pool.md +47 -0
- package/benchmarks/results/profile-input-major.json +347 -0
- package/benchmarks/results/profile-input-major.md +12 -0
- package/benchmarks/results/profile-k16.json +347 -0
- package/benchmarks/results/profile-k16.md +12 -0
- package/benchmarks/results/profile-k4.json +347 -0
- package/benchmarks/results/profile-k4.md +12 -0
- package/benchmarks/results/profile-pool-reuse.json +331 -0
- package/benchmarks/results/profile-pool-reuse.md +12 -0
- package/benchmarks/results/profile-precompiled.json +385 -0
- package/benchmarks/results/profile-precompiled.md +47 -0
- package/benchmarks/results/profile-static-channels.json +385 -0
- package/benchmarks/results/profile-static-channels.md +47 -0
- package/benchmarks/results/profile-static-io.json +385 -0
- package/benchmarks/results/profile-static-io.md +47 -0
- package/benchmarks/results/profile-tiled-conv.json +331 -0
- package/benchmarks/results/profile-tiled-conv.md +12 -0
- package/benchmarks/results/profile-tiled-decoder.json +347 -0
- package/benchmarks/results/profile-tiled-decoder.md +12 -0
- package/benchmarks/results/profile-tiled-matmul.json +331 -0
- package/benchmarks/results/profile-tiled-matmul.md +12 -0
- package/benchmarks/results/profile-unfused-decoder.json +379 -0
- package/benchmarks/results/profile-unfused-decoder.md +12 -0
- package/benchmarks/results/profile-unfused-pool.json +347 -0
- package/benchmarks/results/profile-unfused-pool.md +12 -0
- package/benchmarks/results/spatial-auto.json +575 -0
- package/benchmarks/results/spatial-auto.md +61 -0
- package/benchmarks/results/subgroup-smoke.json +1094 -0
- package/benchmarks/results/subgroup-smoke.md +104 -0
- package/benchmarks/results/webnn-smoke.json +739 -0
- package/benchmarks/results/webnn-smoke.md +76 -0
- package/dist/oidn.js +4189 -22603
- package/dist/oidn.umd.cjs +784 -5807
- package/lib/UNet.d.ts +66 -15
- package/lib/UNet.js +162 -257
- 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 +8 -4
- package/lib/backend.js +36 -44
- package/lib/backend.js.map +1 -1
- package/lib/graphOptimizer.d.ts +54 -0
- package/lib/graphOptimizer.js +216 -0
- package/lib/graphOptimizer.js.map +1 -0
- package/lib/main.d.ts +33 -10
- package/lib/main.js +5 -0
- 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 +67 -0
- package/lib/nativeUNet.js +1735 -0
- package/lib/nativeUNet.js.map +1 -0
- package/lib/process.js +38 -35
- 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 +33 -0
- package/lib/tileScheduler.js +86 -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 +9 -5
- package/scripts/inspect-model.mjs +64 -0
- package/src/UNet.ts +236 -339
- package/src/WGPUComputePass.ts +6 -4
- package/src/backend.ts +42 -55
- package/src/graphOptimizer.ts +301 -0
- package/src/main.ts +71 -11
- package/src/modelSpec.ts +414 -0
- package/src/nativeUNet.ts +2256 -0
- package/src/process.ts +38 -36
- package/src/resourceTracker.ts +94 -0
- package/src/tileScheduler.ts +138 -0
- package/src/webnnUNet.ts +812 -0
- package/tests/modelSpec.test.mjs +128 -0
- package/tests/resourceLifecycle.test.mjs +383 -0
- package/tests/tileScheduler.test.mjs +90 -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,23 +1,4 @@
|
|
|
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,
|
|
@@ -25,43 +6,26 @@ import {
|
|
|
25
6
|
hdrTransferFuncCPU,
|
|
26
7
|
hdrTransferFuncInverseCPU
|
|
27
8
|
} from './process';
|
|
28
|
-
import
|
|
29
|
-
|
|
30
|
-
|
|
31
|
-
|
|
32
|
-
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
|
|
38
|
-
|
|
39
|
-
|
|
40
|
-
|
|
41
|
-
|
|
42
|
-
|
|
43
|
-
|
|
44
|
-
|
|
45
|
-
|
|
46
|
-
}
|
|
9
|
+
import {
|
|
10
|
+
DynamicTileController,
|
|
11
|
+
type DynamicTileSetting,
|
|
12
|
+
fitTileDimension,
|
|
13
|
+
OIDN_TILE_ALIGNMENT,
|
|
14
|
+
waitForSubmittedGPUWork
|
|
15
|
+
} from './tileScheduler';
|
|
16
|
+
import {
|
|
17
|
+
detectUNetModelSpec,
|
|
18
|
+
validateUNetModel,
|
|
19
|
+
type UNetModelSpec
|
|
20
|
+
} from './modelSpec';
|
|
21
|
+
import {
|
|
22
|
+
NativeUNetExecutor,
|
|
23
|
+
type NativeUNetKernelSetting,
|
|
24
|
+
type NativeUNetPrecisionSetting
|
|
25
|
+
} from './nativeUNet';
|
|
26
|
+
import { WebNNUNetExecutor } from './webnnUNet';
|
|
47
27
|
|
|
48
|
-
|
|
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
|
-
}
|
|
28
|
+
export type UNetEngineSetting = 'auto' | 'wgsl' | 'webnn';
|
|
65
29
|
|
|
66
30
|
interface HDRImageData {
|
|
67
31
|
data: Float32Array;
|
|
@@ -81,13 +45,24 @@ interface GPUImageDataOutput {
|
|
|
81
45
|
height: number;
|
|
82
46
|
}
|
|
83
47
|
|
|
48
|
+
export interface UNetExecutionStats {
|
|
49
|
+
width: number;
|
|
50
|
+
height: number;
|
|
51
|
+
tileWidth: number;
|
|
52
|
+
tileHeight: number;
|
|
53
|
+
tileCount: number;
|
|
54
|
+
durationMs: number;
|
|
55
|
+
tileTimeMs: {
|
|
56
|
+
min: number;
|
|
57
|
+
median: number;
|
|
58
|
+
mean: number;
|
|
59
|
+
max: number;
|
|
60
|
+
};
|
|
61
|
+
}
|
|
62
|
+
|
|
84
63
|
function roundUp(a: number, b: number) {
|
|
85
64
|
return Math.ceil(a / b) * b;
|
|
86
65
|
}
|
|
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
66
|
|
|
92
67
|
function isGPUImageData(
|
|
93
68
|
data: ImageData | GPUImageData | HDRImageData
|
|
@@ -95,18 +70,7 @@ function isGPUImageData(
|
|
|
95
70
|
return data.data instanceof GPUBuffer || data.data instanceof GPUTexture;
|
|
96
71
|
}
|
|
97
72
|
|
|
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
73
|
class UNet {
|
|
109
|
-
private _tfModel: LayersModel | undefined;
|
|
110
74
|
private _device: GPUDevice | undefined;
|
|
111
75
|
|
|
112
76
|
// TODO calculate the tile size from memory size
|
|
@@ -121,15 +85,18 @@ class UNet {
|
|
|
121
85
|
private _hdr;
|
|
122
86
|
|
|
123
87
|
private _dataProcessGPU?: GPUDataProcess;
|
|
88
|
+
private _nativeExecutor?: NativeUNetExecutor;
|
|
89
|
+
private _webNNExecutor?: WebNNUNetExecutor;
|
|
90
|
+
private _modelSpec: UNetModelSpec;
|
|
91
|
+
private _inputChannels: number;
|
|
92
|
+
private _engine: UNetEngineSetting;
|
|
124
93
|
|
|
125
|
-
private
|
|
126
|
-
|
|
127
|
-
private _tensors = new Map<string, Tensor>();
|
|
128
|
-
private _modelsCache = new Map<string, LayersModel>();
|
|
94
|
+
private _dynamicTileController: DynamicTileController;
|
|
95
|
+
private _lastExecution?: UNetExecutionStats;
|
|
129
96
|
|
|
130
97
|
constructor(
|
|
131
|
-
|
|
132
|
-
|
|
98
|
+
hostTensors: Map<string, HostTensor>,
|
|
99
|
+
backend: { device: GPUDevice; adapterInfo: GPUAdapterInfo },
|
|
133
100
|
opts: {
|
|
134
101
|
/**
|
|
135
102
|
* If use auxiliary data.
|
|
@@ -140,228 +107,137 @@ class UNet {
|
|
|
140
107
|
*/
|
|
141
108
|
hdr?: boolean;
|
|
142
109
|
maxTileSize?: number;
|
|
110
|
+
dynamicTile?: DynamicTileSetting;
|
|
111
|
+
/** Reserved for explicit native WGSL selection. */
|
|
112
|
+
engine?: UNetEngineSetting;
|
|
113
|
+
/** Arithmetic/storage precision used by the native WGSL engine. */
|
|
114
|
+
precision?: NativeUNetPrecisionSetting;
|
|
115
|
+
/** Model-independent convolution kernel selection. */
|
|
116
|
+
kernel?: NativeUNetKernelSetting;
|
|
117
|
+
/** Explicit descriptor for a new OIDN topology not in the built-in registry. */
|
|
118
|
+
modelSpec?: UNetModelSpec;
|
|
143
119
|
} = {}
|
|
144
120
|
) {
|
|
145
121
|
this._aux = opts.aux || false;
|
|
146
122
|
this._hdr = opts.hdr || false;
|
|
123
|
+
this._engine = opts.engine ?? 'auto';
|
|
124
|
+
const modelSpec = opts.modelSpec ?? detectUNetModelSpec(hostTensors);
|
|
125
|
+
const validatedModel = validateUNetModel(hostTensors, modelSpec);
|
|
126
|
+
this._modelSpec = validatedModel.spec;
|
|
127
|
+
this._inputChannels = validatedModel.inputChannels;
|
|
128
|
+
|
|
129
|
+
const expectedInputChannels = this._aux ? 9 : 3;
|
|
130
|
+
if (validatedModel.inputChannels !== expectedInputChannels) {
|
|
131
|
+
throw new Error(
|
|
132
|
+
`OIDN model expects ${validatedModel.inputChannels} input channels, ` +
|
|
133
|
+
`but aux=${this._aux} provides ${expectedInputChannels}`
|
|
134
|
+
);
|
|
135
|
+
}
|
|
147
136
|
|
|
148
|
-
this.
|
|
137
|
+
this._dynamicTileController = new DynamicTileController(
|
|
138
|
+
opts.maxTileSize ?? 512,
|
|
139
|
+
opts.dynamicTile
|
|
140
|
+
);
|
|
149
141
|
|
|
150
|
-
this._device =
|
|
142
|
+
this._device = backend.device;
|
|
143
|
+
if (this._engine === 'webnn') {
|
|
144
|
+
this._webNNExecutor = new WebNNUNetExecutor(
|
|
145
|
+
this._device,
|
|
146
|
+
validatedModel,
|
|
147
|
+
{ precision: opts.precision }
|
|
148
|
+
);
|
|
149
|
+
} else {
|
|
150
|
+
this._nativeExecutor = new NativeUNetExecutor(
|
|
151
|
+
this._device,
|
|
152
|
+
validatedModel,
|
|
153
|
+
{ precision: opts.precision, kernel: opts.kernel }
|
|
154
|
+
);
|
|
155
|
+
}
|
|
151
156
|
}
|
|
152
157
|
|
|
153
158
|
getDevice() {
|
|
154
159
|
return this._device;
|
|
155
160
|
}
|
|
156
161
|
|
|
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'
|
|
162
|
+
/** Completes backend compilation before first interactive use. */
|
|
163
|
+
async prepare() {
|
|
164
|
+
if (this._webNNExecutor) {
|
|
165
|
+
await this._webNNExecutor.prepare();
|
|
166
|
+
const overlap = roundUp(
|
|
167
|
+
this._modelSpec.receptiveField / 2,
|
|
168
|
+
OIDN_TILE_ALIGNMENT
|
|
208
169
|
);
|
|
209
|
-
|
|
210
|
-
|
|
211
|
-
|
|
212
|
-
|
|
213
|
-
|
|
214
|
-
|
|
215
|
-
|
|
216
|
-
|
|
170
|
+
const outputTileEdges = [
|
|
171
|
+
this._dynamicTileController.tileSize,
|
|
172
|
+
this._dynamicTileController.minTileSize
|
|
173
|
+
];
|
|
174
|
+
await this._webNNExecutor.prewarm(
|
|
175
|
+
[...new Set(outputTileEdges)].map((edge) => ({
|
|
176
|
+
width: edge + 2 * overlap,
|
|
177
|
+
height: edge + 2 * overlap
|
|
178
|
+
}))
|
|
217
179
|
);
|
|
218
|
-
|
|
180
|
+
return;
|
|
219
181
|
}
|
|
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;
|
|
182
|
+
await this._nativeExecutor!.prepare();
|
|
233
183
|
}
|
|
234
184
|
|
|
235
|
-
|
|
236
|
-
|
|
237
|
-
|
|
238
|
-
|
|
239
|
-
|
|
240
|
-
|
|
241
|
-
|
|
242
|
-
|
|
243
|
-
|
|
244
|
-
|
|
245
|
-
|
|
246
|
-
|
|
247
|
-
|
|
248
|
-
|
|
249
|
-
|
|
250
|
-
|
|
251
|
-
|
|
185
|
+
getRuntimeInfo() {
|
|
186
|
+
return {
|
|
187
|
+
configuredEngine: this._engine,
|
|
188
|
+
gpuEngine: this._webNNExecutor ? 'webnn' as const : 'wgsl' as const,
|
|
189
|
+
precision: (this._webNNExecutor ?? this._nativeExecutor!).precision,
|
|
190
|
+
kernel: this._nativeExecutor
|
|
191
|
+
? {
|
|
192
|
+
configured: this._nativeExecutor.kernelSetting,
|
|
193
|
+
maxSpatialInputBlocks: this._nativeExecutor.maxSpatialInputBlocks,
|
|
194
|
+
subgroupsAvailable: this._nativeExecutor.subgroupsAvailable
|
|
195
|
+
}
|
|
196
|
+
: undefined,
|
|
197
|
+
webnn: this._webNNExecutor?.support,
|
|
198
|
+
resources: (
|
|
199
|
+
this._webNNExecutor ?? this._nativeExecutor!
|
|
200
|
+
).getResourceInfo(),
|
|
201
|
+
model: this._modelSpec.id,
|
|
202
|
+
modelFamily: this._modelSpec.family,
|
|
203
|
+
inputChannels: this._inputChannels,
|
|
204
|
+
dynamicTile: {
|
|
205
|
+
enabled: this._dynamicTileController.enabled,
|
|
206
|
+
currentTileSize: this._dynamicTileController.tileSize,
|
|
207
|
+
minTileSize: this._dynamicTileController.minTileSize,
|
|
208
|
+
maxTileSize: this._dynamicTileController.maxTileSize,
|
|
209
|
+
targetTileTimeMs: this._dynamicTileController.targetTileTimeMs
|
|
210
|
+
},
|
|
211
|
+
lastExecution: this._lastExecution
|
|
212
|
+
};
|
|
252
213
|
}
|
|
253
214
|
|
|
254
|
-
|
|
255
|
-
|
|
256
|
-
|
|
257
|
-
poolSize: [2, 2],
|
|
258
|
-
strides: [2, 2],
|
|
259
|
-
padding: 'same',
|
|
260
|
-
trainable: false
|
|
261
|
-
});
|
|
262
|
-
// https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L33
|
|
263
|
-
return poolingLayer.apply(source) as SymbolicTensor;
|
|
215
|
+
/** Captures per-node GPU timestamps for the next native tile execution. */
|
|
216
|
+
profileNextExecution() {
|
|
217
|
+
return this._nativeExecutor?.profileNextExecution() ?? false;
|
|
264
218
|
}
|
|
265
219
|
|
|
266
|
-
|
|
267
|
-
|
|
268
|
-
name: source.name + '/upsampling',
|
|
269
|
-
size: [2, 2],
|
|
270
|
-
trainable: false
|
|
271
|
-
});
|
|
272
|
-
return upsamplingLayer.apply(source) as SymbolicTensor;
|
|
220
|
+
getLastExecutionProfile() {
|
|
221
|
+
return this._nativeExecutor?.getLastExecutionProfile();
|
|
273
222
|
}
|
|
274
223
|
|
|
275
|
-
private
|
|
276
|
-
|
|
277
|
-
const pool1 = (x = this._createPooling(
|
|
278
|
-
this._createConv('enc_conv1', x, 'relu')
|
|
279
|
-
));
|
|
280
|
-
const pool2 = (x = this._createPooling(
|
|
281
|
-
this._createConv('enc_conv2', x, 'relu')
|
|
282
|
-
));
|
|
283
|
-
const pool3 = (x = this._createPooling(
|
|
284
|
-
this._createConv('enc_conv3', x, 'relu')
|
|
285
|
-
));
|
|
286
|
-
const pool4 = (x = this._createPooling(
|
|
287
|
-
this._createConv('enc_conv4', x, 'relu')
|
|
288
|
-
));
|
|
289
|
-
x = this._createConv('enc_conv5a', pool4, 'relu');
|
|
290
|
-
x = this._addUpsamplingLayer(this._createConv('enc_conv5b', x, 'relu'));
|
|
291
|
-
|
|
292
|
-
x = this._createConcatConv('dec_conv4a', x, pool3);
|
|
293
|
-
x = this._addUpsamplingLayer(this._createConv('dec_conv4b', x, 'relu'));
|
|
294
|
-
|
|
295
|
-
x = this._createConcatConv('dec_conv3a', x, pool2);
|
|
296
|
-
x = this._addUpsamplingLayer(this._createConv('dec_conv3b', x, 'relu'));
|
|
297
|
-
|
|
298
|
-
x = this._createConcatConv('dec_conv2a', x, pool1);
|
|
299
|
-
x = this._addUpsamplingLayer(this._createConv('dec_conv2b', x, 'relu'));
|
|
300
|
-
|
|
301
|
-
x = this._createConcatConv('dec_conv1a', x, input);
|
|
302
|
-
x = this._createConv('dec_conv1b', x, 'relu');
|
|
303
|
-
x = this._createConv('dec_conv0', x, 'relu');
|
|
304
|
-
|
|
305
|
-
return x;
|
|
306
|
-
}
|
|
224
|
+
private _updateModel(width: number, height: number) {
|
|
225
|
+
const maxTileSize = this._dynamicTileController.tileSize;
|
|
307
226
|
|
|
308
|
-
|
|
309
|
-
let
|
|
310
|
-
const
|
|
311
|
-
this.
|
|
312
|
-
|
|
313
|
-
|
|
314
|
-
|
|
315
|
-
|
|
316
|
-
));
|
|
317
|
-
x = this._createConv('enc_conv3a', x, 'relu');
|
|
318
|
-
const pool3 = (x = this._createPooling(
|
|
319
|
-
this._createConv('enc_conv3b', x, 'relu')
|
|
320
|
-
));
|
|
321
|
-
x = this._createConv('enc_conv4a', x, 'relu');
|
|
322
|
-
const pool4 = (x = this._createPooling(
|
|
323
|
-
this._createConv('enc_conv4b', x, 'relu')
|
|
324
|
-
));
|
|
325
|
-
|
|
326
|
-
x = this._createConv('enc_conv5a', pool4, 'relu');
|
|
327
|
-
x = this._addUpsamplingLayer(this._createConv('enc_conv5b', x, 'relu'));
|
|
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;
|
|
343
|
-
}
|
|
227
|
+
let tileWidth = fitTileDimension(width, maxTileSize);
|
|
228
|
+
let tileHeight = fitTileDimension(height, maxTileSize);
|
|
229
|
+
const defaultTileOverlap = roundUp(
|
|
230
|
+
this._modelSpec.receptiveField / 2,
|
|
231
|
+
OIDN_TILE_ALIGNMENT
|
|
232
|
+
);
|
|
233
|
+
let tileOverlapX = defaultTileOverlap;
|
|
234
|
+
let tileOverlapY = defaultTileOverlap;
|
|
344
235
|
|
|
345
|
-
|
|
346
|
-
|
|
347
|
-
const maxTileSize = this._maxTileSize;
|
|
348
|
-
|
|
349
|
-
let tileWidth = maxTileSize;
|
|
350
|
-
let tileHeight = maxTileSize;
|
|
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
|
-
}
|
|
236
|
+
if (width <= maxTileSize) {
|
|
237
|
+
tileOverlapX = 0;
|
|
359
238
|
}
|
|
360
|
-
if (height
|
|
361
|
-
|
|
362
|
-
if (height <= maxTileSize) {
|
|
363
|
-
tileOverlapY = 0;
|
|
364
|
-
}
|
|
239
|
+
if (height <= maxTileSize) {
|
|
240
|
+
tileOverlapY = 0;
|
|
365
241
|
}
|
|
366
242
|
|
|
367
243
|
// Force width and height has same size. reduce the cache in memory
|
|
@@ -376,16 +252,13 @@ class UNet {
|
|
|
376
252
|
tileWidth !== this._tileWidth ||
|
|
377
253
|
tileHeight !== this._tileHeight ||
|
|
378
254
|
tileOverlapX !== this._tileOverlapX ||
|
|
379
|
-
tileOverlapY !== this._tileOverlapY
|
|
380
|
-
!this._tfModel
|
|
255
|
+
tileOverlapY !== this._tileOverlapY
|
|
381
256
|
) {
|
|
382
257
|
// console.log(tileWidth, tileHeight, tileOverlapX, tileOverlapY);
|
|
383
258
|
this._tileWidth = tileWidth;
|
|
384
259
|
this._tileHeight = tileHeight;
|
|
385
260
|
this._tileOverlapX = tileOverlapX;
|
|
386
261
|
this._tileOverlapY = tileOverlapY;
|
|
387
|
-
|
|
388
|
-
this._buildModel(isLarge);
|
|
389
262
|
}
|
|
390
263
|
}
|
|
391
264
|
|
|
@@ -497,7 +370,7 @@ class UNet {
|
|
|
497
370
|
}
|
|
498
371
|
}
|
|
499
372
|
|
|
500
|
-
private _executeTile(
|
|
373
|
+
private async _executeTile(
|
|
501
374
|
inputData:
|
|
502
375
|
| Float32Array
|
|
503
376
|
| {
|
|
@@ -534,7 +407,8 @@ class UNet {
|
|
|
534
407
|
|
|
535
408
|
const srcTile = new Tile(srcX0, srcY0, srcTileWidth, srcTileHeight);
|
|
536
409
|
|
|
537
|
-
let
|
|
410
|
+
let nativeOutputBuffer: GPUBuffer | undefined;
|
|
411
|
+
let denoisedData: Float32Array | undefined;
|
|
538
412
|
let inputScale = 1;
|
|
539
413
|
const device = this._device!;
|
|
540
414
|
let dataProcessGPU = this._dataProcessGPU;
|
|
@@ -552,11 +426,11 @@ class UNet {
|
|
|
552
426
|
inputScale
|
|
553
427
|
});
|
|
554
428
|
}
|
|
555
|
-
|
|
429
|
+
denoisedData = await (this._webNNExecutor ?? this._nativeExecutor!).executeCPU(
|
|
556
430
|
tileData,
|
|
557
|
-
|
|
558
|
-
|
|
559
|
-
)
|
|
431
|
+
srcTileWidth,
|
|
432
|
+
srcTileHeight
|
|
433
|
+
);
|
|
560
434
|
} else {
|
|
561
435
|
if (!dataProcessGPU) {
|
|
562
436
|
dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(
|
|
@@ -577,33 +451,14 @@ class UNet {
|
|
|
577
451
|
denoiseAlpha
|
|
578
452
|
);
|
|
579
453
|
|
|
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
|
-
}
|
|
454
|
+
nativeOutputBuffer = await (this._webNNExecutor ?? this._nativeExecutor!).execute(
|
|
455
|
+
this._aux ? [color, albedo!, normal!] : [color],
|
|
456
|
+
srcTileWidth,
|
|
457
|
+
srcTileHeight
|
|
458
|
+
);
|
|
603
459
|
}
|
|
604
460
|
|
|
605
461
|
let outBuffer: GPUBuffer;
|
|
606
|
-
const outputTensor = this._tfModel!.predict(tileTensor) as Tensor;
|
|
607
462
|
|
|
608
463
|
const dstWidth = Math.min(dstTileSize.width, width);
|
|
609
464
|
const dstHeight = Math.min(dstTileSize.height, height);
|
|
@@ -612,10 +467,9 @@ class UNet {
|
|
|
612
467
|
dstTile.height = Math.min(dstTile.height, height - dstTile.y);
|
|
613
468
|
|
|
614
469
|
if (inputData instanceof Float32Array) {
|
|
615
|
-
let denoisedData = outputTensor.dataSync();
|
|
616
470
|
if (isHDR) {
|
|
617
471
|
denoisedData = hdrTransferFuncInverseCPU({
|
|
618
|
-
data: denoisedData
|
|
472
|
+
data: denoisedData!,
|
|
619
473
|
channels: 3,
|
|
620
474
|
inputScale
|
|
621
475
|
});
|
|
@@ -625,7 +479,7 @@ class UNet {
|
|
|
625
479
|
outputImageData!,
|
|
626
480
|
srcTile,
|
|
627
481
|
dstTile,
|
|
628
|
-
denoisedData
|
|
482
|
+
denoisedData!,
|
|
629
483
|
srcTileSize.width,
|
|
630
484
|
isHDR
|
|
631
485
|
);
|
|
@@ -641,17 +495,8 @@ class UNet {
|
|
|
641
495
|
}
|
|
642
496
|
} else {
|
|
643
497
|
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
498
|
outBuffer = dataProcessGPU!.inverse(
|
|
654
|
-
|
|
499
|
+
nativeOutputBuffer!,
|
|
655
500
|
inputData.color
|
|
656
501
|
);
|
|
657
502
|
}
|
|
@@ -694,6 +539,9 @@ class UNet {
|
|
|
694
539
|
|
|
695
540
|
const width = color.width;
|
|
696
541
|
const height = color.height;
|
|
542
|
+
const adaptiveTileSize = this._dynamicTileController.tileSize;
|
|
543
|
+
const shouldAdaptTileSize =
|
|
544
|
+
width > adaptiveTileSize || height > adaptiveTileSize;
|
|
697
545
|
this._updateModel(width, height);
|
|
698
546
|
|
|
699
547
|
// TODO should fixed to be hdr when UNet is created.
|
|
@@ -733,14 +581,24 @@ class UNet {
|
|
|
733
581
|
|
|
734
582
|
let aborted = false;
|
|
735
583
|
|
|
736
|
-
const
|
|
584
|
+
const now = () =>
|
|
585
|
+
typeof performance === 'undefined' ? Date.now() : performance.now();
|
|
586
|
+
const executionStartTime = now();
|
|
587
|
+
const tileTimesMs: number[] = [];
|
|
588
|
+
const scheduleNextTile = (callback: () => void) => {
|
|
589
|
+
if (typeof requestAnimationFrame === 'undefined') {
|
|
590
|
+
setTimeout(callback, 0);
|
|
591
|
+
} else {
|
|
592
|
+
requestAnimationFrame(callback);
|
|
593
|
+
}
|
|
594
|
+
};
|
|
595
|
+
|
|
596
|
+
const executeTile = async (i: number, j: number) => {
|
|
737
597
|
if (aborted) {
|
|
738
598
|
return;
|
|
739
599
|
}
|
|
740
|
-
|
|
741
|
-
|
|
742
|
-
ENGINE.startScope();
|
|
743
|
-
resGPUBuffer = this._executeTile(
|
|
600
|
+
const tileStartTime = now();
|
|
601
|
+
const resGPUBuffer = await this._executeTile(
|
|
744
602
|
isGPUImageData(color)
|
|
745
603
|
? {
|
|
746
604
|
color: color.data,
|
|
@@ -757,8 +615,7 @@ class UNet {
|
|
|
757
615
|
hdr,
|
|
758
616
|
denoiseAlpha
|
|
759
617
|
);
|
|
760
|
-
|
|
761
|
-
// }, true);
|
|
618
|
+
if (aborted) return;
|
|
762
619
|
const output = outputImageData || {
|
|
763
620
|
data: resGPUBuffer,
|
|
764
621
|
width,
|
|
@@ -773,18 +630,58 @@ class UNet {
|
|
|
773
630
|
tileCountW * tileCountH
|
|
774
631
|
);
|
|
775
632
|
|
|
776
|
-
|
|
777
|
-
|
|
778
|
-
|
|
779
|
-
|
|
780
|
-
|
|
781
|
-
|
|
633
|
+
const hasNextTile = i + 1 < tileCountW || j + 1 < tileCountH;
|
|
634
|
+
const continueAfterGPUWork = () => {
|
|
635
|
+
tileTimesMs.push(now() - tileStartTime);
|
|
636
|
+
if (aborted) return;
|
|
637
|
+
|
|
638
|
+
if (hasNextTile) {
|
|
639
|
+
scheduleNextTile(() => {
|
|
640
|
+
if (aborted) return;
|
|
641
|
+
if (i + 1 < tileCountW) {
|
|
642
|
+
executeTile(i + 1, j);
|
|
643
|
+
} else if (j + 1 < tileCountH) {
|
|
644
|
+
executeTile(0, j + 1);
|
|
645
|
+
}
|
|
646
|
+
});
|
|
647
|
+
} else {
|
|
648
|
+
const sortedTileTimes = [...tileTimesMs].sort((a, b) => a - b);
|
|
649
|
+
const middle = Math.floor(sortedTileTimes.length / 2);
|
|
650
|
+
const medianTileTime = sortedTileTimes.length % 2
|
|
651
|
+
? sortedTileTimes[middle]
|
|
652
|
+
: (sortedTileTimes[middle - 1] + sortedTileTimes[middle]) / 2;
|
|
653
|
+
this._lastExecution = {
|
|
654
|
+
width,
|
|
655
|
+
height,
|
|
656
|
+
tileWidth,
|
|
657
|
+
tileHeight,
|
|
658
|
+
tileCount: tileCountW * tileCountH,
|
|
659
|
+
durationMs: now() - executionStartTime,
|
|
660
|
+
tileTimeMs: {
|
|
661
|
+
min: sortedTileTimes[0],
|
|
662
|
+
median: medianTileTime,
|
|
663
|
+
mean:
|
|
664
|
+
sortedTileTimes.reduce((sum, value) => sum + value, 0) /
|
|
665
|
+
sortedTileTimes.length,
|
|
666
|
+
max: sortedTileTimes[sortedTileTimes.length - 1]
|
|
667
|
+
}
|
|
668
|
+
};
|
|
669
|
+
// Adapt only from complete executions. Cancelled work is commonly
|
|
670
|
+
// contending with interactive rendering and is not representative.
|
|
671
|
+
if (shouldAdaptTileSize) {
|
|
672
|
+
this._dynamicTileController.observe(tileTimesMs);
|
|
782
673
|
}
|
|
783
|
-
|
|
784
|
-
|
|
785
|
-
|
|
786
|
-
|
|
787
|
-
|
|
674
|
+
// console.log(memory());
|
|
675
|
+
done(output as any);
|
|
676
|
+
}
|
|
677
|
+
};
|
|
678
|
+
|
|
679
|
+
// requestAnimationFrame only throttles JavaScript submission. Waiting
|
|
680
|
+
// for the queue here keeps at most one OIDN tile in flight, so aborting
|
|
681
|
+
// cannot leave a long tail of already-submitted GPU work.
|
|
682
|
+
void waitForSubmittedGPUWork(this._device!.queue).then(
|
|
683
|
+
continueAfterGPUWork
|
|
684
|
+
);
|
|
788
685
|
};
|
|
789
686
|
|
|
790
687
|
executeTile(0, 0);
|
|
@@ -795,9 +692,9 @@ class UNet {
|
|
|
795
692
|
}
|
|
796
693
|
|
|
797
694
|
dispose() {
|
|
798
|
-
this._tfModel?.dispose();
|
|
799
695
|
this._dataProcessGPU?.dispose();
|
|
800
|
-
this.
|
|
696
|
+
this._nativeExecutor?.dispose();
|
|
697
|
+
this._webNNExecutor?.dispose();
|
|
801
698
|
}
|
|
802
699
|
}
|
|
803
700
|
|