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/hdr.ts ADDED
@@ -0,0 +1,398 @@
1
+ // https://github.com/bbbbx/hdr.js/blob/main/src/hdr.ts
2
+ function frexp(v: number): { f: number; e: number } {
3
+ v = Number(v);
4
+ const result = { f: v, e: 0 };
5
+
6
+ if (v !== 0 && Number.isFinite(v)) {
7
+ const absV = Math.abs(v);
8
+ const log2 =
9
+ Math.log2 ||
10
+ function log2(n) {
11
+ return Math.log(n) * Math.LOG2E;
12
+ };
13
+
14
+ // Math.pow(2, -exp) === Infinity when exp <= -1024
15
+ let e = Math.max(-1023, Math.floor(log2(absV)) + 1);
16
+ let f = absV * Math.pow(2.0, -e);
17
+
18
+ while (f >= 1.0) {
19
+ f *= 0.5;
20
+ e++;
21
+ }
22
+ while (f < 0.5) {
23
+ f *= 2;
24
+ e--;
25
+ }
26
+
27
+ if (v < 0) {
28
+ f = -f;
29
+ }
30
+
31
+ result.f = f;
32
+ result.e = e;
33
+ }
34
+
35
+ return result;
36
+ }
37
+
38
+ function ldexp(f: number, e: number) {
39
+ return f * Math.pow(2.0, e);
40
+ }
41
+
42
+ /**
43
+ * Convert 4 byte uint8 buffer to 3 channels float data
44
+ * @param rgbe input uint8 buffer
45
+ * @param float output float data
46
+ */
47
+ function rgbe2float(rgbe: Uint8Array, float: Float32Array) {
48
+ if (rgbe[3] !== 0) {
49
+ const f1 = ldexp(1.0, rgbe[3] - (128 + 8));
50
+ float[0] = rgbe[0] * f1;
51
+ float[1] = rgbe[1] * f1;
52
+ float[2] = rgbe[2] * f1;
53
+ } else {
54
+ float[0] = float[1] = float[2] = 0.0;
55
+ }
56
+ }
57
+
58
+ /**
59
+ * Convert 3 channels float data to 4 byte uint8 buffer
60
+ * @param float input float data
61
+ * @param rgbe output uint8 buffer
62
+ */
63
+ function float2rgbe(float: Float32Array, rgbe: Uint8Array) {
64
+ const red = float[0];
65
+ const green = float[1];
66
+ const blue = float[2];
67
+ const v = Math.max(red, Math.max(green, blue));
68
+ if (v < 1e-32) {
69
+ rgbe[0] = rgbe[1] = rgbe[2] = rgbe[3] = 0;
70
+ } else {
71
+ const { f, e } = frexp(v);
72
+ const s = (f / v) * 256.0;
73
+ rgbe[0] = red * s;
74
+ rgbe[1] = green * s;
75
+ rgbe[2] = blue * s;
76
+ rgbe[3] = e + 128;
77
+ }
78
+ }
79
+
80
+ /**
81
+ * Write float data to RGBE(.hdr) file buffer
82
+ * @param w image width
83
+ * @param h image height
84
+ * @param data float data, RGB 3 channels.
85
+ * @returns file buffer
86
+ */
87
+ function write_hdr(w: number, h: number, data: Float32Array): Uint8Array {
88
+ const s: number[] = [];
89
+ const comp = 3;
90
+ write_hdr_core(s, w, h, comp, data);
91
+ return new Uint8Array(s);
92
+ }
93
+
94
+ function write_hdr_core(
95
+ s: number[],
96
+ x: number,
97
+ y: number,
98
+ comp: number,
99
+ data: Float32Array
100
+ ): void {
101
+ if (y <= 0 || x <= 0 || !data) {
102
+ return;
103
+ }
104
+
105
+ const header =
106
+ '#?RADIANCE\n' +
107
+ 'FORMAT=32-bit_rle_rgbe\n' +
108
+ 'EXPOSURE=1.0\n' +
109
+ '\n' +
110
+ '-Y ' +
111
+ y +
112
+ ' +X ' +
113
+ x +
114
+ '\n';
115
+ header.split('').forEach((c) => {
116
+ const charCode = c.charCodeAt(0);
117
+ s.push(charCode);
118
+ });
119
+
120
+ const scratch = new Uint8Array(x * 4);
121
+ const stbi__flip_vertically_on_write = false;
122
+ for (let i = 0; i < y; i++) {
123
+ write_hdr_scanline(
124
+ s,
125
+ x,
126
+ comp,
127
+ scratch,
128
+ data.subarray(comp * x * (stbi__flip_vertically_on_write ? y - 1 - i : i))
129
+ );
130
+ }
131
+ }
132
+
133
+ function write_hdr_scanline(
134
+ s: number[],
135
+ width: number,
136
+ ncomp: number,
137
+ scratch: Uint8Array,
138
+ scanline: Float32Array
139
+ ): void {
140
+ const scanlineheader = new Uint8Array([2, 2, 0, 0]);
141
+ const rgbe = new Uint8Array(4);
142
+ const linear = new Float32Array(3);
143
+ let x;
144
+
145
+ scanlineheader[2] = (width & 0xff00) >> 8;
146
+ scanlineheader[3] = width & 0x00ff;
147
+
148
+ /* skip RLE for images too small or large */
149
+ if (width < 8 || width >= 32768) {
150
+ for (x = 0; x < width; x++) {
151
+ switch (ncomp) {
152
+ case 4: /* fallthrough */
153
+ case 3:
154
+ linear[2] = scanline[x * ncomp + 2];
155
+ linear[1] = scanline[x * ncomp + 1];
156
+ linear[0] = scanline[x * ncomp + 0];
157
+ break;
158
+ default:
159
+ linear[0] = linear[1] = linear[2] = scanline[x * ncomp + 0];
160
+ break;
161
+ }
162
+ float2rgbe(linear, rgbe);
163
+ s.push(rgbe[0], rgbe[1], rgbe[2], rgbe[3]);
164
+ }
165
+ } else {
166
+ let c, r;
167
+ /* encode into scratch buffer */
168
+ for (x = 0; x < width; x++) {
169
+ switch (ncomp) {
170
+ case 4: /* fallthrough */
171
+ case 3:
172
+ linear[2] = scanline[x * ncomp + 2];
173
+ linear[1] = scanline[x * ncomp + 1];
174
+ linear[0] = scanline[x * ncomp + 0];
175
+ break;
176
+ default:
177
+ linear[0] = linear[1] = linear[2] = scanline[x * ncomp + 0];
178
+ break;
179
+ }
180
+ float2rgbe(linear, rgbe);
181
+ scratch[x + width * 0] = rgbe[0];
182
+ scratch[x + width * 1] = rgbe[1];
183
+ scratch[x + width * 2] = rgbe[2];
184
+ scratch[x + width * 3] = rgbe[3];
185
+ }
186
+
187
+ s.push(
188
+ scanlineheader[0],
189
+ scanlineheader[1],
190
+ scanlineheader[2],
191
+ scanlineheader[3]
192
+ );
193
+
194
+ /* RLE each component separately */
195
+ for (c = 0; c < 4; c++) {
196
+ const comp = scratch.subarray(width * c);
197
+
198
+ x = 0;
199
+ while (x < width) {
200
+ // find first run
201
+ r = x;
202
+ while (r + 2 < width) {
203
+ if (comp[r] === comp[r + 1] && comp[r] === comp[r + 2]) break;
204
+ ++r;
205
+ }
206
+ if (r + 2 >= width) r = width;
207
+ // dump up to first run
208
+ while (x < r) {
209
+ let len = r - x;
210
+ if (len > 128) len = 128;
211
+ write_dump_data(s, len, comp.subarray(x));
212
+ x += len;
213
+ }
214
+ // if there's a run, output it
215
+ if (r + 2 < width) {
216
+ // same test as what we break out of in search loop, so only true if we break'd
217
+ // find next byte after run
218
+ while (r < width && comp[r] == comp[x]) ++r;
219
+ // output run up to r
220
+ while (x < r) {
221
+ let len = r - x;
222
+ if (len > 127) len = 127;
223
+ write_run_data(s, len, comp[x]);
224
+ x += len;
225
+ }
226
+ }
227
+ }
228
+ }
229
+ }
230
+ }
231
+
232
+ function write_dump_data(s: number[], length: number, data: Uint8Array): void {
233
+ const lengthbyte = length & 0xff;
234
+ if (!(length <= 128)) throw new Error('length is greater than 128');
235
+ s.push(lengthbyte);
236
+ for (let i = 0; i < length; i++) {
237
+ s.push(data[i]);
238
+ }
239
+ }
240
+
241
+ function write_run_data(s: number[], length: number, databyte: number): void {
242
+ const lengthbyte = (length + 128) & 0xff;
243
+ if (!(length + 128 <= 255)) throw new Error('length is greater than 128');
244
+ s.push(lengthbyte, databyte);
245
+ }
246
+
247
+ function testMagic(uint8: Uint8Array, signature: string): boolean {
248
+ const l = signature.length;
249
+ for (let i = 0; i < l; i++) {
250
+ if (uint8[i] !== signature.charCodeAt(i)) {
251
+ return false;
252
+ }
253
+ }
254
+ return true;
255
+ }
256
+
257
+ function isHdr(uint8: Uint8Array): boolean {
258
+ let r = testMagic(uint8, '#?RADIANCE\n');
259
+ if (!r) {
260
+ r = testMagic(uint8, '#?RGBE\n');
261
+ }
262
+ return r;
263
+ }
264
+
265
+ const MAX_HEADER_LENGTH = 1024 * 10;
266
+ const MAX_DIMENSIONS = 1 << 24;
267
+
268
+ /**
269
+ * Read a RGBE(.hdr) file from buffer
270
+ * @param uint8 RGBE(.hdr) file buffer
271
+ * @returns Failure reason or resolved data.
272
+ */
273
+ function read_hdr(uint8: Uint8Array):
274
+ | string
275
+ | {
276
+ rgbFloat: Float32Array;
277
+ width: number;
278
+ height: number;
279
+ } {
280
+ let header = '';
281
+ let pos = 0;
282
+
283
+ if (!isHdr(uint8)) {
284
+ return 'Corrupt HDR image.';
285
+ }
286
+
287
+ // read header
288
+ while (!header.match(/\n\n[^\n]+\n/g) && pos < MAX_HEADER_LENGTH) {
289
+ header += String.fromCharCode(uint8[pos++]);
290
+ }
291
+
292
+ // check format
293
+ const format = header.match(/FORMAT=(.*)$/m)?.[1];
294
+ if (format !== '32-bit_rle_rgbe') {
295
+ return 'Unsupported HDR format: ' + format;
296
+ }
297
+
298
+ // parse resolution
299
+ const rez: string[] = header.split(/\n/).reverse()[1].split(' ');
300
+ if (rez[0] !== '-Y' || rez[2] !== '+X') {
301
+ return 'Unsupported HDR format';
302
+ }
303
+ const width = Number.parseFloat(rez[3]);
304
+ const height = Number.parseFloat(rez[1]);
305
+ if (width > MAX_DIMENSIONS || height > MAX_DIMENSIONS) {
306
+ return 'Very large image (corrupt?)';
307
+ }
308
+
309
+ let i, j;
310
+ let c1: number = uint8[pos];
311
+ let c2: number = uint8[pos + 1];
312
+ let len: number = uint8[pos + 2];
313
+
314
+ // not run-length encoded, so we have to actually use THIS data as a decoded
315
+ // pixel (note this can't be a valid pixel--one of RGB must be >= 128)
316
+ const notRLE: boolean = c1 !== 2 || c2 !== 2 || !!(len & 0x80); // not run-length encoded
317
+
318
+ const hdrData = new Float32Array(width * height * 3);
319
+ if (width < 8 || width >= 32768 || notRLE) {
320
+ // 32768: 2^15
321
+ // Read flat data
322
+ for (j = 0; j < height; ++j) {
323
+ for (i = 0; i < width; ++i) {
324
+ const rgbe = uint8.subarray(pos, pos + 4);
325
+ pos += 4;
326
+ const start = (j * width + i) * 3;
327
+ rgbe2float(rgbe, hdrData.subarray(start, start + 3));
328
+ }
329
+ }
330
+ } else {
331
+ // Read RLE-encoded data
332
+ let scanline: Uint8Array | undefined;
333
+ let c1: number;
334
+ let c2: number;
335
+ let len: number;
336
+ for (let j = 0; j < height; j++) {
337
+ c1 = uint8[pos++];
338
+ c2 = uint8[pos++];
339
+ len = uint8[pos++];
340
+ if (c1 !== 2 || c2 !== 2 || len & 0x80) {
341
+ return 'Invalid scanline';
342
+ }
343
+
344
+ len = len << 8;
345
+ len |= uint8[pos++];
346
+ if (len !== width) {
347
+ return 'invalid decoded scanline length';
348
+ }
349
+ if (!scanline) {
350
+ scanline = new Uint8Array(width * 4);
351
+ }
352
+
353
+ let count: number;
354
+ let value: number;
355
+ for (let k = 0; k < 4; k++) {
356
+ let nLeft: number;
357
+ i = 0;
358
+ while ((nLeft = width - i) > 0) {
359
+ count = uint8[pos++];
360
+ if (count > 128) {
361
+ // is RUN
362
+ value = uint8[pos++];
363
+ count -= 128;
364
+ if (count > nLeft) {
365
+ return 'bad RLE data in HDR';
366
+ }
367
+ for (let z = 0; z < count; z++) {
368
+ scanline[i++ * 4 + k] = value;
369
+ }
370
+ } else {
371
+ // is DUMP
372
+ if (count > nLeft) {
373
+ return 'bad RLE data in HDR';
374
+ }
375
+ for (let z = 0; z < count; z++) {
376
+ scanline[i++ * 4 + k] = uint8[pos++];
377
+ }
378
+ }
379
+ }
380
+ }
381
+
382
+ for (let i = 0; i < width; i++) {
383
+ rgbe2float(
384
+ scanline.subarray(i * 4),
385
+ hdrData.subarray((j * width + i) * 3)
386
+ );
387
+ }
388
+ }
389
+ }
390
+
391
+ return {
392
+ rgbFloat: hdrData,
393
+ width: width,
394
+ height: height
395
+ };
396
+ }
397
+
398
+ export { read_hdr as readHDR, write_hdr as writeHDR, float2rgbe, rgbe2float };
package/src/helper.ts ADDED
@@ -0,0 +1,35 @@
1
+ import * as tfjs from '@tensorflow/tfjs-core';
2
+ export function profileAndLogKernelCode(execute: () => void, disabled = true) {
3
+ if (disabled) {
4
+ execute();
5
+ return;
6
+ }
7
+ tfjs
8
+ .profile(() => {
9
+ execute();
10
+ })
11
+ .then((res) => {
12
+ const kernelNames = Array.from(
13
+ new Set(
14
+ res.kernels.map((k) => k.name).filter((name) => !name.endsWith('_op'))
15
+ )
16
+ );
17
+ function nameToConfig(name: string) {
18
+ return `${name[0].toLowerCase()}${name.slice(1)}Config`;
19
+ }
20
+ const importCode = kernelNames.map(
21
+ (name) =>
22
+ `import { ${nameToConfig(
23
+ name
24
+ )} } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/${name}';`
25
+ );
26
+ const configCode = kernelNames.map((name) => `${nameToConfig(name)},`);
27
+ const code = `
28
+ ${importCode.join('\n')}
29
+ const kernelConfigs: KernelConfig[] = [
30
+ ${configCode.join('\n')}
31
+ ]
32
+ `;
33
+ console.log(code);
34
+ });
35
+ }
package/src/kernels.ts ADDED
@@ -0,0 +1,31 @@
1
+ import {
2
+ KernelConfig,
3
+ registerKernel
4
+ } from '@tensorflow/tfjs-core/dist/kernel_registry';
5
+
6
+ import { mirrorPadConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/MirrorPad';
7
+ import { padV2Config } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/PadV2';
8
+ import { sliceConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/Slice';
9
+ import { fusedConv2DConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/FusedConv2D';
10
+ import { maxPoolConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/MaxPool';
11
+ import { resizeNearestNeighborConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/ResizeNearestNeighbor';
12
+ import { concatConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/Concat';
13
+ import { identityConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/Identity';
14
+
15
+ const kernelConfigs: KernelConfig[] = [
16
+ mirrorPadConfig,
17
+ padV2Config,
18
+ sliceConfig,
19
+ fusedConv2DConfig,
20
+ maxPoolConfig,
21
+ resizeNearestNeighborConfig,
22
+ concatConfig,
23
+ identityConfig
24
+ ];
25
+
26
+ for (const kernelConfig of kernelConfigs) {
27
+ registerKernel({
28
+ ...kernelConfig,
29
+ backendName: 'webgpu-oidn'
30
+ });
31
+ }
package/src/main.ts ADDED
@@ -0,0 +1,42 @@
1
+ import { WebGPUBackend } from '@tensorflow/tfjs-backend-webgpu/dist/base';
2
+ import { parseTZA } from './tza';
3
+ import UNet from './UNet';
4
+ import { initWebGPUBackend, initWebGPUBackendWithDevice } from './backend';
5
+
6
+ export { parseTZA, UNet };
7
+
8
+ export async function initUNetFromBuffer(
9
+ tzaBuffer: ArrayBuffer,
10
+ backendParams?: { device: GPUDevice; adapterInfo: GPUAdapterInfo },
11
+ opts?: {
12
+ aux?: boolean;
13
+ hdr?: boolean;
14
+ maxTileSize?: number;
15
+ }
16
+ ) {
17
+ const backend = await (backendParams
18
+ ? initWebGPUBackendWithDevice(
19
+ backendParams.device,
20
+ backendParams.adapterInfo
21
+ )
22
+ : initWebGPUBackend());
23
+ const tensors = parseTZA(tzaBuffer);
24
+ const unet = new UNet(tensors, backend!, opts);
25
+ return unet;
26
+ }
27
+
28
+ export async function initUNetFromURL(
29
+ modelPath: string,
30
+ backendParams?: { device: GPUDevice; adapterInfo: GPUAdapterInfo },
31
+ opts?: {
32
+ aux?: boolean;
33
+ hdr?: boolean;
34
+ maxTileSize?: number;
35
+ }
36
+ ) {
37
+ return fetch(modelPath)
38
+ .then((res) => res.arrayBuffer())
39
+ .then((ab) => {
40
+ return initUNetFromBuffer(ab, backendParams, opts);
41
+ });
42
+ }