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.
- package/CHANGELOG.md +45 -0
- package/README.md +83 -19
- package/dist/oidn.js +3239 -2642
- package/dist/oidn.umd.cjs +512 -296
- package/lib/UNet.d.ts +54 -20
- package/lib/UNet.js +194 -118
- package/lib/UNet.js.map +1 -1
- package/lib/backend.d.ts +1 -8
- package/lib/backend.js +1 -9
- package/lib/backend.js.map +1 -1
- package/lib/finalRgbShader.d.ts +13 -0
- package/lib/finalRgbShader.js +160 -0
- package/lib/finalRgbShader.js.map +1 -0
- package/lib/graphOptimizer.js +1 -2
- package/lib/graphOptimizer.js.map +1 -1
- package/lib/hdrTransfer.d.ts +14 -0
- package/lib/hdrTransfer.js +61 -0
- package/lib/hdrTransfer.js.map +1 -0
- package/lib/main.d.ts +16 -7
- package/lib/main.js +4 -5
- package/lib/main.js.map +1 -1
- package/lib/nativeUNet.d.ts +39 -3
- package/lib/nativeUNet.js +449 -120
- package/lib/nativeUNet.js.map +1 -1
- package/lib/process.d.ts +5 -11
- package/lib/process.js +35 -49
- package/lib/process.js.map +1 -1
- package/lib/tileScheduler.d.ts +32 -4
- package/lib/tileScheduler.js +133 -20
- package/lib/tileScheduler.js.map +1 -1
- package/package.json +9 -2
- package/src/UNet.ts +287 -158
- package/src/backend.ts +1 -14
- package/src/finalRgbShader.ts +186 -0
- package/src/graphOptimizer.ts +1 -2
- package/src/hdrTransfer.ts +88 -0
- package/src/main.ts +28 -13
- package/src/nativeUNet.ts +515 -116
- package/src/process.ts +43 -70
- package/src/tileScheduler.ts +216 -24
- package/benchmarks/compare.mjs +0 -651
- package/benchmarks/leak.mjs +0 -255
- package/benchmarks/results/before-spatial.json +0 -391
- package/benchmarks/results/before-spatial.md +0 -47
- package/benchmarks/results/int8-scan.json +0 -2007
- package/benchmarks/results/int8-scan.md +0 -160
- package/benchmarks/results/int8-w8a8-scan.json +0 -2007
- package/benchmarks/results/int8-w8a8-scan.md +0 -160
- package/benchmarks/results/int8-weight-channel.json +0 -1413
- package/benchmarks/results/int8-weight-channel.md +0 -118
- package/benchmarks/results/int8-weight-only.json +0 -1437
- package/benchmarks/results/int8-weight-only.md +0 -118
- package/benchmarks/results/kernel-webnn-final.json +0 -1115
- package/benchmarks/results/kernel-webnn-final.md +0 -104
- package/benchmarks/results/latest-optimized.json +0 -375
- package/benchmarks/results/latest-optimized.md +0 -47
- package/benchmarks/results/latest.json +0 -391
- package/benchmarks/results/latest.md +0 -47
- package/benchmarks/results/profile-baseline.json +0 -331
- package/benchmarks/results/profile-baseline.md +0 -12
- package/benchmarks/results/profile-conv2x.json +0 -331
- package/benchmarks/results/profile-conv2x.md +0 -12
- package/benchmarks/results/profile-fast-init.json +0 -385
- package/benchmarks/results/profile-fast-init.md +0 -47
- package/benchmarks/results/profile-fp16-fma.json +0 -369
- package/benchmarks/results/profile-fp16-fma.md +0 -47
- package/benchmarks/results/profile-fp16-tiled-decoder.json +0 -369
- package/benchmarks/results/profile-fp16-tiled-decoder.md +0 -47
- package/benchmarks/results/profile-fp16-tiled-encoder.json +0 -369
- package/benchmarks/results/profile-fp16-tiled-encoder.md +0 -47
- package/benchmarks/results/profile-fp16-unfused-pool.json +0 -385
- package/benchmarks/results/profile-fp16-unfused-pool.md +0 -47
- package/benchmarks/results/profile-input-major.json +0 -347
- package/benchmarks/results/profile-input-major.md +0 -12
- package/benchmarks/results/profile-k16.json +0 -347
- package/benchmarks/results/profile-k16.md +0 -12
- package/benchmarks/results/profile-k4.json +0 -347
- package/benchmarks/results/profile-k4.md +0 -12
- package/benchmarks/results/profile-pool-reuse.json +0 -331
- package/benchmarks/results/profile-pool-reuse.md +0 -12
- package/benchmarks/results/profile-precompiled.json +0 -385
- package/benchmarks/results/profile-precompiled.md +0 -47
- package/benchmarks/results/profile-static-channels.json +0 -385
- package/benchmarks/results/profile-static-channels.md +0 -47
- package/benchmarks/results/profile-static-io.json +0 -385
- package/benchmarks/results/profile-static-io.md +0 -47
- package/benchmarks/results/profile-tiled-conv.json +0 -331
- package/benchmarks/results/profile-tiled-conv.md +0 -12
- package/benchmarks/results/profile-tiled-decoder.json +0 -347
- package/benchmarks/results/profile-tiled-decoder.md +0 -12
- package/benchmarks/results/profile-tiled-matmul.json +0 -331
- package/benchmarks/results/profile-tiled-matmul.md +0 -12
- package/benchmarks/results/profile-unfused-decoder.json +0 -379
- package/benchmarks/results/profile-unfused-decoder.md +0 -12
- package/benchmarks/results/profile-unfused-pool.json +0 -347
- package/benchmarks/results/profile-unfused-pool.md +0 -12
- package/benchmarks/results/spatial-auto.json +0 -575
- package/benchmarks/results/spatial-auto.md +0 -61
- package/benchmarks/results/subgroup-smoke.json +0 -1094
- package/benchmarks/results/subgroup-smoke.md +0 -104
- package/benchmarks/results/webnn-smoke.json +0 -739
- package/benchmarks/results/webnn-smoke.md +0 -76
- package/scripts/inspect-model.mjs +0 -64
- package/tests/modelSpec.test.mjs +0 -128
- package/tests/resourceLifecycle.test.mjs +0 -383
- 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
|
-
|
|
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
|
+
}
|
package/src/graphOptimizer.ts
CHANGED
|
@@ -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
|
|
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
|
|
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
|
|
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`
|
|
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
|
|
96
|
+
backendParams?: { device: GPUDevice },
|
|
77
97
|
opts?: UNetOptions
|
|
78
98
|
) {
|
|
79
|
-
const
|
|
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,
|
|
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
|
|
108
|
+
backendParams?: { device: GPUDevice },
|
|
94
109
|
opts?: UNetOptions
|
|
95
110
|
) {
|
|
96
111
|
return fetch(modelPath)
|