oidn-web 0.3.4 → 0.4.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 (122) hide show
  1. package/CHANGELOG.md +49 -0
  2. package/README.md +140 -4
  3. package/benchmarks/compare.mjs +651 -0
  4. package/benchmarks/leak.mjs +255 -0
  5. package/benchmarks/results/before-spatial.json +391 -0
  6. package/benchmarks/results/before-spatial.md +47 -0
  7. package/benchmarks/results/int8-scan.json +2007 -0
  8. package/benchmarks/results/int8-scan.md +160 -0
  9. package/benchmarks/results/int8-w8a8-scan.json +2007 -0
  10. package/benchmarks/results/int8-w8a8-scan.md +160 -0
  11. package/benchmarks/results/int8-weight-channel.json +1413 -0
  12. package/benchmarks/results/int8-weight-channel.md +118 -0
  13. package/benchmarks/results/int8-weight-only.json +1437 -0
  14. package/benchmarks/results/int8-weight-only.md +118 -0
  15. package/benchmarks/results/kernel-webnn-final.json +1115 -0
  16. package/benchmarks/results/kernel-webnn-final.md +104 -0
  17. package/benchmarks/results/latest-optimized.json +375 -0
  18. package/benchmarks/results/latest-optimized.md +47 -0
  19. package/benchmarks/results/latest.json +391 -0
  20. package/benchmarks/results/latest.md +47 -0
  21. package/benchmarks/results/profile-baseline.json +331 -0
  22. package/benchmarks/results/profile-baseline.md +12 -0
  23. package/benchmarks/results/profile-conv2x.json +331 -0
  24. package/benchmarks/results/profile-conv2x.md +12 -0
  25. package/benchmarks/results/profile-fast-init.json +385 -0
  26. package/benchmarks/results/profile-fast-init.md +47 -0
  27. package/benchmarks/results/profile-fp16-fma.json +369 -0
  28. package/benchmarks/results/profile-fp16-fma.md +47 -0
  29. package/benchmarks/results/profile-fp16-tiled-decoder.json +369 -0
  30. package/benchmarks/results/profile-fp16-tiled-decoder.md +47 -0
  31. package/benchmarks/results/profile-fp16-tiled-encoder.json +369 -0
  32. package/benchmarks/results/profile-fp16-tiled-encoder.md +47 -0
  33. package/benchmarks/results/profile-fp16-unfused-pool.json +385 -0
  34. package/benchmarks/results/profile-fp16-unfused-pool.md +47 -0
  35. package/benchmarks/results/profile-input-major.json +347 -0
  36. package/benchmarks/results/profile-input-major.md +12 -0
  37. package/benchmarks/results/profile-k16.json +347 -0
  38. package/benchmarks/results/profile-k16.md +12 -0
  39. package/benchmarks/results/profile-k4.json +347 -0
  40. package/benchmarks/results/profile-k4.md +12 -0
  41. package/benchmarks/results/profile-pool-reuse.json +331 -0
  42. package/benchmarks/results/profile-pool-reuse.md +12 -0
  43. package/benchmarks/results/profile-precompiled.json +385 -0
  44. package/benchmarks/results/profile-precompiled.md +47 -0
  45. package/benchmarks/results/profile-static-channels.json +385 -0
  46. package/benchmarks/results/profile-static-channels.md +47 -0
  47. package/benchmarks/results/profile-static-io.json +385 -0
  48. package/benchmarks/results/profile-static-io.md +47 -0
  49. package/benchmarks/results/profile-tiled-conv.json +331 -0
  50. package/benchmarks/results/profile-tiled-conv.md +12 -0
  51. package/benchmarks/results/profile-tiled-decoder.json +347 -0
  52. package/benchmarks/results/profile-tiled-decoder.md +12 -0
  53. package/benchmarks/results/profile-tiled-matmul.json +331 -0
  54. package/benchmarks/results/profile-tiled-matmul.md +12 -0
  55. package/benchmarks/results/profile-unfused-decoder.json +379 -0
  56. package/benchmarks/results/profile-unfused-decoder.md +12 -0
  57. package/benchmarks/results/profile-unfused-pool.json +347 -0
  58. package/benchmarks/results/profile-unfused-pool.md +12 -0
  59. package/benchmarks/results/spatial-auto.json +575 -0
  60. package/benchmarks/results/spatial-auto.md +61 -0
  61. package/benchmarks/results/subgroup-smoke.json +1094 -0
  62. package/benchmarks/results/subgroup-smoke.md +104 -0
  63. package/benchmarks/results/webnn-smoke.json +739 -0
  64. package/benchmarks/results/webnn-smoke.md +76 -0
  65. package/dist/oidn.js +4189 -22603
  66. package/dist/oidn.umd.cjs +784 -5807
  67. package/lib/UNet.d.ts +66 -15
  68. package/lib/UNet.js +162 -257
  69. package/lib/UNet.js.map +1 -1
  70. package/lib/WGPUComputePass.d.ts +1 -1
  71. package/lib/WGPUComputePass.js +6 -4
  72. package/lib/WGPUComputePass.js.map +1 -1
  73. package/lib/backend.d.ts +8 -4
  74. package/lib/backend.js +36 -44
  75. package/lib/backend.js.map +1 -1
  76. package/lib/graphOptimizer.d.ts +54 -0
  77. package/lib/graphOptimizer.js +216 -0
  78. package/lib/graphOptimizer.js.map +1 -0
  79. package/lib/main.d.ts +33 -10
  80. package/lib/main.js +5 -0
  81. package/lib/main.js.map +1 -1
  82. package/lib/modelSpec.d.ts +80 -0
  83. package/lib/modelSpec.js +270 -0
  84. package/lib/modelSpec.js.map +1 -0
  85. package/lib/nativeUNet.d.ts +67 -0
  86. package/lib/nativeUNet.js +1735 -0
  87. package/lib/nativeUNet.js.map +1 -0
  88. package/lib/process.js +38 -35
  89. package/lib/process.js.map +1 -1
  90. package/lib/resourceTracker.d.ts +26 -0
  91. package/lib/resourceTracker.js +65 -0
  92. package/lib/resourceTracker.js.map +1 -0
  93. package/lib/tileScheduler.d.ts +33 -0
  94. package/lib/tileScheduler.js +86 -0
  95. package/lib/tileScheduler.js.map +1 -0
  96. package/lib/webnnUNet.d.ts +52 -0
  97. package/lib/webnnUNet.js +535 -0
  98. package/lib/webnnUNet.js.map +1 -0
  99. package/package.json +9 -5
  100. package/scripts/inspect-model.mjs +64 -0
  101. package/src/UNet.ts +236 -339
  102. package/src/WGPUComputePass.ts +6 -4
  103. package/src/backend.ts +42 -55
  104. package/src/graphOptimizer.ts +301 -0
  105. package/src/main.ts +71 -11
  106. package/src/modelSpec.ts +414 -0
  107. package/src/nativeUNet.ts +2256 -0
  108. package/src/process.ts +38 -36
  109. package/src/resourceTracker.ts +94 -0
  110. package/src/tileScheduler.ts +138 -0
  111. package/src/webnnUNet.ts +812 -0
  112. package/tests/modelSpec.test.mjs +128 -0
  113. package/tests/resourceLifecycle.test.mjs +383 -0
  114. package/tests/tileScheduler.test.mjs +90 -0
  115. package/lib/helper.d.ts +0 -4
  116. package/lib/helper.js +0 -33
  117. package/lib/helper.js.map +0 -1
  118. package/lib/kernels.d.ts +0 -1
  119. package/lib/kernels.js +0 -26
  120. package/lib/kernels.js.map +0 -1
  121. package/src/helper.ts +0 -43
  122. package/src/kernels.ts +0 -31
@@ -0,0 +1,2256 @@
1
+ import { Float16Array } from '@petamoriken/float16';
2
+ import {
3
+ optimizeModelGraph,
4
+ planModelExecution,
5
+ type ExecutableModelNode,
6
+ type ModelExecutionPlan,
7
+ type ModelValueShape
8
+ } from './graphOptimizer.js';
9
+ import type {
10
+ Conv2DNodeSpec,
11
+ UNetModelGraph,
12
+ ValidatedConvTensor,
13
+ ValidatedUNetModel
14
+ } from './modelSpec';
15
+ import type { HostTensor } from './tza';
16
+ import {
17
+ OIDNResourceTracker,
18
+ type OIDNResourceSnapshot
19
+ } from './resourceTracker.js';
20
+
21
+ export type NativeUNetPrecision = 'fp32' | 'fp16';
22
+ export type NativeUNetPrecisionSetting = NativeUNetPrecision | 'auto';
23
+ export type NativeUNetKernel =
24
+ | 'direct'
25
+ | 'implicit-gemm'
26
+ | 'spatial'
27
+ | 'subgroup';
28
+ export type NativeUNetKernelSetting = NativeUNetKernel | 'auto';
29
+
30
+ export interface NativeUNetOptions {
31
+ precision?: NativeUNetPrecisionSetting;
32
+ /**
33
+ * Convolution kernel selection. `auto` uses a model-independent capability
34
+ * heuristic and falls back to the direct kernel when a tile does not fit.
35
+ */
36
+ kernel?: NativeUNetKernelSetting;
37
+ /** Maximum number of shape-dependent activation plans retained. */
38
+ shapeCacheSize?: number;
39
+ }
40
+
41
+ export interface NativeUNetLayerTiming {
42
+ id: string;
43
+ durationMs: number;
44
+ }
45
+
46
+ export interface NativeUNetExecutionProfile {
47
+ totalMs: number;
48
+ layers: NativeUNetLayerTiming[];
49
+ }
50
+
51
+ interface PackedConvBuffers {
52
+ weights: GPUBuffer;
53
+ bias: GPUBuffer;
54
+ }
55
+
56
+ interface NativePipelineSpec {
57
+ key: string;
58
+ code: string;
59
+ kernel?: NativeUNetKernel;
60
+ }
61
+
62
+ interface ActivationSlot {
63
+ buffer: GPUBuffer;
64
+ capacity: number;
65
+ activeValue?: string;
66
+ }
67
+
68
+ interface CachedExecution {
69
+ plan: ModelExecutionPlan;
70
+ valueBuffers: Map<string, GPUBuffer>;
71
+ slots: ActivationSlot[];
72
+ nodeBindings: GPUBindGroup[];
73
+ nodePipelines: GPUComputePipeline[];
74
+ nodeKernels: NativeUNetKernel[];
75
+ inputPipeline: GPUComputePipeline;
76
+ inputUniform: GPUBuffer;
77
+ ownedBuffers: GPUBuffer[];
78
+ cpuInputBuffers?: GPUBuffer[];
79
+ cpuReadbackBuffer?: GPUBuffer;
80
+ lastUsed: number;
81
+ }
82
+
83
+ interface SharedPipelineCache {
84
+ ready: Map<string, GPUComputePipeline>;
85
+ pending: Map<string, Promise<GPUComputePipeline>>;
86
+ }
87
+
88
+ const pipelineCachesByDevice = new WeakMap<GPUDevice, SharedPipelineCache>();
89
+
90
+ function sharedPipelineCache(device: GPUDevice) {
91
+ let cache = pipelineCachesByDevice.get(device);
92
+ if (!cache) {
93
+ cache = { ready: new Map(), pending: new Map() };
94
+ pipelineCachesByDevice.set(device, cache);
95
+ }
96
+ return cache;
97
+ }
98
+
99
+ const WORKGROUP_SIZE = 8;
100
+ const TILED_CONV_WORKGROUP = 8;
101
+ const TILED_CONV_ROWS_PER_THREAD = 4;
102
+ const TILED_CONV_M =
103
+ TILED_CONV_WORKGROUP * TILED_CONV_ROWS_PER_THREAD;
104
+ const TILED_CONV_N_BLOCKS = TILED_CONV_WORKGROUP;
105
+ const TILED_CONV_K_BLOCKS = 8;
106
+ const SPATIAL_CONV_WORKGROUP = 8;
107
+ const SPATIAL_CONV_PATCH = SPATIAL_CONV_WORKGROUP + 2;
108
+
109
+ function roundUp(value: number, alignment: number) {
110
+ return Math.ceil(value / alignment) * alignment;
111
+ }
112
+
113
+ function blocksForChannels(channels: number) {
114
+ return Math.ceil(channels / 4);
115
+ }
116
+
117
+ function activationByteSize(
118
+ shape: ModelValueShape,
119
+ bytesPerScalar: number
120
+ ) {
121
+ return (
122
+ shape.width *
123
+ shape.height *
124
+ blocksForChannels(shape.channels) *
125
+ 4 *
126
+ bytesPerScalar
127
+ );
128
+ }
129
+
130
+ let float16ToFloat32Lookup: Float32Array | undefined;
131
+
132
+ function halfBitsToNumber(bits: number) {
133
+ const sign = bits & 0x8000 ? -1 : 1;
134
+ const exponent = (bits >>> 10) & 0x1f;
135
+ const fraction = bits & 0x3ff;
136
+ if (exponent === 0) {
137
+ return sign * fraction * 2 ** -24;
138
+ }
139
+ if (exponent === 0x1f) {
140
+ return fraction === 0 ? sign * Infinity : NaN;
141
+ }
142
+ return sign * (1 + fraction / 1024) * 2 ** (exponent - 15);
143
+ }
144
+
145
+ function halfLookup() {
146
+ if (!float16ToFloat32Lookup) {
147
+ float16ToFloat32Lookup = new Float32Array(1 << 16);
148
+ for (let bits = 0; bits < float16ToFloat32Lookup.length; bits++) {
149
+ float16ToFloat32Lookup[bits] = halfBitsToNumber(bits);
150
+ }
151
+ }
152
+ return float16ToFloat32Lookup;
153
+ }
154
+
155
+ function tensorFloat32Values(tensor: HostTensor): Float32Array {
156
+ if (tensor.desc.dataType === 'Float32') {
157
+ return new Float32Array(
158
+ tensor.data.buffer,
159
+ tensor.data.byteOffset,
160
+ tensor.data.byteLength / 4
161
+ );
162
+ }
163
+ const bits = new Uint16Array(
164
+ tensor.data.buffer,
165
+ tensor.data.byteOffset,
166
+ tensor.data.byteLength / 2
167
+ );
168
+ const values = new Float32Array(bits.length);
169
+ if (bits.length < 4096) {
170
+ for (let index = 0; index < bits.length; index++) {
171
+ values[index] = halfBitsToNumber(bits[index]);
172
+ }
173
+ } else {
174
+ const lookup = halfLookup();
175
+ for (let index = 0; index < bits.length; index++) {
176
+ values[index] = lookup[bits[index]];
177
+ }
178
+ }
179
+ return values;
180
+ }
181
+
182
+ function createMappedBuffer(
183
+ device: GPUDevice,
184
+ label: string,
185
+ data: ArrayBufferView,
186
+ usage: GPUBufferUsageFlags
187
+ ) {
188
+ const size = roundUp(data.byteLength, 4);
189
+ const buffer = device.createBuffer({
190
+ label,
191
+ size,
192
+ usage,
193
+ mappedAtCreation: true
194
+ });
195
+ new Uint8Array(buffer.getMappedRange()).set(
196
+ new Uint8Array(data.buffer, data.byteOffset, data.byteLength)
197
+ );
198
+ buffer.unmap();
199
+ return buffer;
200
+ }
201
+
202
+ function createUniformBuffer(
203
+ device: GPUDevice,
204
+ label: string,
205
+ values: readonly number[]
206
+ ) {
207
+ const data = new Uint32Array(roundUp(values.length, 4));
208
+ data.set(values);
209
+ return createMappedBuffer(device, label, data, GPUBufferUsage.UNIFORM);
210
+ }
211
+
212
+ function packConvTensors(
213
+ device: GPUDevice,
214
+ id: string,
215
+ tensors: ValidatedConvTensor,
216
+ precision: NativeUNetPrecision
217
+ ): PackedConvBuffers {
218
+ const inputBlocks = blocksForChannels(tensors.inputChannels);
219
+ const outputBlocks = blocksForChannels(tensors.outputChannels);
220
+ const packedWeightCount =
221
+ outputBlocks *
222
+ tensors.kernelHeight *
223
+ tensors.kernelWidth *
224
+ inputBlocks *
225
+ 4 *
226
+ 4;
227
+ const canCopyHalfBits =
228
+ precision === 'fp16' && tensors.weight.desc.dataType === 'Float16';
229
+ const packedWeights = canCopyHalfBits
230
+ ? new Uint16Array(packedWeightCount)
231
+ : precision === 'fp16'
232
+ ? new Float16Array(packedWeightCount)
233
+ : new Float32Array(packedWeightCount);
234
+ const sourceWeights = canCopyHalfBits
235
+ ? new Uint16Array(
236
+ tensors.weight.data.buffer,
237
+ tensors.weight.data.byteOffset,
238
+ tensors.weight.data.byteLength / 2
239
+ )
240
+ : tensorFloat32Values(tensors.weight);
241
+
242
+ for (let outputBlock = 0; outputBlock < outputBlocks; outputBlock++) {
243
+ for (let y = 0; y < tensors.kernelHeight; y++) {
244
+ for (let x = 0; x < tensors.kernelWidth; x++) {
245
+ for (let inputBlock = 0; inputBlock < inputBlocks; inputBlock++) {
246
+ for (let outputLane = 0; outputLane < 4; outputLane++) {
247
+ const outputChannel = outputBlock * 4 + outputLane;
248
+ for (let inputLane = 0; inputLane < 4; inputLane++) {
249
+ const inputChannel = inputBlock * 4 + inputLane;
250
+ const packedBlock =
251
+ ((((outputBlock * tensors.kernelHeight + y) *
252
+ tensors.kernelWidth +
253
+ x) *
254
+ inputBlocks +
255
+ inputBlock) *
256
+ 16);
257
+ // One vec4 contains the four output lanes for an input lane.
258
+ // Direct and tiled kernels share this layout, keeping model
259
+ // descriptors independent from the selected kernel.
260
+ const packedIndex =
261
+ packedBlock + inputLane * 4 + outputLane;
262
+ if (
263
+ outputChannel < tensors.outputChannels &&
264
+ inputChannel < tensors.inputChannels
265
+ ) {
266
+ const sourceIndex =
267
+ ((outputChannel * tensors.inputChannels + inputChannel) *
268
+ tensors.kernelHeight +
269
+ y) *
270
+ tensors.kernelWidth +
271
+ x;
272
+ packedWeights[packedIndex] = sourceWeights[sourceIndex];
273
+ }
274
+ }
275
+ }
276
+ }
277
+ }
278
+ }
279
+ }
280
+
281
+ // Bias stays f32 even for half activations because convolution accumulates
282
+ // in f32. Padded lanes are zero and never escape the final three channels.
283
+ const packedBias = new Float32Array(outputBlocks * 4);
284
+ packedBias.set(tensorFloat32Values(tensors.bias));
285
+
286
+ const weights = createMappedBuffer(
287
+ device,
288
+ `oidn/${id}/weights/${precision}`,
289
+ packedWeights,
290
+ GPUBufferUsage.STORAGE
291
+ );
292
+ try {
293
+ return {
294
+ weights,
295
+ bias: createMappedBuffer(
296
+ device,
297
+ `oidn/${id}/bias`,
298
+ packedBias,
299
+ GPUBufferUsage.STORAGE
300
+ )
301
+ };
302
+ } catch (error) {
303
+ weights.destroy();
304
+ throw error;
305
+ }
306
+ }
307
+
308
+ function storageVecType(precision: NativeUNetPrecision) {
309
+ return precision === 'fp16' ? 'vec4<f16>' : 'vec4<f32>';
310
+ }
311
+
312
+ function shaderPreamble(precision: NativeUNetPrecision) {
313
+ return precision === 'fp16' ? 'enable f16;\n' : '';
314
+ }
315
+
316
+ function storeExpression(
317
+ expression: string,
318
+ outputPrecision: NativeUNetPrecision
319
+ ) {
320
+ return outputPrecision === 'fp16'
321
+ ? `vec4<f16>(${expression})`
322
+ : expression;
323
+ }
324
+
325
+ function activationExpression(
326
+ expression: string,
327
+ activation: Conv2DNodeSpec['activation']
328
+ ) {
329
+ return activation === 'relu'
330
+ ? `max(${expression}, vec4<f32>(0.0))`
331
+ : expression;
332
+ }
333
+
334
+ function accumulationCode(
335
+ inputExpression: string,
336
+ weightBase: string,
337
+ precision: NativeUNetPrecision
338
+ ) {
339
+ if (precision === 'fp32') {
340
+ return /* wgsl */ `
341
+ let inputValue = vec4<f32>(${inputExpression});
342
+ let weightBase = ${weightBase};
343
+ acc = fma(vec4<f32>(weights[weightBase]), vec4<f32>(inputValue.x), acc);
344
+ acc = fma(vec4<f32>(weights[weightBase + 1u]), vec4<f32>(inputValue.y), acc);
345
+ acc = fma(vec4<f32>(weights[weightBase + 2u]), vec4<f32>(inputValue.z), acc);
346
+ acc = fma(vec4<f32>(weights[weightBase + 3u]), vec4<f32>(inputValue.w), acc);
347
+ `;
348
+ }
349
+ return /* wgsl */ `
350
+ let inputValue = vec4<f16>(${inputExpression});
351
+ let weightBase = ${weightBase};
352
+ var partial = vec4<f16>(0.0h);
353
+ partial = fma(weights[weightBase], vec4<f16>(inputValue.x), partial);
354
+ partial = fma(weights[weightBase + 1u], vec4<f16>(inputValue.y), partial);
355
+ partial = fma(weights[weightBase + 2u], vec4<f16>(inputValue.z), partial);
356
+ partial = fma(weights[weightBase + 3u], vec4<f16>(inputValue.w), partial);
357
+ acc += vec4<f32>(partial);
358
+ `;
359
+ }
360
+
361
+ function subgroupAccumulationCode(
362
+ inputExpression: string,
363
+ weightBase: string,
364
+ precision: NativeUNetPrecision
365
+ ) {
366
+ const inputType = precision === 'fp16' ? 'vec4<f16>' : 'vec4<f32>';
367
+ const accumulator = precision === 'fp16'
368
+ ? `var partial = vec4<f16>(0.0h);
369
+ partial = fma(subgroupBroadcast(weights[weightBase], 0u), ${inputType}(inputValue.x), partial);
370
+ partial = fma(subgroupBroadcast(weights[weightBase + 1u], 0u), ${inputType}(inputValue.y), partial);
371
+ partial = fma(subgroupBroadcast(weights[weightBase + 2u], 0u), ${inputType}(inputValue.z), partial);
372
+ partial = fma(subgroupBroadcast(weights[weightBase + 3u], 0u), ${inputType}(inputValue.w), partial);
373
+ acc += vec4<f32>(partial);`
374
+ : `acc = fma(subgroupBroadcast(weights[weightBase], 0u), vec4<f32>(inputValue.x), acc);
375
+ acc = fma(subgroupBroadcast(weights[weightBase + 1u], 0u), vec4<f32>(inputValue.y), acc);
376
+ acc = fma(subgroupBroadcast(weights[weightBase + 2u], 0u), vec4<f32>(inputValue.z), acc);
377
+ acc = fma(subgroupBroadcast(weights[weightBase + 3u], 0u), vec4<f32>(inputValue.w), acc);`;
378
+ return /* wgsl */ `
379
+ let inputValue = ${inputType}(${inputExpression});
380
+ let weightBase = ${weightBase};
381
+ ${accumulator}
382
+ `;
383
+ }
384
+
385
+ function tiledAccumulationCode(precision: NativeUNetPrecision) {
386
+ if (precision === 'fp16') {
387
+ return /* wgsl */ `
388
+ var partial: array<vec4<f16>, ${TILED_CONV_ROWS_PER_THREAD}>;
389
+ for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
390
+ partial[row] = vec4<f16>(0.0h);
391
+ }
392
+ for (var tileK = 0u; tileK < ${TILED_CONV_K_BLOCKS}u; tileK++) {
393
+ let weightBase =
394
+ (tileK * ${TILED_CONV_N_BLOCKS}u + localId.x) * 4u;
395
+ for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
396
+ let tileSpatial =
397
+ localId.y * ${TILED_CONV_ROWS_PER_THREAD}u + row;
398
+ let inputValue =
399
+ inputTile[tileSpatial * ${TILED_CONV_K_BLOCKS}u + tileK];
400
+ partial[row] = fma(
401
+ weightTile[weightBase],
402
+ vec4<f16>(inputValue.x),
403
+ partial[row]
404
+ );
405
+ partial[row] = fma(
406
+ weightTile[weightBase + 1u],
407
+ vec4<f16>(inputValue.y),
408
+ partial[row]
409
+ );
410
+ partial[row] = fma(
411
+ weightTile[weightBase + 2u],
412
+ vec4<f16>(inputValue.z),
413
+ partial[row]
414
+ );
415
+ partial[row] = fma(
416
+ weightTile[weightBase + 3u],
417
+ vec4<f16>(inputValue.w),
418
+ partial[row]
419
+ );
420
+ }
421
+ }
422
+ for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
423
+ acc[row] += vec4<f32>(partial[row]);
424
+ }
425
+ `;
426
+ }
427
+ return /* wgsl */ `
428
+ for (var tileK = 0u; tileK < ${TILED_CONV_K_BLOCKS}u; tileK++) {
429
+ let weightBase =
430
+ (tileK * ${TILED_CONV_N_BLOCKS}u + localId.x) * 4u;
431
+ for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
432
+ let tileSpatial =
433
+ localId.y * ${TILED_CONV_ROWS_PER_THREAD}u + row;
434
+ let inputValue =
435
+ inputTile[tileSpatial * ${TILED_CONV_K_BLOCKS}u + tileK];
436
+ acc[row] = fma(
437
+ weightTile[weightBase],
438
+ vec4<f32>(inputValue.x),
439
+ acc[row]
440
+ );
441
+ acc[row] = fma(
442
+ weightTile[weightBase + 1u],
443
+ vec4<f32>(inputValue.y),
444
+ acc[row]
445
+ );
446
+ acc[row] = fma(
447
+ weightTile[weightBase + 2u],
448
+ vec4<f32>(inputValue.z),
449
+ acc[row]
450
+ );
451
+ acc[row] = fma(
452
+ weightTile[weightBase + 3u],
453
+ vec4<f32>(inputValue.w),
454
+ acc[row]
455
+ );
456
+ }
457
+ }
458
+ `;
459
+ }
460
+
461
+ function createConvShader(
462
+ precision: NativeUNetPrecision,
463
+ outputPrecision: NativeUNetPrecision,
464
+ activation: Conv2DNodeSpec['activation'],
465
+ inputBlocks: number,
466
+ outputBlocks: number
467
+ ) {
468
+ const inputType = storageVecType(precision);
469
+ const weightType = storageVecType(precision);
470
+ const outputType = storageVecType(outputPrecision);
471
+ const stored = storeExpression(
472
+ activationExpression('acc', activation),
473
+ outputPrecision
474
+ );
475
+
476
+ return /* wgsl */ `${shaderPreamble(precision)}
477
+ struct Params {
478
+ inputWidth: u32,
479
+ inputHeight: u32,
480
+ outputWidth: u32,
481
+ outputHeight: u32,
482
+ inputBlocks: u32,
483
+ outputBlocks: u32,
484
+ }
485
+
486
+ @group(0) @binding(0) var<storage, read> inputData: array<${inputType}>;
487
+ @group(0) @binding(1) var<storage, read> weights: array<${weightType}>;
488
+ @group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
489
+ @group(0) @binding(3) var<storage, read_write> outputData: array<${outputType}>;
490
+ @group(0) @binding(4) var<uniform> params: Params;
491
+
492
+ @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
493
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
494
+ if (gid.x >= params.outputWidth || gid.y >= params.outputHeight || gid.z >= ${outputBlocks}u) {
495
+ return;
496
+ }
497
+ var acc = bias[gid.z];
498
+ for (var ky = 0u; ky < 3u; ky++) {
499
+ let inputY = i32(gid.y) + i32(ky) - 1;
500
+ if (inputY < 0 || inputY >= i32(params.inputHeight)) { continue; }
501
+ for (var kx = 0u; kx < 3u; kx++) {
502
+ let inputX = i32(gid.x) + i32(kx) - 1;
503
+ if (inputX < 0 || inputX >= i32(params.inputWidth)) { continue; }
504
+ let pixelBase = (u32(inputY) * params.inputWidth + u32(inputX)) * ${inputBlocks}u;
505
+ for (var inputBlock = 0u; inputBlock < ${inputBlocks}u; inputBlock++) {
506
+ ${accumulationCode(
507
+ 'inputData[pixelBase + inputBlock]',
508
+ `((((gid.z * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`,
509
+ precision
510
+ )}
511
+ }
512
+ }
513
+ }
514
+ let outputIndex = (gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
515
+ outputData[outputIndex] = ${stored};
516
+ }
517
+ `;
518
+ }
519
+
520
+ /** Direct convolution with subgroup-wide weight broadcast. */
521
+ function createSubgroupConvShader(
522
+ precision: NativeUNetPrecision,
523
+ outputPrecision: NativeUNetPrecision,
524
+ activation: Conv2DNodeSpec['activation'],
525
+ inputBlocks: number,
526
+ outputBlocks: number
527
+ ) {
528
+ const inputType = storageVecType(precision);
529
+ const weightType = storageVecType(precision);
530
+ const outputType = storageVecType(outputPrecision);
531
+ const stored = storeExpression(
532
+ activationExpression('acc', activation),
533
+ outputPrecision
534
+ );
535
+ return /* wgsl */ `${shaderPreamble(precision)}
536
+ enable subgroups;
537
+ struct Params {
538
+ inputWidth: u32,
539
+ inputHeight: u32,
540
+ outputWidth: u32,
541
+ outputHeight: u32,
542
+ inputBlocks: u32,
543
+ outputBlocks: u32,
544
+ }
545
+ @group(0) @binding(0) var<storage, read> inputData: array<${inputType}>;
546
+ @group(0) @binding(1) var<storage, read> weights: array<${weightType}>;
547
+ @group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
548
+ @group(0) @binding(3) var<storage, read_write> outputData: array<${outputType}>;
549
+ @group(0) @binding(4) var<uniform> params: Params;
550
+
551
+ @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
552
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
553
+ let outputInBounds =
554
+ gid.x < params.outputWidth && gid.y < params.outputHeight;
555
+ var acc = bias[gid.z];
556
+ for (var ky = 0u; ky < 3u; ky++) {
557
+ let inputY = i32(gid.y) + i32(ky) - 1;
558
+ let clampedY = u32(clamp(inputY, 0, i32(params.inputHeight) - 1));
559
+ for (var kx = 0u; kx < 3u; kx++) {
560
+ let inputX = i32(gid.x) + i32(kx) - 1;
561
+ let clampedX = u32(clamp(inputX, 0, i32(params.inputWidth) - 1));
562
+ let inputInBounds =
563
+ outputInBounds && inputX >= 0 && inputY >= 0 &&
564
+ inputX < i32(params.inputWidth) && inputY < i32(params.inputHeight);
565
+ let pixelBase =
566
+ (clampedY * params.inputWidth + clampedX) * ${inputBlocks}u;
567
+ for (var inputBlock = 0u; inputBlock < ${inputBlocks}u; inputBlock++) {
568
+ ${subgroupAccumulationCode(
569
+ `select(${inputType}(0.0), inputData[pixelBase + inputBlock], inputInBounds)`,
570
+ `((((gid.z * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`,
571
+ precision
572
+ )}
573
+ }
574
+ }
575
+ }
576
+ if (outputInBounds) {
577
+ let outputIndex =
578
+ (gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
579
+ outputData[outputIndex] = ${stored};
580
+ }
581
+ }
582
+ `;
583
+ }
584
+
585
+ /**
586
+ * A 2D convolution tile which loads the complete 3x3 halo into workgroup
587
+ * memory once. Kernel choice only depends on the operation shape, precision,
588
+ * and device limits; it is deliberately independent of OIDN model names.
589
+ */
590
+ function createSpatialConvShader(
591
+ precision: NativeUNetPrecision,
592
+ outputPrecision: NativeUNetPrecision,
593
+ activation: Conv2DNodeSpec['activation'],
594
+ inputBlocks: number,
595
+ outputBlocks: number
596
+ ) {
597
+ const inputType = storageVecType(precision);
598
+ const outputType = storageVecType(outputPrecision);
599
+ const stored = storeExpression(
600
+ activationExpression('acc', activation),
601
+ outputPrecision
602
+ );
603
+ const patchValues =
604
+ SPATIAL_CONV_PATCH * SPATIAL_CONV_PATCH * inputBlocks;
605
+ const workgroupThreads =
606
+ SPATIAL_CONV_WORKGROUP * SPATIAL_CONV_WORKGROUP;
607
+
608
+ return /* wgsl */ `${shaderPreamble(precision)}
609
+ struct Params {
610
+ inputWidth: u32,
611
+ inputHeight: u32,
612
+ outputWidth: u32,
613
+ outputHeight: u32,
614
+ inputBlocks: u32,
615
+ outputBlocks: u32,
616
+ }
617
+ @group(0) @binding(0) var<storage, read> inputData: array<${inputType}>;
618
+ @group(0) @binding(1) var<storage, read> weights: array<${inputType}>;
619
+ @group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
620
+ @group(0) @binding(3) var<storage, read_write> outputData: array<${outputType}>;
621
+ @group(0) @binding(4) var<uniform> params: Params;
622
+
623
+ var<workgroup> inputPatch: array<${inputType}, ${patchValues}>;
624
+
625
+ @compute @workgroup_size(${SPATIAL_CONV_WORKGROUP}, ${SPATIAL_CONV_WORKGROUP}, 1)
626
+ fn main(
627
+ @builtin(local_invocation_id) localId: vec3<u32>,
628
+ @builtin(workgroup_id) workgroupId: vec3<u32>
629
+ ) {
630
+ let localLinear =
631
+ localId.y * ${SPATIAL_CONV_WORKGROUP}u + localId.x;
632
+ for (
633
+ var loadIndex = localLinear;
634
+ loadIndex < ${patchValues}u;
635
+ loadIndex += ${workgroupThreads}u
636
+ ) {
637
+ let patchPixel = loadIndex / ${inputBlocks}u;
638
+ let inputBlock = loadIndex % ${inputBlocks}u;
639
+ let patchX = patchPixel % ${SPATIAL_CONV_PATCH}u;
640
+ let patchY = patchPixel / ${SPATIAL_CONV_PATCH}u;
641
+ let inputX =
642
+ i32(workgroupId.x * ${SPATIAL_CONV_WORKGROUP}u + patchX) - 1;
643
+ let inputY =
644
+ i32(workgroupId.y * ${SPATIAL_CONV_WORKGROUP}u + patchY) - 1;
645
+ var value = ${inputType}(0.0);
646
+ if (
647
+ inputX >= 0 && inputX < i32(params.inputWidth) &&
648
+ inputY >= 0 && inputY < i32(params.inputHeight)
649
+ ) {
650
+ let inputIndex =
651
+ (u32(inputY) * params.inputWidth + u32(inputX)) *
652
+ ${inputBlocks}u + inputBlock;
653
+ value = inputData[inputIndex];
654
+ }
655
+ inputPatch[loadIndex] = value;
656
+ }
657
+ workgroupBarrier();
658
+
659
+ let outputX =
660
+ workgroupId.x * ${SPATIAL_CONV_WORKGROUP}u + localId.x;
661
+ let outputY =
662
+ workgroupId.y * ${SPATIAL_CONV_WORKGROUP}u + localId.y;
663
+ let outputBlock = workgroupId.z;
664
+ if (
665
+ outputX >= params.outputWidth || outputY >= params.outputHeight ||
666
+ outputBlock >= ${outputBlocks}u
667
+ ) {
668
+ return;
669
+ }
670
+
671
+ var acc = bias[outputBlock];
672
+ for (var ky = 0u; ky < 3u; ky++) {
673
+ for (var kx = 0u; kx < 3u; kx++) {
674
+ let patchBase =
675
+ ((localId.y + ky) * ${SPATIAL_CONV_PATCH}u + localId.x + kx) *
676
+ ${inputBlocks}u;
677
+ for (var inputBlock = 0u; inputBlock < ${inputBlocks}u; inputBlock++) {
678
+ ${accumulationCode(
679
+ 'inputPatch[patchBase + inputBlock]',
680
+ `((((outputBlock * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`,
681
+ precision
682
+ )}
683
+ }
684
+ }
685
+ }
686
+ let outputIndex =
687
+ (outputY * params.outputWidth + outputX) * ${outputBlocks}u + outputBlock;
688
+ outputData[outputIndex] = ${stored};
689
+ }
690
+ `;
691
+ }
692
+
693
+ /**
694
+ * Implicit-GEMM convolution for FP32. This follows the proven packed WebGPU
695
+ * shape: one 8x8 workgroup computes 32 spatial rows by 8 vec4 output blocks,
696
+ * with four output rows per thread and an eight-vec4 K tile.
697
+ */
698
+ function createTiledConvShader(
699
+ precision: NativeUNetPrecision,
700
+ outputPrecision: NativeUNetPrecision,
701
+ activation: Conv2DNodeSpec['activation'],
702
+ inputBlocks: number,
703
+ outputBlocks: number
704
+ ) {
705
+ const inputType = storageVecType(precision);
706
+ const outputType = storageVecType(outputPrecision);
707
+ const stored = storeExpression(
708
+ activationExpression('acc', activation),
709
+ outputPrecision
710
+ );
711
+ const workgroupThreads =
712
+ TILED_CONV_WORKGROUP * TILED_CONV_WORKGROUP;
713
+ const inputTileValues = TILED_CONV_M * TILED_CONV_K_BLOCKS;
714
+ const weightTileValues =
715
+ TILED_CONV_K_BLOCKS * TILED_CONV_N_BLOCKS * 4;
716
+ return /* wgsl */ `${shaderPreamble(precision)}
717
+ struct Params {
718
+ inputWidth: u32,
719
+ inputHeight: u32,
720
+ outputWidth: u32,
721
+ outputHeight: u32,
722
+ inputBlocks: u32,
723
+ outputBlocks: u32,
724
+ }
725
+ @group(0) @binding(0) var<storage, read> inputData: array<${inputType}>;
726
+ @group(0) @binding(1) var<storage, read> weights: array<${inputType}>;
727
+ @group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
728
+ @group(0) @binding(3) var<storage, read_write> outputData: array<${outputType}>;
729
+ @group(0) @binding(4) var<uniform> params: Params;
730
+
731
+ var<workgroup> inputTile: array<${inputType}, ${inputTileValues}>;
732
+ var<workgroup> weightTile: array<${inputType}, ${weightTileValues}>;
733
+
734
+ @compute @workgroup_size(${TILED_CONV_WORKGROUP}, ${TILED_CONV_WORKGROUP}, 1)
735
+ fn main(
736
+ @builtin(local_invocation_id) localId: vec3<u32>,
737
+ @builtin(workgroup_id) workgroupId: vec3<u32>
738
+ ) {
739
+ let spatialBase =
740
+ workgroupId.x * ${TILED_CONV_M}u +
741
+ localId.y * ${TILED_CONV_ROWS_PER_THREAD}u;
742
+ let outputBlock =
743
+ workgroupId.y * ${TILED_CONV_N_BLOCKS}u + localId.x;
744
+ let spatialCount = params.outputWidth * params.outputHeight;
745
+ var acc: array<vec4<f32>, ${TILED_CONV_ROWS_PER_THREAD}>;
746
+ if (outputBlock < ${outputBlocks}u) {
747
+ for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
748
+ acc[row] = bias[outputBlock];
749
+ }
750
+ }
751
+
752
+ let localLinear =
753
+ localId.y * ${TILED_CONV_WORKGROUP}u + localId.x;
754
+ let totalK = ${inputBlocks * 9}u;
755
+ for (var kBase = 0u; kBase < totalK; kBase += ${TILED_CONV_K_BLOCKS}u) {
756
+ for (
757
+ var loadIndex = localLinear;
758
+ loadIndex < ${inputTileValues}u;
759
+ loadIndex += ${workgroupThreads}u
760
+ ) {
761
+ let tileSpatial = loadIndex / ${TILED_CONV_K_BLOCKS}u;
762
+ let tileK = loadIndex % ${TILED_CONV_K_BLOCKS}u;
763
+ let inputSpatialIndex =
764
+ workgroupId.x * ${TILED_CONV_M}u + tileSpatial;
765
+ let kIndex = kBase + tileK;
766
+ var value = ${inputType}(0.0);
767
+ if (inputSpatialIndex < spatialCount && kIndex < totalK) {
768
+ let outputY = inputSpatialIndex / params.outputWidth;
769
+ let outputX = inputSpatialIndex % params.outputWidth;
770
+ let inputBlock = kIndex % ${inputBlocks}u;
771
+ let kernelIndex = kIndex / ${inputBlocks}u;
772
+ let kernelY = kernelIndex / 3u;
773
+ let kernelX = kernelIndex % 3u;
774
+ let inputY = i32(outputY) + i32(kernelY) - 1;
775
+ let inputX = i32(outputX) + i32(kernelX) - 1;
776
+ if (
777
+ inputY >= 0 && inputY < i32(params.inputHeight) &&
778
+ inputX >= 0 && inputX < i32(params.inputWidth)
779
+ ) {
780
+ let inputIndex =
781
+ (u32(inputY) * params.inputWidth + u32(inputX)) *
782
+ ${inputBlocks}u + inputBlock;
783
+ value = inputData[inputIndex];
784
+ }
785
+ }
786
+ inputTile[loadIndex] = value;
787
+ }
788
+
789
+ for (
790
+ var loadIndex = localLinear;
791
+ loadIndex < ${weightTileValues}u;
792
+ loadIndex += ${workgroupThreads}u
793
+ ) {
794
+ let tileK = loadIndex / ${TILED_CONV_N_BLOCKS * 4}u;
795
+ let outputRemainder = loadIndex % ${TILED_CONV_N_BLOCKS * 4}u;
796
+ let tileOutputBlock = outputRemainder / 4u;
797
+ let outputLane = outputRemainder % 4u;
798
+ let loadedOutputBlock =
799
+ workgroupId.y * ${TILED_CONV_N_BLOCKS}u + tileOutputBlock;
800
+ let kIndex = kBase + tileK;
801
+ var value = ${inputType}(0.0);
802
+ if (loadedOutputBlock < ${outputBlocks}u && kIndex < totalK) {
803
+ let inputBlock = kIndex % ${inputBlocks}u;
804
+ let kernelIndex = kIndex / ${inputBlocks}u;
805
+ let kernelY = kernelIndex / 3u;
806
+ let kernelX = kernelIndex % 3u;
807
+ let weightIndex =
808
+ ((((loadedOutputBlock * 3u + kernelY) * 3u + kernelX) *
809
+ ${inputBlocks}u + inputBlock) * 4u + outputLane);
810
+ value = weights[weightIndex];
811
+ }
812
+ weightTile[loadIndex] = value;
813
+ }
814
+
815
+ workgroupBarrier();
816
+ ${tiledAccumulationCode(precision)}
817
+ workgroupBarrier();
818
+ }
819
+
820
+ if (outputBlock < ${outputBlocks}u) {
821
+ for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
822
+ let spatialIndex = spatialBase + row;
823
+ if (spatialIndex < spatialCount) {
824
+ let outputIndex = spatialIndex * ${outputBlocks}u + outputBlock;
825
+ outputData[outputIndex] = ${stored.replaceAll('acc', 'acc[row]')};
826
+ }
827
+ }
828
+ }
829
+ }
830
+ `;
831
+ }
832
+
833
+ function createFusedConvPoolShader(
834
+ precision: NativeUNetPrecision,
835
+ activation: Conv2DNodeSpec['activation'],
836
+ inputBlocks: number,
837
+ outputBlocks: number
838
+ ) {
839
+ const valueType = storageVecType(precision);
840
+ const activated = activationExpression('acc', activation);
841
+ return /* wgsl */ `${shaderPreamble(precision)}
842
+ struct Params {
843
+ inputWidth: u32,
844
+ inputHeight: u32,
845
+ outputWidth: u32,
846
+ outputHeight: u32,
847
+ inputBlocks: u32,
848
+ outputBlocks: u32,
849
+ }
850
+ @group(0) @binding(0) var<storage, read> inputData: array<${valueType}>;
851
+ @group(0) @binding(1) var<storage, read> weights: array<${valueType}>;
852
+ @group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
853
+ @group(0) @binding(3) var<storage, read_write> outputData: array<${valueType}>;
854
+ @group(0) @binding(4) var<uniform> params: Params;
855
+
856
+ @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
857
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
858
+ if (gid.x >= params.outputWidth || gid.y >= params.outputHeight || gid.z >= ${outputBlocks}u) {
859
+ return;
860
+ }
861
+ var pooled = vec4<f32>(-3.402823466e+38);
862
+ for (var py = 0u; py < 2u; py++) {
863
+ let centerY = gid.y * 2u + py;
864
+ if (centerY >= params.inputHeight) { continue; }
865
+ for (var px = 0u; px < 2u; px++) {
866
+ let centerX = gid.x * 2u + px;
867
+ if (centerX >= params.inputWidth) { continue; }
868
+ var acc = bias[gid.z];
869
+ for (var ky = 0u; ky < 3u; ky++) {
870
+ let inputY = i32(centerY) + i32(ky) - 1;
871
+ if (inputY < 0 || inputY >= i32(params.inputHeight)) { continue; }
872
+ for (var kx = 0u; kx < 3u; kx++) {
873
+ let inputX = i32(centerX) + i32(kx) - 1;
874
+ if (inputX < 0 || inputX >= i32(params.inputWidth)) { continue; }
875
+ let pixelBase = (u32(inputY) * params.inputWidth + u32(inputX)) * ${inputBlocks}u;
876
+ for (var inputBlock = 0u; inputBlock < ${inputBlocks}u; inputBlock++) {
877
+ ${accumulationCode(
878
+ 'inputData[pixelBase + inputBlock]',
879
+ `((((gid.z * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`,
880
+ precision
881
+ )}
882
+ }
883
+ }
884
+ }
885
+ pooled = max(pooled, ${activated});
886
+ }
887
+ }
888
+ let outputIndex = (gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
889
+ outputData[outputIndex] = ${storeExpression('pooled', precision)};
890
+ }
891
+ `;
892
+ }
893
+
894
+ function createMaxPoolShader(
895
+ precision: NativeUNetPrecision,
896
+ outputBlocks: number
897
+ ) {
898
+ const valueType = storageVecType(precision);
899
+ return /* wgsl */ `${shaderPreamble(precision)}
900
+ struct Params {
901
+ inputWidth: u32,
902
+ inputHeight: u32,
903
+ outputWidth: u32,
904
+ outputHeight: u32,
905
+ outputBlocks: u32,
906
+ }
907
+ @group(0) @binding(0) var<storage, read> inputData: array<${valueType}>;
908
+ @group(0) @binding(1) var<storage, read_write> outputData: array<${valueType}>;
909
+ @group(0) @binding(2) var<uniform> params: Params;
910
+
911
+ @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
912
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
913
+ if (
914
+ gid.x >= params.outputWidth ||
915
+ gid.y >= params.outputHeight ||
916
+ gid.z >= ${outputBlocks}u
917
+ ) {
918
+ return;
919
+ }
920
+ var pooled = vec4<f32>(-3.402823466e+38);
921
+ for (var py = 0u; py < 2u; py++) {
922
+ let inputY = gid.y * 2u + py;
923
+ if (inputY >= params.inputHeight) { continue; }
924
+ for (var px = 0u; px < 2u; px++) {
925
+ let inputX = gid.x * 2u + px;
926
+ if (inputX >= params.inputWidth) { continue; }
927
+ let inputIndex =
928
+ (inputY * params.inputWidth + inputX) * ${outputBlocks}u + gid.z;
929
+ pooled = max(pooled, vec4<f32>(inputData[inputIndex]));
930
+ }
931
+ }
932
+ let outputIndex =
933
+ (gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
934
+ outputData[outputIndex] = ${storeExpression('pooled', precision)};
935
+ }
936
+ `;
937
+ }
938
+
939
+ function createTiledDecoderShader(
940
+ precision: NativeUNetPrecision,
941
+ activation: Conv2DNodeSpec['activation'],
942
+ sourceBlocks: readonly [number, number],
943
+ upsampledSource: 0 | 1,
944
+ outputBlocks: number
945
+ ) {
946
+ const valueType = storageVecType(precision);
947
+ const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
948
+ const stored = storeExpression(
949
+ activationExpression('acc[row]', activation),
950
+ precision
951
+ );
952
+ const workgroupThreads =
953
+ TILED_CONV_WORKGROUP * TILED_CONV_WORKGROUP;
954
+ const inputTileValues = TILED_CONV_M * TILED_CONV_K_BLOCKS;
955
+ const weightTileValues =
956
+ TILED_CONV_K_BLOCKS * TILED_CONV_N_BLOCKS * 4;
957
+ const sourceRead = (source: 0 | 1, blockExpression: string) => {
958
+ const sourceX =
959
+ source === upsampledSource
960
+ ? 'u32(inputX) / 2u'
961
+ : 'u32(inputX)';
962
+ const sourceY =
963
+ source === upsampledSource
964
+ ? 'u32(inputY) / 2u'
965
+ : 'u32(inputY)';
966
+ return /* wgsl */ `
967
+ {
968
+ let sourceBlock = ${blockExpression};
969
+ let sourceX = ${sourceX};
970
+ let sourceY = ${sourceY};
971
+ let sourceIndex =
972
+ (sourceY * params.source${source}Width + sourceX) *
973
+ ${sourceBlocks[source]}u + sourceBlock;
974
+ value = input${source}[sourceIndex];
975
+ }`;
976
+ };
977
+ return /* wgsl */ `${shaderPreamble(precision)}
978
+ struct Params {
979
+ outputWidth: u32,
980
+ outputHeight: u32,
981
+ outputBlocks: u32,
982
+ inputBlocks: u32,
983
+ source0Width: u32,
984
+ source0Height: u32,
985
+ source1Width: u32,
986
+ source1Height: u32,
987
+ }
988
+ @group(0) @binding(0) var<storage, read> input0: array<${valueType}>;
989
+ @group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
990
+ @group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
991
+ @group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
992
+ @group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
993
+ @group(0) @binding(5) var<uniform> params: Params;
994
+
995
+ var<workgroup> inputTile: array<${valueType}, ${inputTileValues}>;
996
+ var<workgroup> weightTile: array<${valueType}, ${weightTileValues}>;
997
+
998
+ @compute @workgroup_size(${TILED_CONV_WORKGROUP}, ${TILED_CONV_WORKGROUP}, 1)
999
+ fn main(
1000
+ @builtin(local_invocation_id) localId: vec3<u32>,
1001
+ @builtin(workgroup_id) workgroupId: vec3<u32>
1002
+ ) {
1003
+ let spatialBase =
1004
+ workgroupId.x * ${TILED_CONV_M}u +
1005
+ localId.y * ${TILED_CONV_ROWS_PER_THREAD}u;
1006
+ let outputBlock =
1007
+ workgroupId.y * ${TILED_CONV_N_BLOCKS}u + localId.x;
1008
+ let spatialCount = params.outputWidth * params.outputHeight;
1009
+ var acc: array<vec4<f32>, ${TILED_CONV_ROWS_PER_THREAD}>;
1010
+ if (outputBlock < ${outputBlocks}u) {
1011
+ for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
1012
+ acc[row] = bias[outputBlock];
1013
+ }
1014
+ }
1015
+
1016
+ let localLinear =
1017
+ localId.y * ${TILED_CONV_WORKGROUP}u + localId.x;
1018
+ let totalK = ${inputBlocks * 9}u;
1019
+ for (var kBase = 0u; kBase < totalK; kBase += ${TILED_CONV_K_BLOCKS}u) {
1020
+ for (
1021
+ var loadIndex = localLinear;
1022
+ loadIndex < ${inputTileValues}u;
1023
+ loadIndex += ${workgroupThreads}u
1024
+ ) {
1025
+ let tileSpatial = loadIndex / ${TILED_CONV_K_BLOCKS}u;
1026
+ let tileK = loadIndex % ${TILED_CONV_K_BLOCKS}u;
1027
+ let outputSpatialIndex =
1028
+ workgroupId.x * ${TILED_CONV_M}u + tileSpatial;
1029
+ let kIndex = kBase + tileK;
1030
+ var value = ${valueType}(0.0);
1031
+ if (outputSpatialIndex < spatialCount && kIndex < totalK) {
1032
+ let outputY = outputSpatialIndex / params.outputWidth;
1033
+ let outputX = outputSpatialIndex % params.outputWidth;
1034
+ let inputBlock = kIndex % ${inputBlocks}u;
1035
+ let kernelIndex = kIndex / ${inputBlocks}u;
1036
+ let kernelY = kernelIndex / 3u;
1037
+ let kernelX = kernelIndex % 3u;
1038
+ let inputY = i32(outputY) + i32(kernelY) - 1;
1039
+ let inputX = i32(outputX) + i32(kernelX) - 1;
1040
+ if (
1041
+ inputY >= 0 && inputY < i32(params.outputHeight) &&
1042
+ inputX >= 0 && inputX < i32(params.outputWidth)
1043
+ ) {
1044
+ if (inputBlock < ${sourceBlocks[0]}u) {
1045
+ ${sourceRead(0, 'inputBlock')}
1046
+ } else {
1047
+ ${sourceRead(1, `inputBlock - ${sourceBlocks[0]}u`)}
1048
+ }
1049
+ }
1050
+ }
1051
+ inputTile[loadIndex] = value;
1052
+ }
1053
+
1054
+ for (
1055
+ var loadIndex = localLinear;
1056
+ loadIndex < ${weightTileValues}u;
1057
+ loadIndex += ${workgroupThreads}u
1058
+ ) {
1059
+ let tileK = loadIndex / ${TILED_CONV_N_BLOCKS * 4}u;
1060
+ let outputRemainder = loadIndex % ${TILED_CONV_N_BLOCKS * 4}u;
1061
+ let tileOutputBlock = outputRemainder / 4u;
1062
+ let outputLane = outputRemainder % 4u;
1063
+ let loadedOutputBlock =
1064
+ workgroupId.y * ${TILED_CONV_N_BLOCKS}u + tileOutputBlock;
1065
+ let kIndex = kBase + tileK;
1066
+ var value = ${valueType}(0.0);
1067
+ if (loadedOutputBlock < ${outputBlocks}u && kIndex < totalK) {
1068
+ let inputBlock = kIndex % ${inputBlocks}u;
1069
+ let kernelIndex = kIndex / ${inputBlocks}u;
1070
+ let kernelY = kernelIndex / 3u;
1071
+ let kernelX = kernelIndex % 3u;
1072
+ let weightIndex =
1073
+ ((((loadedOutputBlock * 3u + kernelY) * 3u + kernelX) *
1074
+ ${inputBlocks}u + inputBlock) * 4u + outputLane);
1075
+ value = weights[weightIndex];
1076
+ }
1077
+ weightTile[loadIndex] = value;
1078
+ }
1079
+
1080
+ workgroupBarrier();
1081
+ ${tiledAccumulationCode(precision)}
1082
+ workgroupBarrier();
1083
+ }
1084
+
1085
+ if (outputBlock < ${outputBlocks}u) {
1086
+ for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
1087
+ let spatialIndex = spatialBase + row;
1088
+ if (spatialIndex < spatialCount) {
1089
+ let outputIndex = spatialIndex * ${outputBlocks}u + outputBlock;
1090
+ outputData[outputIndex] = ${stored};
1091
+ }
1092
+ }
1093
+ }
1094
+ }
1095
+ `;
1096
+ }
1097
+
1098
+ function createFusedDecoderShader(
1099
+ precision: NativeUNetPrecision,
1100
+ activation: Conv2DNodeSpec['activation'],
1101
+ sourceBlocks: readonly [number, number],
1102
+ upsampledSource: 0 | 1,
1103
+ outputBlocks: number
1104
+ ) {
1105
+ const valueType = storageVecType(precision);
1106
+ const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
1107
+ const sourceCode = (source: 0 | 1, blockOffset: number) => {
1108
+ const isUpsampled = source === upsampledSource;
1109
+ return /* wgsl */ `
1110
+ {
1111
+ let sourceX = ${isUpsampled ? 'u32(inputX) / 2u' : 'u32(inputX)'};
1112
+ let sourceY = ${isUpsampled ? 'u32(inputY) / 2u' : 'u32(inputY)'};
1113
+ let sourcePixelBase = (sourceY * params.source${source}Width + sourceX) * ${sourceBlocks[source]}u;
1114
+ for (var sourceBlock = 0u; sourceBlock < ${sourceBlocks[source]}u; sourceBlock++) {
1115
+ let inputBlock = ${blockOffset}u + sourceBlock;
1116
+ ${accumulationCode(
1117
+ `input${source}[sourcePixelBase + sourceBlock]`,
1118
+ `((((gid.z * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`,
1119
+ precision
1120
+ )}
1121
+ }
1122
+ }
1123
+ `;
1124
+ };
1125
+ const stored = storeExpression(
1126
+ activationExpression('acc', activation),
1127
+ precision
1128
+ );
1129
+ return /* wgsl */ `${shaderPreamble(precision)}
1130
+ struct Params {
1131
+ outputWidth: u32,
1132
+ outputHeight: u32,
1133
+ outputBlocks: u32,
1134
+ inputBlocks: u32,
1135
+ source0Width: u32,
1136
+ source0Height: u32,
1137
+ source1Width: u32,
1138
+ source1Height: u32,
1139
+ }
1140
+ @group(0) @binding(0) var<storage, read> input0: array<${valueType}>;
1141
+ @group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
1142
+ @group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
1143
+ @group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
1144
+ @group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
1145
+ @group(0) @binding(5) var<uniform> params: Params;
1146
+
1147
+ @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
1148
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
1149
+ if (gid.x >= params.outputWidth || gid.y >= params.outputHeight || gid.z >= ${outputBlocks}u) {
1150
+ return;
1151
+ }
1152
+ var acc = bias[gid.z];
1153
+ for (var ky = 0u; ky < 3u; ky++) {
1154
+ let inputY = i32(gid.y) + i32(ky) - 1;
1155
+ if (inputY < 0 || inputY >= i32(params.outputHeight)) { continue; }
1156
+ for (var kx = 0u; kx < 3u; kx++) {
1157
+ let inputX = i32(gid.x) + i32(kx) - 1;
1158
+ if (inputX < 0 || inputX >= i32(params.outputWidth)) { continue; }
1159
+ ${sourceCode(0, 0)}
1160
+ ${sourceCode(1, sourceBlocks[0])}
1161
+ }
1162
+ }
1163
+ let outputIndex = (gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
1164
+ outputData[outputIndex] = ${stored};
1165
+ }
1166
+ `;
1167
+ }
1168
+
1169
+ function createSpatialDecoderShader(
1170
+ precision: NativeUNetPrecision,
1171
+ activation: Conv2DNodeSpec['activation'],
1172
+ sourceBlocks: readonly [number, number],
1173
+ upsampledSource: 0 | 1,
1174
+ outputBlocks: number
1175
+ ) {
1176
+ const valueType = storageVecType(precision);
1177
+ const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
1178
+ const patchValues =
1179
+ SPATIAL_CONV_PATCH * SPATIAL_CONV_PATCH * inputBlocks;
1180
+ const workgroupThreads =
1181
+ SPATIAL_CONV_WORKGROUP * SPATIAL_CONV_WORKGROUP;
1182
+ const stored = storeExpression(
1183
+ activationExpression('acc', activation),
1184
+ precision
1185
+ );
1186
+ const sourceRead = (source: 0 | 1, blockExpression: string) => {
1187
+ const sourceX =
1188
+ source === upsampledSource ? 'u32(inputX) / 2u' : 'u32(inputX)';
1189
+ const sourceY =
1190
+ source === upsampledSource ? 'u32(inputY) / 2u' : 'u32(inputY)';
1191
+ return /* wgsl */ `
1192
+ let sourceBlock = ${blockExpression};
1193
+ let sourceIndex =
1194
+ (${sourceY} * params.source${source}Width + ${sourceX}) *
1195
+ ${sourceBlocks[source]}u + sourceBlock;
1196
+ value = input${source}[sourceIndex];`;
1197
+ };
1198
+
1199
+ return /* wgsl */ `${shaderPreamble(precision)}
1200
+ struct Params {
1201
+ outputWidth: u32,
1202
+ outputHeight: u32,
1203
+ outputBlocks: u32,
1204
+ inputBlocks: u32,
1205
+ source0Width: u32,
1206
+ source0Height: u32,
1207
+ source1Width: u32,
1208
+ source1Height: u32,
1209
+ }
1210
+ @group(0) @binding(0) var<storage, read> input0: array<${valueType}>;
1211
+ @group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
1212
+ @group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
1213
+ @group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
1214
+ @group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
1215
+ @group(0) @binding(5) var<uniform> params: Params;
1216
+
1217
+ var<workgroup> inputPatch: array<${valueType}, ${patchValues}>;
1218
+
1219
+ @compute @workgroup_size(${SPATIAL_CONV_WORKGROUP}, ${SPATIAL_CONV_WORKGROUP}, 1)
1220
+ fn main(
1221
+ @builtin(local_invocation_id) localId: vec3<u32>,
1222
+ @builtin(workgroup_id) workgroupId: vec3<u32>
1223
+ ) {
1224
+ let localLinear =
1225
+ localId.y * ${SPATIAL_CONV_WORKGROUP}u + localId.x;
1226
+ for (
1227
+ var loadIndex = localLinear;
1228
+ loadIndex < ${patchValues}u;
1229
+ loadIndex += ${workgroupThreads}u
1230
+ ) {
1231
+ let patchPixel = loadIndex / ${inputBlocks}u;
1232
+ let inputBlock = loadIndex % ${inputBlocks}u;
1233
+ let patchX = patchPixel % ${SPATIAL_CONV_PATCH}u;
1234
+ let patchY = patchPixel / ${SPATIAL_CONV_PATCH}u;
1235
+ let inputX =
1236
+ i32(workgroupId.x * ${SPATIAL_CONV_WORKGROUP}u + patchX) - 1;
1237
+ let inputY =
1238
+ i32(workgroupId.y * ${SPATIAL_CONV_WORKGROUP}u + patchY) - 1;
1239
+ var value = ${valueType}(0.0);
1240
+ if (
1241
+ inputX >= 0 && inputX < i32(params.outputWidth) &&
1242
+ inputY >= 0 && inputY < i32(params.outputHeight)
1243
+ ) {
1244
+ if (inputBlock < ${sourceBlocks[0]}u) {
1245
+ ${sourceRead(0, 'inputBlock')}
1246
+ } else {
1247
+ ${sourceRead(1, `inputBlock - ${sourceBlocks[0]}u`)}
1248
+ }
1249
+ }
1250
+ inputPatch[loadIndex] = value;
1251
+ }
1252
+ workgroupBarrier();
1253
+
1254
+ let outputX =
1255
+ workgroupId.x * ${SPATIAL_CONV_WORKGROUP}u + localId.x;
1256
+ let outputY =
1257
+ workgroupId.y * ${SPATIAL_CONV_WORKGROUP}u + localId.y;
1258
+ let outputBlock = workgroupId.z;
1259
+ if (
1260
+ outputX >= params.outputWidth || outputY >= params.outputHeight ||
1261
+ outputBlock >= ${outputBlocks}u
1262
+ ) {
1263
+ return;
1264
+ }
1265
+
1266
+ var acc = bias[outputBlock];
1267
+ for (var ky = 0u; ky < 3u; ky++) {
1268
+ for (var kx = 0u; kx < 3u; kx++) {
1269
+ let patchBase =
1270
+ ((localId.y + ky) * ${SPATIAL_CONV_PATCH}u + localId.x + kx) *
1271
+ ${inputBlocks}u;
1272
+ for (var inputBlock = 0u; inputBlock < ${inputBlocks}u; inputBlock++) {
1273
+ ${accumulationCode(
1274
+ 'inputPatch[patchBase + inputBlock]',
1275
+ `((((outputBlock * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`,
1276
+ precision
1277
+ )}
1278
+ }
1279
+ }
1280
+ }
1281
+ let outputIndex =
1282
+ (outputY * params.outputWidth + outputX) * ${outputBlocks}u + outputBlock;
1283
+ outputData[outputIndex] = ${stored};
1284
+ }
1285
+ `;
1286
+ }
1287
+
1288
+ function createInputPackShader(
1289
+ precision: NativeUNetPrecision,
1290
+ sourceCount: number
1291
+ ) {
1292
+ const outputType = storageVecType(precision);
1293
+ const inputBindings = Array.from(
1294
+ { length: sourceCount },
1295
+ (_, index) =>
1296
+ `@group(0) @binding(${index}) var<storage, read> input${index}: array<vec4<f32>>;`
1297
+ ).join('\n');
1298
+ const readBranches = Array.from({ length: sourceCount }, (_, index) => {
1299
+ const firstChannel = index * 3;
1300
+ return `if (channel < ${firstChannel + 3}u) { return input${index}[pixel][channel - ${firstChannel}u]; }`;
1301
+ }).join('\n ');
1302
+ const outputBinding = sourceCount;
1303
+ const paramsBinding = sourceCount + 1;
1304
+ return /* wgsl */ `${shaderPreamble(precision)}
1305
+ struct Params {
1306
+ width: u32,
1307
+ height: u32,
1308
+ outputBlocks: u32,
1309
+ inputChannels: u32,
1310
+ }
1311
+ ${inputBindings}
1312
+ @group(0) @binding(${outputBinding}) var<storage, read_write> outputData: array<${outputType}>;
1313
+ @group(0) @binding(${paramsBinding}) var<uniform> params: Params;
1314
+
1315
+ fn readChannel(pixel: u32, channel: u32) -> f32 {
1316
+ ${readBranches}
1317
+ return 0.0;
1318
+ }
1319
+
1320
+ @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
1321
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
1322
+ if (gid.x >= params.width || gid.y >= params.height || gid.z >= params.outputBlocks) {
1323
+ return;
1324
+ }
1325
+ let pixel = gid.y * params.width + gid.x;
1326
+ let firstChannel = gid.z * 4u;
1327
+ let value = vec4<f32>(
1328
+ readChannel(pixel, firstChannel),
1329
+ readChannel(pixel, firstChannel + 1u),
1330
+ readChannel(pixel, firstChannel + 2u),
1331
+ readChannel(pixel, firstChannel + 3u)
1332
+ );
1333
+ outputData[pixel * params.outputBlocks + gid.z] = ${storeExpression('value', precision)};
1334
+ }
1335
+ `;
1336
+ }
1337
+
1338
+ function executableInputs(node: ExecutableModelNode): readonly string[] {
1339
+ if (node.op === 'concat') return node.inputs;
1340
+ if (node.op === 'fusedUpsampleConcatConv2d') {
1341
+ return node.inputs.map((input) => input.value);
1342
+ }
1343
+ return [node.input];
1344
+ }
1345
+
1346
+ export function resolveNativeUNetPrecision(
1347
+ device: GPUDevice,
1348
+ requested: NativeUNetPrecisionSetting = 'auto'
1349
+ ): NativeUNetPrecision {
1350
+ const hasShaderF16 = device.features.has('shader-f16');
1351
+ if (requested === 'fp16' && !hasShaderF16) {
1352
+ throw new Error(
1353
+ 'OIDN FP16 was requested but the GPUDevice does not have shader-f16 enabled'
1354
+ );
1355
+ }
1356
+ if (requested === 'auto') return hasShaderF16 ? 'fp16' : 'fp32';
1357
+ return requested;
1358
+ }
1359
+
1360
+ /** Native, model-driven OIDN U-Net executor. */
1361
+ export class NativeUNetExecutor {
1362
+ readonly precision: NativeUNetPrecision;
1363
+ readonly kernelSetting: NativeUNetKernelSetting;
1364
+ readonly maxSpatialInputBlocks: number;
1365
+ readonly subgroupsAvailable: boolean;
1366
+
1367
+ private _model: UNetModelGraph;
1368
+ private _packedConvs = new Map<string, PackedConvBuffers>();
1369
+ private _pipelineCache: Map<string, GPUComputePipeline>;
1370
+ private _pipelinePromises: Map<string, Promise<GPUComputePipeline>>;
1371
+ private _executionCache = new Map<string, CachedExecution>();
1372
+ private _retiredExecutions = new Set<CachedExecution>();
1373
+ private _clock = 0;
1374
+ private _shapeCacheSize: number;
1375
+ private _profileNextExecution = false;
1376
+ private _lastExecutionProfile?: Promise<NativeUNetExecutionProfile>;
1377
+ private _profileOperations = 0;
1378
+ private _resources = new OIDNResourceTracker();
1379
+ private _disposed = false;
1380
+
1381
+ constructor(
1382
+ private _device: GPUDevice,
1383
+ model: ValidatedUNetModel,
1384
+ options: NativeUNetOptions = {}
1385
+ ) {
1386
+ const pipelineCache = sharedPipelineCache(_device);
1387
+ this._pipelineCache = pipelineCache.ready;
1388
+ this._pipelinePromises = pipelineCache.pending;
1389
+ this.precision = resolveNativeUNetPrecision(
1390
+ _device,
1391
+ options.precision ?? 'auto'
1392
+ );
1393
+ this.kernelSetting = options.kernel ?? 'auto';
1394
+ this.subgroupsAvailable = _device.features.has(
1395
+ 'subgroups' as GPUFeatureName
1396
+ );
1397
+ this.maxSpatialInputBlocks =
1398
+ this.precision === 'fp16' &&
1399
+ _device.limits.maxComputeInvocationsPerWorkgroup >=
1400
+ SPATIAL_CONV_WORKGROUP * SPATIAL_CONV_WORKGROUP &&
1401
+ _device.limits.maxComputeWorkgroupSizeX >= SPATIAL_CONV_WORKGROUP &&
1402
+ _device.limits.maxComputeWorkgroupSizeY >= SPATIAL_CONV_WORKGROUP
1403
+ ? Math.floor(
1404
+ _device.limits.maxComputeWorkgroupStorageSize /
1405
+ (SPATIAL_CONV_PATCH * SPATIAL_CONV_PATCH * 4 * 2)
1406
+ )
1407
+ : 0;
1408
+ this._shapeCacheSize = Math.max(1, options.shapeCacheSize ?? 2);
1409
+
1410
+ if (
1411
+ model.inputChannels % 3 !== 0 ||
1412
+ model.inputChannels < 3 ||
1413
+ model.inputChannels > 9
1414
+ ) {
1415
+ throw new Error(
1416
+ `Native OIDN expects 3, 6, or 9 input channels, got ${model.inputChannels}`
1417
+ );
1418
+ }
1419
+ this._model = {
1420
+ spec: model.spec,
1421
+ inputChannels: model.inputChannels,
1422
+ outputChannels: model.outputChannels,
1423
+ channelsByValue: new Map(model.channelsByValue),
1424
+ convChannels: new Map(model.convChannels)
1425
+ };
1426
+ for (const node of optimizeModelGraph(this._model, {
1427
+ fuseConvPool: false
1428
+ }).nodes) {
1429
+ if (
1430
+ node.op !== 'conv2d' &&
1431
+ node.op !== 'maxPool2d' &&
1432
+ node.op !== 'fusedConvReluMaxPool2d' &&
1433
+ node.op !== 'fusedUpsampleConcatConv2d'
1434
+ ) {
1435
+ throw new Error(
1436
+ `Native OIDN descriptor ${model.spec.id} leaves unsupported ` +
1437
+ `${node.op} node ${node.id} after graph optimization`
1438
+ );
1439
+ }
1440
+ }
1441
+ try {
1442
+ for (const [id, tensors] of model.convTensors) {
1443
+ const packed = packConvTensors(_device, id, tensors, this.precision);
1444
+ this._resources.track('gpu-buffer', packed.weights);
1445
+ this._resources.track('gpu-buffer', packed.bias);
1446
+ this._packedConvs.set(id, packed);
1447
+ }
1448
+ } catch (error) {
1449
+ for (const packed of this._packedConvs.values()) {
1450
+ this._releaseBuffer(packed.weights);
1451
+ this._releaseBuffer(packed.bias);
1452
+ }
1453
+ this._packedConvs.clear();
1454
+ throw error;
1455
+ }
1456
+ }
1457
+
1458
+ private _pipeline(key: string, code: string) {
1459
+ let pipeline = this._pipelineCache.get(key);
1460
+ if (!pipeline) {
1461
+ pipeline = this._device.createComputePipeline({
1462
+ label: `oidn/${key}`,
1463
+ layout: 'auto',
1464
+ compute: {
1465
+ module: this._device.createShaderModule({
1466
+ label: `oidn/${key}`,
1467
+ code
1468
+ }),
1469
+ entryPoint: 'main'
1470
+ }
1471
+ });
1472
+ this._pipelineCache.set(key, pipeline);
1473
+ }
1474
+ return pipeline;
1475
+ }
1476
+
1477
+ private _pipelineAsync(key: string, code: string) {
1478
+ const ready = this._pipelineCache.get(key);
1479
+ if (ready) return Promise.resolve(ready);
1480
+ const pending = this._pipelinePromises.get(key);
1481
+ if (pending) return pending;
1482
+ const promise = this._device.createComputePipelineAsync({
1483
+ label: `oidn/${key}`,
1484
+ layout: 'auto',
1485
+ compute: {
1486
+ module: this._device.createShaderModule({
1487
+ label: `oidn/${key}`,
1488
+ code
1489
+ }),
1490
+ entryPoint: 'main'
1491
+ }
1492
+ }).then((pipeline) => {
1493
+ this._pipelineCache.set(key, pipeline);
1494
+ this._pipelinePromises.delete(key);
1495
+ return pipeline;
1496
+ }, (error) => {
1497
+ this._pipelinePromises.delete(key);
1498
+ throw error;
1499
+ });
1500
+ this._pipelinePromises.set(key, promise);
1501
+ return promise;
1502
+ }
1503
+
1504
+ private _nodePipelineSpec(
1505
+ node: ExecutableModelNode,
1506
+ isFinal: boolean
1507
+ ): NativePipelineSpec {
1508
+ if (node.op === 'conv2d') {
1509
+ const outputPrecision = isFinal ? 'fp32' : this.precision;
1510
+ const inputBlocks = blocksForChannels(
1511
+ this._model.convChannels.get(node.id)!.inputChannels
1512
+ );
1513
+ const outputBlocks = blocksForChannels(
1514
+ this._model.convChannels.get(node.id)!.outputChannels
1515
+ );
1516
+ const kernel = this._selectConvKernel(inputBlocks, isFinal);
1517
+ const key =
1518
+ `conv-${kernel}/${this.precision}/` +
1519
+ `${outputPrecision}/${node.activation}/` +
1520
+ `in${inputBlocks}/out${outputBlocks}`;
1521
+ return {
1522
+ key,
1523
+ kernel,
1524
+ code: kernel === 'implicit-gemm'
1525
+ ? createTiledConvShader(
1526
+ this.precision,
1527
+ outputPrecision,
1528
+ node.activation,
1529
+ inputBlocks,
1530
+ outputBlocks
1531
+ )
1532
+ : kernel === 'spatial'
1533
+ ? createSpatialConvShader(
1534
+ this.precision,
1535
+ outputPrecision,
1536
+ node.activation,
1537
+ inputBlocks,
1538
+ outputBlocks
1539
+ )
1540
+ : kernel === 'subgroup'
1541
+ ? createSubgroupConvShader(
1542
+ this.precision,
1543
+ outputPrecision,
1544
+ node.activation,
1545
+ inputBlocks,
1546
+ outputBlocks
1547
+ )
1548
+ : createConvShader(
1549
+ this.precision,
1550
+ outputPrecision,
1551
+ node.activation,
1552
+ inputBlocks,
1553
+ outputBlocks
1554
+ )
1555
+ };
1556
+ }
1557
+ if (node.op === 'maxPool2d') {
1558
+ const outputBlocks = blocksForChannels(
1559
+ this._model.channelsByValue.get(node.id)!
1560
+ );
1561
+ const key = `max-pool/${this.precision}/out${outputBlocks}`;
1562
+ return {
1563
+ key,
1564
+ kernel: 'direct',
1565
+ code: createMaxPoolShader(this.precision, outputBlocks)
1566
+ };
1567
+ }
1568
+ if (node.op === 'fusedConvReluMaxPool2d') {
1569
+ const inputBlocks = blocksForChannels(
1570
+ this._model.convChannels.get(node.conv.id)!.inputChannels
1571
+ );
1572
+ const outputBlocks = blocksForChannels(
1573
+ this._model.convChannels.get(node.conv.id)!.outputChannels
1574
+ );
1575
+ const key =
1576
+ `conv-pool/${this.precision}/${node.conv.activation}/` +
1577
+ `in${inputBlocks}/out${outputBlocks}`;
1578
+ return {
1579
+ key,
1580
+ kernel: 'direct',
1581
+ code: createFusedConvPoolShader(
1582
+ this.precision,
1583
+ node.conv.activation,
1584
+ inputBlocks,
1585
+ outputBlocks
1586
+ )
1587
+ };
1588
+ }
1589
+ if (node.op === 'fusedUpsampleConcatConv2d') {
1590
+ if (node.inputs.length !== 2) {
1591
+ throw new Error(`Native fused decoder ${node.id} requires two inputs`);
1592
+ }
1593
+ const sourceBlocks = node.inputs.map((input) =>
1594
+ blocksForChannels(this._model.channelsByValue.get(input.value)!)
1595
+ ) as [number, number];
1596
+ const upsampledSource = node.inputs.findIndex((input) => input.upsample);
1597
+ if (upsampledSource !== 0 && upsampledSource !== 1) {
1598
+ throw new Error(`Native fused decoder ${node.id} has no upsample input`);
1599
+ }
1600
+ const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
1601
+ const selectedKernel = this._selectConvKernel(inputBlocks, false);
1602
+ // The subgroup broadcast path currently targets the common standalone
1603
+ // convolution layout; fused decoder reads use the direct kernel.
1604
+ const kernel = selectedKernel === 'subgroup'
1605
+ ? 'direct'
1606
+ : selectedKernel;
1607
+ const outputBlocks = blocksForChannels(
1608
+ this._model.convChannels.get(node.conv.id)!.outputChannels
1609
+ );
1610
+ const key =
1611
+ `decoder-${kernel}/${this.precision}/` +
1612
+ `${node.conv.activation}/` +
1613
+ `${sourceBlocks.join('+')}/out${outputBlocks}/up${upsampledSource}`;
1614
+ return {
1615
+ key,
1616
+ kernel,
1617
+ code: kernel === 'implicit-gemm'
1618
+ ? createTiledDecoderShader(
1619
+ this.precision,
1620
+ node.conv.activation,
1621
+ sourceBlocks,
1622
+ upsampledSource,
1623
+ outputBlocks
1624
+ )
1625
+ : kernel === 'spatial'
1626
+ ? createSpatialDecoderShader(
1627
+ this.precision,
1628
+ node.conv.activation,
1629
+ sourceBlocks,
1630
+ upsampledSource,
1631
+ outputBlocks
1632
+ )
1633
+ : createFusedDecoderShader(
1634
+ this.precision,
1635
+ node.conv.activation,
1636
+ sourceBlocks,
1637
+ upsampledSource,
1638
+ outputBlocks
1639
+ )
1640
+ };
1641
+ }
1642
+ throw new Error(
1643
+ `Native OIDN does not implement unfused ${node.op} node ${node.id}`
1644
+ );
1645
+ }
1646
+
1647
+ private _selectConvKernel(
1648
+ inputBlocks: number,
1649
+ isFinal: boolean
1650
+ ): NativeUNetKernel {
1651
+ const spatialFits =
1652
+ this.precision === 'fp16' &&
1653
+ inputBlocks <= this.maxSpatialInputBlocks;
1654
+ if (this.kernelSetting === 'direct') return 'direct';
1655
+ if (this.kernelSetting === 'spatial') {
1656
+ return spatialFits ? 'spatial' : 'direct';
1657
+ }
1658
+ if (this.kernelSetting === 'implicit-gemm') {
1659
+ return isFinal ? 'direct' : 'implicit-gemm';
1660
+ }
1661
+ if (this.kernelSetting === 'subgroup') {
1662
+ return this.subgroupsAvailable ? 'subgroup' : 'direct';
1663
+ }
1664
+ if (this.precision === 'fp32' && !isFinal) return 'implicit-gemm';
1665
+ return 'direct';
1666
+ }
1667
+
1668
+ private _nodePipeline(node: ExecutableModelNode, isFinal: boolean) {
1669
+ const { key, code } = this._nodePipelineSpec(node, isFinal);
1670
+ return this._pipeline(key, code);
1671
+ }
1672
+
1673
+ /** Compiles all shape-independent kernels before the model reports ready. */
1674
+ async prepare() {
1675
+ if (this._disposed) throw new Error('Native OIDN executor is disposed');
1676
+ const graph = optimizeModelGraph(this._model, { fuseConvPool: false });
1677
+ const sourceCount = this._model.inputChannels / 3;
1678
+ const specs: NativePipelineSpec[] = [
1679
+ {
1680
+ key: `pack/${this.precision}/${sourceCount}`,
1681
+ code: createInputPackShader(this.precision, sourceCount)
1682
+ },
1683
+ ...graph.nodes.map((node) =>
1684
+ this._nodePipelineSpec(node, node.id === graph.spec.output)
1685
+ )
1686
+ ];
1687
+ await Promise.all(
1688
+ specs.map(({ key, code }) => this._pipelineAsync(key, code))
1689
+ );
1690
+ }
1691
+
1692
+ private _createExecution(width: number, height: number): CachedExecution {
1693
+ const plan = planModelExecution(this._model, width, height, {
1694
+ fuseConvPool: false
1695
+ });
1696
+ const valueBuffers = new Map<string, GPUBuffer>();
1697
+ const slots: ActivationSlot[] = [];
1698
+ const lastUses = new Map<string, number>();
1699
+ plan.nodes.forEach((node, index) => {
1700
+ for (const input of executableInputs(node)) lastUses.set(input, index);
1701
+ });
1702
+ lastUses.set(plan.spec.output, plan.nodes.length);
1703
+ const createdBuffers: GPUBuffer[] = [];
1704
+ const own = (buffer: GPUBuffer) => {
1705
+ createdBuffers.push(buffer);
1706
+ return this._resources.track('gpu-buffer', buffer);
1707
+ };
1708
+
1709
+ try {
1710
+ const allocate = (
1711
+ value: string,
1712
+ shape: ModelValueShape,
1713
+ bytesPerScalar: number,
1714
+ index: number
1715
+ ) => {
1716
+ for (const slot of slots) {
1717
+ if (
1718
+ slot.activeValue &&
1719
+ (lastUses.get(slot.activeValue) ?? -1) < index
1720
+ ) {
1721
+ slot.activeValue = undefined;
1722
+ }
1723
+ }
1724
+ const requiredSize = activationByteSize(shape, bytesPerScalar);
1725
+ let slot = slots
1726
+ .filter((candidate) => !candidate.activeValue && candidate.capacity >= requiredSize)
1727
+ .sort((a, b) => a.capacity - b.capacity)[0];
1728
+ if (!slot) {
1729
+ const buffer = own(this._device.createBuffer({
1730
+ label: `oidn/activation/${width}x${height}/${slots.length}`,
1731
+ size: roundUp(requiredSize, 4),
1732
+ usage:
1733
+ GPUBufferUsage.STORAGE |
1734
+ GPUBufferUsage.COPY_SRC |
1735
+ GPUBufferUsage.COPY_DST
1736
+ }));
1737
+ slot = { buffer, capacity: requiredSize };
1738
+ slots.push(slot);
1739
+ }
1740
+ slot.activeValue = value;
1741
+ valueBuffers.set(value, slot.buffer);
1742
+ };
1743
+
1744
+ allocate(
1745
+ plan.spec.input,
1746
+ plan.inputShape,
1747
+ this.precision === 'fp16' ? 2 : 4,
1748
+ -1
1749
+ );
1750
+ plan.plannedNodes.forEach(({ node, outputShape }, index) => {
1751
+ const isFinal = node.id === plan.spec.output;
1752
+ allocate(
1753
+ node.id,
1754
+ outputShape,
1755
+ isFinal || this.precision === 'fp32' ? 4 : 2,
1756
+ index
1757
+ );
1758
+ });
1759
+
1760
+ const inputSourceCount = this._model.inputChannels / 3;
1761
+ const inputKey = `pack/${this.precision}/${inputSourceCount}`;
1762
+ const inputPipeline = this._pipeline(
1763
+ inputKey,
1764
+ createInputPackShader(this.precision, inputSourceCount)
1765
+ );
1766
+ const nodePipelines: GPUComputePipeline[] = [];
1767
+ const nodeKernels: NativeUNetKernel[] = [];
1768
+ const nodeBindings: GPUBindGroup[] = [];
1769
+ const ownedBuffers: GPUBuffer[] = [];
1770
+
1771
+ plan.plannedNodes.forEach(({ node, outputShape }, index) => {
1772
+ const isFinal = node.id === plan.spec.output;
1773
+ const pipelineSpec = this._nodePipelineSpec(node, isFinal);
1774
+ const pipeline = this._pipeline(pipelineSpec.key, pipelineSpec.code);
1775
+ nodePipelines.push(pipeline);
1776
+ nodeKernels.push(pipelineSpec.kernel ?? 'direct');
1777
+ const output = valueBuffers.get(node.id)!;
1778
+ let entries: GPUBindGroupEntry[];
1779
+ let uniformValues: number[];
1780
+ let convId: string;
1781
+
1782
+ if (node.op === 'conv2d') {
1783
+ const inputShape = plan.valueShapes.get(node.input)!;
1784
+ convId = node.id;
1785
+ uniformValues = [
1786
+ inputShape.width,
1787
+ inputShape.height,
1788
+ outputShape.width,
1789
+ outputShape.height,
1790
+ blocksForChannels(inputShape.channels),
1791
+ blocksForChannels(outputShape.channels)
1792
+ ];
1793
+ const packed = this._packedConvs.get(convId)!;
1794
+ entries = [
1795
+ { binding: 0, resource: { buffer: valueBuffers.get(node.input)! } },
1796
+ { binding: 1, resource: { buffer: packed.weights } },
1797
+ { binding: 2, resource: { buffer: packed.bias } },
1798
+ { binding: 3, resource: { buffer: output } }
1799
+ ];
1800
+ } else if (node.op === 'maxPool2d') {
1801
+ const inputShape = plan.valueShapes.get(node.input)!;
1802
+ uniformValues = [
1803
+ inputShape.width,
1804
+ inputShape.height,
1805
+ outputShape.width,
1806
+ outputShape.height,
1807
+ blocksForChannels(outputShape.channels)
1808
+ ];
1809
+ entries = [
1810
+ { binding: 0, resource: { buffer: valueBuffers.get(node.input)! } },
1811
+ { binding: 1, resource: { buffer: output } }
1812
+ ];
1813
+ } else if (node.op === 'fusedConvReluMaxPool2d') {
1814
+ const inputShape = plan.valueShapes.get(node.input)!;
1815
+ convId = node.conv.id;
1816
+ uniformValues = [
1817
+ inputShape.width,
1818
+ inputShape.height,
1819
+ outputShape.width,
1820
+ outputShape.height,
1821
+ blocksForChannels(inputShape.channels),
1822
+ blocksForChannels(outputShape.channels)
1823
+ ];
1824
+ const packed = this._packedConvs.get(convId)!;
1825
+ entries = [
1826
+ { binding: 0, resource: { buffer: valueBuffers.get(node.input)! } },
1827
+ { binding: 1, resource: { buffer: packed.weights } },
1828
+ { binding: 2, resource: { buffer: packed.bias } },
1829
+ { binding: 3, resource: { buffer: output } }
1830
+ ];
1831
+ } else if (node.op === 'fusedUpsampleConcatConv2d') {
1832
+ convId = node.conv.id;
1833
+ const firstShape = plan.valueShapes.get(node.inputs[0].value)!;
1834
+ const secondShape = plan.valueShapes.get(node.inputs[1].value)!;
1835
+ uniformValues = [
1836
+ outputShape.width,
1837
+ outputShape.height,
1838
+ blocksForChannels(outputShape.channels),
1839
+ blocksForChannels(this._model.convChannels.get(convId)!.inputChannels),
1840
+ firstShape.width,
1841
+ firstShape.height,
1842
+ secondShape.width,
1843
+ secondShape.height
1844
+ ];
1845
+ const packed = this._packedConvs.get(convId)!;
1846
+ entries = [
1847
+ {
1848
+ binding: 0,
1849
+ resource: { buffer: valueBuffers.get(node.inputs[0].value)! }
1850
+ },
1851
+ {
1852
+ binding: 1,
1853
+ resource: { buffer: valueBuffers.get(node.inputs[1].value)! }
1854
+ },
1855
+ { binding: 2, resource: { buffer: packed.weights } },
1856
+ { binding: 3, resource: { buffer: packed.bias } },
1857
+ { binding: 4, resource: { buffer: output } }
1858
+ ];
1859
+ } else {
1860
+ throw new Error(`Unexpected native node ${node.op}`);
1861
+ }
1862
+
1863
+ const uniform = own(
1864
+ createUniformBuffer(
1865
+ this._device,
1866
+ `oidn/${node.id}/params/${width}x${height}`,
1867
+ uniformValues
1868
+ )
1869
+ );
1870
+ ownedBuffers.push(uniform);
1871
+ entries.push({ binding: entries.length, resource: { buffer: uniform } });
1872
+ nodeBindings.push(
1873
+ this._device.createBindGroup({
1874
+ label: `oidn/${node.id}/bindings`,
1875
+ layout: pipeline.getBindGroupLayout(0),
1876
+ entries
1877
+ })
1878
+ );
1879
+ });
1880
+
1881
+ const inputUniform = own(
1882
+ createUniformBuffer(
1883
+ this._device,
1884
+ `oidn/input/params/${width}x${height}`,
1885
+ [
1886
+ width,
1887
+ height,
1888
+ blocksForChannels(this._model.inputChannels),
1889
+ this._model.inputChannels
1890
+ ]
1891
+ )
1892
+ );
1893
+ ownedBuffers.push(inputUniform);
1894
+
1895
+ return {
1896
+ plan,
1897
+ valueBuffers,
1898
+ slots,
1899
+ nodeBindings,
1900
+ nodePipelines,
1901
+ nodeKernels,
1902
+ inputPipeline,
1903
+ inputUniform,
1904
+ ownedBuffers,
1905
+ lastUsed: ++this._clock
1906
+ };
1907
+ } catch (error) {
1908
+ for (const buffer of createdBuffers) this._releaseBuffer(buffer);
1909
+ throw error;
1910
+ }
1911
+ }
1912
+
1913
+ private _execution(width: number, height: number) {
1914
+ if (this._disposed) throw new Error('Native OIDN executor is disposed');
1915
+ const key = `${width}x${height}`;
1916
+ let execution = this._executionCache.get(key);
1917
+ if (!execution) {
1918
+ execution = this._createExecution(width, height);
1919
+ this._executionCache.set(key, execution);
1920
+ if (this._executionCache.size > this._shapeCacheSize) {
1921
+ const oldest = [...this._executionCache.entries()]
1922
+ .filter(([candidateKey]) => candidateKey !== key)
1923
+ .sort((a, b) => a[1].lastUsed - b[1].lastUsed)[0];
1924
+ if (oldest) {
1925
+ this._executionCache.delete(oldest[0]);
1926
+ // Commands using an evicted plan may still be submitted. Defer actual
1927
+ // destruction until all work currently on the shared queue completes.
1928
+ this._retiredExecutions.add(oldest[1]);
1929
+ void this._device.queue.onSubmittedWorkDone()
1930
+ .catch(() => undefined)
1931
+ .then(() => {
1932
+ this._retiredExecutions.delete(oldest[1]);
1933
+ this._destroyExecution(oldest[1]);
1934
+ });
1935
+ }
1936
+ }
1937
+ }
1938
+ execution.lastUsed = ++this._clock;
1939
+ return execution;
1940
+ }
1941
+
1942
+ /** Captures per-pass GPU timestamps for the next execute call when supported. */
1943
+ profileNextExecution() {
1944
+ if (!this._device.features.has('timestamp-query')) return false;
1945
+ this._profileNextExecution = true;
1946
+ return true;
1947
+ }
1948
+
1949
+ getLastExecutionProfile() {
1950
+ return this._lastExecutionProfile;
1951
+ }
1952
+
1953
+ execute(inputBuffers: readonly GPUBuffer[], width: number, height: number) {
1954
+ const sourceCount = this._model.inputChannels / 3;
1955
+ if (inputBuffers.length !== sourceCount) {
1956
+ throw new Error(
1957
+ `Native OIDN expected ${sourceCount} input buffers, got ${inputBuffers.length}`
1958
+ );
1959
+ }
1960
+ const execution = this._execution(width, height);
1961
+ const profileLabels = [
1962
+ 'input-pack',
1963
+ ...execution.plan.nodes.map((node) => node.id)
1964
+ ];
1965
+ const shouldProfile =
1966
+ this._profileNextExecution &&
1967
+ this._device.features.has('timestamp-query');
1968
+ this._profileNextExecution = false;
1969
+ const queryCount = profileLabels.length * 2;
1970
+ const querySet = shouldProfile
1971
+ ? this._resources.track(
1972
+ 'gpu-query-set',
1973
+ this._device.createQuerySet({ type: 'timestamp', count: queryCount })
1974
+ )
1975
+ : undefined;
1976
+ const queryBufferSize = queryCount * 8;
1977
+ let queryResolveBuffer: GPUBuffer | undefined;
1978
+ let queryReadbackBuffer: GPUBuffer | undefined;
1979
+ try {
1980
+ queryResolveBuffer = shouldProfile
1981
+ ? this._resources.track(
1982
+ 'gpu-buffer',
1983
+ this._device.createBuffer({
1984
+ label: `oidn/profile/resolve/${width}x${height}`,
1985
+ size: queryBufferSize,
1986
+ usage: GPUBufferUsage.QUERY_RESOLVE | GPUBufferUsage.COPY_SRC
1987
+ })
1988
+ )
1989
+ : undefined;
1990
+ queryReadbackBuffer = shouldProfile
1991
+ ? this._resources.track(
1992
+ 'gpu-buffer',
1993
+ this._device.createBuffer({
1994
+ label: `oidn/profile/readback/${width}x${height}`,
1995
+ size: queryBufferSize,
1996
+ usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.MAP_READ
1997
+ })
1998
+ )
1999
+ : undefined;
2000
+ } catch (error) {
2001
+ this._releaseBuffer(queryReadbackBuffer);
2002
+ this._releaseBuffer(queryResolveBuffer);
2003
+ this._releaseQuerySet(querySet);
2004
+ throw error;
2005
+ }
2006
+ const passDescriptor = (label: string, index: number) => ({
2007
+ label,
2008
+ ...(querySet
2009
+ ? {
2010
+ timestampWrites: {
2011
+ querySet,
2012
+ beginningOfPassWriteIndex: index * 2,
2013
+ endOfPassWriteIndex: index * 2 + 1
2014
+ }
2015
+ }
2016
+ : {})
2017
+ });
2018
+ try {
2019
+ const encoder = this._device.createCommandEncoder({
2020
+ label: `oidn/native/${width}x${height}`
2021
+ });
2022
+
2023
+ const inputEntries: GPUBindGroupEntry[] = inputBuffers.map(
2024
+ (buffer, binding) => ({ binding, resource: { buffer } })
2025
+ );
2026
+ inputEntries.push({
2027
+ binding: sourceCount,
2028
+ resource: { buffer: execution.valueBuffers.get(execution.plan.spec.input)! }
2029
+ });
2030
+ inputEntries.push({
2031
+ binding: sourceCount + 1,
2032
+ resource: { buffer: execution.inputUniform }
2033
+ });
2034
+ const inputBindings = this._device.createBindGroup({
2035
+ label: 'oidn/input/bindings',
2036
+ layout: execution.inputPipeline.getBindGroupLayout(0),
2037
+ entries: inputEntries
2038
+ });
2039
+ {
2040
+ const pass = encoder.beginComputePass(
2041
+ passDescriptor('oidn/input-pack', 0)
2042
+ );
2043
+ pass.setPipeline(execution.inputPipeline);
2044
+ pass.setBindGroup(0, inputBindings);
2045
+ pass.dispatchWorkgroups(
2046
+ Math.ceil(width / WORKGROUP_SIZE),
2047
+ Math.ceil(height / WORKGROUP_SIZE),
2048
+ blocksForChannels(this._model.inputChannels)
2049
+ );
2050
+ pass.end();
2051
+ }
2052
+
2053
+ execution.plan.plannedNodes.forEach(({ node, outputShape }, index) => {
2054
+ const pass = encoder.beginComputePass(
2055
+ passDescriptor(`oidn/${execution.plan.nodes[index].id}`, index + 1)
2056
+ );
2057
+ pass.setPipeline(execution.nodePipelines[index]);
2058
+ pass.setBindGroup(0, execution.nodeBindings[index]);
2059
+ if (execution.nodeKernels[index] === 'implicit-gemm') {
2060
+ pass.dispatchWorkgroups(
2061
+ Math.ceil(
2062
+ (outputShape.width * outputShape.height) / TILED_CONV_M
2063
+ ),
2064
+ Math.ceil(
2065
+ blocksForChannels(outputShape.channels) / TILED_CONV_N_BLOCKS
2066
+ ),
2067
+ 1
2068
+ );
2069
+ } else {
2070
+ pass.dispatchWorkgroups(
2071
+ Math.ceil(outputShape.width / WORKGROUP_SIZE),
2072
+ Math.ceil(outputShape.height / WORKGROUP_SIZE),
2073
+ blocksForChannels(outputShape.channels)
2074
+ );
2075
+ }
2076
+ pass.end();
2077
+ });
2078
+
2079
+ if (querySet) {
2080
+ encoder.resolveQuerySet(
2081
+ querySet,
2082
+ 0,
2083
+ queryCount,
2084
+ queryResolveBuffer!,
2085
+ 0
2086
+ );
2087
+ encoder.copyBufferToBuffer(
2088
+ queryResolveBuffer!,
2089
+ 0,
2090
+ queryReadbackBuffer!,
2091
+ 0,
2092
+ queryBufferSize
2093
+ );
2094
+ }
2095
+
2096
+ this._device.queue.submit([encoder.finish()]);
2097
+ } catch (error) {
2098
+ this._releaseBuffer(queryReadbackBuffer);
2099
+ this._releaseBuffer(queryResolveBuffer);
2100
+ this._releaseQuerySet(querySet);
2101
+ throw error;
2102
+ }
2103
+ if (querySet) {
2104
+ this._profileOperations++;
2105
+ this._lastExecutionProfile = (async () => {
2106
+ try {
2107
+ await queryReadbackBuffer!.mapAsync(GPUMapMode.READ);
2108
+ const timestamps = new BigUint64Array(
2109
+ queryReadbackBuffer!.getMappedRange()
2110
+ );
2111
+ const layers = profileLabels.map((id, index) => ({
2112
+ id,
2113
+ durationMs:
2114
+ Number(timestamps[index * 2 + 1] - timestamps[index * 2]) /
2115
+ 1_000_000
2116
+ }));
2117
+ return {
2118
+ totalMs: layers.reduce(
2119
+ (sum, layer) => sum + layer.durationMs,
2120
+ 0
2121
+ ),
2122
+ layers
2123
+ };
2124
+ } finally {
2125
+ if (queryReadbackBuffer!.mapState === 'mapped') {
2126
+ queryReadbackBuffer!.unmap();
2127
+ }
2128
+ this._releaseQuerySet(querySet);
2129
+ this._releaseBuffer(queryResolveBuffer);
2130
+ this._releaseBuffer(queryReadbackBuffer);
2131
+ this._profileOperations--;
2132
+ }
2133
+ })();
2134
+ }
2135
+ return execution.valueBuffers.get(execution.plan.spec.output)!;
2136
+ }
2137
+
2138
+ /** Compatibility path for ImageData/HDR arrays without TensorFlow.js. */
2139
+ async executeCPU(
2140
+ interleavedInput: Float32Array,
2141
+ width: number,
2142
+ height: number
2143
+ ): Promise<Float32Array> {
2144
+ const expectedLength = width * height * this._model.inputChannels;
2145
+ if (interleavedInput.length !== expectedLength) {
2146
+ throw new Error(
2147
+ `Native OIDN CPU input has ${interleavedInput.length} values, expected ${expectedLength}`
2148
+ );
2149
+ }
2150
+ const execution = this._execution(width, height);
2151
+ const sourceCount = this._model.inputChannels / 3;
2152
+ const pixelCount = width * height;
2153
+ if (!execution.cpuInputBuffers) {
2154
+ execution.cpuInputBuffers = Array.from({ length: sourceCount }, (_, index) => {
2155
+ const buffer = this._device.createBuffer({
2156
+ label: `oidn/cpu-input/${width}x${height}/${index}`,
2157
+ size: pixelCount * 16,
2158
+ usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST
2159
+ });
2160
+ this._resources.track('gpu-buffer', buffer);
2161
+ execution.ownedBuffers.push(buffer);
2162
+ return buffer;
2163
+ });
2164
+ execution.cpuReadbackBuffer = this._device.createBuffer({
2165
+ label: `oidn/cpu-readback/${width}x${height}`,
2166
+ size: pixelCount * 16,
2167
+ usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.MAP_READ
2168
+ });
2169
+ this._resources.track('gpu-buffer', execution.cpuReadbackBuffer);
2170
+ execution.ownedBuffers.push(execution.cpuReadbackBuffer);
2171
+ }
2172
+
2173
+ for (let source = 0; source < sourceCount; source++) {
2174
+ const upload = new Float32Array(pixelCount * 4);
2175
+ for (let pixel = 0; pixel < pixelCount; pixel++) {
2176
+ const inputOffset =
2177
+ pixel * this._model.inputChannels + source * 3;
2178
+ const outputOffset = pixel * 4;
2179
+ upload[outputOffset] = interleavedInput[inputOffset];
2180
+ upload[outputOffset + 1] = interleavedInput[inputOffset + 1];
2181
+ upload[outputOffset + 2] = interleavedInput[inputOffset + 2];
2182
+ }
2183
+ this._device.queue.writeBuffer(
2184
+ execution.cpuInputBuffers[source],
2185
+ 0,
2186
+ upload
2187
+ );
2188
+ }
2189
+
2190
+ const output = this.execute(execution.cpuInputBuffers, width, height);
2191
+ const encoder = this._device.createCommandEncoder({
2192
+ label: `oidn/cpu-readback/${width}x${height}`
2193
+ });
2194
+ encoder.copyBufferToBuffer(
2195
+ output,
2196
+ 0,
2197
+ execution.cpuReadbackBuffer!,
2198
+ 0,
2199
+ pixelCount * 16
2200
+ );
2201
+ this._device.queue.submit([encoder.finish()]);
2202
+ await execution.cpuReadbackBuffer!.mapAsync(GPUMapMode.READ);
2203
+ const rgba = new Float32Array(
2204
+ execution.cpuReadbackBuffer!.getMappedRange()
2205
+ );
2206
+ const rgb = new Float32Array(pixelCount * 3);
2207
+ for (let pixel = 0; pixel < pixelCount; pixel++) {
2208
+ rgb[pixel * 3] = rgba[pixel * 4];
2209
+ rgb[pixel * 3 + 1] = rgba[pixel * 4 + 1];
2210
+ rgb[pixel * 3 + 2] = rgba[pixel * 4 + 2];
2211
+ }
2212
+ execution.cpuReadbackBuffer!.unmap();
2213
+ return rgb;
2214
+ }
2215
+
2216
+ private _releaseBuffer(buffer: GPUBuffer | undefined) {
2217
+ this._resources.release('gpu-buffer', buffer, () => buffer!.destroy());
2218
+ }
2219
+
2220
+ private _releaseQuerySet(querySet: GPUQuerySet | undefined) {
2221
+ this._resources.release(
2222
+ 'gpu-query-set',
2223
+ querySet,
2224
+ () => querySet!.destroy()
2225
+ );
2226
+ }
2227
+
2228
+ private _destroyExecution(execution: CachedExecution) {
2229
+ execution.slots.forEach((slot) => this._releaseBuffer(slot.buffer));
2230
+ execution.ownedBuffers.forEach((buffer) => this._releaseBuffer(buffer));
2231
+ }
2232
+
2233
+ getResourceInfo(): OIDNResourceSnapshot {
2234
+ return this._resources.snapshot(
2235
+ this._retiredExecutions.size + this._profileOperations
2236
+ );
2237
+ }
2238
+
2239
+ dispose() {
2240
+ if (this._disposed) return;
2241
+ this._disposed = true;
2242
+ for (const packed of this._packedConvs.values()) {
2243
+ this._releaseBuffer(packed.weights);
2244
+ this._releaseBuffer(packed.bias);
2245
+ }
2246
+ this._packedConvs.clear();
2247
+ for (const execution of this._executionCache.values()) {
2248
+ this._destroyExecution(execution);
2249
+ }
2250
+ this._executionCache.clear();
2251
+ for (const execution of this._retiredExecutions) {
2252
+ this._destroyExecution(execution);
2253
+ }
2254
+ this._retiredExecutions.clear();
2255
+ }
2256
+ }