oidn-web 0.4.0 → 0.5.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (106) hide show
  1. package/CHANGELOG.md +45 -0
  2. package/README.md +83 -19
  3. package/dist/oidn.js +3239 -2642
  4. package/dist/oidn.umd.cjs +512 -296
  5. package/lib/UNet.d.ts +54 -20
  6. package/lib/UNet.js +194 -118
  7. package/lib/UNet.js.map +1 -1
  8. package/lib/backend.d.ts +1 -8
  9. package/lib/backend.js +1 -9
  10. package/lib/backend.js.map +1 -1
  11. package/lib/finalRgbShader.d.ts +13 -0
  12. package/lib/finalRgbShader.js +160 -0
  13. package/lib/finalRgbShader.js.map +1 -0
  14. package/lib/graphOptimizer.js +1 -2
  15. package/lib/graphOptimizer.js.map +1 -1
  16. package/lib/hdrTransfer.d.ts +14 -0
  17. package/lib/hdrTransfer.js +61 -0
  18. package/lib/hdrTransfer.js.map +1 -0
  19. package/lib/main.d.ts +16 -7
  20. package/lib/main.js +4 -5
  21. package/lib/main.js.map +1 -1
  22. package/lib/nativeUNet.d.ts +39 -3
  23. package/lib/nativeUNet.js +449 -120
  24. package/lib/nativeUNet.js.map +1 -1
  25. package/lib/process.d.ts +5 -11
  26. package/lib/process.js +35 -49
  27. package/lib/process.js.map +1 -1
  28. package/lib/tileScheduler.d.ts +32 -4
  29. package/lib/tileScheduler.js +133 -20
  30. package/lib/tileScheduler.js.map +1 -1
  31. package/package.json +9 -2
  32. package/src/UNet.ts +287 -158
  33. package/src/backend.ts +1 -14
  34. package/src/finalRgbShader.ts +186 -0
  35. package/src/graphOptimizer.ts +1 -2
  36. package/src/hdrTransfer.ts +88 -0
  37. package/src/main.ts +28 -13
  38. package/src/nativeUNet.ts +515 -116
  39. package/src/process.ts +43 -70
  40. package/src/tileScheduler.ts +216 -24
  41. package/benchmarks/compare.mjs +0 -651
  42. package/benchmarks/leak.mjs +0 -255
  43. package/benchmarks/results/before-spatial.json +0 -391
  44. package/benchmarks/results/before-spatial.md +0 -47
  45. package/benchmarks/results/int8-scan.json +0 -2007
  46. package/benchmarks/results/int8-scan.md +0 -160
  47. package/benchmarks/results/int8-w8a8-scan.json +0 -2007
  48. package/benchmarks/results/int8-w8a8-scan.md +0 -160
  49. package/benchmarks/results/int8-weight-channel.json +0 -1413
  50. package/benchmarks/results/int8-weight-channel.md +0 -118
  51. package/benchmarks/results/int8-weight-only.json +0 -1437
  52. package/benchmarks/results/int8-weight-only.md +0 -118
  53. package/benchmarks/results/kernel-webnn-final.json +0 -1115
  54. package/benchmarks/results/kernel-webnn-final.md +0 -104
  55. package/benchmarks/results/latest-optimized.json +0 -375
  56. package/benchmarks/results/latest-optimized.md +0 -47
  57. package/benchmarks/results/latest.json +0 -391
  58. package/benchmarks/results/latest.md +0 -47
  59. package/benchmarks/results/profile-baseline.json +0 -331
  60. package/benchmarks/results/profile-baseline.md +0 -12
  61. package/benchmarks/results/profile-conv2x.json +0 -331
  62. package/benchmarks/results/profile-conv2x.md +0 -12
  63. package/benchmarks/results/profile-fast-init.json +0 -385
  64. package/benchmarks/results/profile-fast-init.md +0 -47
  65. package/benchmarks/results/profile-fp16-fma.json +0 -369
  66. package/benchmarks/results/profile-fp16-fma.md +0 -47
  67. package/benchmarks/results/profile-fp16-tiled-decoder.json +0 -369
  68. package/benchmarks/results/profile-fp16-tiled-decoder.md +0 -47
  69. package/benchmarks/results/profile-fp16-tiled-encoder.json +0 -369
  70. package/benchmarks/results/profile-fp16-tiled-encoder.md +0 -47
  71. package/benchmarks/results/profile-fp16-unfused-pool.json +0 -385
  72. package/benchmarks/results/profile-fp16-unfused-pool.md +0 -47
  73. package/benchmarks/results/profile-input-major.json +0 -347
  74. package/benchmarks/results/profile-input-major.md +0 -12
  75. package/benchmarks/results/profile-k16.json +0 -347
  76. package/benchmarks/results/profile-k16.md +0 -12
  77. package/benchmarks/results/profile-k4.json +0 -347
  78. package/benchmarks/results/profile-k4.md +0 -12
  79. package/benchmarks/results/profile-pool-reuse.json +0 -331
  80. package/benchmarks/results/profile-pool-reuse.md +0 -12
  81. package/benchmarks/results/profile-precompiled.json +0 -385
  82. package/benchmarks/results/profile-precompiled.md +0 -47
  83. package/benchmarks/results/profile-static-channels.json +0 -385
  84. package/benchmarks/results/profile-static-channels.md +0 -47
  85. package/benchmarks/results/profile-static-io.json +0 -385
  86. package/benchmarks/results/profile-static-io.md +0 -47
  87. package/benchmarks/results/profile-tiled-conv.json +0 -331
  88. package/benchmarks/results/profile-tiled-conv.md +0 -12
  89. package/benchmarks/results/profile-tiled-decoder.json +0 -347
  90. package/benchmarks/results/profile-tiled-decoder.md +0 -12
  91. package/benchmarks/results/profile-tiled-matmul.json +0 -331
  92. package/benchmarks/results/profile-tiled-matmul.md +0 -12
  93. package/benchmarks/results/profile-unfused-decoder.json +0 -379
  94. package/benchmarks/results/profile-unfused-decoder.md +0 -12
  95. package/benchmarks/results/profile-unfused-pool.json +0 -347
  96. package/benchmarks/results/profile-unfused-pool.md +0 -12
  97. package/benchmarks/results/spatial-auto.json +0 -575
  98. package/benchmarks/results/spatial-auto.md +0 -61
  99. package/benchmarks/results/subgroup-smoke.json +0 -1094
  100. package/benchmarks/results/subgroup-smoke.md +0 -104
  101. package/benchmarks/results/webnn-smoke.json +0 -739
  102. package/benchmarks/results/webnn-smoke.md +0 -76
  103. package/scripts/inspect-model.mjs +0 -64
  104. package/tests/modelSpec.test.mjs +0 -128
  105. package/tests/resourceLifecycle.test.mjs +0 -383
  106. package/tests/tileScheduler.test.mjs +0 -90
package/lib/nativeUNet.js CHANGED
@@ -1,4 +1,5 @@
1
1
  import { Float16Array } from '@petamoriken/float16';
2
+ import { createFinalRgbShader, sharedMemoryBytes as finalRgbSharedMemoryBytes } from './finalRgbShader.js';
2
3
  import { optimizeModelGraph, planModelExecution } from './graphOptimizer.js';
3
4
  import { OIDNResourceTracker } from './resourceTracker.js';
4
5
  const pipelineCachesByDevice = new WeakMap();
@@ -11,13 +12,42 @@ function sharedPipelineCache(device) {
11
12
  return cache;
12
13
  }
13
14
  const WORKGROUP_SIZE = 8;
14
- const TILED_CONV_WORKGROUP = 8;
15
- const TILED_CONV_ROWS_PER_THREAD = 4;
16
- const TILED_CONV_M = TILED_CONV_WORKGROUP * TILED_CONV_ROWS_PER_THREAD;
17
- const TILED_CONV_N_BLOCKS = TILED_CONV_WORKGROUP;
18
15
  const TILED_CONV_K_BLOCKS = 8;
19
16
  const SPATIAL_CONV_WORKGROUP = 8;
20
17
  const SPATIAL_CONV_PATCH = SPATIAL_CONV_WORKGROUP + 2;
18
+ function gemmTile(options) {
19
+ const [workgroupX, workgroupY] = options.workgroupSize;
20
+ return {
21
+ workgroupX,
22
+ workgroupY,
23
+ rowsPerThread: options.rowsPerThread,
24
+ tileM: workgroupY * options.rowsPerThread,
25
+ tileNBlocks: workgroupX,
26
+ // Keep the half partial accumulation grouping independent of tile tuning.
27
+ tileKBlocks: TILED_CONV_K_BLOCKS
28
+ };
29
+ }
30
+ // Holes change only workgroup addresses. Cooperative loads still visit the
31
+ // original dense logical tiles and never read the unused padding elements.
32
+ function gemmSharedSizes(options) {
33
+ const tile = gemmTile(options);
34
+ const paddedInput = options.sharedLayout === 'padded' || options.sharedLayout === 'padded-input';
35
+ const paddedWeights = options.sharedLayout === 'padded' || options.sharedLayout === 'padded-weights';
36
+ return {
37
+ input: tile.tileM * tile.tileKBlocks + (paddedInput ? tile.workgroupY : 0),
38
+ weights: tile.tileKBlocks * tile.tileNBlocks * (paddedWeights ? 5 : 4)
39
+ };
40
+ }
41
+ function gemmSharedInputIndex(index, options) {
42
+ return options.sharedLayout === 'padded' || options.sharedLayout === 'padded-input'
43
+ ? `(${index}) + (${index}) / ${options.rowsPerThread * TILED_CONV_K_BLOCKS}u`
44
+ : index;
45
+ }
46
+ function gemmSharedWeightIndex(index, options) {
47
+ return options.sharedLayout === 'padded' || options.sharedLayout === 'padded-weights'
48
+ ? `(${index}) + (${index}) / 4u`
49
+ : index;
50
+ }
21
51
  function roundUp(value, alignment) {
22
52
  return Math.ceil(value / alignment) * alignment;
23
53
  }
@@ -89,7 +119,7 @@ function createUniformBuffer(device, label, values) {
89
119
  data.set(values);
90
120
  return createMappedBuffer(device, label, data, GPUBufferUsage.UNIFORM);
91
121
  }
92
- function packConvTensors(device, id, tensors, precision) {
122
+ function packConvTensors(device, id, tensors, precision, weightLayout) {
93
123
  const inputBlocks = blocksForChannels(tensors.inputChannels);
94
124
  const outputBlocks = blocksForChannels(tensors.outputChannels);
95
125
  const packedWeightCount = outputBlocks *
@@ -115,15 +145,18 @@ function packConvTensors(device, id, tensors, precision) {
115
145
  const outputChannel = outputBlock * 4 + outputLane;
116
146
  for (let inputLane = 0; inputLane < 4; inputLane++) {
117
147
  const inputChannel = inputBlock * 4 + inputLane;
118
- const packedBlock = ((((outputBlock * tensors.kernelHeight + y) *
119
- tensors.kernelWidth +
120
- x) *
121
- inputBlocks +
122
- inputBlock) *
123
- 16);
148
+ const packedBlock = weightLayout === 'k-major'
149
+ ? ((((y * tensors.kernelWidth + x) * inputBlocks + inputBlock) *
150
+ outputBlocks + outputBlock) * 16)
151
+ : ((((outputBlock * tensors.kernelHeight + y) *
152
+ tensors.kernelWidth +
153
+ x) *
154
+ inputBlocks +
155
+ inputBlock) *
156
+ 16);
124
157
  // One vec4 contains the four output lanes for an input lane.
125
- // Direct and tiled kernels share this layout, keeping model
126
- // descriptors independent from the selected kernel.
158
+ // Both layouts preserve the vec4's output lanes and the
159
+ // convolution's reduction order; only the vec4 address changes.
127
160
  const packedIndex = packedBlock + inputLane * 4 + outputLane;
128
161
  if (outputChannel < tensors.outputChannels &&
129
162
  inputChannel < tensors.inputChannels) {
@@ -144,10 +177,11 @@ function packConvTensors(device, id, tensors, precision) {
144
177
  // in f32. Padded lanes are zero and never escape the final three channels.
145
178
  const packedBias = new Float32Array(outputBlocks * 4);
146
179
  packedBias.set(tensorFloat32Values(tensors.bias));
147
- const weights = createMappedBuffer(device, `oidn/${id}/weights/${precision}`, packedWeights, GPUBufferUsage.STORAGE);
180
+ const weights = createMappedBuffer(device, `oidn/${id}/weights/${precision}/${weightLayout}`, packedWeights, GPUBufferUsage.STORAGE);
148
181
  try {
149
182
  return {
150
183
  weights,
184
+ weightLayout,
151
185
  bias: createMappedBuffer(device, `oidn/${id}/bias`, packedBias, GPUBufferUsage.STORAGE)
152
186
  };
153
187
  }
@@ -213,21 +247,52 @@ let weightBase = ${weightBase};
213
247
  ${accumulator}
214
248
  `;
215
249
  }
216
- function tiledAccumulationCode(precision) {
250
+ function gemmPackedLoad(precision, gemm, weights) {
251
+ return precision === 'fp16' &&
252
+ (gemm.loadMode === 'packed-all' ||
253
+ (weights && gemm.loadMode === 'packed-weights'));
254
+ }
255
+ // vec2<u32> reinterprets exactly the same eight bytes as vec4<f16>.
256
+ function gemmLoadExpression(expression, packed) {
257
+ return packed ? `bitcast<vec4<f16>>(${expression})` : expression;
258
+ }
259
+ function tiledAccumulationCode(precision, gemm) {
260
+ const { rowsPerThread, tileNBlocks, tileKBlocks } = gemmTile(gemm);
261
+ const inputIndex = gemmSharedInputIndex(`tileSpatial * ${tileKBlocks}u + tileK`, gemm);
262
+ const weightBase = `(tileK * ${tileNBlocks}u + localId.x) * ${gemm.sharedLayout === 'padded' || gemm.sharedLayout === 'padded-weights' ? 5 : 4}u`;
263
+ if (gemm.accumulationOrder === 'row-major') {
264
+ const valueType = storageVecType(precision);
265
+ const target = precision === 'fp16' ? 'partial' : 'acc[row]';
266
+ return /* wgsl */ `
267
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
268
+ ${precision === 'fp16' ? 'var partial = vec4<f16>(0.0h);' : ''}
269
+ let tileSpatial = localId.y * ${rowsPerThread}u + row;
270
+ for (var tileK = 0u; tileK < ${tileKBlocks}u; tileK++) {
271
+ let weightBase = ${weightBase};
272
+ let inputValue = inputTile[${inputIndex}];
273
+ ${target} = fma(weightTile[weightBase], ${valueType}(inputValue.x), ${target});
274
+ ${target} = fma(weightTile[weightBase + 1u], ${valueType}(inputValue.y), ${target});
275
+ ${target} = fma(weightTile[weightBase + 2u], ${valueType}(inputValue.z), ${target});
276
+ ${target} = fma(weightTile[weightBase + 3u], ${valueType}(inputValue.w), ${target});
277
+ }
278
+ ${precision === 'fp16' ? 'acc[row] += vec4<f32>(partial);' : ''}
279
+ }
280
+ `;
281
+ }
217
282
  if (precision === 'fp16') {
218
283
  return /* wgsl */ `
219
- var partial: array<vec4<f16>, ${TILED_CONV_ROWS_PER_THREAD}>;
220
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
284
+ var partial: array<vec4<f16>, ${rowsPerThread}>;
285
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
221
286
  partial[row] = vec4<f16>(0.0h);
222
287
  }
223
- for (var tileK = 0u; tileK < ${TILED_CONV_K_BLOCKS}u; tileK++) {
288
+ for (var tileK = 0u; tileK < ${tileKBlocks}u; tileK++) {
224
289
  let weightBase =
225
- (tileK * ${TILED_CONV_N_BLOCKS}u + localId.x) * 4u;
226
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
290
+ ${weightBase};
291
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
227
292
  let tileSpatial =
228
- localId.y * ${TILED_CONV_ROWS_PER_THREAD}u + row;
293
+ localId.y * ${rowsPerThread}u + row;
229
294
  let inputValue =
230
- inputTile[tileSpatial * ${TILED_CONV_K_BLOCKS}u + tileK];
295
+ inputTile[${inputIndex}];
231
296
  partial[row] = fma(
232
297
  weightTile[weightBase],
233
298
  vec4<f16>(inputValue.x),
@@ -250,20 +315,20 @@ function tiledAccumulationCode(precision) {
250
315
  );
251
316
  }
252
317
  }
253
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
318
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
254
319
  acc[row] += vec4<f32>(partial[row]);
255
320
  }
256
321
  `;
257
322
  }
258
323
  return /* wgsl */ `
259
- for (var tileK = 0u; tileK < ${TILED_CONV_K_BLOCKS}u; tileK++) {
324
+ for (var tileK = 0u; tileK < ${tileKBlocks}u; tileK++) {
260
325
  let weightBase =
261
- (tileK * ${TILED_CONV_N_BLOCKS}u + localId.x) * 4u;
262
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
326
+ ${weightBase};
327
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
263
328
  let tileSpatial =
264
- localId.y * ${TILED_CONV_ROWS_PER_THREAD}u + row;
329
+ localId.y * ${rowsPerThread}u + row;
265
330
  let inputValue =
266
- inputTile[tileSpatial * ${TILED_CONV_K_BLOCKS}u + tileK];
331
+ inputTile[${inputIndex}];
267
332
  acc[row] = fma(
268
333
  weightTile[weightBase],
269
334
  vec4<f32>(inputValue.x),
@@ -474,18 +539,92 @@ fn main(
474
539
  }
475
540
  `;
476
541
  }
542
+ /** Precompute spatial predicates once, then advance the original flattened K order. */
543
+ function tiledAddressCacheCode(inputBlocks, decoder, gemm, sources) {
544
+ const { workgroupX, workgroupY, tileM, tileKBlocks } = gemmTile(gemm);
545
+ const width = decoder ? 'params.outputWidth' : 'params.inputWidth';
546
+ const height = decoder ? 'params.outputHeight' : 'params.inputHeight';
547
+ const baseOffset = gemm.addressMode === 'base-offset';
548
+ // Decoder parity travels in the high mask bits, avoiding extra cached arrays.
549
+ const baseAssignments = sources
550
+ ? sources.blocks.map((blocks, source) => `cachedBase${source}[load] = (${source === sources.upsampled ? 'y / 2u' : 'y'} * params.source${source}Width + ${source === sources.upsampled ? 'x / 2u' : 'x'}) * ${blocks}u;`).join('\n ')
551
+ : `cachedBase0[load] = (y * params.inputWidth + x) * ${inputBlocks}u;`;
552
+ const loads = tileM * tileKBlocks /
553
+ (workgroupX * workgroupY);
554
+ return /* wgsl */ `
555
+ ${baseOffset ? `var cachedBase0: array<u32, ${loads}>;
556
+ ${sources ? `var cachedBase1: array<u32, ${loads}>;` : ''}` : `var cachedX: array<u32, ${loads}>;
557
+ var cachedY: array<u32, ${loads}>;`}
558
+ var cachedMask: array<u32, ${loads}>;
559
+ for (var load = 0u; load < ${loads}u; load++) {
560
+ let loadIndex = localLinear + load * ${workgroupX * workgroupY}u;
561
+ let spatial = workgroupId.x * ${tileM}u + loadIndex / ${tileKBlocks}u;
562
+ let x = spatial % params.outputWidth;
563
+ let y = spatial / params.outputWidth;
564
+ ${baseOffset ? baseAssignments : `cachedX[load] = x;
565
+ cachedY[load] = y;`}
566
+ let columns =
567
+ select(0u, 0x049u, x > 0u && x - 1u < ${width}) |
568
+ select(0u, 0x092u, x < ${width}) |
569
+ select(0u, 0x124u, x + 1u < ${width});
570
+ let rows =
571
+ select(0u, 0x007u, y > 0u && y - 1u < ${height}) |
572
+ select(0u, 0x038u, y < ${height}) |
573
+ select(0u, 0x1c0u, y + 1u < ${height});
574
+ cachedMask[load] = select(0u, columns & rows, spatial < spatialCount)${baseOffset && sources ? ' | ((x & 1u) << 9u) | ((y & 1u) << 10u)' : ''};
575
+ }
576
+ var channelBlock = (localLinear % ${tileKBlocks}u) % ${inputBlocks}u;
577
+ var filterPosition = (localLinear % ${tileKBlocks}u) / ${inputBlocks}u;
578
+ `;
579
+ }
580
+ function tiledAddressAdvanceCode(inputBlocks, gemm) {
581
+ const { tileKBlocks } = gemmTile(gemm);
582
+ // The remainder is less than inputBlocks, so at most one carry is needed.
583
+ return /* wgsl */ `
584
+ channelBlock += ${tileKBlocks % inputBlocks}u;
585
+ let carry = channelBlock >= ${inputBlocks}u;
586
+ channelBlock -= select(0u, ${inputBlocks}u, carry);
587
+ filterPosition += ${Math.floor(tileKBlocks / inputBlocks)}u + select(0u, 1u, carry);
588
+ `;
589
+ }
590
+ function tiledIncrementalInputCode(valueType, readValue, gemm) {
591
+ const { workgroupX, workgroupY, tileM, tileKBlocks } = gemmTile(gemm);
592
+ const loads = tileM * tileKBlocks /
593
+ (workgroupX * workgroupY);
594
+ return /* wgsl */ `
595
+ let inputBlock = channelBlock;
596
+ let kernelY = filterPosition / 3u;
597
+ let kernelX = filterPosition % 3u;
598
+ for (var load = 0u; load < ${loads}u; load++) {
599
+ let loadIndex = localLinear + load * ${workgroupX * workgroupY}u;
600
+ var value = ${valueType}(0.0);
601
+ if (filterPosition < 9u && (cachedMask[load] & (1u << filterPosition)) != 0u) {
602
+ ${gemm.addressMode === 'base-offset' ? '' : `let inputX = i32(cachedX[load]) + i32(kernelX) - 1;
603
+ let inputY = i32(cachedY[load]) + i32(kernelY) - 1;`}
604
+ ${readValue}
605
+ }
606
+ inputTile[${gemmSharedInputIndex('loadIndex', gemm)}] = value;
607
+ }
608
+ `;
609
+ }
477
610
  /**
478
- * Implicit-GEMM convolution for FP32. This follows the proven packed WebGPU
479
- * shape: one 8x8 workgroup computes 32 spatial rows by 8 vec4 output blocks,
480
- * with four output rows per thread and an eight-vec4 K tile.
611
+ * Implicit-GEMM convolution for FP16/FP32. This follows the packed WebGPU
612
+ * shape: configurable spatial rows and vec4 output blocks per workgroup.
613
+ * The eight-vec4 K tile stays fixed to preserve FP16 rounding across variants.
481
614
  */
482
- function createTiledConvShader(precision, outputPrecision, activation, inputBlocks, outputBlocks) {
615
+ function createTiledConvShader(precision, outputPrecision, activation, inputBlocks, outputBlocks, gemm) {
616
+ const { workgroupX, workgroupY, rowsPerThread, tileM, tileNBlocks, tileKBlocks } = gemmTile(gemm);
617
+ const { addressMode, weightLayout } = gemm;
483
618
  const inputType = storageVecType(precision);
484
619
  const outputType = storageVecType(outputPrecision);
620
+ const packedInput = gemmPackedLoad(precision, gemm, false);
621
+ const packedWeights = gemmPackedLoad(precision, gemm, true);
622
+ const inputStorageType = packedInput ? 'vec2<u32>' : inputType;
623
+ const weightStorageType = packedWeights ? 'vec2<u32>' : inputType;
485
624
  const stored = storeExpression(activationExpression('acc', activation), outputPrecision);
486
- const workgroupThreads = TILED_CONV_WORKGROUP * TILED_CONV_WORKGROUP;
487
- const inputTileValues = TILED_CONV_M * TILED_CONV_K_BLOCKS;
488
- const weightTileValues = TILED_CONV_K_BLOCKS * TILED_CONV_N_BLOCKS * 4;
625
+ const workgroupThreads = workgroupX * workgroupY;
626
+ const inputTileValues = tileM * tileKBlocks;
627
+ const weightTileValues = tileKBlocks * tileNBlocks * 4;
489
628
  return /* wgsl */ `${shaderPreamble(precision)}
490
629
  struct Params {
491
630
  inputWidth: u32,
@@ -495,46 +634,53 @@ struct Params {
495
634
  inputBlocks: u32,
496
635
  outputBlocks: u32,
497
636
  }
498
- @group(0) @binding(0) var<storage, read> inputData: array<${inputType}>;
499
- @group(0) @binding(1) var<storage, read> weights: array<${inputType}>;
637
+ @group(0) @binding(0) var<storage, read> inputData: array<${inputStorageType}>;
638
+ @group(0) @binding(1) var<storage, read> weights: array<${weightStorageType}>;
500
639
  @group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
501
640
  @group(0) @binding(3) var<storage, read_write> outputData: array<${outputType}>;
502
641
  @group(0) @binding(4) var<uniform> params: Params;
503
642
 
504
- var<workgroup> inputTile: array<${inputType}, ${inputTileValues}>;
505
- var<workgroup> weightTile: array<${inputType}, ${weightTileValues}>;
643
+ var<workgroup> inputTile: array<${inputType}, ${gemmSharedSizes(gemm).input}>;
644
+ var<workgroup> weightTile: array<${inputType}, ${gemmSharedSizes(gemm).weights}>;
506
645
 
507
- @compute @workgroup_size(${TILED_CONV_WORKGROUP}, ${TILED_CONV_WORKGROUP}, 1)
646
+ @compute @workgroup_size(${workgroupX}, ${workgroupY}, 1)
508
647
  fn main(
509
648
  @builtin(local_invocation_id) localId: vec3<u32>,
510
649
  @builtin(workgroup_id) workgroupId: vec3<u32>
511
650
  ) {
512
651
  let spatialBase =
513
- workgroupId.x * ${TILED_CONV_M}u +
514
- localId.y * ${TILED_CONV_ROWS_PER_THREAD}u;
652
+ workgroupId.x * ${tileM}u +
653
+ localId.y * ${rowsPerThread}u;
515
654
  let outputBlock =
516
- workgroupId.y * ${TILED_CONV_N_BLOCKS}u + localId.x;
655
+ workgroupId.y * ${tileNBlocks}u + localId.x;
517
656
  let spatialCount = params.outputWidth * params.outputHeight;
518
- var acc: array<vec4<f32>, ${TILED_CONV_ROWS_PER_THREAD}>;
657
+ var acc: array<vec4<f32>, ${rowsPerThread}>;
519
658
  if (outputBlock < ${outputBlocks}u) {
520
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
659
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
521
660
  acc[row] = bias[outputBlock];
522
661
  }
523
662
  }
524
663
 
525
664
  let localLinear =
526
- localId.y * ${TILED_CONV_WORKGROUP}u + localId.x;
665
+ localId.y * ${workgroupX}u + localId.x;
666
+ ${addressMode !== 'analytic' ? tiledAddressCacheCode(inputBlocks, false, gemm) : ''}
527
667
  let totalK = ${inputBlocks * 9}u;
528
- for (var kBase = 0u; kBase < totalK; kBase += ${TILED_CONV_K_BLOCKS}u) {
668
+ for (var kBase = 0u; kBase < totalK; kBase += ${tileKBlocks}u) {
669
+ ${addressMode !== 'analytic' ? tiledIncrementalInputCode(inputType, `
670
+ ${addressMode === 'base-offset' ? `let offset = (i32(kernelY) - 1) * i32(params.inputWidth) + i32(kernelX) - 1;
671
+ let inputIndex = cachedBase0[load] + u32(offset * ${inputBlocks}) + inputBlock;` : `let inputIndex = (u32(inputY) * params.inputWidth + u32(inputX)) *
672
+ ${inputBlocks}u + inputBlock;`}
673
+ value = ${gemmLoadExpression('inputData[inputIndex]', packedInput)};
674
+ `, gemm) : /* wgsl */ `
529
675
  for (
530
676
  var loadIndex = localLinear;
531
677
  loadIndex < ${inputTileValues}u;
532
678
  loadIndex += ${workgroupThreads}u
533
679
  ) {
534
- let tileSpatial = loadIndex / ${TILED_CONV_K_BLOCKS}u;
535
- let tileK = loadIndex % ${TILED_CONV_K_BLOCKS}u;
680
+ let tileSpatial = loadIndex / ${tileKBlocks}u;
681
+ let tileK = loadIndex % ${tileKBlocks}u;
536
682
  let inputSpatialIndex =
537
- workgroupId.x * ${TILED_CONV_M}u + tileSpatial;
683
+ workgroupId.x * ${tileM}u + tileSpatial;
538
684
  let kIndex = kBase + tileK;
539
685
  var value = ${inputType}(0.0);
540
686
  if (inputSpatialIndex < spatialCount && kIndex < totalK) {
@@ -553,26 +699,32 @@ fn main(
553
699
  let inputIndex =
554
700
  (u32(inputY) * params.inputWidth + u32(inputX)) *
555
701
  ${inputBlocks}u + inputBlock;
556
- value = inputData[inputIndex];
702
+ value = ${gemmLoadExpression('inputData[inputIndex]', packedInput)};
557
703
  }
558
704
  }
559
- inputTile[loadIndex] = value;
705
+ inputTile[${gemmSharedInputIndex('loadIndex', gemm)}] = value;
560
706
  }
707
+ `}
561
708
 
562
709
  for (
563
710
  var loadIndex = localLinear;
564
711
  loadIndex < ${weightTileValues}u;
565
712
  loadIndex += ${workgroupThreads}u
566
713
  ) {
567
- let tileK = loadIndex / ${TILED_CONV_N_BLOCKS * 4}u;
568
- let outputRemainder = loadIndex % ${TILED_CONV_N_BLOCKS * 4}u;
714
+ let tileK = loadIndex / ${tileNBlocks * 4}u;
715
+ let outputRemainder = loadIndex % ${tileNBlocks * 4}u;
569
716
  let tileOutputBlock = outputRemainder / 4u;
570
717
  let outputLane = outputRemainder % 4u;
571
718
  let loadedOutputBlock =
572
- workgroupId.y * ${TILED_CONV_N_BLOCKS}u + tileOutputBlock;
719
+ workgroupId.y * ${tileNBlocks}u + tileOutputBlock;
573
720
  let kIndex = kBase + tileK;
574
721
  var value = ${inputType}(0.0);
575
722
  if (loadedOutputBlock < ${outputBlocks}u && kIndex < totalK) {
723
+ ${weightLayout === 'k-major' ? /* wgsl */ `
724
+ let weightIndex = (kIndex * ${outputBlocks}u + loadedOutputBlock) * 4u + outputLane;
725
+ ` : addressMode !== 'analytic' ? /* wgsl */ `
726
+ let weightIndex = (loadedOutputBlock * ${inputBlocks * 9}u + kIndex) * 4u + outputLane;
727
+ ` : /* wgsl */ `
576
728
  let inputBlock = kIndex % ${inputBlocks}u;
577
729
  let kernelIndex = kIndex / ${inputBlocks}u;
578
730
  let kernelY = kernelIndex / 3u;
@@ -580,18 +732,20 @@ fn main(
580
732
  let weightIndex =
581
733
  ((((loadedOutputBlock * 3u + kernelY) * 3u + kernelX) *
582
734
  ${inputBlocks}u + inputBlock) * 4u + outputLane);
583
- value = weights[weightIndex];
735
+ `}
736
+ value = ${gemmLoadExpression('weights[weightIndex]', packedWeights)};
584
737
  }
585
- weightTile[loadIndex] = value;
738
+ weightTile[${gemmSharedWeightIndex('loadIndex', gemm)}] = value;
586
739
  }
587
740
 
588
741
  workgroupBarrier();
589
- ${tiledAccumulationCode(precision)}
742
+ ${tiledAccumulationCode(precision, gemm)}
590
743
  workgroupBarrier();
744
+ ${addressMode !== 'analytic' ? tiledAddressAdvanceCode(inputBlocks, gemm) : ''}
591
745
  }
592
746
 
593
747
  if (outputBlock < ${outputBlocks}u) {
594
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
748
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
595
749
  let spatialIndex = spatialBase + row;
596
750
  if (spatialIndex < spatialCount) {
597
751
  let outputIndex = spatialIndex * ${outputBlocks}u + outputBlock;
@@ -653,7 +807,7 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
653
807
  }
654
808
  `;
655
809
  }
656
- function createMaxPoolShader(precision, outputBlocks) {
810
+ function createMaxPoolShader(precision, outputBlocks, coalesced) {
657
811
  const valueType = storageVecType(precision);
658
812
  return /* wgsl */ `${shaderPreamble(precision)}
659
813
  struct Params {
@@ -668,11 +822,17 @@ struct Params {
668
822
  @group(0) @binding(2) var<uniform> params: Params;
669
823
 
670
824
  @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
671
- fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
825
+ fn main(
826
+ @builtin(global_invocation_id) gid: vec3<u32>${coalesced ? ',\n @builtin(num_workgroups) groupCount: vec3<u32>' : ''}
827
+ ) {
828
+ ${coalesced ? `let linearX = gid.x + gid.z * groupCount.x * ${WORKGROUP_SIZE}u;
829
+ let outputBlock = linearX % ${outputBlocks}u;
830
+ let outputX = linearX / ${outputBlocks}u;` : `let outputBlock = gid.z;
831
+ let outputX = gid.x;`}
672
832
  if (
673
- gid.x >= params.outputWidth ||
833
+ outputX >= params.outputWidth ||
674
834
  gid.y >= params.outputHeight ||
675
- gid.z >= ${outputBlocks}u
835
+ outputBlock >= ${outputBlocks}u
676
836
  ) {
677
837
  return;
678
838
  }
@@ -681,26 +841,33 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
681
841
  let inputY = gid.y * 2u + py;
682
842
  if (inputY >= params.inputHeight) { continue; }
683
843
  for (var px = 0u; px < 2u; px++) {
684
- let inputX = gid.x * 2u + px;
844
+ let inputX = outputX * 2u + px;
685
845
  if (inputX >= params.inputWidth) { continue; }
686
846
  let inputIndex =
687
- (inputY * params.inputWidth + inputX) * ${outputBlocks}u + gid.z;
847
+ (inputY * params.inputWidth + inputX) * ${outputBlocks}u + outputBlock;
688
848
  pooled = max(pooled, vec4<f32>(inputData[inputIndex]));
689
849
  }
690
850
  }
691
851
  let outputIndex =
692
- (gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
852
+ (gid.y * params.outputWidth + outputX) * ${outputBlocks}u + outputBlock;
693
853
  outputData[outputIndex] = ${storeExpression('pooled', precision)};
694
854
  }
695
855
  `;
696
856
  }
697
- function createTiledDecoderShader(precision, activation, sourceBlocks, upsampledSource, outputBlocks) {
857
+ function createTiledDecoderShader(precision, outputPrecision, activation, sourceBlocks, upsampledSource, outputBlocks, gemm) {
858
+ const { workgroupX, workgroupY, rowsPerThread, tileM, tileNBlocks, tileKBlocks } = gemmTile(gemm);
859
+ const { addressMode, weightLayout } = gemm;
698
860
  const valueType = storageVecType(precision);
861
+ const outputType = storageVecType(outputPrecision);
862
+ const packedInput = gemmPackedLoad(precision, gemm, false);
863
+ const packedWeights = gemmPackedLoad(precision, gemm, true);
864
+ const inputStorageType = packedInput ? 'vec2<u32>' : valueType;
865
+ const weightStorageType = packedWeights ? 'vec2<u32>' : valueType;
699
866
  const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
700
- const stored = storeExpression(activationExpression('acc[row]', activation), precision);
701
- const workgroupThreads = TILED_CONV_WORKGROUP * TILED_CONV_WORKGROUP;
702
- const inputTileValues = TILED_CONV_M * TILED_CONV_K_BLOCKS;
703
- const weightTileValues = TILED_CONV_K_BLOCKS * TILED_CONV_N_BLOCKS * 4;
867
+ const stored = storeExpression(activationExpression('acc[row]', activation), outputPrecision);
868
+ const workgroupThreads = workgroupX * workgroupY;
869
+ const inputTileValues = tileM * tileKBlocks;
870
+ const weightTileValues = tileKBlocks * tileNBlocks * 4;
704
871
  const sourceRead = (source, blockExpression) => {
705
872
  const sourceX = source === upsampledSource
706
873
  ? 'u32(inputX) / 2u'
@@ -711,12 +878,15 @@ function createTiledDecoderShader(precision, activation, sourceBlocks, upsampled
711
878
  return /* wgsl */ `
712
879
  {
713
880
  let sourceBlock = ${blockExpression};
714
- let sourceX = ${sourceX};
881
+ ${addressMode === 'base-offset' ? `let dx = ${source === upsampledSource ? '(i32((cachedMask[load] >> 9u) & 1u) + i32(kernelX) - 1) >> 1u' : 'i32(kernelX) - 1'};
882
+ let dy = ${source === upsampledSource ? '(i32((cachedMask[load] >> 10u) & 1u) + i32(kernelY) - 1) >> 1u' : 'i32(kernelY) - 1'};
883
+ let offset = (dy * i32(params.source${source}Width) + dx) * ${sourceBlocks[source]};
884
+ let sourceIndex = cachedBase${source}[load] + u32(offset) + sourceBlock;` : `let sourceX = ${sourceX};
715
885
  let sourceY = ${sourceY};
716
886
  let sourceIndex =
717
887
  (sourceY * params.source${source}Width + sourceX) *
718
- ${sourceBlocks[source]}u + sourceBlock;
719
- value = input${source}[sourceIndex];
888
+ ${sourceBlocks[source]}u + sourceBlock;`}
889
+ value = ${gemmLoadExpression(`input${source}[sourceIndex]`, packedInput)};
720
890
  }`;
721
891
  };
722
892
  return /* wgsl */ `${shaderPreamble(precision)}
@@ -730,47 +900,61 @@ struct Params {
730
900
  source1Width: u32,
731
901
  source1Height: u32,
732
902
  }
733
- @group(0) @binding(0) var<storage, read> input0: array<${valueType}>;
734
- @group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
735
- @group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
903
+ @group(0) @binding(0) var<storage, read> input0: array<${inputStorageType}>;
904
+ @group(0) @binding(1) var<storage, read> input1: array<${inputStorageType}>;
905
+ @group(0) @binding(2) var<storage, read> weights: array<${weightStorageType}>;
736
906
  @group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
737
- @group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
907
+ @group(0) @binding(4) var<storage, read_write> outputData: array<${outputType}>;
738
908
  @group(0) @binding(5) var<uniform> params: Params;
739
909
 
740
- var<workgroup> inputTile: array<${valueType}, ${inputTileValues}>;
741
- var<workgroup> weightTile: array<${valueType}, ${weightTileValues}>;
910
+ var<workgroup> inputTile: array<${valueType}, ${gemmSharedSizes(gemm).input}>;
911
+ var<workgroup> weightTile: array<${valueType}, ${gemmSharedSizes(gemm).weights}>;
742
912
 
743
- @compute @workgroup_size(${TILED_CONV_WORKGROUP}, ${TILED_CONV_WORKGROUP}, 1)
913
+ @compute @workgroup_size(${workgroupX}, ${workgroupY}, 1)
744
914
  fn main(
745
915
  @builtin(local_invocation_id) localId: vec3<u32>,
746
916
  @builtin(workgroup_id) workgroupId: vec3<u32>
747
917
  ) {
748
918
  let spatialBase =
749
- workgroupId.x * ${TILED_CONV_M}u +
750
- localId.y * ${TILED_CONV_ROWS_PER_THREAD}u;
919
+ workgroupId.x * ${tileM}u +
920
+ localId.y * ${rowsPerThread}u;
751
921
  let outputBlock =
752
- workgroupId.y * ${TILED_CONV_N_BLOCKS}u + localId.x;
922
+ workgroupId.y * ${tileNBlocks}u + localId.x;
753
923
  let spatialCount = params.outputWidth * params.outputHeight;
754
- var acc: array<vec4<f32>, ${TILED_CONV_ROWS_PER_THREAD}>;
924
+ var acc: array<vec4<f32>, ${rowsPerThread}>;
755
925
  if (outputBlock < ${outputBlocks}u) {
756
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
926
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
757
927
  acc[row] = bias[outputBlock];
758
928
  }
759
929
  }
760
930
 
761
931
  let localLinear =
762
- localId.y * ${TILED_CONV_WORKGROUP}u + localId.x;
932
+ localId.y * ${workgroupX}u + localId.x;
933
+ ${addressMode !== 'analytic' ? tiledAddressCacheCode(inputBlocks, true, gemm, { blocks: sourceBlocks, upsampled: upsampledSource }) : ''}
763
934
  let totalK = ${inputBlocks * 9}u;
764
- for (var kBase = 0u; kBase < totalK; kBase += ${TILED_CONV_K_BLOCKS}u) {
935
+ for (var kBase = 0u; kBase < totalK; kBase += ${tileKBlocks}u) {
936
+ ${addressMode !== 'analytic' ? (gemm.decoderLoad === 'source-first' ? `
937
+ if (channelBlock < ${sourceBlocks[0]}u) {
938
+ ${tiledIncrementalInputCode(valueType, sourceRead(0, 'inputBlock'), gemm)}
939
+ } else {
940
+ ${tiledIncrementalInputCode(valueType, sourceRead(1, `inputBlock - ${sourceBlocks[0]}u`), gemm)}
941
+ }
942
+ ` : tiledIncrementalInputCode(valueType, `
943
+ if (inputBlock < ${sourceBlocks[0]}u) {
944
+ ${sourceRead(0, 'inputBlock')}
945
+ } else {
946
+ ${sourceRead(1, `inputBlock - ${sourceBlocks[0]}u`)}
947
+ }
948
+ `, gemm)) : /* wgsl */ `
765
949
  for (
766
950
  var loadIndex = localLinear;
767
951
  loadIndex < ${inputTileValues}u;
768
952
  loadIndex += ${workgroupThreads}u
769
953
  ) {
770
- let tileSpatial = loadIndex / ${TILED_CONV_K_BLOCKS}u;
771
- let tileK = loadIndex % ${TILED_CONV_K_BLOCKS}u;
954
+ let tileSpatial = loadIndex / ${tileKBlocks}u;
955
+ let tileK = loadIndex % ${tileKBlocks}u;
772
956
  let outputSpatialIndex =
773
- workgroupId.x * ${TILED_CONV_M}u + tileSpatial;
957
+ workgroupId.x * ${tileM}u + tileSpatial;
774
958
  let kIndex = kBase + tileK;
775
959
  var value = ${valueType}(0.0);
776
960
  if (outputSpatialIndex < spatialCount && kIndex < totalK) {
@@ -793,23 +977,29 @@ fn main(
793
977
  }
794
978
  }
795
979
  }
796
- inputTile[loadIndex] = value;
980
+ inputTile[${gemmSharedInputIndex('loadIndex', gemm)}] = value;
797
981
  }
982
+ `}
798
983
 
799
984
  for (
800
985
  var loadIndex = localLinear;
801
986
  loadIndex < ${weightTileValues}u;
802
987
  loadIndex += ${workgroupThreads}u
803
988
  ) {
804
- let tileK = loadIndex / ${TILED_CONV_N_BLOCKS * 4}u;
805
- let outputRemainder = loadIndex % ${TILED_CONV_N_BLOCKS * 4}u;
989
+ let tileK = loadIndex / ${tileNBlocks * 4}u;
990
+ let outputRemainder = loadIndex % ${tileNBlocks * 4}u;
806
991
  let tileOutputBlock = outputRemainder / 4u;
807
992
  let outputLane = outputRemainder % 4u;
808
993
  let loadedOutputBlock =
809
- workgroupId.y * ${TILED_CONV_N_BLOCKS}u + tileOutputBlock;
994
+ workgroupId.y * ${tileNBlocks}u + tileOutputBlock;
810
995
  let kIndex = kBase + tileK;
811
996
  var value = ${valueType}(0.0);
812
997
  if (loadedOutputBlock < ${outputBlocks}u && kIndex < totalK) {
998
+ ${weightLayout === 'k-major' ? /* wgsl */ `
999
+ let weightIndex = (kIndex * ${outputBlocks}u + loadedOutputBlock) * 4u + outputLane;
1000
+ ` : addressMode !== 'analytic' ? /* wgsl */ `
1001
+ let weightIndex = (loadedOutputBlock * ${inputBlocks * 9}u + kIndex) * 4u + outputLane;
1002
+ ` : /* wgsl */ `
813
1003
  let inputBlock = kIndex % ${inputBlocks}u;
814
1004
  let kernelIndex = kIndex / ${inputBlocks}u;
815
1005
  let kernelY = kernelIndex / 3u;
@@ -817,18 +1007,20 @@ fn main(
817
1007
  let weightIndex =
818
1008
  ((((loadedOutputBlock * 3u + kernelY) * 3u + kernelX) *
819
1009
  ${inputBlocks}u + inputBlock) * 4u + outputLane);
820
- value = weights[weightIndex];
1010
+ `}
1011
+ value = ${gemmLoadExpression('weights[weightIndex]', packedWeights)};
821
1012
  }
822
- weightTile[loadIndex] = value;
1013
+ weightTile[${gemmSharedWeightIndex('loadIndex', gemm)}] = value;
823
1014
  }
824
1015
 
825
1016
  workgroupBarrier();
826
- ${tiledAccumulationCode(precision)}
1017
+ ${tiledAccumulationCode(precision, gemm)}
827
1018
  workgroupBarrier();
1019
+ ${addressMode !== 'analytic' ? tiledAddressAdvanceCode(inputBlocks, gemm) : ''}
828
1020
  }
829
1021
 
830
1022
  if (outputBlock < ${outputBlocks}u) {
831
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
1023
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
832
1024
  let spatialIndex = spatialBase + row;
833
1025
  if (spatialIndex < spatialCount) {
834
1026
  let outputIndex = spatialIndex * ${outputBlocks}u + outputBlock;
@@ -839,8 +1031,9 @@ fn main(
839
1031
  }
840
1032
  `;
841
1033
  }
842
- function createFusedDecoderShader(precision, activation, sourceBlocks, upsampledSource, outputBlocks) {
1034
+ function createFusedDecoderShader(precision, outputPrecision, activation, sourceBlocks, upsampledSource, outputBlocks) {
843
1035
  const valueType = storageVecType(precision);
1036
+ const outputType = storageVecType(outputPrecision);
844
1037
  const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
845
1038
  const sourceCode = (source, blockOffset) => {
846
1039
  const isUpsampled = source === upsampledSource;
@@ -856,7 +1049,7 @@ function createFusedDecoderShader(precision, activation, sourceBlocks, upsampled
856
1049
  }
857
1050
  `;
858
1051
  };
859
- const stored = storeExpression(activationExpression('acc', activation), precision);
1052
+ const stored = storeExpression(activationExpression('acc', activation), outputPrecision);
860
1053
  return /* wgsl */ `${shaderPreamble(precision)}
861
1054
  struct Params {
862
1055
  outputWidth: u32,
@@ -872,7 +1065,7 @@ struct Params {
872
1065
  @group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
873
1066
  @group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
874
1067
  @group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
875
- @group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
1068
+ @group(0) @binding(4) var<storage, read_write> outputData: array<${outputType}>;
876
1069
  @group(0) @binding(5) var<uniform> params: Params;
877
1070
 
878
1071
  @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
@@ -896,12 +1089,13 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
896
1089
  }
897
1090
  `;
898
1091
  }
899
- function createSpatialDecoderShader(precision, activation, sourceBlocks, upsampledSource, outputBlocks) {
1092
+ function createSpatialDecoderShader(precision, outputPrecision, activation, sourceBlocks, upsampledSource, outputBlocks) {
900
1093
  const valueType = storageVecType(precision);
1094
+ const outputType = storageVecType(outputPrecision);
901
1095
  const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
902
1096
  const patchValues = SPATIAL_CONV_PATCH * SPATIAL_CONV_PATCH * inputBlocks;
903
1097
  const workgroupThreads = SPATIAL_CONV_WORKGROUP * SPATIAL_CONV_WORKGROUP;
904
- const stored = storeExpression(activationExpression('acc', activation), precision);
1098
+ const stored = storeExpression(activationExpression('acc', activation), outputPrecision);
905
1099
  const sourceRead = (source, blockExpression) => {
906
1100
  const sourceX = source === upsampledSource ? 'u32(inputX) / 2u' : 'u32(inputX)';
907
1101
  const sourceY = source === upsampledSource ? 'u32(inputY) / 2u' : 'u32(inputY)';
@@ -927,7 +1121,7 @@ struct Params {
927
1121
  @group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
928
1122
  @group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
929
1123
  @group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
930
- @group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
1124
+ @group(0) @binding(4) var<storage, read_write> outputData: array<${outputType}>;
931
1125
  @group(0) @binding(5) var<uniform> params: Params;
932
1126
 
933
1127
  var<workgroup> inputPatch: array<${valueType}, ${patchValues}>;
@@ -1060,9 +1254,11 @@ export class NativeUNetExecutor {
1060
1254
  _device;
1061
1255
  precision;
1062
1256
  kernelSetting;
1257
+ gemm;
1063
1258
  maxSpatialInputBlocks;
1064
1259
  subgroupsAvailable;
1065
1260
  _model;
1261
+ _gemmByOutputBlocks = new Map();
1066
1262
  _packedConvs = new Map();
1067
1263
  _pipelineCache;
1068
1264
  _pipelinePromises;
@@ -1082,6 +1278,60 @@ export class NativeUNetExecutor {
1082
1278
  this._pipelinePromises = pipelineCache.pending;
1083
1279
  this.precision = resolveNativeUNetPrecision(_device, options.precision ?? 'auto');
1084
1280
  this.kernelSetting = options.kernel ?? 'auto';
1281
+ const workgroupSize = options.gemm?.workgroupSize ?? [8, 8];
1282
+ this.gemm = Object.freeze({
1283
+ tilePolicy: options.gemm?.tilePolicy ?? (options.gemm?.workgroupSize ? 'fixed' : 'output-aligned'),
1284
+ decoderLoad: options.gemm?.decoderLoad ?? 'per-load',
1285
+ poolLayout: options.gemm?.poolLayout ?? 'channels',
1286
+ sharedLayout: options.gemm?.sharedLayout ?? (this.precision === 'fp16' ? 'padded-input' : 'padded'),
1287
+ accumulationOrder: options.gemm?.accumulationOrder ?? 'k-major',
1288
+ finalLayer: options.gemm?.finalLayer ?? 'shared-auto',
1289
+ loadMode: options.gemm?.loadMode ?? 'native',
1290
+ addressMode: options.gemm?.addressMode ?? 'incremental',
1291
+ weightLayout: options.gemm?.weightLayout ?? 'k-major',
1292
+ rowsPerThread: options.gemm?.rowsPerThread ?? 8,
1293
+ workgroupSize: Object.freeze([workgroupSize[0], workgroupSize[1]])
1294
+ });
1295
+ if (!['fixed', 'output-aligned'].includes(this.gemm.tilePolicy)) {
1296
+ throw new Error(`Unsupported GEMM tile policy: ${this.gemm.tilePolicy}`);
1297
+ }
1298
+ if (!['per-load', 'source-first'].includes(this.gemm.decoderLoad)) {
1299
+ throw new Error(`Unsupported GEMM decoder load: ${this.gemm.decoderLoad}`);
1300
+ }
1301
+ if (!['spatial', 'channels'].includes(this.gemm.poolLayout)) {
1302
+ throw new Error(`Unsupported GEMM pool layout: ${this.gemm.poolLayout}`);
1303
+ }
1304
+ if (!['linear', 'padded', 'padded-input', 'padded-weights'].includes(this.gemm.sharedLayout)) {
1305
+ throw new Error(`Unsupported GEMM shared layout: ${this.gemm.sharedLayout}`);
1306
+ }
1307
+ if (!['k-major', 'row-major'].includes(this.gemm.accumulationOrder)) {
1308
+ throw new Error(`Unsupported GEMM accumulation order: ${this.gemm.accumulationOrder}`);
1309
+ }
1310
+ if (!['direct', 'shared-input', 'shared-input-weights', 'shared-auto'].includes(this.gemm.finalLayer)) {
1311
+ throw new Error(`Unsupported GEMM final layer: ${this.gemm.finalLayer}`);
1312
+ }
1313
+ if (!['native', 'packed-weights', 'packed-all'].includes(this.gemm.loadMode)) {
1314
+ throw new Error(`Unsupported GEMM load mode: ${this.gemm.loadMode}`);
1315
+ }
1316
+ if (!['analytic', 'incremental', 'base-offset'].includes(this.gemm.addressMode)) {
1317
+ throw new Error(`Unsupported GEMM address mode: ${this.gemm.addressMode}`);
1318
+ }
1319
+ if (!['output-major', 'k-major'].includes(this.gemm.weightLayout)) {
1320
+ throw new Error(`Unsupported GEMM weight layout: ${this.gemm.weightLayout}`);
1321
+ }
1322
+ const tile = gemmTile(this.gemm);
1323
+ const tileStorageBytes = (gemmSharedSizes(this.gemm).input + gemmSharedSizes(this.gemm).weights) *
1324
+ 4 * (this.precision === 'fp16' ? 2 : 4);
1325
+ if (![2, 4, 8].includes(tile.rowsPerThread) ||
1326
+ ![4, 8, 16].includes(tile.workgroupX) ||
1327
+ ![4, 8].includes(tile.workgroupY) ||
1328
+ workgroupSize.length !== 2 ||
1329
+ tile.workgroupX > _device.limits.maxComputeWorkgroupSizeX ||
1330
+ tile.workgroupY > _device.limits.maxComputeWorkgroupSizeY ||
1331
+ tile.workgroupX * tile.workgroupY > _device.limits.maxComputeInvocationsPerWorkgroup ||
1332
+ tileStorageBytes > _device.limits.maxComputeWorkgroupStorageSize) {
1333
+ throw new Error('Unsupported GEMM tile configuration for this GPUDevice');
1334
+ }
1085
1335
  this.subgroupsAvailable = _device.features.has('subgroups');
1086
1336
  this.maxSpatialInputBlocks =
1087
1337
  this.precision === 'fp16' &&
@@ -1118,7 +1368,14 @@ export class NativeUNetExecutor {
1118
1368
  }
1119
1369
  try {
1120
1370
  for (const [id, tensors] of model.convTensors) {
1121
- const packed = packConvTensors(_device, id, tensors, this.precision);
1371
+ // Kernel selection is shape-independent. Pack only the active layout
1372
+ // for this executor, keeping direct/spatial/subgroup and final weights
1373
+ // in their original ABI without duplicating GPU allocations.
1374
+ const kernel = this._selectConvKernel(blocksForChannels(tensors.inputChannels), id === model.spec.output);
1375
+ const weightLayout = kernel === 'implicit-gemm'
1376
+ ? this.gemm.weightLayout
1377
+ : 'output-major';
1378
+ const packed = packConvTensors(_device, id, tensors, this.precision, weightLayout);
1122
1379
  this._resources.track('gpu-buffer', packed.weights);
1123
1380
  this._resources.track('gpu-buffer', packed.bias);
1124
1381
  this._packedConvs.set(id, packed);
@@ -1184,15 +1441,38 @@ export class NativeUNetExecutor {
1184
1441
  const outputPrecision = isFinal ? 'fp32' : this.precision;
1185
1442
  const inputBlocks = blocksForChannels(this._model.convChannels.get(node.id).inputChannels);
1186
1443
  const outputBlocks = blocksForChannels(this._model.convChannels.get(node.id).outputChannels);
1444
+ const gemm = this._gemmForOutput(outputBlocks);
1187
1445
  const kernel = this._selectConvKernel(inputBlocks, isFinal);
1446
+ const cacheFinalWeights = gemm.finalLayer === 'shared-input-weights' ||
1447
+ (gemm.finalLayer === 'shared-auto' &&
1448
+ finalRgbSharedMemoryBytes(this.precision, inputBlocks, true) <=
1449
+ this._device.limits.maxComputeWorkgroupStorageSize);
1450
+ if (isFinal && this._model.convChannels.get(node.id).outputChannels === 3 &&
1451
+ (this.kernelSetting === 'auto' || this.kernelSetting === 'implicit-gemm') &&
1452
+ gemm.finalLayer !== 'direct' &&
1453
+ this._device.limits.maxComputeWorkgroupSizeX >= 8 &&
1454
+ this._device.limits.maxComputeWorkgroupSizeY >= 8 &&
1455
+ this._device.limits.maxComputeInvocationsPerWorkgroup >= 64 &&
1456
+ finalRgbSharedMemoryBytes(this.precision, inputBlocks, cacheFinalWeights) <=
1457
+ this._device.limits.maxComputeWorkgroupStorageSize) {
1458
+ return {
1459
+ key: `conv-final-rgb/${this.precision}/${node.activation}/in${inputBlocks}/weights-${cacheFinalWeights ? 'shared' : 'storage'}`,
1460
+ kernel: 'direct',
1461
+ code: createFinalRgbShader(this.precision, node.activation, inputBlocks, cacheFinalWeights)
1462
+ };
1463
+ }
1188
1464
  const key = `conv-${kernel}/${this.precision}/` +
1189
1465
  `${outputPrecision}/${node.activation}/` +
1190
- `in${inputBlocks}/out${outputBlocks}`;
1466
+ `in${inputBlocks}/out${outputBlocks}` +
1467
+ (kernel === 'implicit-gemm'
1468
+ ? `/address-${gemm.addressMode}/weights-${gemm.weightLayout}-v1` +
1469
+ `/tile-${gemm.workgroupSize.join('x')}-r${gemm.rowsPerThread}/loads-${gemm.loadMode}/shared-${gemm.sharedLayout}/acc-${gemm.accumulationOrder}`
1470
+ : '');
1191
1471
  return {
1192
1472
  key,
1193
1473
  kernel,
1194
1474
  code: kernel === 'implicit-gemm'
1195
- ? createTiledConvShader(this.precision, outputPrecision, node.activation, inputBlocks, outputBlocks)
1475
+ ? createTiledConvShader(this.precision, outputPrecision, node.activation, inputBlocks, outputBlocks, gemm)
1196
1476
  : kernel === 'spatial'
1197
1477
  ? createSpatialConvShader(this.precision, outputPrecision, node.activation, inputBlocks, outputBlocks)
1198
1478
  : kernel === 'subgroup'
@@ -1202,11 +1482,12 @@ export class NativeUNetExecutor {
1202
1482
  }
1203
1483
  if (node.op === 'maxPool2d') {
1204
1484
  const outputBlocks = blocksForChannels(this._model.channelsByValue.get(node.id));
1205
- const key = `max-pool/${this.precision}/out${outputBlocks}`;
1485
+ const coalesced = this._coalescedPool();
1486
+ const key = `max-pool/${this.precision}/out${outputBlocks}/${coalesced ? 'channels' : 'spatial'}`;
1206
1487
  return {
1207
1488
  key,
1208
1489
  kernel: 'direct',
1209
- code: createMaxPoolShader(this.precision, outputBlocks)
1490
+ code: createMaxPoolShader(this.precision, outputBlocks, coalesced)
1210
1491
  };
1211
1492
  }
1212
1493
  if (node.op === 'fusedConvReluMaxPool2d') {
@@ -1230,28 +1511,63 @@ export class NativeUNetExecutor {
1230
1511
  throw new Error(`Native fused decoder ${node.id} has no upsample input`);
1231
1512
  }
1232
1513
  const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
1233
- const selectedKernel = this._selectConvKernel(inputBlocks, false);
1514
+ const outputPrecision = isFinal ? 'fp32' : this.precision;
1515
+ const selectedKernel = this._selectConvKernel(inputBlocks, isFinal);
1234
1516
  // The subgroup broadcast path currently targets the common standalone
1235
1517
  // convolution layout; fused decoder reads use the direct kernel.
1236
1518
  const kernel = selectedKernel === 'subgroup'
1237
1519
  ? 'direct'
1238
1520
  : selectedKernel;
1239
1521
  const outputBlocks = blocksForChannels(this._model.convChannels.get(node.conv.id).outputChannels);
1240
- const key = `decoder-${kernel}/${this.precision}/` +
1522
+ const gemm = this._gemmForOutput(outputBlocks);
1523
+ const key = `decoder-${kernel}/${this.precision}/${outputPrecision}/` +
1241
1524
  `${node.conv.activation}/` +
1242
- `${sourceBlocks.join('+')}/out${outputBlocks}/up${upsampledSource}`;
1525
+ `${sourceBlocks.join('+')}/out${outputBlocks}/up${upsampledSource}` +
1526
+ (kernel === 'implicit-gemm'
1527
+ ? `/address-${gemm.addressMode}/weights-${gemm.weightLayout}-v1` +
1528
+ `/tile-${gemm.workgroupSize.join('x')}-r${gemm.rowsPerThread}/loads-${gemm.loadMode}/shared-${gemm.sharedLayout}/acc-${gemm.accumulationOrder}/decoder-${gemm.decoderLoad}`
1529
+ : '');
1243
1530
  return {
1244
1531
  key,
1245
1532
  kernel,
1246
1533
  code: kernel === 'implicit-gemm'
1247
- ? createTiledDecoderShader(this.precision, node.conv.activation, sourceBlocks, upsampledSource, outputBlocks)
1534
+ ? createTiledDecoderShader(this.precision, outputPrecision, node.conv.activation, sourceBlocks, upsampledSource, outputBlocks, gemm)
1248
1535
  : kernel === 'spatial'
1249
- ? createSpatialDecoderShader(this.precision, node.conv.activation, sourceBlocks, upsampledSource, outputBlocks)
1250
- : createFusedDecoderShader(this.precision, node.conv.activation, sourceBlocks, upsampledSource, outputBlocks)
1536
+ ? createSpatialDecoderShader(this.precision, outputPrecision, node.conv.activation, sourceBlocks, upsampledSource, outputBlocks)
1537
+ : createFusedDecoderShader(this.precision, outputPrecision, node.conv.activation, sourceBlocks, upsampledSource, outputBlocks)
1251
1538
  };
1252
1539
  }
1253
1540
  throw new Error(`Native OIDN does not implement unfused ${node.op} node ${node.id}`);
1254
1541
  }
1542
+ // Wider output tiles reuse each input load across twice as many channels.
1543
+ // Require full output blocks and keep the validated base tile as fallback.
1544
+ _gemmForOutput(outputBlocks) {
1545
+ if (this.gemm.tilePolicy !== 'output-aligned' ||
1546
+ this.gemm.workgroupSize[0] !== 8 || outputBlocks <= 0 || outputBlocks % 16 !== 0) {
1547
+ return this.gemm;
1548
+ }
1549
+ const cached = this._gemmByOutputBlocks.get(outputBlocks);
1550
+ if (cached)
1551
+ return cached;
1552
+ const candidate = Object.freeze({
1553
+ ...this.gemm,
1554
+ workgroupSize: Object.freeze([16, this.gemm.workgroupSize[1]])
1555
+ });
1556
+ const tile = gemmTile(candidate);
1557
+ const sizes = gemmSharedSizes(candidate);
1558
+ const limits = this._device.limits;
1559
+ const fits = tile.workgroupX <= limits.maxComputeWorkgroupSizeX &&
1560
+ tile.workgroupY <= limits.maxComputeWorkgroupSizeY &&
1561
+ tile.workgroupX * tile.workgroupY <= limits.maxComputeInvocationsPerWorkgroup &&
1562
+ (sizes.input + sizes.weights) * 4 * (this.precision === 'fp16' ? 2 : 4) <= limits.maxComputeWorkgroupStorageSize;
1563
+ const selected = fits ? candidate : this.gemm;
1564
+ this._gemmByOutputBlocks.set(outputBlocks, selected);
1565
+ return selected;
1566
+ }
1567
+ _coalescedPool() {
1568
+ return this.gemm.poolLayout === 'channels' &&
1569
+ (this.kernelSetting === 'auto' || this.kernelSetting === 'implicit-gemm');
1570
+ }
1255
1571
  _selectConvKernel(inputBlocks, isFinal) {
1256
1572
  const spatialFits = this.precision === 'fp16' &&
1257
1573
  inputBlocks <= this.maxSpatialInputBlocks;
@@ -1266,7 +1582,7 @@ export class NativeUNetExecutor {
1266
1582
  if (this.kernelSetting === 'subgroup') {
1267
1583
  return this.subgroupsAvailable ? 'subgroup' : 'direct';
1268
1584
  }
1269
- if (this.precision === 'fp32' && !isFinal)
1585
+ if (!isFinal)
1270
1586
  return 'implicit-gemm';
1271
1587
  return 'direct';
1272
1588
  }
@@ -1502,6 +1818,12 @@ export class NativeUNetExecutor {
1502
1818
  execution.lastUsed = ++this._clock;
1503
1819
  return execution;
1504
1820
  }
1821
+ /** Allocates shape-dependent execution resources before interactive use. */
1822
+ prewarm(shapes) {
1823
+ for (const shape of shapes) {
1824
+ this._execution(shape.width, shape.height);
1825
+ }
1826
+ }
1505
1827
  /** Captures per-pass GPU timestamps for the next execute call when supported. */
1506
1828
  profileNextExecution() {
1507
1829
  if (!this._device.features.has('timestamp-query'))
@@ -1595,8 +1917,15 @@ export class NativeUNetExecutor {
1595
1917
  const pass = encoder.beginComputePass(passDescriptor(`oidn/${execution.plan.nodes[index].id}`, index + 1));
1596
1918
  pass.setPipeline(execution.nodePipelines[index]);
1597
1919
  pass.setBindGroup(0, execution.nodeBindings[index]);
1598
- if (execution.nodeKernels[index] === 'implicit-gemm') {
1599
- pass.dispatchWorkgroups(Math.ceil((outputShape.width * outputShape.height) / TILED_CONV_M), Math.ceil(blocksForChannels(outputShape.channels) / TILED_CONV_N_BLOCKS), 1);
1920
+ if (node.op === 'maxPool2d' && this._coalescedPool()) {
1921
+ const groupsX = Math.ceil(outputShape.width * blocksForChannels(outputShape.channels) / WORKGROUP_SIZE);
1922
+ const maxGroups = this._device.limits.maxComputeWorkgroupsPerDimension;
1923
+ // Wide channel rows continue in Z instead of exceeding the X limit.
1924
+ pass.dispatchWorkgroups(Math.min(groupsX, maxGroups), Math.ceil(outputShape.height / WORKGROUP_SIZE), Math.ceil(groupsX / maxGroups));
1925
+ }
1926
+ else if (execution.nodeKernels[index] === 'implicit-gemm') {
1927
+ const tile = gemmTile(this._gemmForOutput(blocksForChannels(outputShape.channels)));
1928
+ pass.dispatchWorkgroups(Math.ceil((outputShape.width * outputShape.height) / tile.tileM), Math.ceil(blocksForChannels(outputShape.channels) / tile.tileNBlocks), 1);
1600
1929
  }
1601
1930
  else {
1602
1931
  pass.dispatchWorkgroups(Math.ceil(outputShape.width / WORKGROUP_SIZE), Math.ceil(outputShape.height / WORKGROUP_SIZE), blocksForChannels(outputShape.channels));
@@ -1644,7 +1973,7 @@ export class NativeUNetExecutor {
1644
1973
  }
1645
1974
  return execution.valueBuffers.get(execution.plan.spec.output);
1646
1975
  }
1647
- /** Compatibility path for ImageData/HDR arrays without TensorFlow.js. */
1976
+ /** Executes interleaved CPU image data through the native GPU runtime. */
1648
1977
  async executeCPU(interleavedInput, width, height) {
1649
1978
  const expectedLength = width * height * this._model.inputChannels;
1650
1979
  if (interleavedInput.length !== expectedLength) {