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.
Files changed (52) hide show
  1. package/LICENSE +21 -0
  2. package/README.md +2 -0
  3. package/dist/oidn.mjs +22909 -0
  4. package/dist/oidn.umd.js +5905 -0
  5. package/lib/UNet.d.ts +54 -0
  6. package/lib/UNet.js +467 -0
  7. package/lib/UNet.js.map +1 -0
  8. package/lib/WGPUComputePass.d.ts +53 -0
  9. package/lib/WGPUComputePass.js +220 -0
  10. package/lib/WGPUComputePass.js.map +1 -0
  11. package/lib/WGPUFullQuadPass.d.ts +51 -0
  12. package/lib/WGPUFullQuadPass.js +261 -0
  13. package/lib/WGPUFullQuadPass.js.map +1 -0
  14. package/lib/backend.d.ts +5 -0
  15. package/lib/backend.js +41 -0
  16. package/lib/backend.js.map +1 -0
  17. package/lib/hdr.d.ts +31 -0
  18. package/lib/hdr.js +340 -0
  19. package/lib/hdr.js.map +1 -0
  20. package/lib/helper.d.ts +1 -0
  21. package/lib/helper.js +27 -0
  22. package/lib/helper.js.map +1 -0
  23. package/lib/kernels.d.ts +1 -0
  24. package/lib/kernels.js +26 -0
  25. package/lib/kernels.js.map +1 -0
  26. package/lib/main.d.ts +20 -0
  27. package/lib/main.js +20 -0
  28. package/lib/main.js.map +1 -0
  29. package/lib/process.d.ts +39 -0
  30. package/lib/process.js +309 -0
  31. package/lib/process.js.map +1 -0
  32. package/lib/tza.d.ts +15 -0
  33. package/lib/tza.js +114 -0
  34. package/lib/tza.js.map +1 -0
  35. package/package.json +26 -0
  36. package/src/UNet.ts +708 -0
  37. package/src/WGPUComputePass.ts +318 -0
  38. package/src/WGPUFullQuadPass.ts +348 -0
  39. package/src/backend.ts +53 -0
  40. package/src/hdr.ts +398 -0
  41. package/src/helper.ts +35 -0
  42. package/src/kernels.ts +31 -0
  43. package/src/main.ts +42 -0
  44. package/src/process.ts +362 -0
  45. package/src/tza.ts +136 -0
  46. package/weights/.gitattributes +1 -0
  47. package/weights/LICENSE.txt +202 -0
  48. package/weights/README.md +7 -0
  49. package/weights/rt_hdr.tza +0 -0
  50. package/weights/rt_hdr_alb_nrm.tza +0 -0
  51. package/weights/rt_ldr.tza +0 -0
  52. 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;