oidn-web 0.3.5 → 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 +94 -0
- package/README.md +208 -8
- package/dist/oidn.js +4699 -22516
- package/dist/oidn.umd.cjs +989 -5796
- package/lib/UNet.d.ts +111 -26
- package/lib/UNet.js +310 -329
- package/lib/UNet.js.map +1 -1
- package/lib/WGPUComputePass.d.ts +1 -1
- package/lib/WGPUComputePass.js +6 -4
- package/lib/WGPUComputePass.js.map +1 -1
- package/lib/backend.d.ts +1 -4
- package/lib/backend.js +28 -44
- 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.d.ts +54 -0
- package/lib/graphOptimizer.js +215 -0
- package/lib/graphOptimizer.js.map +1 -0
- 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 +43 -11
- package/lib/main.js +9 -5
- package/lib/main.js.map +1 -1
- package/lib/modelSpec.d.ts +80 -0
- package/lib/modelSpec.js +270 -0
- package/lib/modelSpec.js.map +1 -0
- package/lib/nativeUNet.d.ts +103 -0
- package/lib/nativeUNet.js +2064 -0
- package/lib/nativeUNet.js.map +1 -0
- package/lib/process.d.ts +5 -11
- package/lib/process.js +38 -49
- package/lib/process.js.map +1 -1
- package/lib/resourceTracker.d.ts +26 -0
- package/lib/resourceTracker.js +65 -0
- package/lib/resourceTracker.js.map +1 -0
- package/lib/tileScheduler.d.ts +61 -0
- package/lib/tileScheduler.js +199 -0
- package/lib/tileScheduler.js.map +1 -0
- package/lib/webnnUNet.d.ts +52 -0
- package/lib/webnnUNet.js +535 -0
- package/lib/webnnUNet.js.map +1 -0
- package/package.json +16 -5
- package/src/UNet.ts +463 -437
- package/src/WGPUComputePass.ts +6 -4
- package/src/backend.ts +33 -59
- package/src/finalRgbShader.ts +186 -0
- package/src/graphOptimizer.ts +300 -0
- package/src/hdrTransfer.ts +88 -0
- package/src/main.ts +95 -20
- package/src/modelSpec.ts +414 -0
- package/src/nativeUNet.ts +2655 -0
- package/src/process.ts +46 -71
- package/src/resourceTracker.ts +94 -0
- package/src/tileScheduler.ts +330 -0
- package/src/webnnUNet.ts +812 -0
- package/lib/helper.d.ts +0 -4
- package/lib/helper.js +0 -33
- package/lib/helper.js.map +0 -1
- package/lib/kernels.d.ts +0 -1
- package/lib/kernels.js +0 -26
- package/lib/kernels.js.map +0 -1
- package/src/helper.ts +0 -43
- package/src/kernels.ts +0 -31
package/src/WGPUComputePass.ts
CHANGED
|
@@ -152,13 +152,15 @@ export class WGPUComputePass<I extends string, O extends string> {
|
|
|
152
152
|
return this._outputBuffers[name].buffer;
|
|
153
153
|
}
|
|
154
154
|
|
|
155
|
-
dispose() {
|
|
155
|
+
dispose(destroyOutputBuffers = true) {
|
|
156
156
|
Object.keys(this._uniformBuffers).forEach((key) => {
|
|
157
157
|
(this._uniformBuffers as any)[key].destroy();
|
|
158
158
|
});
|
|
159
|
-
|
|
160
|
-
(this._outputBuffers
|
|
161
|
-
|
|
159
|
+
if (destroyOutputBuffers) {
|
|
160
|
+
Object.keys(this._outputBuffers).forEach((key) => {
|
|
161
|
+
(this._outputBuffers as any)[key].buffer.destroy();
|
|
162
|
+
});
|
|
163
|
+
}
|
|
162
164
|
}
|
|
163
165
|
|
|
164
166
|
private _createBuffer(params: WGPUComputePassOutput) {
|
package/src/backend.ts
CHANGED
|
@@ -1,62 +1,36 @@
|
|
|
1
|
-
import { WebGPUBackend } from '@tensorflow/tfjs-backend-webgpu/dist/base';
|
|
2
|
-
import { ENGINE } from '@tensorflow/tfjs-core/dist/engine';
|
|
3
|
-
|
|
4
|
-
import './kernels';
|
|
5
|
-
|
|
6
1
|
export async function initWebGPUBackend() {
|
|
7
|
-
|
|
8
|
-
|
|
9
|
-
|
|
10
|
-
|
|
11
|
-
|
|
12
|
-
|
|
13
|
-
|
|
14
|
-
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
|
|
18
|
-
|
|
19
|
-
if (adapter.features.has('bgra8unorm-storage')) {
|
|
20
|
-
requiredFeatures.push(['bgra8unorm-storage']);
|
|
21
|
-
}
|
|
22
|
-
deviceDescriptor.requiredFeatures =
|
|
23
|
-
requiredFeatures as Iterable<GPUFeatureName>;
|
|
24
|
-
|
|
25
|
-
const adapterLimits = adapter.limits;
|
|
26
|
-
deviceDescriptor.requiredLimits = {
|
|
27
|
-
maxComputeWorkgroupStorageSize:
|
|
28
|
-
adapterLimits.maxComputeWorkgroupStorageSize,
|
|
29
|
-
maxComputeWorkgroupsPerDimension:
|
|
30
|
-
adapterLimits.maxComputeWorkgroupsPerDimension,
|
|
31
|
-
maxStorageBufferBindingSize: adapterLimits.maxStorageBufferBindingSize,
|
|
32
|
-
maxBufferSize: adapterLimits.maxBufferSize,
|
|
33
|
-
maxComputeWorkgroupSizeX: adapterLimits.maxComputeWorkgroupSizeX,
|
|
34
|
-
maxComputeInvocationsPerWorkgroup:
|
|
35
|
-
adapterLimits.maxComputeInvocationsPerWorkgroup
|
|
36
|
-
};
|
|
37
|
-
const device = await adapter.requestDevice(deviceDescriptor);
|
|
38
|
-
const adapterInfo =
|
|
39
|
-
// requestAdapterInfo is deprecated
|
|
40
|
-
// @ts-ignore
|
|
41
|
-
adapter!.info ?? (await adapter!.requestAdapterInfo?.());
|
|
42
|
-
|
|
43
|
-
return initWebGPUBackendWithDevice(device, adapterInfo);
|
|
44
|
-
} catch (e) {}
|
|
45
|
-
}
|
|
46
|
-
|
|
47
|
-
export async function initWebGPUBackendWithDevice(
|
|
48
|
-
device: GPUDevice,
|
|
49
|
-
adapter: GPUAdapterInfo
|
|
50
|
-
) {
|
|
51
|
-
// TODO multiple device and adapter in one backend
|
|
52
|
-
let backend = ENGINE.findBackend('webgpu-oidn');
|
|
53
|
-
if (backend != null) {
|
|
54
|
-
return backend as WebGPUBackend;
|
|
2
|
+
if (!navigator.gpu) throw new Error('WebGPU is not available');
|
|
3
|
+
const gpuDescriptor: GPURequestAdapterOptions = {
|
|
4
|
+
powerPreference: 'high-performance'
|
|
5
|
+
};
|
|
6
|
+
|
|
7
|
+
const adapter = await navigator.gpu.requestAdapter(gpuDescriptor);
|
|
8
|
+
if (!adapter) throw new Error('No WebGPU adapter is available');
|
|
9
|
+
const deviceDescriptor: GPUDeviceDescriptor = {};
|
|
10
|
+
|
|
11
|
+
const requiredFeatures: GPUFeatureName[] = [];
|
|
12
|
+
if (adapter.features.has('timestamp-query')) {
|
|
13
|
+
requiredFeatures.push('timestamp-query');
|
|
55
14
|
}
|
|
56
|
-
|
|
57
|
-
|
|
58
|
-
|
|
59
|
-
|
|
60
|
-
|
|
61
|
-
|
|
15
|
+
if (adapter.features.has('bgra8unorm-storage')) {
|
|
16
|
+
requiredFeatures.push('bgra8unorm-storage');
|
|
17
|
+
}
|
|
18
|
+
if (adapter.features.has('shader-f16')) {
|
|
19
|
+
requiredFeatures.push('shader-f16');
|
|
20
|
+
}
|
|
21
|
+
deviceDescriptor.requiredFeatures = requiredFeatures;
|
|
22
|
+
|
|
23
|
+
const adapterLimits = adapter.limits;
|
|
24
|
+
deviceDescriptor.requiredLimits = {
|
|
25
|
+
maxComputeWorkgroupStorageSize:
|
|
26
|
+
adapterLimits.maxComputeWorkgroupStorageSize,
|
|
27
|
+
maxComputeWorkgroupsPerDimension:
|
|
28
|
+
adapterLimits.maxComputeWorkgroupsPerDimension,
|
|
29
|
+
maxStorageBufferBindingSize: adapterLimits.maxStorageBufferBindingSize,
|
|
30
|
+
maxBufferSize: adapterLimits.maxBufferSize,
|
|
31
|
+
maxComputeWorkgroupSizeX: adapterLimits.maxComputeWorkgroupSizeX,
|
|
32
|
+
maxComputeInvocationsPerWorkgroup:
|
|
33
|
+
adapterLimits.maxComputeInvocationsPerWorkgroup
|
|
34
|
+
};
|
|
35
|
+
return adapter.requestDevice(deviceDescriptor);
|
|
62
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
|
+
}
|
|
@@ -0,0 +1,300 @@
|
|
|
1
|
+
import type {
|
|
2
|
+
Conv2DNodeSpec,
|
|
3
|
+
MaxPool2DNodeSpec,
|
|
4
|
+
ModelNodeSpec,
|
|
5
|
+
UNetModelSpec,
|
|
6
|
+
Upsample2DNodeSpec,
|
|
7
|
+
UNetModelGraph
|
|
8
|
+
} from './modelSpec';
|
|
9
|
+
|
|
10
|
+
export interface FusedConvPoolNodeSpec {
|
|
11
|
+
op: 'fusedConvReluMaxPool2d';
|
|
12
|
+
id: string;
|
|
13
|
+
input: string;
|
|
14
|
+
conv: Conv2DNodeSpec;
|
|
15
|
+
pool: MaxPool2DNodeSpec;
|
|
16
|
+
}
|
|
17
|
+
|
|
18
|
+
export interface FusedUpsampleConcatConvNodeSpec {
|
|
19
|
+
op: 'fusedUpsampleConcatConv2d';
|
|
20
|
+
id: string;
|
|
21
|
+
/** Inputs stay in concat order because that order selects weight channels. */
|
|
22
|
+
inputs: readonly {
|
|
23
|
+
value: string;
|
|
24
|
+
upsample?: Upsample2DNodeSpec;
|
|
25
|
+
}[];
|
|
26
|
+
conv: Conv2DNodeSpec;
|
|
27
|
+
}
|
|
28
|
+
|
|
29
|
+
export type ExecutableModelNode =
|
|
30
|
+
| ModelNodeSpec
|
|
31
|
+
| FusedConvPoolNodeSpec
|
|
32
|
+
| FusedUpsampleConcatConvNodeSpec;
|
|
33
|
+
|
|
34
|
+
export interface GraphOptimizationOptions {
|
|
35
|
+
fuseConvPool?: boolean;
|
|
36
|
+
fuseUpsampleConcatConv?: boolean;
|
|
37
|
+
}
|
|
38
|
+
|
|
39
|
+
export interface OptimizedModelGraph {
|
|
40
|
+
spec: UNetModelSpec;
|
|
41
|
+
nodes: readonly ExecutableModelNode[];
|
|
42
|
+
fusions: {
|
|
43
|
+
convPool: number;
|
|
44
|
+
upsampleConcatConv: number;
|
|
45
|
+
};
|
|
46
|
+
}
|
|
47
|
+
|
|
48
|
+
function nodeInputs(node: ModelNodeSpec): readonly string[] {
|
|
49
|
+
return node.op === 'concat' ? node.inputs : [node.input];
|
|
50
|
+
}
|
|
51
|
+
|
|
52
|
+
function buildConsumers(spec: UNetModelSpec) {
|
|
53
|
+
const consumers = new Map<string, ModelNodeSpec[]>();
|
|
54
|
+
for (const node of spec.nodes) {
|
|
55
|
+
for (const input of nodeInputs(node)) {
|
|
56
|
+
const list = consumers.get(input) ?? [];
|
|
57
|
+
list.push(node);
|
|
58
|
+
consumers.set(input, list);
|
|
59
|
+
}
|
|
60
|
+
}
|
|
61
|
+
return consumers;
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
function onlyConsumer<T extends ModelNodeSpec['op']>(
|
|
65
|
+
consumers: ReadonlyMap<string, ModelNodeSpec[]>,
|
|
66
|
+
value: string,
|
|
67
|
+
op: T
|
|
68
|
+
): Extract<ModelNodeSpec, { op: T }> | undefined {
|
|
69
|
+
const list = consumers.get(value);
|
|
70
|
+
if (list?.length !== 1 || list[0].op !== op) return undefined;
|
|
71
|
+
return list[0] as Extract<ModelNodeSpec, { op: T }>;
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
/**
|
|
75
|
+
* Applies topology-only fusions. It never depends on a particular OIDN model
|
|
76
|
+
* name, so new descriptors automatically benefit from known graph patterns.
|
|
77
|
+
*/
|
|
78
|
+
export function optimizeModelGraph(
|
|
79
|
+
validated: UNetModelGraph,
|
|
80
|
+
options: GraphOptimizationOptions = {}
|
|
81
|
+
): OptimizedModelGraph {
|
|
82
|
+
const spec = validated.spec;
|
|
83
|
+
const consumers = buildConsumers(spec);
|
|
84
|
+
const nodesById = new Map(spec.nodes.map((node) => [node.id, node]));
|
|
85
|
+
const eliminated = new Set<string>();
|
|
86
|
+
const fusedAt = new Map<string, ExecutableModelNode>();
|
|
87
|
+
let convPool = 0;
|
|
88
|
+
let upsampleConcatConv = 0;
|
|
89
|
+
|
|
90
|
+
if (options.fuseConvPool !== false) {
|
|
91
|
+
for (const node of spec.nodes) {
|
|
92
|
+
if (node.op !== 'conv2d' || node.activation !== 'relu') continue;
|
|
93
|
+
const pool = onlyConsumer(consumers, node.id, 'maxPool2d');
|
|
94
|
+
if (!pool || pool.size !== 2 || pool.stride !== 2) continue;
|
|
95
|
+
|
|
96
|
+
eliminated.add(node.id);
|
|
97
|
+
fusedAt.set(pool.id, {
|
|
98
|
+
op: 'fusedConvReluMaxPool2d',
|
|
99
|
+
id: pool.id,
|
|
100
|
+
input: node.input,
|
|
101
|
+
conv: node,
|
|
102
|
+
pool
|
|
103
|
+
});
|
|
104
|
+
convPool++;
|
|
105
|
+
}
|
|
106
|
+
}
|
|
107
|
+
|
|
108
|
+
if (options.fuseUpsampleConcatConv !== false) {
|
|
109
|
+
for (const node of spec.nodes) {
|
|
110
|
+
if (node.op !== 'conv2d') continue;
|
|
111
|
+
const concat = nodesById.get(node.input);
|
|
112
|
+
if (concat?.op !== 'concat' || concat.inputs.length !== 2) continue;
|
|
113
|
+
if (onlyConsumer(consumers, concat.id, 'conv2d') !== node) continue;
|
|
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
|
+
const firstInputChannels = validated.channelsByValue.get(concat.inputs[0]);
|
|
117
|
+
if (firstInputChannels === undefined || firstInputChannels % 4 !== 0) {
|
|
118
|
+
continue;
|
|
119
|
+
}
|
|
120
|
+
|
|
121
|
+
const inputs = concat.inputs.map((value) => {
|
|
122
|
+
const candidate = nodesById.get(value);
|
|
123
|
+
if (
|
|
124
|
+
candidate?.op === 'upsample2d' &&
|
|
125
|
+
candidate.scale === 2 &&
|
|
126
|
+
candidate.mode === 'nearest' &&
|
|
127
|
+
onlyConsumer(consumers, candidate.id, 'concat') === concat
|
|
128
|
+
) {
|
|
129
|
+
return { value: candidate.input, upsample: candidate };
|
|
130
|
+
}
|
|
131
|
+
return { value };
|
|
132
|
+
});
|
|
133
|
+
const upsampleCount = inputs.filter((input) => input.upsample).length;
|
|
134
|
+
if (upsampleCount !== 1) continue;
|
|
135
|
+
|
|
136
|
+
eliminated.add(concat.id);
|
|
137
|
+
for (const value of concat.inputs) {
|
|
138
|
+
const candidate = nodesById.get(value);
|
|
139
|
+
if (candidate?.op === 'upsample2d') eliminated.add(candidate.id);
|
|
140
|
+
}
|
|
141
|
+
fusedAt.set(node.id, {
|
|
142
|
+
op: 'fusedUpsampleConcatConv2d',
|
|
143
|
+
id: node.id,
|
|
144
|
+
inputs,
|
|
145
|
+
conv: node
|
|
146
|
+
});
|
|
147
|
+
upsampleConcatConv++;
|
|
148
|
+
}
|
|
149
|
+
}
|
|
150
|
+
|
|
151
|
+
const nodes: ExecutableModelNode[] = [];
|
|
152
|
+
for (const node of spec.nodes) {
|
|
153
|
+
const fused = fusedAt.get(node.id);
|
|
154
|
+
if (fused) {
|
|
155
|
+
nodes.push(fused);
|
|
156
|
+
} else if (!eliminated.has(node.id)) {
|
|
157
|
+
nodes.push(node);
|
|
158
|
+
}
|
|
159
|
+
}
|
|
160
|
+
|
|
161
|
+
return {
|
|
162
|
+
spec,
|
|
163
|
+
nodes,
|
|
164
|
+
fusions: { convPool, upsampleConcatConv }
|
|
165
|
+
};
|
|
166
|
+
}
|
|
167
|
+
|
|
168
|
+
export interface ModelValueShape {
|
|
169
|
+
width: number;
|
|
170
|
+
height: number;
|
|
171
|
+
channels: number;
|
|
172
|
+
}
|
|
173
|
+
|
|
174
|
+
export interface PlannedModelNode {
|
|
175
|
+
node: ExecutableModelNode;
|
|
176
|
+
outputShape: ModelValueShape;
|
|
177
|
+
/** Last planned node that reads the output; output itself uses nodes.length. */
|
|
178
|
+
lastUse: number;
|
|
179
|
+
}
|
|
180
|
+
|
|
181
|
+
export interface ModelExecutionPlan extends OptimizedModelGraph {
|
|
182
|
+
inputShape: ModelValueShape;
|
|
183
|
+
valueShapes: ReadonlyMap<string, ModelValueShape>;
|
|
184
|
+
plannedNodes: readonly PlannedModelNode[];
|
|
185
|
+
}
|
|
186
|
+
|
|
187
|
+
function executableInputs(node: ExecutableModelNode): readonly string[] {
|
|
188
|
+
if (node.op === 'concat') return node.inputs;
|
|
189
|
+
if (node.op === 'fusedUpsampleConcatConv2d') {
|
|
190
|
+
return node.inputs.map((input) => input.value);
|
|
191
|
+
}
|
|
192
|
+
return [node.input];
|
|
193
|
+
}
|
|
194
|
+
|
|
195
|
+
function sameSpatialShape(
|
|
196
|
+
left: ModelValueShape,
|
|
197
|
+
right: ModelValueShape
|
|
198
|
+
): boolean {
|
|
199
|
+
return left.width === right.width && left.height === right.height;
|
|
200
|
+
}
|
|
201
|
+
|
|
202
|
+
/** Resolve all runtime shapes and value lifetimes before allocating GPU data. */
|
|
203
|
+
export function planModelExecution(
|
|
204
|
+
validated: UNetModelGraph,
|
|
205
|
+
width: number,
|
|
206
|
+
height: number,
|
|
207
|
+
options?: GraphOptimizationOptions
|
|
208
|
+
): ModelExecutionPlan {
|
|
209
|
+
if (!Number.isInteger(width) || width <= 0 || !Number.isInteger(height) || height <= 0) {
|
|
210
|
+
throw new Error(`Invalid model input size ${width}x${height}`);
|
|
211
|
+
}
|
|
212
|
+
|
|
213
|
+
const graph = optimizeModelGraph(validated, options);
|
|
214
|
+
const inputShape = { width, height, channels: validated.inputChannels };
|
|
215
|
+
const valueShapes = new Map<string, ModelValueShape>([
|
|
216
|
+
[validated.spec.input, inputShape]
|
|
217
|
+
]);
|
|
218
|
+
const outputShapes: ModelValueShape[] = [];
|
|
219
|
+
|
|
220
|
+
const shapeOf = (value: string, nodeId: string) => {
|
|
221
|
+
const shape = valueShapes.get(value);
|
|
222
|
+
if (!shape) throw new Error(`Planned node ${nodeId} reads missing value ${value}`);
|
|
223
|
+
return shape;
|
|
224
|
+
};
|
|
225
|
+
|
|
226
|
+
for (const node of graph.nodes) {
|
|
227
|
+
let outputShape: ModelValueShape;
|
|
228
|
+
if (node.op === 'conv2d') {
|
|
229
|
+
const input = shapeOf(node.input, node.id);
|
|
230
|
+
outputShape = {
|
|
231
|
+
width: input.width,
|
|
232
|
+
height: input.height,
|
|
233
|
+
channels: validated.convChannels.get(node.id)!.outputChannels
|
|
234
|
+
};
|
|
235
|
+
} else if (node.op === 'maxPool2d') {
|
|
236
|
+
const input = shapeOf(node.input, node.id);
|
|
237
|
+
outputShape = {
|
|
238
|
+
width: Math.ceil(input.width / 2),
|
|
239
|
+
height: Math.ceil(input.height / 2),
|
|
240
|
+
channels: input.channels
|
|
241
|
+
};
|
|
242
|
+
} else if (node.op === 'upsample2d') {
|
|
243
|
+
const input = shapeOf(node.input, node.id);
|
|
244
|
+
outputShape = {
|
|
245
|
+
width: input.width * 2,
|
|
246
|
+
height: input.height * 2,
|
|
247
|
+
channels: input.channels
|
|
248
|
+
};
|
|
249
|
+
} else if (node.op === 'concat') {
|
|
250
|
+
const inputs = node.inputs.map((value) => shapeOf(value, node.id));
|
|
251
|
+
if (inputs.some((shape) => !sameSpatialShape(shape, inputs[0]))) {
|
|
252
|
+
throw new Error(`Concat ${node.id} has mismatched spatial shapes`);
|
|
253
|
+
}
|
|
254
|
+
outputShape = {
|
|
255
|
+
width: inputs[0].width,
|
|
256
|
+
height: inputs[0].height,
|
|
257
|
+
channels: inputs.reduce((sum, shape) => sum + shape.channels, 0)
|
|
258
|
+
};
|
|
259
|
+
} else if (node.op === 'fusedConvReluMaxPool2d') {
|
|
260
|
+
const input = shapeOf(node.input, node.id);
|
|
261
|
+
outputShape = {
|
|
262
|
+
width: Math.ceil(input.width / 2),
|
|
263
|
+
height: Math.ceil(input.height / 2),
|
|
264
|
+
channels: validated.convChannels.get(node.conv.id)!.outputChannels
|
|
265
|
+
};
|
|
266
|
+
} else {
|
|
267
|
+
const inputs = node.inputs.map((input) => {
|
|
268
|
+
const source = shapeOf(input.value, node.id);
|
|
269
|
+
return input.upsample
|
|
270
|
+
? { ...source, width: source.width * 2, height: source.height * 2 }
|
|
271
|
+
: source;
|
|
272
|
+
});
|
|
273
|
+
if (inputs.some((shape) => !sameSpatialShape(shape, inputs[0]))) {
|
|
274
|
+
throw new Error(`Fused decoder ${node.id} has mismatched spatial shapes`);
|
|
275
|
+
}
|
|
276
|
+
outputShape = {
|
|
277
|
+
width: inputs[0].width,
|
|
278
|
+
height: inputs[0].height,
|
|
279
|
+
channels: validated.convChannels.get(node.conv.id)!.outputChannels
|
|
280
|
+
};
|
|
281
|
+
}
|
|
282
|
+
|
|
283
|
+
valueShapes.set(node.id, outputShape);
|
|
284
|
+
outputShapes.push(outputShape);
|
|
285
|
+
}
|
|
286
|
+
|
|
287
|
+
const lastUses = new Map<string, number>();
|
|
288
|
+
graph.nodes.forEach((node, index) => {
|
|
289
|
+
for (const input of executableInputs(node)) lastUses.set(input, index);
|
|
290
|
+
});
|
|
291
|
+
lastUses.set(validated.spec.output, graph.nodes.length);
|
|
292
|
+
|
|
293
|
+
const plannedNodes = graph.nodes.map((node, index) => ({
|
|
294
|
+
node,
|
|
295
|
+
outputShape: outputShapes[index],
|
|
296
|
+
lastUse: lastUses.get(node.id) ?? index
|
|
297
|
+
}));
|
|
298
|
+
|
|
299
|
+
return { ...graph, inputShape, valueShapes, plannedNodes };
|
|
300
|
+
}
|
|
@@ -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
|
+
}
|