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.
Files changed (65) hide show
  1. package/CHANGELOG.md +94 -0
  2. package/README.md +208 -8
  3. package/dist/oidn.js +4699 -22516
  4. package/dist/oidn.umd.cjs +989 -5796
  5. package/lib/UNet.d.ts +111 -26
  6. package/lib/UNet.js +310 -329
  7. package/lib/UNet.js.map +1 -1
  8. package/lib/WGPUComputePass.d.ts +1 -1
  9. package/lib/WGPUComputePass.js +6 -4
  10. package/lib/WGPUComputePass.js.map +1 -1
  11. package/lib/backend.d.ts +1 -4
  12. package/lib/backend.js +28 -44
  13. package/lib/backend.js.map +1 -1
  14. package/lib/finalRgbShader.d.ts +13 -0
  15. package/lib/finalRgbShader.js +160 -0
  16. package/lib/finalRgbShader.js.map +1 -0
  17. package/lib/graphOptimizer.d.ts +54 -0
  18. package/lib/graphOptimizer.js +215 -0
  19. package/lib/graphOptimizer.js.map +1 -0
  20. package/lib/hdrTransfer.d.ts +14 -0
  21. package/lib/hdrTransfer.js +61 -0
  22. package/lib/hdrTransfer.js.map +1 -0
  23. package/lib/main.d.ts +43 -11
  24. package/lib/main.js +9 -5
  25. package/lib/main.js.map +1 -1
  26. package/lib/modelSpec.d.ts +80 -0
  27. package/lib/modelSpec.js +270 -0
  28. package/lib/modelSpec.js.map +1 -0
  29. package/lib/nativeUNet.d.ts +103 -0
  30. package/lib/nativeUNet.js +2064 -0
  31. package/lib/nativeUNet.js.map +1 -0
  32. package/lib/process.d.ts +5 -11
  33. package/lib/process.js +38 -49
  34. package/lib/process.js.map +1 -1
  35. package/lib/resourceTracker.d.ts +26 -0
  36. package/lib/resourceTracker.js +65 -0
  37. package/lib/resourceTracker.js.map +1 -0
  38. package/lib/tileScheduler.d.ts +61 -0
  39. package/lib/tileScheduler.js +199 -0
  40. package/lib/tileScheduler.js.map +1 -0
  41. package/lib/webnnUNet.d.ts +52 -0
  42. package/lib/webnnUNet.js +535 -0
  43. package/lib/webnnUNet.js.map +1 -0
  44. package/package.json +16 -5
  45. package/src/UNet.ts +463 -437
  46. package/src/WGPUComputePass.ts +6 -4
  47. package/src/backend.ts +33 -59
  48. package/src/finalRgbShader.ts +186 -0
  49. package/src/graphOptimizer.ts +300 -0
  50. package/src/hdrTransfer.ts +88 -0
  51. package/src/main.ts +95 -20
  52. package/src/modelSpec.ts +414 -0
  53. package/src/nativeUNet.ts +2655 -0
  54. package/src/process.ts +46 -71
  55. package/src/resourceTracker.ts +94 -0
  56. package/src/tileScheduler.ts +330 -0
  57. package/src/webnnUNet.ts +812 -0
  58. package/lib/helper.d.ts +0 -4
  59. package/lib/helper.js +0 -33
  60. package/lib/helper.js.map +0 -1
  61. package/lib/kernels.d.ts +0 -1
  62. package/lib/kernels.js +0 -26
  63. package/lib/kernels.js.map +0 -1
  64. package/src/helper.ts +0 -43
  65. package/src/kernels.ts +0 -31
@@ -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
- Object.keys(this._outputBuffers).forEach((key) => {
160
- (this._outputBuffers as any)[key].buffer.destroy();
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
- try {
8
- const gpuDescriptor: GPURequestAdapterOptions = {
9
- powerPreference: 'high-performance'
10
- };
11
-
12
- const adapter = (await navigator.gpu.requestAdapter(gpuDescriptor))!;
13
- const deviceDescriptor: GPUDeviceDescriptor = {};
14
-
15
- const requiredFeatures = [];
16
- if (adapter.features.has('timestamp-query')) {
17
- requiredFeatures.push('timestamp-query');
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
- backend = new WebGPUBackend(device, adapter);
58
- ENGINE.registerBackend('webgpu-oidn', () => backend);
59
- await ENGINE.setBackend('webgpu-oidn');
60
-
61
- return backend as WebGPUBackend;
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
+ }