oidn-web 0.4.0 → 0.5.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 (106) hide show
  1. package/CHANGELOG.md +45 -0
  2. package/README.md +83 -19
  3. package/dist/oidn.js +3239 -2642
  4. package/dist/oidn.umd.cjs +512 -296
  5. package/lib/UNet.d.ts +54 -20
  6. package/lib/UNet.js +194 -118
  7. package/lib/UNet.js.map +1 -1
  8. package/lib/backend.d.ts +1 -8
  9. package/lib/backend.js +1 -9
  10. package/lib/backend.js.map +1 -1
  11. package/lib/finalRgbShader.d.ts +13 -0
  12. package/lib/finalRgbShader.js +160 -0
  13. package/lib/finalRgbShader.js.map +1 -0
  14. package/lib/graphOptimizer.js +1 -2
  15. package/lib/graphOptimizer.js.map +1 -1
  16. package/lib/hdrTransfer.d.ts +14 -0
  17. package/lib/hdrTransfer.js +61 -0
  18. package/lib/hdrTransfer.js.map +1 -0
  19. package/lib/main.d.ts +16 -7
  20. package/lib/main.js +4 -5
  21. package/lib/main.js.map +1 -1
  22. package/lib/nativeUNet.d.ts +39 -3
  23. package/lib/nativeUNet.js +449 -120
  24. package/lib/nativeUNet.js.map +1 -1
  25. package/lib/process.d.ts +5 -11
  26. package/lib/process.js +35 -49
  27. package/lib/process.js.map +1 -1
  28. package/lib/tileScheduler.d.ts +32 -4
  29. package/lib/tileScheduler.js +133 -20
  30. package/lib/tileScheduler.js.map +1 -1
  31. package/package.json +9 -2
  32. package/src/UNet.ts +287 -158
  33. package/src/backend.ts +1 -14
  34. package/src/finalRgbShader.ts +186 -0
  35. package/src/graphOptimizer.ts +1 -2
  36. package/src/hdrTransfer.ts +88 -0
  37. package/src/main.ts +28 -13
  38. package/src/nativeUNet.ts +515 -116
  39. package/src/process.ts +43 -70
  40. package/src/tileScheduler.ts +216 -24
  41. package/benchmarks/compare.mjs +0 -651
  42. package/benchmarks/leak.mjs +0 -255
  43. package/benchmarks/results/before-spatial.json +0 -391
  44. package/benchmarks/results/before-spatial.md +0 -47
  45. package/benchmarks/results/int8-scan.json +0 -2007
  46. package/benchmarks/results/int8-scan.md +0 -160
  47. package/benchmarks/results/int8-w8a8-scan.json +0 -2007
  48. package/benchmarks/results/int8-w8a8-scan.md +0 -160
  49. package/benchmarks/results/int8-weight-channel.json +0 -1413
  50. package/benchmarks/results/int8-weight-channel.md +0 -118
  51. package/benchmarks/results/int8-weight-only.json +0 -1437
  52. package/benchmarks/results/int8-weight-only.md +0 -118
  53. package/benchmarks/results/kernel-webnn-final.json +0 -1115
  54. package/benchmarks/results/kernel-webnn-final.md +0 -104
  55. package/benchmarks/results/latest-optimized.json +0 -375
  56. package/benchmarks/results/latest-optimized.md +0 -47
  57. package/benchmarks/results/latest.json +0 -391
  58. package/benchmarks/results/latest.md +0 -47
  59. package/benchmarks/results/profile-baseline.json +0 -331
  60. package/benchmarks/results/profile-baseline.md +0 -12
  61. package/benchmarks/results/profile-conv2x.json +0 -331
  62. package/benchmarks/results/profile-conv2x.md +0 -12
  63. package/benchmarks/results/profile-fast-init.json +0 -385
  64. package/benchmarks/results/profile-fast-init.md +0 -47
  65. package/benchmarks/results/profile-fp16-fma.json +0 -369
  66. package/benchmarks/results/profile-fp16-fma.md +0 -47
  67. package/benchmarks/results/profile-fp16-tiled-decoder.json +0 -369
  68. package/benchmarks/results/profile-fp16-tiled-decoder.md +0 -47
  69. package/benchmarks/results/profile-fp16-tiled-encoder.json +0 -369
  70. package/benchmarks/results/profile-fp16-tiled-encoder.md +0 -47
  71. package/benchmarks/results/profile-fp16-unfused-pool.json +0 -385
  72. package/benchmarks/results/profile-fp16-unfused-pool.md +0 -47
  73. package/benchmarks/results/profile-input-major.json +0 -347
  74. package/benchmarks/results/profile-input-major.md +0 -12
  75. package/benchmarks/results/profile-k16.json +0 -347
  76. package/benchmarks/results/profile-k16.md +0 -12
  77. package/benchmarks/results/profile-k4.json +0 -347
  78. package/benchmarks/results/profile-k4.md +0 -12
  79. package/benchmarks/results/profile-pool-reuse.json +0 -331
  80. package/benchmarks/results/profile-pool-reuse.md +0 -12
  81. package/benchmarks/results/profile-precompiled.json +0 -385
  82. package/benchmarks/results/profile-precompiled.md +0 -47
  83. package/benchmarks/results/profile-static-channels.json +0 -385
  84. package/benchmarks/results/profile-static-channels.md +0 -47
  85. package/benchmarks/results/profile-static-io.json +0 -385
  86. package/benchmarks/results/profile-static-io.md +0 -47
  87. package/benchmarks/results/profile-tiled-conv.json +0 -331
  88. package/benchmarks/results/profile-tiled-conv.md +0 -12
  89. package/benchmarks/results/profile-tiled-decoder.json +0 -347
  90. package/benchmarks/results/profile-tiled-decoder.md +0 -12
  91. package/benchmarks/results/profile-tiled-matmul.json +0 -331
  92. package/benchmarks/results/profile-tiled-matmul.md +0 -12
  93. package/benchmarks/results/profile-unfused-decoder.json +0 -379
  94. package/benchmarks/results/profile-unfused-decoder.md +0 -12
  95. package/benchmarks/results/profile-unfused-pool.json +0 -347
  96. package/benchmarks/results/profile-unfused-pool.md +0 -12
  97. package/benchmarks/results/spatial-auto.json +0 -575
  98. package/benchmarks/results/spatial-auto.md +0 -61
  99. package/benchmarks/results/subgroup-smoke.json +0 -1094
  100. package/benchmarks/results/subgroup-smoke.md +0 -104
  101. package/benchmarks/results/webnn-smoke.json +0 -739
  102. package/benchmarks/results/webnn-smoke.md +0 -76
  103. package/scripts/inspect-model.mjs +0 -64
  104. package/tests/modelSpec.test.mjs +0 -128
  105. package/tests/resourceLifecycle.test.mjs +0 -383
  106. package/tests/tileScheduler.test.mjs +0 -90
package/src/backend.ts CHANGED
@@ -32,18 +32,5 @@ export async function initWebGPUBackend() {
32
32
  maxComputeInvocationsPerWorkgroup:
33
33
  adapterLimits.maxComputeInvocationsPerWorkgroup
34
34
  };
35
- const device = await adapter.requestDevice(deviceDescriptor);
36
- const adapterInfo =
37
- // requestAdapterInfo is deprecated
38
- // @ts-ignore
39
- adapter.info ?? (await adapter.requestAdapterInfo?.());
40
-
41
- return initWebGPUBackendWithDevice(device, adapterInfo);
42
- }
43
-
44
- export async function initWebGPUBackendWithDevice(
45
- device: GPUDevice,
46
- adapterInfo: GPUAdapterInfo
47
- ) {
48
- return { device, adapterInfo };
35
+ return adapter.requestDevice(deviceDescriptor);
49
36
  }
@@ -0,0 +1,186 @@
1
+ export type FinalRgbPrecision = 'fp16' | 'fp32';
2
+ export type FinalRgbActivation = 'relu' | 'identity';
3
+
4
+ const WORKGROUP_SIZE = 8;
5
+ const PATCH_SIZE = WORKGROUP_SIZE + 2;
6
+ const KERNEL_ELEMENTS = 3 * 3;
7
+
8
+ /** Shared storage required by the 8x8 final-RGB workgroup. */
9
+ export function sharedMemoryBytes(
10
+ precision: FinalRgbPrecision,
11
+ inputBlocks: number,
12
+ cacheWeights = false
13
+ ) {
14
+ const bytesPerScalar = precision === 'fp16' ? 2 : 4;
15
+ const inputVec4Count = PATCH_SIZE * PATCH_SIZE * inputBlocks;
16
+ const weightVec4Count = cacheWeights
17
+ ? KERNEL_ELEMENTS * inputBlocks * 4
18
+ : 0;
19
+ const vec4Count = inputVec4Count + weightVec4Count;
20
+ return vec4Count * 4 * bytesPerScalar;
21
+ }
22
+
23
+ function storageVecType(precision: FinalRgbPrecision) {
24
+ return precision === 'fp16' ? 'vec4<f16>' : 'vec4<f32>';
25
+ }
26
+
27
+ function shaderPreamble(precision: FinalRgbPrecision) {
28
+ return precision === 'fp16' ? 'enable f16;\n' : '';
29
+ }
30
+
31
+ function accumulationCode(
32
+ precision: FinalRgbPrecision,
33
+ inputExpression: string,
34
+ weightExpression: string
35
+ ) {
36
+ const weight = (lane: number) =>
37
+ `${weightExpression}[weightBase + ${lane}u]`;
38
+ if (precision === 'fp16') {
39
+ return /* wgsl */ `
40
+ let inputValue = vec4<f16>(${inputExpression});
41
+ var partial = vec4<f16>(0.0h);
42
+ partial = fma(${weight(0)}, vec4<f16>(inputValue.x), partial);
43
+ partial = fma(${weight(1)}, vec4<f16>(inputValue.y), partial);
44
+ partial = fma(${weight(2)}, vec4<f16>(inputValue.z), partial);
45
+ partial = fma(${weight(3)}, vec4<f16>(inputValue.w), partial);
46
+ acc += vec4<f32>(partial);
47
+ `;
48
+ }
49
+ return /* wgsl */ `
50
+ let inputValue = vec4<f32>(${inputExpression});
51
+ acc = fma(vec4<f32>(${weight(0)}), vec4<f32>(inputValue.x), acc);
52
+ acc = fma(vec4<f32>(${weight(1)}), vec4<f32>(inputValue.y), acc);
53
+ acc = fma(vec4<f32>(${weight(2)}), vec4<f32>(inputValue.z), acc);
54
+ acc = fma(vec4<f32>(${weight(3)}), vec4<f32>(inputValue.w), acc);
55
+ `;
56
+ }
57
+
58
+ /**
59
+ * Builds the final three-channel same-padded convolution shader.
60
+ *
61
+ * The bind-group layout and Params block intentionally match createConvShader:
62
+ * input, weights, bias, output, then uniform params. Weights retain the
63
+ * output-major packed ABI, with four vec4 values per (kernel position, input
64
+ * block), one for each input lane.
65
+ */
66
+ export function createFinalRgbShader(
67
+ precision: FinalRgbPrecision,
68
+ activation: FinalRgbActivation,
69
+ inputBlocks: number,
70
+ cacheWeights = false
71
+ ) {
72
+ if (!Number.isInteger(inputBlocks) || inputBlocks < 1) {
73
+ throw new Error(`Final RGB shader requires positive input blocks, got ${inputBlocks}`);
74
+ }
75
+ const inputType = storageVecType(precision);
76
+ const outputType = 'vec4<f32>';
77
+ const inputTileValues = PATCH_SIZE * PATCH_SIZE * inputBlocks;
78
+ const weightTileValues = KERNEL_ELEMENTS * inputBlocks * 4;
79
+ const stored = activation === 'relu'
80
+ ? 'max(acc, vec4<f32>(0.0))'
81
+ : 'acc';
82
+ const accumulation = accumulationCode(
83
+ precision,
84
+ 'inputTile[patchBase + inputBlock]',
85
+ cacheWeights ? 'weightTile' : 'weights'
86
+ );
87
+
88
+ return /* wgsl */ `${shaderPreamble(precision)}
89
+ struct Params {
90
+ inputWidth: u32,
91
+ inputHeight: u32,
92
+ outputWidth: u32,
93
+ outputHeight: u32,
94
+ inputBlocks: u32,
95
+ outputBlocks: u32,
96
+ }
97
+
98
+ @group(0) @binding(0) var<storage, read> inputData: array<${inputType}>;
99
+ @group(0) @binding(1) var<storage, read> weights: array<${inputType}>;
100
+ @group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
101
+ @group(0) @binding(3) var<storage, read_write> outputData: array<${outputType}>;
102
+ @group(0) @binding(4) var<uniform> params: Params;
103
+
104
+ var<workgroup> inputTile: array<${inputType}, ${inputTileValues}>;
105
+ ${cacheWeights
106
+ ? `var<workgroup> weightTile: array<${inputType}, ${weightTileValues}>;`
107
+ : ''}
108
+
109
+ @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
110
+ fn main(
111
+ @builtin(local_invocation_id) localId: vec3<u32>,
112
+ @builtin(global_invocation_id) gid: vec3<u32>,
113
+ @builtin(workgroup_id) workgroupId: vec3<u32>
114
+ ) {
115
+ let localLinear = localId.y * ${WORKGROUP_SIZE}u + localId.x;
116
+ for (
117
+ var loadIndex = localLinear;
118
+ loadIndex < ${inputTileValues}u;
119
+ loadIndex += ${WORKGROUP_SIZE * WORKGROUP_SIZE}u
120
+ ) {
121
+ let tilePixel = loadIndex / ${inputBlocks}u;
122
+ let inputBlock = loadIndex % ${inputBlocks}u;
123
+ let tileX = tilePixel % ${PATCH_SIZE}u;
124
+ let tileY = tilePixel / ${PATCH_SIZE}u;
125
+ let inputX = i32(workgroupId.x * ${WORKGROUP_SIZE}u + tileX) - 1;
126
+ let inputY = i32(workgroupId.y * ${WORKGROUP_SIZE}u + tileY) - 1;
127
+ var value = ${inputType}(0.0);
128
+ if (
129
+ inputX >= 0 && inputX < i32(params.inputWidth) &&
130
+ inputY >= 0 && inputY < i32(params.inputHeight)
131
+ ) {
132
+ let inputIndex =
133
+ (u32(inputY) * params.inputWidth + u32(inputX)) *
134
+ ${inputBlocks}u + inputBlock;
135
+ value = inputData[inputIndex];
136
+ }
137
+ inputTile[loadIndex] = value;
138
+ }
139
+
140
+ ${cacheWeights ? ` for (
141
+ var loadIndex = localLinear;
142
+ loadIndex < ${weightTileValues}u;
143
+ loadIndex += ${WORKGROUP_SIZE * WORKGROUP_SIZE}u
144
+ ) {
145
+ weightTile[loadIndex] = weights[loadIndex];
146
+ }
147
+
148
+ ` : ''} // Out-of-range invocations must reach this barrier before returning.
149
+ workgroupBarrier();
150
+
151
+ let outputInBounds =
152
+ gid.x < params.outputWidth &&
153
+ gid.y < params.outputHeight &&
154
+ gid.z < params.outputBlocks;
155
+ if (!outputInBounds) {
156
+ return;
157
+ }
158
+
159
+ var acc = bias[gid.z];
160
+ for (var ky = 0u; ky < 3u; ky++) {
161
+ let inputY = i32(gid.y) + i32(ky) - 1;
162
+ if (inputY < 0 || inputY >= i32(params.inputHeight)) {
163
+ continue;
164
+ }
165
+ for (var kx = 0u; kx < 3u; kx++) {
166
+ let inputX = i32(gid.x) + i32(kx) - 1;
167
+ if (inputX < 0 || inputX >= i32(params.inputWidth)) {
168
+ continue;
169
+ }
170
+ let patchBase =
171
+ ((localId.y + ky) * ${PATCH_SIZE}u + localId.x + kx) *
172
+ ${inputBlocks}u;
173
+ for (var inputBlock = 0u; inputBlock < ${inputBlocks}u; inputBlock++) {
174
+ let weightBase =
175
+ ((ky * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u;
176
+ ${accumulation}
177
+ }
178
+ }
179
+ }
180
+
181
+ let outputIndex =
182
+ (gid.y * params.outputWidth + gid.x) * ${1}u + gid.z;
183
+ outputData[outputIndex] = ${stored};
184
+ }
185
+ `;
186
+ }
@@ -112,8 +112,7 @@ export function optimizeModelGraph(
112
112
  if (concat?.op !== 'concat' || concat.inputs.length !== 2) continue;
113
113
  if (onlyConsumer(consumers, concat.id, 'conv2d') !== node) continue;
114
114
  // The native blocked layout can remove concat only when the source
115
- // boundary is also a vec4 boundary. Other graphs keep the generic ops
116
- // and can use the compatibility engine until a scalar-tail kernel exists.
115
+ // boundary is also a vec4 boundary. Other graphs keep the generic ops.
117
116
  const firstInputChannels = validated.channelsByValue.get(concat.inputs[0]);
118
117
  if (firstInputChannels === undefined || firstInputChannels % 4 !== 0) {
119
118
  continue;
@@ -0,0 +1,88 @@
1
+ /** HDR transfer functions used by the OIDN input and output processing passes. */
2
+ export type HDRTransfer = 'pu' | 'log';
3
+
4
+ const a = 1.41283765e3;
5
+ const b = 1.64593172;
6
+ const c = 4.31384981e-1;
7
+ const d = -2.94139609e-3;
8
+ const e = 1.92653254e-1;
9
+ const f = 6.26026094e-3;
10
+ const g = 9.98620152e-1;
11
+ const y0 = 1.5794576e-6;
12
+ const y1 = 3.22087631e-2;
13
+ const x0 = 2.23151711e-3;
14
+ const x1 = 3.70974749e-1;
15
+ const yMax = 65504;
16
+ const puXMax = puForward(yMax);
17
+ const puNormScale = 1 / puXMax;
18
+ const puRcpNormScale = puXMax;
19
+ const logXMax = Math.log(yMax + 1);
20
+ const logNormScale = 1 / logXMax;
21
+
22
+ function puForward(y: number) {
23
+ if (y <= y0) return a * y;
24
+ if (y <= y1) return b * Math.pow(y, c) + d;
25
+ return e * Math.log(y + f) + g;
26
+ }
27
+
28
+ function puInverse(x: number) {
29
+ if (x <= x0) return x / a;
30
+ if (x <= x1) return Math.pow((x - d) / b, 1 / c);
31
+ return Math.exp((x - g) / e) - f;
32
+ }
33
+
34
+ function forward(y: number, transfer: HDRTransfer) {
35
+ return transfer === 'log'
36
+ ? Math.log(y + 1) * logNormScale
37
+ : puForward(y) * puNormScale;
38
+ }
39
+
40
+ function inverse(x: number, transfer: HDRTransfer) {
41
+ return transfer === 'log'
42
+ ? Math.exp(x * logXMax) - 1
43
+ : puInverse(x * puRcpNormScale);
44
+ }
45
+
46
+ export function hdrTransferFuncCPU({
47
+ data,
48
+ channels,
49
+ inputScale,
50
+ transfer = 'pu'
51
+ }: {
52
+ data: Float32Array;
53
+ channels: number;
54
+ inputScale: number;
55
+ transfer?: HDRTransfer;
56
+ }) {
57
+ const newData = new Float32Array(data);
58
+ for (let i = 0; i < newData.length; i += channels) {
59
+ for (let channel = 0; channel < 3; channel++) {
60
+ newData[i + channel] = forward(
61
+ newData[i + channel] * inputScale,
62
+ transfer
63
+ );
64
+ }
65
+ }
66
+ return newData;
67
+ }
68
+
69
+ export function hdrTransferFuncInverseCPU({
70
+ data,
71
+ channels,
72
+ inputScale,
73
+ transfer = 'pu'
74
+ }: {
75
+ data: Float32Array;
76
+ channels: number;
77
+ inputScale: number;
78
+ transfer?: HDRTransfer;
79
+ }) {
80
+ const newData = new Float32Array(data);
81
+ const outputScale = 1 / inputScale;
82
+ for (let i = 0; i < newData.length; i += channels) {
83
+ for (let channel = 0; channel < 3; channel++) {
84
+ newData[i + channel] = inverse(newData[i + channel], transfer) * outputScale;
85
+ }
86
+ }
87
+ return newData;
88
+ }
package/src/main.ts CHANGED
@@ -1,16 +1,25 @@
1
1
  import { parseTZA } from './tza';
2
2
  import UNet from './UNet';
3
3
  import type { UNetEngineSetting, UNetExecutionStats } from './UNet';
4
- import { initWebGPUBackend, initWebGPUBackendWithDevice } from './backend';
4
+ import { initWebGPUBackend } from './backend';
5
5
  import type { DynamicTileSetting } from './tileScheduler';
6
6
  import type { UNetModelSpec } from './modelSpec';
7
7
  import type {
8
+ NativeUNetGemmOptions,
8
9
  NativeUNetKernelSetting,
9
10
  NativeUNetPrecisionSetting
10
11
  } from './nativeUNet';
12
+ import type { HDRTransfer } from './process';
11
13
 
12
14
  export { parseTZA, UNet };
13
- export type { DynamicTileOptions, DynamicTileSetting } from './tileScheduler';
15
+ export { planTileGrid } from './tileScheduler';
16
+ export type {
17
+ DynamicTileOptions,
18
+ DynamicTileSetting,
19
+ PlannedTile,
20
+ TilePlan,
21
+ TileRect
22
+ } from './tileScheduler';
14
23
  export {
15
24
  detectUNetModelSpec,
16
25
  OIDN_UNET_LARGE_SPEC,
@@ -36,6 +45,8 @@ export {
36
45
  resolveNativeUNetPrecision
37
46
  } from './nativeUNet';
38
47
  export type {
48
+ NativeUNetGemmWorkgroup,
49
+ NativeUNetGemmOptions,
39
50
  NativeUNetExecutionProfile,
40
51
  NativeUNetKernel,
41
52
  NativeUNetKernelSetting,
@@ -48,21 +59,30 @@ export type {
48
59
  export interface UNetOptions {
49
60
  aux?: boolean;
50
61
  hdr?: boolean;
62
+ /** HDR transfer function expected by the trained model. Defaults to PU. */
63
+ hdrTransfer?: HDRTransfer;
51
64
  /** Hard upper bound for an output tile edge. Defaults to 512. */
52
65
  maxTileSize?: number;
53
66
  /** Adaptive GPU-time-based tile sizing. Enabled by default. */
54
67
  dynamicTile?: DynamicTileSetting;
55
- /** `auto` uses stable WGSL; `webnn` opts into the experimental WebNN backend. */
68
+ /** `auto` uses native WGSL; `webnn` opts into the WebNN backend. */
56
69
  engine?: UNetEngineSetting;
57
70
  /** `auto` selects FP16 when shader-f16 was enabled on the GPUDevice. */
58
71
  precision?: NativeUNetPrecisionSetting;
59
- /** `auto` selects kernels from precision, operation shape, and GPU limits. */
72
+ /** `auto` uses implicit GEMM for FP16/FP32 convolutions, except the direct output layer. */
60
73
  kernel?: NativeUNetKernelSetting;
74
+ /**
75
+ * Experimental implicit-GEMM tuning for native WGSL execution. These knobs
76
+ * exist for benchmarking and may change or be removed in any release.
77
+ * @experimental
78
+ */
79
+ gemm?: NativeUNetGemmOptions;
61
80
  /** Versioned topology descriptor for future/custom OIDN TZA models. */
62
81
  modelSpec?: UNetModelSpec;
63
82
  }
64
83
 
65
84
  export type { UNetEngineSetting, UNetExecutionStats } from './UNet';
85
+ export type { HDRTransfer } from './process';
66
86
  export { WebNNUNetExecutor } from './webnnUNet';
67
87
  export type { WebNNRuntimeSupport, WebNNUNetOptions } from './webnnUNet';
68
88
  export type {
@@ -73,24 +93,19 @@ export type {
73
93
 
74
94
  export async function initUNetFromBuffer(
75
95
  tzaBuffer: ArrayBuffer,
76
- backendParams?: { device: GPUDevice; adapterInfo: GPUAdapterInfo },
96
+ backendParams?: { device: GPUDevice },
77
97
  opts?: UNetOptions
78
98
  ) {
79
- const backend = await (backendParams
80
- ? initWebGPUBackendWithDevice(
81
- backendParams.device,
82
- backendParams.adapterInfo
83
- )
84
- : initWebGPUBackend());
99
+ const device = backendParams?.device ?? await initWebGPUBackend();
85
100
  const tensors = parseTZA(tzaBuffer);
86
- const unet = new UNet(tensors, backend, opts);
101
+ const unet = new UNet(tensors, device, opts);
87
102
  await unet.prepare();
88
103
  return unet;
89
104
  }
90
105
 
91
106
  export async function initUNetFromURL(
92
107
  modelPath: string,
93
- backendParams?: { device: GPUDevice; adapterInfo: GPUAdapterInfo },
108
+ backendParams?: { device: GPUDevice },
94
109
  opts?: UNetOptions
95
110
  ) {
96
111
  return fetch(modelPath)