oidn-web 0.2.2 → 0.3.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/src/UNet.ts CHANGED
@@ -7,6 +7,7 @@ import { mirrorPad } from '@tensorflow/tfjs-core/dist/ops/mirror_pad';
7
7
  import { pad4d } from '@tensorflow/tfjs-core/dist/ops/pad4d';
8
8
  import { slice4d } from '@tensorflow/tfjs-core/dist/ops/slice4d';
9
9
  import { concat4d } from '@tensorflow/tfjs-core/dist/ops/concat_4d';
10
+ import { ENGINE } from '@tensorflow/tfjs-core/dist/engine';
10
11
  import {
11
12
  Conv2D,
12
13
  UpSampling2D
@@ -25,7 +26,8 @@ import {
25
26
  hdrTransferFuncInverseCPU
26
27
  } from './process';
27
28
  import type { WebGPUBackend } from '@tensorflow/tfjs-backend-webgpu';
28
- // import { profileAndLogKernelCode } from './helper';
29
+
30
+ // import { profileAndLogKernelCode, memory } from './helper';
29
31
 
30
32
  function getTensorData(
31
33
  ubytes: Uint8Array,
@@ -68,6 +70,12 @@ interface HDRImageData {
68
70
  }
69
71
 
70
72
  interface GPUImageData {
73
+ data: GPUBuffer | GPUTexture;
74
+ width: number;
75
+ height: number;
76
+ }
77
+
78
+ interface GPUImageDataOutput {
71
79
  data: GPUBuffer;
72
80
  width: number;
73
81
  height: number;
@@ -84,7 +92,7 @@ function roundUp2(a: number, b: number, c: number) {
84
92
  function isGPUImageData(
85
93
  data: ImageData | GPUImageData | HDRImageData
86
94
  ): data is GPUImageData {
87
- return data.data instanceof GPUBuffer;
95
+ return data.data instanceof GPUBuffer || data.data instanceof GPUTexture;
88
96
  }
89
97
 
90
98
  const receptiveField = 174; // receptive field in pixels
@@ -109,18 +117,27 @@ class UNet {
109
117
  private _tileOverlapX = 0;
110
118
  private _tileOverlapY = 0;
111
119
 
112
- private _aux = false;
113
- private _hdr = false;
120
+ private _aux;
121
+ private _hdr;
114
122
 
115
123
  private _dataProcessGPU?: GPUDataProcess;
116
124
 
117
125
  private _maxTileSize;
118
126
 
127
+ private _tensors = new Map<string, Tensor>();
128
+ private _modelsCache = new Map<string, LayersModel>();
129
+
119
130
  constructor(
120
- private _tensors: Map<string, HostTensor>,
131
+ private _hostTensors: Map<string, HostTensor>,
121
132
  private _backend: WebGPUBackend,
122
133
  opts: {
134
+ /**
135
+ * If use auxiliary data.
136
+ */
123
137
  aux?: boolean;
138
+ /**
139
+ * If input is HDR image.
140
+ */
124
141
  hdr?: boolean;
125
142
  maxTileSize?: number;
126
143
  } = {}
@@ -128,7 +145,7 @@ class UNet {
128
145
  this._aux = opts.aux || false;
129
146
  this._hdr = opts.hdr || false;
130
147
 
131
- this._maxTileSize = opts.maxTileSize ?? 512;
148
+ this._maxTileSize = roundUp(opts.maxTileSize ?? 512, 2);
132
149
 
133
150
  this._device = this._backend.device;
134
151
  }
@@ -141,19 +158,30 @@ class UNet {
141
158
  const aux = this._aux;
142
159
  const channels = 3 + (aux ? 6 : 0);
143
160
  const tileSize = this._getTileSizeWithOverlap();
161
+ const cache = this._modelsCache;
162
+ const key = [tileSize.width, tileSize.height].join(',');
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
+ }
144
173
 
145
- // TODO input process transferFunc
146
- // TODO input shape
147
174
  const input = TFInput({
175
+ name: 'input',
148
176
  shape: [tileSize.height, tileSize.width, channels],
149
177
  dtype: 'float32'
150
178
  });
151
179
 
152
180
  this._tfModel = new LayersModel({
153
181
  inputs: [input],
154
- // TODO output process transferFunc
155
182
  outputs: isLarge ? this._addNetLarge(input) : this._addNet(input)
156
183
  });
184
+ cache.set(key, this._tfModel);
157
185
  }
158
186
 
159
187
  private _createConv(
@@ -161,21 +189,34 @@ class UNet {
161
189
  source: SymbolicTensor,
162
190
  activation?: 'relu'
163
191
  ) {
164
- const unetWeightTensor = this._tensors.get(name + '.weight')!;
165
- const unetBiasTensor = this._tensors.get(name + '.bias')!;
166
- const weightDims = unetWeightTensor.desc.dims;
167
- const weightTensor = tensor(
168
- changeWeightShapes(
169
- getTensorData(unetWeightTensor.data, unetWeightTensor.desc.dataType),
170
- weightDims
171
- ),
172
- [weightDims[2], weightDims[3], weightDims[1], weightDims[0]],
173
- 'float32'
174
- );
175
- const biasTensor = tensor1d(
176
- getTensorData(unetBiasTensor.data, unetBiasTensor.desc.dataType),
177
- 'float32'
178
- );
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'
208
+ );
209
+
210
+ tensors.set(weightTensorName, weightTensor);
211
+ }
212
+ if (!biasTensor) {
213
+ const unetBiasTensor = this._hostTensors.get(name + '.bias')!;
214
+ biasTensor = tensor1d(
215
+ getTensorData(unetBiasTensor.data, unetBiasTensor.desc.dataType),
216
+ 'float32'
217
+ );
218
+ tensors.set(biasTensorName, biasTensor);
219
+ }
179
220
  // TODO whats the purpose of padded dims ?
180
221
  const convLayer = new Conv2D({
181
222
  name,
@@ -196,36 +237,39 @@ class UNet {
196
237
  source1: SymbolicTensor,
197
238
  source2: SymbolicTensor
198
239
  ) {
240
+ const concatLayer = new Concatenate({
241
+ name: name + '/concat',
242
+ trainable: false,
243
+ axis: 3
244
+ });
199
245
  //https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L40
200
246
  return this._createConv(
201
247
  name,
202
248
  // Concat on the channel
203
- new Concatenate({ trainable: false, axis: 3 }).apply([
204
- // convLayer.apply(source2) as SymbolicTensor,
205
- source1,
206
- source2
207
- ]) as SymbolicTensor,
249
+ concatLayer.apply([source1, source2]) as SymbolicTensor,
208
250
  'relu'
209
251
  ) as SymbolicTensor;
210
252
  }
211
253
 
212
254
  private _createPooling(source: SymbolicTensor) {
213
- // https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L33
214
- return new MaxPooling2D({
255
+ const poolingLayer = new MaxPooling2D({
215
256
  name: source.name + '/pooling',
216
257
  poolSize: [2, 2],
217
258
  strides: [2, 2],
218
259
  padding: 'same',
219
260
  trainable: false
220
- }).apply(source) as SymbolicTensor;
261
+ });
262
+ // https://github.com/RenderKit/oidn/blob/713ec7838ba650f99e0a896549c0dca5eeb3652d/training/model.py#L33
263
+ return poolingLayer.apply(source) as SymbolicTensor;
221
264
  }
222
265
 
223
266
  private _addUpsamplingLayer(source: SymbolicTensor) {
224
- return new UpSampling2D({
267
+ const upsamplingLayer = new UpSampling2D({
225
268
  name: source.name + '/upsampling',
226
269
  size: [2, 2],
227
270
  trainable: false
228
- }).apply(source) as SymbolicTensor;
271
+ });
272
+ return upsamplingLayer.apply(source) as SymbolicTensor;
229
273
  }
230
274
 
231
275
  private _addNet(input: SymbolicTensor) {
@@ -308,14 +352,26 @@ class UNet {
308
352
  let tileOverlapY = isLarge ? defaultTileOverlapLarge : defaultTileOverlap;
309
353
 
310
354
  if (width < maxTileSize + defaultTileOverlap * 2) {
311
- tileWidth = roundUp(width, tileAlignment);
312
- tileOverlapX = 0;
355
+ tileWidth = roundUp(width, maxTileSize / 2);
356
+ if (width <= maxTileSize) {
357
+ tileOverlapX = 0;
358
+ }
313
359
  }
314
360
  if (height < maxTileSize + defaultTileOverlap * 2) {
315
- tileHeight = roundUp(height, tileAlignment);
316
- tileOverlapY = 0;
361
+ tileHeight = roundUp(height, maxTileSize / 2);
362
+ if (height <= maxTileSize) {
363
+ tileOverlapY = 0;
364
+ }
317
365
  }
318
366
 
367
+ // Force width and height has same size. reduce the cache in memory
368
+ const tileSize = Math.max(tileWidth, tileHeight);
369
+ const tileOverlap = Math.max(tileOverlapX, tileOverlapY);
370
+ tileWidth = tileSize;
371
+ tileHeight = tileSize;
372
+ tileOverlapX = tileOverlap;
373
+ tileOverlapY = tileOverlap;
374
+
319
375
  if (
320
376
  tileWidth !== this._tileWidth ||
321
377
  tileHeight !== this._tileHeight ||
@@ -323,15 +379,12 @@ class UNet {
323
379
  tileOverlapY !== this._tileOverlapY ||
324
380
  !this._tfModel
325
381
  ) {
382
+ // console.log(tileWidth, tileHeight, tileOverlapX, tileOverlapY);
326
383
  this._tileWidth = tileWidth;
327
384
  this._tileHeight = tileHeight;
328
385
  this._tileOverlapX = tileOverlapX;
329
386
  this._tileOverlapY = tileOverlapY;
330
387
 
331
- if (this._tfModel) {
332
- this._tfModel.dispose();
333
- }
334
-
335
388
  this._buildModel(isLarge);
336
389
  }
337
390
  }
@@ -448,10 +501,10 @@ class UNet {
448
501
  inputData:
449
502
  | Float32Array
450
503
  | {
451
- color: GPUBuffer;
504
+ color: GPUBuffer | GPUTexture;
452
505
  // TODO optional
453
- albedo: GPUBuffer;
454
- normal: GPUBuffer;
506
+ albedo?: GPUBuffer | GPUTexture;
507
+ normal?: GPUBuffer | GPUTexture;
455
508
  },
456
509
  outputTileData: ImageData | HDRImageData | undefined,
457
510
  outputImageData: ImageData | HDRImageData | undefined,
@@ -459,7 +512,8 @@ class UNet {
459
512
  j: number,
460
513
  width: number,
461
514
  height: number,
462
- isHDR: boolean
515
+ isHDR: boolean,
516
+ denoiseAlpha?: boolean
463
517
  ) {
464
518
  const channels = this._aux ? 9 : 3;
465
519
  const tileOverlapX = this._tileOverlapX;
@@ -475,10 +529,9 @@ class UNet {
475
529
  let srcY1 = Math.min(srcY0 + srcTileSize.height, height);
476
530
  srcY0 = Math.max(srcY1 - srcTileSize.height, 0);
477
531
 
478
- const srcTileWidth = Math.min(srcTileSize.width, width);
479
- const srcTileHeight = Math.min(srcTileSize.height, height);
480
- const needsResize =
481
- width < dstTileSize.width || height < dstTileSize.height;
532
+ const srcTileWidth = srcTileSize.width;
533
+ const srcTileHeight = srcTileSize.height;
534
+
482
535
  const srcTile = new Tile(srcX0, srcY0, srcTileWidth, srcTileHeight);
483
536
 
484
537
  let tileTensor!: Tensor;
@@ -491,11 +544,11 @@ class UNet {
491
544
  if (isHDR) {
492
545
  inputScale = avgLogLum({
493
546
  data: tileData,
494
- channels: 9
547
+ channels
495
548
  });
496
549
  tileData = hdrTransferFuncCPU({
497
550
  data: tileData,
498
- channels: 9,
551
+ channels,
499
552
  inputScale
500
553
  });
501
554
  }
@@ -505,11 +558,11 @@ class UNet {
505
558
  'float32'
506
559
  ) as Tensor4D;
507
560
  } else {
508
- if (!isHDR) {
509
- throw new Error('Only hdr is supported for webgpu data.');
510
- }
511
561
  if (!dataProcessGPU) {
512
- dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(device);
562
+ dataProcessGPU = this._dataProcessGPU = new GPUDataProcess(
563
+ device,
564
+ isHDR
565
+ );
513
566
  }
514
567
  dataProcessGPU.setImageSize(width, height);
515
568
  dataProcessGPU.setInputTile(srcTile);
@@ -519,43 +572,38 @@ class UNet {
519
572
  }
520
573
  const { color, albedo, normal } = dataProcessGPU.forward(
521
574
  inputData.color,
522
- inputData.albedo,
523
- inputData.normal
575
+ this._aux ? inputData.albedo : undefined,
576
+ this._aux ? inputData.normal : undefined,
577
+ denoiseAlpha
524
578
  );
525
- const shape = [1, srcTileHeight, srcTileWidth, 4] as any;
526
- const auxTensors = [color, albedo, normal].map((buffer) => {
527
- const tmp = tensor({ buffer, zeroCopy: true }, shape) as Tensor4D;
579
+
580
+ const createTensor = (buffer: GPUBuffer) => {
581
+ const tmp = tensor({ buffer, zeroCopy: true }, [
582
+ 1,
583
+ srcTileHeight,
584
+ srcTileWidth,
585
+ 4
586
+ ]) as Tensor4D;
528
587
  const ret = slice4d(
529
588
  tmp,
530
589
  [0, 0, 0, 0],
531
590
  [1, srcTileHeight, srcTileWidth, 3]
532
591
  );
533
- tmp.dispose();
534
592
  return ret;
535
- });
593
+ };
536
594
 
537
- tileTensor = concat4d(auxTensors, 3);
538
- // TODO reuse tensors?
539
- auxTensors.forEach((t) => t.dispose());
540
- }
541
- // We need resize if input size is smaller than tile size. And is rounded up.
542
- if (needsResize) {
543
- const rawTileTensor = tileTensor;
544
- tileTensor = mirrorPad(
545
- rawTileTensor,
546
- [
547
- [0, 0],
548
- [0, srcTileSize.height - height],
549
- [0, srcTileSize.width - width],
550
- [0, 0]
551
- ],
552
- 'reflect'
553
- );
554
- rawTileTensor.dispose();
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
+ }
555
603
  }
556
604
 
605
+ let outBuffer: GPUBuffer;
557
606
  const outputTensor = this._tfModel!.predict(tileTensor) as Tensor;
558
- tileTensor.dispose();
559
607
 
560
608
  const dstWidth = Math.min(dstTileSize.width, width);
561
609
  const dstHeight = Math.min(dstTileSize.height, height);
@@ -591,8 +639,6 @@ class UNet {
591
639
  }
592
640
  }
593
641
  }
594
-
595
- outputTensor.dispose();
596
642
  } else {
597
643
  dataProcessGPU!.setOutputTile(dstTile, srcTile);
598
644
  // IMPORTANT
@@ -604,14 +650,12 @@ class UNet {
604
650
  [0, 0],
605
651
  [0, 1]
606
652
  ]);
607
- const outBuffer = dataProcessGPU!.inverse(
653
+ outBuffer = dataProcessGPU!.inverse(
608
654
  outputTensor4Channnels.dataToGPU().buffer!,
609
655
  inputData.color
610
656
  );
611
- outputTensor.dispose();
612
- outputTensor4Channnels.dispose();
613
- return outBuffer;
614
657
  }
658
+ return outBuffer!;
615
659
  }
616
660
 
617
661
  tileExecute<T extends ImageData | HDRImageData | GPUImageData>({
@@ -619,15 +663,20 @@ class UNet {
619
663
  albedo,
620
664
  normal,
621
665
  done,
622
- progress
666
+ progress,
667
+ denoiseAlpha
623
668
  }: {
624
669
  color: T;
625
670
  albedo?: ImageData | GPUImageData;
626
671
  normal?: ImageData | GPUImageData;
627
- done: (outputData: T) => void;
672
+ /**
673
+ * If denoise alpha channel. Otherwise denoise RGB channels.
674
+ */
675
+ denoiseAlpha?: boolean;
676
+ done: (outputData: T extends GPUImageData ? GPUImageDataOutput : T) => void;
628
677
  progress?: (
629
- outputData: T,
630
- tileData: T | undefined,
678
+ outputData: T extends GPUImageData ? GPUImageDataOutput : T,
679
+ tileData: (T extends GPUImageData ? GPUImageDataOutput : T) | undefined,
631
680
  tile: Tile,
632
681
  currentIdx: number,
633
682
  totalIdx: number
@@ -655,8 +704,8 @@ class UNet {
655
704
  if (!isGPUImageData(color)) {
656
705
  rawData = this._processImageData(
657
706
  color,
658
- albedo as ImageData,
659
- normal as ImageData,
707
+ albedo as ImageData | undefined,
708
+ normal as ImageData | undefined,
660
709
  hdr
661
710
  );
662
711
  }
@@ -690,12 +739,13 @@ class UNet {
690
739
  }
691
740
  let resGPUBuffer;
692
741
  // profileAndLogKernelCode(() => {
742
+ ENGINE.startScope();
693
743
  resGPUBuffer = this._executeTile(
694
744
  isGPUImageData(color)
695
745
  ? {
696
746
  color: color.data,
697
- albedo: (albedo as GPUImageData).data,
698
- normal: (normal as GPUImageData).data
747
+ albedo: (albedo as GPUImageData | undefined)?.data,
748
+ normal: (normal as GPUImageData | undefined)?.data
699
749
  }
700
750
  : rawData,
701
751
  outputTileData,
@@ -704,8 +754,10 @@ class UNet {
704
754
  j,
705
755
  width,
706
756
  height,
707
- hdr
757
+ hdr,
758
+ denoiseAlpha
708
759
  );
760
+ ENGINE.endScope();
709
761
  // }, true);
710
762
  const output = outputImageData || {
711
763
  data: resGPUBuffer,
@@ -713,9 +765,9 @@ class UNet {
713
765
  height
714
766
  };
715
767
  progress?.(
716
- output as T,
768
+ output as any,
717
769
  // Is undefined if using webgpu buffer
718
- outputTileData as T | undefined,
770
+ outputTileData as any,
719
771
  new Tile(i * tileWidth, j * tileHeight, tileWidth, tileHeight),
720
772
  i + j * tileCountW,
721
773
  tileCountW * tileCountH
@@ -730,7 +782,8 @@ class UNet {
730
782
  }
731
783
  });
732
784
  } else {
733
- done(output as T);
785
+ // console.log(memory());
786
+ done(output as any);
734
787
  }
735
788
  };
736
789
 
@@ -744,6 +797,7 @@ class UNet {
744
797
  dispose() {
745
798
  this._tfModel?.dispose();
746
799
  this._dataProcessGPU?.dispose();
800
+ this._tensors.forEach((tensor) => tensor.dispose());
747
801
  }
748
802
  }
749
803