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