oidn-web 0.1.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/LICENSE +21 -0
- package/README.md +2 -0
- package/dist/oidn.mjs +22909 -0
- package/dist/oidn.umd.js +5905 -0
- package/lib/UNet.d.ts +54 -0
- package/lib/UNet.js +467 -0
- package/lib/UNet.js.map +1 -0
- package/lib/WGPUComputePass.d.ts +53 -0
- package/lib/WGPUComputePass.js +220 -0
- package/lib/WGPUComputePass.js.map +1 -0
- package/lib/WGPUFullQuadPass.d.ts +51 -0
- package/lib/WGPUFullQuadPass.js +261 -0
- package/lib/WGPUFullQuadPass.js.map +1 -0
- package/lib/backend.d.ts +5 -0
- package/lib/backend.js +41 -0
- package/lib/backend.js.map +1 -0
- package/lib/hdr.d.ts +31 -0
- package/lib/hdr.js +340 -0
- package/lib/hdr.js.map +1 -0
- package/lib/helper.d.ts +1 -0
- package/lib/helper.js +27 -0
- package/lib/helper.js.map +1 -0
- package/lib/kernels.d.ts +1 -0
- package/lib/kernels.js +26 -0
- package/lib/kernels.js.map +1 -0
- package/lib/main.d.ts +20 -0
- package/lib/main.js +20 -0
- package/lib/main.js.map +1 -0
- package/lib/process.d.ts +39 -0
- package/lib/process.js +309 -0
- package/lib/process.js.map +1 -0
- package/lib/tza.d.ts +15 -0
- package/lib/tza.js +114 -0
- package/lib/tza.js.map +1 -0
- package/package.json +26 -0
- package/src/UNet.ts +708 -0
- package/src/WGPUComputePass.ts +318 -0
- package/src/WGPUFullQuadPass.ts +348 -0
- package/src/backend.ts +53 -0
- package/src/hdr.ts +398 -0
- package/src/helper.ts +35 -0
- package/src/kernels.ts +31 -0
- package/src/main.ts +42 -0
- package/src/process.ts +362 -0
- package/src/tza.ts +136 -0
- package/weights/.gitattributes +1 -0
- package/weights/LICENSE.txt +202 -0
- package/weights/README.md +7 -0
- package/weights/rt_hdr.tza +0 -0
- package/weights/rt_hdr_alb_nrm.tza +0 -0
- package/weights/rt_ldr.tza +0 -0
- package/weights/rt_ldr_alb_nrm.tza +0 -0
package/src/UNet.ts
ADDED
|
@@ -0,0 +1,708 @@
|
|
|
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 {
|
|
11
|
+
Conv2D,
|
|
12
|
+
UpSampling2D
|
|
13
|
+
} from '@tensorflow/tfjs-layers/dist/layers/convolutional';
|
|
14
|
+
import { MaxPooling2D } from '@tensorflow/tfjs-layers/dist/layers/pooling';
|
|
15
|
+
import { Concatenate } from '@tensorflow/tfjs-layers/dist/layers/merge';
|
|
16
|
+
import { LayersModel } from '@tensorflow/tfjs-layers/dist/engine/training';
|
|
17
|
+
import { Input as TFInput } from '@tensorflow/tfjs-layers/dist/engine/input_layer';
|
|
18
|
+
import { HostTensor } from './tza';
|
|
19
|
+
import { Float16Array } from '@petamoriken/float16';
|
|
20
|
+
import {
|
|
21
|
+
GPUDataProcess,
|
|
22
|
+
Tile,
|
|
23
|
+
avgLogLum,
|
|
24
|
+
hdrTransferFuncCPU,
|
|
25
|
+
hdrTransferFuncInverseCPU
|
|
26
|
+
} from './process';
|
|
27
|
+
import type { WebGPUBackend } from '@tensorflow/tfjs-backend-webgpu';
|
|
28
|
+
// import { profileAndLogKernelCode } from './helper';
|
|
29
|
+
|
|
30
|
+
function getTensorData(
|
|
31
|
+
ubytes: Uint8Array,
|
|
32
|
+
type: HostTensor['desc']['dataType']
|
|
33
|
+
) {
|
|
34
|
+
const buffer = ubytes.buffer;
|
|
35
|
+
if (type === 'Float32') {
|
|
36
|
+
return new Float32Array(ubytes.buffer);
|
|
37
|
+
}
|
|
38
|
+
const float16Data = new Float16Array(buffer);
|
|
39
|
+
const float32Data = new Float32Array(float16Data.length);
|
|
40
|
+
for (let i = 0; i < float32Data.length; ++i) {
|
|
41
|
+
float32Data[i] = float16Data[i];
|
|
42
|
+
}
|
|
43
|
+
return float32Data;
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
function changeWeightShapes(weightData: Float32Array, dims: number[]) {
|
|
47
|
+
const [O, C, H, W] = dims;
|
|
48
|
+
const reorderedWeightData = new Float32Array(weightData.length);
|
|
49
|
+
for (let o = 0; o < O; ++o) {
|
|
50
|
+
for (let c = 0; c < C; ++c) {
|
|
51
|
+
for (let h = 0; h < H; ++h) {
|
|
52
|
+
for (let w = 0; w < W; ++w) {
|
|
53
|
+
// Change OCHW to HWCO
|
|
54
|
+
const idx = o * C * H * W + c * H * W + h * W + w;
|
|
55
|
+
const idx2 = h * W * C * O + w * C * O + c * O + o;
|
|
56
|
+
reorderedWeightData[idx2] = weightData[idx];
|
|
57
|
+
}
|
|
58
|
+
}
|
|
59
|
+
}
|
|
60
|
+
}
|
|
61
|
+
return reorderedWeightData;
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
interface HDRImageData {
|
|
65
|
+
data: Float32Array;
|
|
66
|
+
width: number;
|
|
67
|
+
height: number;
|
|
68
|
+
}
|
|
69
|
+
|
|
70
|
+
interface GPUImageData {
|
|
71
|
+
data: GPUBuffer;
|
|
72
|
+
width: number;
|
|
73
|
+
height: number;
|
|
74
|
+
}
|
|
75
|
+
|
|
76
|
+
function roundUp(a: number, b: number) {
|
|
77
|
+
return Math.ceil(a / b) * b;
|
|
78
|
+
}
|
|
79
|
+
// Returns the smallest integer larger than or equal to a which has remainder c when divided by b
|
|
80
|
+
function roundUp2(a: number, b: number, c: number) {
|
|
81
|
+
return Math.ceil((a - c) / b) * b + c;
|
|
82
|
+
}
|
|
83
|
+
|
|
84
|
+
function isGPUImageData(
|
|
85
|
+
data: ImageData | GPUImageData | HDRImageData
|
|
86
|
+
): data is GPUImageData {
|
|
87
|
+
return data.data instanceof GPUBuffer;
|
|
88
|
+
}
|
|
89
|
+
|
|
90
|
+
const receptiveField = 174; // receptive field in pixels
|
|
91
|
+
// TODO metal is 32?
|
|
92
|
+
const minTileAlignment = 1;
|
|
93
|
+
|
|
94
|
+
const tileAlignment = 16; // required spatial alignment in pixels (padding may be necessary)
|
|
95
|
+
|
|
96
|
+
const defaultTileOverlap = roundUp(receptiveField / 2, tileAlignment);
|
|
97
|
+
class UNet {
|
|
98
|
+
private _tfModel: LayersModel | undefined;
|
|
99
|
+
private _device: GPUDevice | undefined;
|
|
100
|
+
|
|
101
|
+
// TODO calculate the tile size from memory size
|
|
102
|
+
// https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/core/unet_filter.cpp#L287
|
|
103
|
+
private _tileWidth = 0;
|
|
104
|
+
private _tileHeight = 0;
|
|
105
|
+
|
|
106
|
+
private _tileOverlapX = 0;
|
|
107
|
+
private _tileOverlapY = 0;
|
|
108
|
+
|
|
109
|
+
private _aux = false;
|
|
110
|
+
private _hdr = false;
|
|
111
|
+
|
|
112
|
+
private _dataProcessGPU?: GPUDataProcess;
|
|
113
|
+
|
|
114
|
+
private _maxTileSize;
|
|
115
|
+
|
|
116
|
+
constructor(
|
|
117
|
+
private _tensors: Map<string, HostTensor>,
|
|
118
|
+
private _backend: WebGPUBackend,
|
|
119
|
+
opts: {
|
|
120
|
+
aux?: boolean;
|
|
121
|
+
hdr?: boolean;
|
|
122
|
+
maxTileSize?: number;
|
|
123
|
+
} = {}
|
|
124
|
+
) {
|
|
125
|
+
this._aux = opts.aux || false;
|
|
126
|
+
this._hdr = opts.hdr || false;
|
|
127
|
+
|
|
128
|
+
this._maxTileSize = opts.maxTileSize ?? 512;
|
|
129
|
+
|
|
130
|
+
this._device = this._backend.device;
|
|
131
|
+
}
|
|
132
|
+
|
|
133
|
+
private _createConv(
|
|
134
|
+
name: string,
|
|
135
|
+
source: SymbolicTensor,
|
|
136
|
+
activation?: 'relu'
|
|
137
|
+
) {
|
|
138
|
+
const unetWeightTensor = this._tensors.get(name + '.weight')!;
|
|
139
|
+
const unetBiasTensor = this._tensors.get(name + '.bias')!;
|
|
140
|
+
const weightDims = unetWeightTensor.desc.dims;
|
|
141
|
+
const weightTensor = tensor(
|
|
142
|
+
changeWeightShapes(
|
|
143
|
+
getTensorData(unetWeightTensor.data, unetWeightTensor.desc.dataType),
|
|
144
|
+
weightDims
|
|
145
|
+
),
|
|
146
|
+
[weightDims[2], weightDims[3], weightDims[1], weightDims[0]],
|
|
147
|
+
'float32'
|
|
148
|
+
);
|
|
149
|
+
const biasTensor = tensor1d(
|
|
150
|
+
getTensorData(unetBiasTensor.data, unetBiasTensor.desc.dataType),
|
|
151
|
+
'float32'
|
|
152
|
+
);
|
|
153
|
+
// TODO whats the purpose of padded dims ?
|
|
154
|
+
const convLayer = new Conv2D({
|
|
155
|
+
name,
|
|
156
|
+
filters: unetWeightTensor.desc.dims[0],
|
|
157
|
+
kernelSize: unetWeightTensor.desc.dims.slice(2, 4) as [number, number],
|
|
158
|
+
useBias: true,
|
|
159
|
+
activation,
|
|
160
|
+
padding: 'same',
|
|
161
|
+
weights: [weightTensor, biasTensor],
|
|
162
|
+
trainable: false
|
|
163
|
+
});
|
|
164
|
+
|
|
165
|
+
return convLayer.apply(source) as SymbolicTensor;
|
|
166
|
+
}
|
|
167
|
+
|
|
168
|
+
private _createConcatConv(
|
|
169
|
+
name: string,
|
|
170
|
+
source1: SymbolicTensor,
|
|
171
|
+
source2: SymbolicTensor
|
|
172
|
+
) {
|
|
173
|
+
//https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L40
|
|
174
|
+
return this._createConv(
|
|
175
|
+
name,
|
|
176
|
+
// Concat on the channel
|
|
177
|
+
new Concatenate({ trainable: false, axis: 3 }).apply([
|
|
178
|
+
// convLayer.apply(source2) as SymbolicTensor,
|
|
179
|
+
source1,
|
|
180
|
+
source2
|
|
181
|
+
]) as SymbolicTensor,
|
|
182
|
+
'relu'
|
|
183
|
+
) as SymbolicTensor;
|
|
184
|
+
}
|
|
185
|
+
|
|
186
|
+
private _createPooling(source: SymbolicTensor) {
|
|
187
|
+
// https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L33
|
|
188
|
+
return new MaxPooling2D({
|
|
189
|
+
name: source.name + '/pooling',
|
|
190
|
+
poolSize: [2, 2],
|
|
191
|
+
strides: [2, 2],
|
|
192
|
+
padding: 'same',
|
|
193
|
+
trainable: false
|
|
194
|
+
}).apply(source) as SymbolicTensor;
|
|
195
|
+
}
|
|
196
|
+
|
|
197
|
+
private _addUpsamplingLayer(source: SymbolicTensor) {
|
|
198
|
+
return new UpSampling2D({
|
|
199
|
+
name: source.name + '/upsampling',
|
|
200
|
+
size: [2, 2],
|
|
201
|
+
trainable: false
|
|
202
|
+
}).apply(source) as SymbolicTensor;
|
|
203
|
+
}
|
|
204
|
+
|
|
205
|
+
getDevice() {
|
|
206
|
+
return this._device;
|
|
207
|
+
}
|
|
208
|
+
|
|
209
|
+
buildModel() {
|
|
210
|
+
const aux = this._aux;
|
|
211
|
+
const channels = 3 + (aux ? 6 : 0);
|
|
212
|
+
const tileSize = this._getTileSizeWithOverlap();
|
|
213
|
+
|
|
214
|
+
// TODO input process transferFunc
|
|
215
|
+
// TODO input shape
|
|
216
|
+
const input = TFInput({
|
|
217
|
+
shape: [tileSize.height, tileSize.width, channels],
|
|
218
|
+
dtype: 'float32'
|
|
219
|
+
});
|
|
220
|
+
|
|
221
|
+
const encConv0 = this._createConv('enc_conv0', input, 'relu');
|
|
222
|
+
const pool1 = this._createPooling(
|
|
223
|
+
this._createConv('enc_conv1', encConv0, 'relu')
|
|
224
|
+
);
|
|
225
|
+
const pool2 = this._createPooling(
|
|
226
|
+
this._createConv('enc_conv2', pool1, 'relu')
|
|
227
|
+
);
|
|
228
|
+
const pool3 = this._createPooling(
|
|
229
|
+
this._createConv('enc_conv3', pool2, 'relu')
|
|
230
|
+
);
|
|
231
|
+
const pool4 = this._createPooling(
|
|
232
|
+
this._createConv('enc_conv4', pool3, 'relu')
|
|
233
|
+
);
|
|
234
|
+
const encConv5a = this._createConv('enc_conv5a', pool4, 'relu');
|
|
235
|
+
const upsample4 = this._addUpsamplingLayer(
|
|
236
|
+
this._createConv('enc_conv5b', encConv5a, 'relu')
|
|
237
|
+
);
|
|
238
|
+
const decConv4a = this._createConcatConv('dec_conv4a', upsample4, pool3);
|
|
239
|
+
const upsample3 = this._addUpsamplingLayer(
|
|
240
|
+
this._createConv('dec_conv4b', decConv4a, 'relu')
|
|
241
|
+
);
|
|
242
|
+
const decConv3a = this._createConcatConv('dec_conv3a', upsample3, pool2);
|
|
243
|
+
const upsample2 = this._addUpsamplingLayer(
|
|
244
|
+
this._createConv('dec_conv3b', decConv3a, 'relu')
|
|
245
|
+
);
|
|
246
|
+
const decConv2a = this._createConcatConv('dec_conv2a', upsample2, pool1);
|
|
247
|
+
const upsample1 = this._addUpsamplingLayer(
|
|
248
|
+
this._createConv('dec_conv2b', decConv2a, 'relu')
|
|
249
|
+
);
|
|
250
|
+
const decConv1a = this._createConcatConv('dec_conv1a', upsample1, input);
|
|
251
|
+
const decConv1b = this._createConv('dec_conv1b', decConv1a, 'relu');
|
|
252
|
+
const decConv0 = this._createConv('dec_conv0', decConv1b, 'relu');
|
|
253
|
+
|
|
254
|
+
this._tfModel = new LayersModel({
|
|
255
|
+
inputs: [input],
|
|
256
|
+
// TODO output process transferFunc
|
|
257
|
+
outputs: decConv0
|
|
258
|
+
});
|
|
259
|
+
}
|
|
260
|
+
|
|
261
|
+
private _updateModel(width: number, height: number) {
|
|
262
|
+
const maxTileSize = this._maxTileSize;
|
|
263
|
+
let tileWidth = maxTileSize;
|
|
264
|
+
let tileHeight = maxTileSize;
|
|
265
|
+
let tileOverlapX = defaultTileOverlap;
|
|
266
|
+
let tileOverlapY = defaultTileOverlap;
|
|
267
|
+
|
|
268
|
+
if (width < maxTileSize + defaultTileOverlap * 2) {
|
|
269
|
+
tileWidth = roundUp(width, tileAlignment);
|
|
270
|
+
tileOverlapX = 0;
|
|
271
|
+
}
|
|
272
|
+
if (height < maxTileSize + defaultTileOverlap * 2) {
|
|
273
|
+
tileHeight = roundUp(height, tileAlignment);
|
|
274
|
+
tileOverlapY = 0;
|
|
275
|
+
}
|
|
276
|
+
|
|
277
|
+
if (
|
|
278
|
+
tileWidth !== this._tileWidth ||
|
|
279
|
+
tileHeight !== this._tileHeight ||
|
|
280
|
+
tileOverlapX !== this._tileOverlapX ||
|
|
281
|
+
tileOverlapY !== this._tileOverlapY ||
|
|
282
|
+
!this._tfModel
|
|
283
|
+
) {
|
|
284
|
+
this._tileWidth = tileWidth;
|
|
285
|
+
this._tileHeight = tileHeight;
|
|
286
|
+
this._tileOverlapX = tileOverlapX;
|
|
287
|
+
this._tileOverlapY = tileOverlapY;
|
|
288
|
+
|
|
289
|
+
if (this._tfModel) {
|
|
290
|
+
this._tfModel.dispose();
|
|
291
|
+
}
|
|
292
|
+
|
|
293
|
+
this.buildModel();
|
|
294
|
+
}
|
|
295
|
+
}
|
|
296
|
+
|
|
297
|
+
private _getTileSizeWithOverlap() {
|
|
298
|
+
return {
|
|
299
|
+
width: this._tileWidth + 2 * this._tileOverlapX,
|
|
300
|
+
height: this._tileHeight + 2 * this._tileOverlapY
|
|
301
|
+
};
|
|
302
|
+
}
|
|
303
|
+
|
|
304
|
+
private _processImageData(
|
|
305
|
+
color: ImageData | HDRImageData,
|
|
306
|
+
albedo: ImageData | undefined,
|
|
307
|
+
normal: ImageData | undefined,
|
|
308
|
+
isHDR: boolean
|
|
309
|
+
) {
|
|
310
|
+
const rawData = color.data;
|
|
311
|
+
const pixelsCount = rawData.length / 4;
|
|
312
|
+
const channels = this._aux ? 9 : 3;
|
|
313
|
+
const tensorData = new Float32Array(pixelsCount * channels);
|
|
314
|
+
|
|
315
|
+
if ((albedo && !normal) || (normal && !albedo)) {
|
|
316
|
+
throw new Error('Normal map and albedo map are both required');
|
|
317
|
+
}
|
|
318
|
+
if (albedo && normal) {
|
|
319
|
+
if (
|
|
320
|
+
albedo.width !== normal.width ||
|
|
321
|
+
albedo.height !== normal.height ||
|
|
322
|
+
color.width !== albedo.width ||
|
|
323
|
+
color.height !== albedo.height
|
|
324
|
+
) {
|
|
325
|
+
throw new Error('Image size mismatch');
|
|
326
|
+
}
|
|
327
|
+
}
|
|
328
|
+
|
|
329
|
+
const albedoData = albedo?.data;
|
|
330
|
+
const normalData = normal?.data;
|
|
331
|
+
for (let i = 0; i < rawData.length; i += 4) {
|
|
332
|
+
const i2 = (i / 4) * channels;
|
|
333
|
+
|
|
334
|
+
for (let c = 0; c < 3; c++) {
|
|
335
|
+
if (isHDR) {
|
|
336
|
+
tensorData[i2 + c] = rawData[i + c];
|
|
337
|
+
} else {
|
|
338
|
+
tensorData[i2 + c] = rawData[i + c] / 255;
|
|
339
|
+
}
|
|
340
|
+
if (albedoData) {
|
|
341
|
+
tensorData[i2 + c + 3] = albedoData[i + c] / 255;
|
|
342
|
+
}
|
|
343
|
+
if (normalData) {
|
|
344
|
+
tensorData[i2 + c + 6] = normalData[i + c] / 255;
|
|
345
|
+
}
|
|
346
|
+
}
|
|
347
|
+
}
|
|
348
|
+
|
|
349
|
+
return tensorData;
|
|
350
|
+
}
|
|
351
|
+
|
|
352
|
+
private _readTile(
|
|
353
|
+
data: Float32Array,
|
|
354
|
+
channels: number,
|
|
355
|
+
srcTile: Tile,
|
|
356
|
+
width: number
|
|
357
|
+
) {
|
|
358
|
+
const tileData = new Float32Array(
|
|
359
|
+
srcTile.width * srcTile.height * channels
|
|
360
|
+
);
|
|
361
|
+
for (let y = 0; y < srcTile.height; y++) {
|
|
362
|
+
for (let x = 0; x < srcTile.width; x++) {
|
|
363
|
+
const i2 = ((y + srcTile.y) * width + (x + srcTile.x)) * channels;
|
|
364
|
+
const i1 = (y * srcTile.width + x) * channels;
|
|
365
|
+
|
|
366
|
+
for (let c = 0; c < channels; c++) {
|
|
367
|
+
tileData[i1 + c] = data[i2 + c];
|
|
368
|
+
}
|
|
369
|
+
}
|
|
370
|
+
}
|
|
371
|
+
return tileData;
|
|
372
|
+
}
|
|
373
|
+
|
|
374
|
+
private _writeTile(
|
|
375
|
+
imageData: ImageData | HDRImageData,
|
|
376
|
+
srcTile: Tile,
|
|
377
|
+
dstTile: Tile,
|
|
378
|
+
srcTileData: Float32Array,
|
|
379
|
+
srcWidth: number,
|
|
380
|
+
isHDR: boolean
|
|
381
|
+
) {
|
|
382
|
+
const { data: outImageData, width } = imageData;
|
|
383
|
+
const dx = dstTile.x - srcTile.x;
|
|
384
|
+
const dy = dstTile.y - srcTile.y;
|
|
385
|
+
for (let y = 0; y < dstTile.height; y++) {
|
|
386
|
+
for (let x = 0; x < dstTile.width; x++) {
|
|
387
|
+
const i1 = ((y + dy) * srcWidth + x + dx) * 3;
|
|
388
|
+
const i2 = ((y + dstTile.y) * width + (x + dstTile.x)) * 4;
|
|
389
|
+
|
|
390
|
+
for (let c = 0; c < 3; c++) {
|
|
391
|
+
if (isHDR) {
|
|
392
|
+
outImageData[i2 + c] = srcTileData[i1 + c];
|
|
393
|
+
} else {
|
|
394
|
+
outImageData[i2 + c] = Math.min(
|
|
395
|
+
Math.max(srcTileData[i1 + c] * 255, 0),
|
|
396
|
+
255
|
|
397
|
+
);
|
|
398
|
+
}
|
|
399
|
+
}
|
|
400
|
+
imageData.data[i2 + 3] = isHDR ? 1 : 255;
|
|
401
|
+
}
|
|
402
|
+
}
|
|
403
|
+
}
|
|
404
|
+
|
|
405
|
+
private _executeTile(
|
|
406
|
+
inputData:
|
|
407
|
+
| Float32Array
|
|
408
|
+
| {
|
|
409
|
+
color: GPUBuffer;
|
|
410
|
+
// TODO optional
|
|
411
|
+
albedo: GPUBuffer;
|
|
412
|
+
normal: GPUBuffer;
|
|
413
|
+
},
|
|
414
|
+
outputTileData: ImageData | HDRImageData | undefined,
|
|
415
|
+
outputImageData: ImageData | HDRImageData | undefined,
|
|
416
|
+
i: number,
|
|
417
|
+
j: number,
|
|
418
|
+
width: number,
|
|
419
|
+
height: number,
|
|
420
|
+
isHDR: boolean
|
|
421
|
+
) {
|
|
422
|
+
const channels = this._aux ? 9 : 3;
|
|
423
|
+
const tileOverlapX = this._tileOverlapX;
|
|
424
|
+
const tileOverlapY = this._tileOverlapY;
|
|
425
|
+
let srcTileSize = this._getTileSizeWithOverlap();
|
|
426
|
+
let dstTileSize = { width: this._tileWidth, height: this._tileHeight };
|
|
427
|
+
|
|
428
|
+
let srcX0 = i > 0 ? i * dstTileSize.width - tileOverlapX : 0;
|
|
429
|
+
let srcX1 = Math.min(srcX0 + srcTileSize.width, width);
|
|
430
|
+
srcX0 = Math.max(srcX1 - srcTileSize.width, 0);
|
|
431
|
+
|
|
432
|
+
let srcY0 = j > 0 ? j * dstTileSize.height - tileOverlapY : 0;
|
|
433
|
+
let srcY1 = Math.min(srcY0 + srcTileSize.height, height);
|
|
434
|
+
srcY0 = Math.max(srcY1 - srcTileSize.height, 0);
|
|
435
|
+
|
|
436
|
+
const srcTileWidth = Math.min(srcTileSize.width, width);
|
|
437
|
+
const srcTileHeight = Math.min(srcTileSize.height, height);
|
|
438
|
+
const needsResize =
|
|
439
|
+
width < dstTileSize.width || height < dstTileSize.height;
|
|
440
|
+
const srcTile = new Tile(srcX0, srcY0, srcTileWidth, srcTileHeight);
|
|
441
|
+
|
|
442
|
+
let tileTensor!: Tensor;
|
|
443
|
+
let inputScale = 1;
|
|
444
|
+
const device = this._device!;
|
|
445
|
+
let dataProcessGPU = this._dataProcessGPU;
|
|
446
|
+
|
|
447
|
+
if (inputData instanceof Float32Array) {
|
|
448
|
+
let tileData = this._readTile(inputData, channels, srcTile, width);
|
|
449
|
+
if (isHDR) {
|
|
450
|
+
inputScale = avgLogLum({
|
|
451
|
+
data: tileData,
|
|
452
|
+
channels: 9
|
|
453
|
+
});
|
|
454
|
+
tileData = hdrTransferFuncCPU({
|
|
455
|
+
data: tileData,
|
|
456
|
+
channels: 9,
|
|
457
|
+
inputScale
|
|
458
|
+
});
|
|
459
|
+
}
|
|
460
|
+
tileTensor = tensor(
|
|
461
|
+
tileData,
|
|
462
|
+
[1, srcTileHeight, srcTileWidth, channels],
|
|
463
|
+
'float32'
|
|
464
|
+
) as Tensor4D;
|
|
465
|
+
} else {
|
|
466
|
+
if (!isHDR) {
|
|
467
|
+
throw new Error('Only hdr is supported for webgpu data.');
|
|
468
|
+
}
|
|
469
|
+
if (!dataProcessGPU) {
|
|
470
|
+
dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(device);
|
|
471
|
+
}
|
|
472
|
+
dataProcessGPU.setImageSize(width, height);
|
|
473
|
+
dataProcessGPU.setInputTile(srcTile);
|
|
474
|
+
// Display the noisy input instead of prev denoised result
|
|
475
|
+
if (i === 0 && j === 0) {
|
|
476
|
+
dataProcessGPU.copyInputDataToOutput(inputData.color);
|
|
477
|
+
}
|
|
478
|
+
const { color, albedo, normal } = dataProcessGPU.forward(
|
|
479
|
+
inputData.color,
|
|
480
|
+
inputData.albedo,
|
|
481
|
+
inputData.normal
|
|
482
|
+
);
|
|
483
|
+
const shape = [1, srcTileHeight, srcTileWidth, 4] as any;
|
|
484
|
+
|
|
485
|
+
tileTensor = concat4d(
|
|
486
|
+
[color, albedo, normal].map((buffer) => {
|
|
487
|
+
const tmp = tensor({ buffer, zeroCopy: true }, shape) as Tensor4D;
|
|
488
|
+
const ret = slice4d(
|
|
489
|
+
tmp,
|
|
490
|
+
[0, 0, 0, 0],
|
|
491
|
+
[1, srcTileHeight, srcTileWidth, 3]
|
|
492
|
+
);
|
|
493
|
+
tmp.dispose();
|
|
494
|
+
return ret;
|
|
495
|
+
}),
|
|
496
|
+
3
|
|
497
|
+
);
|
|
498
|
+
}
|
|
499
|
+
// We need resize if input size is smaller than tile size. And is rounded up.
|
|
500
|
+
if (needsResize) {
|
|
501
|
+
const rawTileTensor = tileTensor;
|
|
502
|
+
tileTensor = mirrorPad(
|
|
503
|
+
rawTileTensor,
|
|
504
|
+
[
|
|
505
|
+
[0, 0],
|
|
506
|
+
[0, srcTileSize.height - height],
|
|
507
|
+
[0, srcTileSize.width - width],
|
|
508
|
+
[0, 0]
|
|
509
|
+
],
|
|
510
|
+
'reflect'
|
|
511
|
+
);
|
|
512
|
+
rawTileTensor.dispose();
|
|
513
|
+
}
|
|
514
|
+
|
|
515
|
+
const outputTensor = this._tfModel!.predict(tileTensor) as Tensor;
|
|
516
|
+
tileTensor.dispose();
|
|
517
|
+
|
|
518
|
+
const dstWidth = Math.min(dstTileSize.width, width);
|
|
519
|
+
const dstHeight = Math.min(dstTileSize.height, height);
|
|
520
|
+
const dstTile = new Tile(i * dstWidth, j * dstHeight, dstWidth, dstHeight);
|
|
521
|
+
dstTile.width = Math.min(dstTile.width, width - dstTile.x);
|
|
522
|
+
dstTile.height = Math.min(dstTile.height, height - dstTile.y);
|
|
523
|
+
|
|
524
|
+
if (inputData instanceof Float32Array) {
|
|
525
|
+
let denoisedData = outputTensor.dataSync();
|
|
526
|
+
if (isHDR) {
|
|
527
|
+
denoisedData = hdrTransferFuncInverseCPU({
|
|
528
|
+
data: denoisedData as Float32Array,
|
|
529
|
+
channels: 3,
|
|
530
|
+
inputScale
|
|
531
|
+
});
|
|
532
|
+
}
|
|
533
|
+
|
|
534
|
+
this._writeTile(
|
|
535
|
+
outputImageData!,
|
|
536
|
+
srcTile,
|
|
537
|
+
dstTile,
|
|
538
|
+
denoisedData as Float32Array,
|
|
539
|
+
srcTileSize.width,
|
|
540
|
+
isHDR
|
|
541
|
+
);
|
|
542
|
+
|
|
543
|
+
for (let y = 0; y < dstHeight; y++) {
|
|
544
|
+
for (let x = 0; x < dstWidth; x++) {
|
|
545
|
+
const i1 = (y * dstWidth + x) * 4;
|
|
546
|
+
const i2 = ((y + dstTile.y) * width + (x + dstTile.x)) * 4;
|
|
547
|
+
for (let c = 0; c < 4; c++) {
|
|
548
|
+
outputTileData!.data[i1 + c] = outputImageData!.data[i2 + c];
|
|
549
|
+
}
|
|
550
|
+
}
|
|
551
|
+
}
|
|
552
|
+
|
|
553
|
+
outputTensor.dispose();
|
|
554
|
+
} else {
|
|
555
|
+
dataProcessGPU!.setOutputTile(dstTile, srcTile);
|
|
556
|
+
// IMPORTANT
|
|
557
|
+
// storage buffer has alignment. that 3 channels still needs 16 bytes data.
|
|
558
|
+
// So we need to pad it to 4 channels.
|
|
559
|
+
const outputTensor4Channnels = pad4d(outputTensor as Tensor4D, [
|
|
560
|
+
[0, 0],
|
|
561
|
+
[0, 0],
|
|
562
|
+
[0, 0],
|
|
563
|
+
[0, 1]
|
|
564
|
+
]);
|
|
565
|
+
const outBuffer = dataProcessGPU!.inverse(
|
|
566
|
+
outputTensor4Channnels.dataToGPU().buffer!,
|
|
567
|
+
inputData.color
|
|
568
|
+
);
|
|
569
|
+
outputTensor.dispose();
|
|
570
|
+
outputTensor4Channnels.dispose();
|
|
571
|
+
return outBuffer;
|
|
572
|
+
}
|
|
573
|
+
}
|
|
574
|
+
|
|
575
|
+
progressiveExecute<T extends ImageData | HDRImageData | GPUImageData>({
|
|
576
|
+
color,
|
|
577
|
+
albedo,
|
|
578
|
+
normal,
|
|
579
|
+
done,
|
|
580
|
+
progress
|
|
581
|
+
}: {
|
|
582
|
+
color: T;
|
|
583
|
+
albedo?: ImageData | GPUImageData;
|
|
584
|
+
normal?: ImageData | GPUImageData;
|
|
585
|
+
done: (outputData: T) => void;
|
|
586
|
+
progress?: (
|
|
587
|
+
outputData: T,
|
|
588
|
+
tileData: T | undefined,
|
|
589
|
+
tile: Tile,
|
|
590
|
+
currentIdx: number,
|
|
591
|
+
totalIdx: number
|
|
592
|
+
) => void;
|
|
593
|
+
}): () => void {
|
|
594
|
+
if (this._aux && (!albedo || !normal)) {
|
|
595
|
+
throw new Error('Normal map and albedo map are both required');
|
|
596
|
+
}
|
|
597
|
+
|
|
598
|
+
if (!this._aux) {
|
|
599
|
+
if (albedo || normal) {
|
|
600
|
+
throw new Error('Normal map and albedo map are not required');
|
|
601
|
+
}
|
|
602
|
+
}
|
|
603
|
+
|
|
604
|
+
const width = color.width;
|
|
605
|
+
const height = color.height;
|
|
606
|
+
this._updateModel(width, height);
|
|
607
|
+
|
|
608
|
+
// TODO should fixed to be hdr when UNet is created.
|
|
609
|
+
// weights of hdr and ldr is different
|
|
610
|
+
|
|
611
|
+
const hdr = this._hdr || false;
|
|
612
|
+
let rawData: Float32Array;
|
|
613
|
+
if (!isGPUImageData(color)) {
|
|
614
|
+
rawData = this._processImageData(
|
|
615
|
+
color,
|
|
616
|
+
albedo as ImageData,
|
|
617
|
+
normal as ImageData,
|
|
618
|
+
hdr
|
|
619
|
+
);
|
|
620
|
+
}
|
|
621
|
+
const tileWidth = this._tileWidth;
|
|
622
|
+
const tileHeight = this._tileHeight;
|
|
623
|
+
const tileCountH = Math.ceil(height / tileHeight);
|
|
624
|
+
const tileCountW = Math.ceil(width / tileWidth);
|
|
625
|
+
|
|
626
|
+
function makeImageData(width: number, height: number) {
|
|
627
|
+
return hdr
|
|
628
|
+
? {
|
|
629
|
+
data: new Float32Array(width * height * 4),
|
|
630
|
+
width,
|
|
631
|
+
height
|
|
632
|
+
}
|
|
633
|
+
: new ImageData(width, height);
|
|
634
|
+
}
|
|
635
|
+
|
|
636
|
+
const outputImageData = isGPUImageData(color)
|
|
637
|
+
? undefined
|
|
638
|
+
: makeImageData(width, height);
|
|
639
|
+
const outputTileData = isGPUImageData(color)
|
|
640
|
+
? undefined
|
|
641
|
+
: makeImageData(Math.min(tileWidth, width), Math.min(tileHeight, height));
|
|
642
|
+
|
|
643
|
+
let aborted = false;
|
|
644
|
+
|
|
645
|
+
const executeTile = (i: number, j: number) => {
|
|
646
|
+
if (aborted) {
|
|
647
|
+
return;
|
|
648
|
+
}
|
|
649
|
+
let resGPUBuffer;
|
|
650
|
+
// profileAndLogKernelCode(() => {
|
|
651
|
+
resGPUBuffer = this._executeTile(
|
|
652
|
+
isGPUImageData(color)
|
|
653
|
+
? {
|
|
654
|
+
color: color.data,
|
|
655
|
+
albedo: (albedo as GPUImageData).data,
|
|
656
|
+
normal: (normal as GPUImageData).data
|
|
657
|
+
}
|
|
658
|
+
: rawData,
|
|
659
|
+
outputTileData,
|
|
660
|
+
outputImageData,
|
|
661
|
+
i,
|
|
662
|
+
j,
|
|
663
|
+
width,
|
|
664
|
+
height,
|
|
665
|
+
hdr
|
|
666
|
+
);
|
|
667
|
+
// }, true);
|
|
668
|
+
const output = outputImageData || {
|
|
669
|
+
data: resGPUBuffer,
|
|
670
|
+
width,
|
|
671
|
+
height
|
|
672
|
+
};
|
|
673
|
+
progress?.(
|
|
674
|
+
output as T,
|
|
675
|
+
// Is undefined if using webgpu buffer
|
|
676
|
+
outputTileData as T | undefined,
|
|
677
|
+
new Tile(i * tileWidth, j * tileHeight, tileWidth, tileHeight),
|
|
678
|
+
i + j * tileCountW,
|
|
679
|
+
tileCountW * tileCountH
|
|
680
|
+
);
|
|
681
|
+
|
|
682
|
+
if (i + 1 < tileCountW || j + 1 < tileCountH) {
|
|
683
|
+
requestAnimationFrame(() => {
|
|
684
|
+
if (i + 1 < tileCountW) {
|
|
685
|
+
executeTile(i + 1, j);
|
|
686
|
+
} else if (j + 1 < tileCountH) {
|
|
687
|
+
executeTile(0, j + 1);
|
|
688
|
+
}
|
|
689
|
+
});
|
|
690
|
+
} else {
|
|
691
|
+
done(output as T);
|
|
692
|
+
}
|
|
693
|
+
};
|
|
694
|
+
|
|
695
|
+
executeTile(0, 0);
|
|
696
|
+
|
|
697
|
+
return () => {
|
|
698
|
+
aborted = true;
|
|
699
|
+
};
|
|
700
|
+
}
|
|
701
|
+
|
|
702
|
+
dispose() {
|
|
703
|
+
this._tfModel?.dispose();
|
|
704
|
+
this._dataProcessGPU?.dispose();
|
|
705
|
+
}
|
|
706
|
+
}
|
|
707
|
+
|
|
708
|
+
export default UNet;
|