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/src/nativeUNet.ts CHANGED
@@ -1,4 +1,5 @@
1
1
  import { Float16Array } from '@petamoriken/float16';
2
+ import { createFinalRgbShader, sharedMemoryBytes as finalRgbSharedMemoryBytes } from './finalRgbShader.js';
2
3
  import {
3
4
  optimizeModelGraph,
4
5
  planModelExecution,
@@ -26,14 +27,42 @@ export type NativeUNetKernel =
26
27
  | 'spatial'
27
28
  | 'subgroup';
28
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
+ }
29
56
 
30
57
  export interface NativeUNetOptions {
31
58
  precision?: NativeUNetPrecisionSetting;
32
59
  /**
33
- * Convolution kernel selection. `auto` uses a model-independent capability
34
- * heuristic and falls back to the direct kernel when a tile does not fit.
60
+ * Convolution kernel selection. `auto` uses implicit GEMM for FP16/FP32
61
+ * convolutions and the direct kernel for the final output layer.
35
62
  */
36
63
  kernel?: NativeUNetKernelSetting;
64
+ /** Optional implicit-GEMM tuning; other convolution kernels ignore it. */
65
+ gemm?: NativeUNetGemmOptions;
37
66
  /** Maximum number of shape-dependent activation plans retained. */
38
67
  shapeCacheSize?: number;
39
68
  }
@@ -51,6 +80,7 @@ export interface NativeUNetExecutionProfile {
51
80
  interface PackedConvBuffers {
52
81
  weights: GPUBuffer;
53
82
  bias: GPUBuffer;
83
+ weightLayout: NonNullable<NativeUNetGemmOptions['weightLayout']>;
54
84
  }
55
85
 
56
86
  interface NativePipelineSpec {
@@ -97,15 +127,47 @@ function sharedPipelineCache(device: GPUDevice) {
97
127
  }
98
128
 
99
129
  const WORKGROUP_SIZE = 8;
100
- const TILED_CONV_WORKGROUP = 8;
101
- const TILED_CONV_ROWS_PER_THREAD = 4;
102
- const TILED_CONV_M =
103
- TILED_CONV_WORKGROUP * TILED_CONV_ROWS_PER_THREAD;
104
- const TILED_CONV_N_BLOCKS = TILED_CONV_WORKGROUP;
105
130
  const TILED_CONV_K_BLOCKS = 8;
106
131
  const SPATIAL_CONV_WORKGROUP = 8;
107
132
  const SPATIAL_CONV_PATCH = SPATIAL_CONV_WORKGROUP + 2;
108
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
+
109
171
  function roundUp(value: number, alignment: number) {
110
172
  return Math.ceil(value / alignment) * alignment;
111
173
  }
@@ -213,7 +275,8 @@ function packConvTensors(
213
275
  device: GPUDevice,
214
276
  id: string,
215
277
  tensors: ValidatedConvTensor,
216
- precision: NativeUNetPrecision
278
+ precision: NativeUNetPrecision,
279
+ weightLayout: NonNullable<NativeUNetGemmOptions['weightLayout']>
217
280
  ): PackedConvBuffers {
218
281
  const inputBlocks = blocksForChannels(tensors.inputChannels);
219
282
  const outputBlocks = blocksForChannels(tensors.outputChannels);
@@ -247,16 +310,18 @@ function packConvTensors(
247
310
  const outputChannel = outputBlock * 4 + outputLane;
248
311
  for (let inputLane = 0; inputLane < 4; inputLane++) {
249
312
  const inputChannel = inputBlock * 4 + inputLane;
250
- const packedBlock =
251
- ((((outputBlock * tensors.kernelHeight + y) *
313
+ const packedBlock = weightLayout === 'k-major'
314
+ ? ((((y * tensors.kernelWidth + x) * inputBlocks + inputBlock) *
315
+ outputBlocks + outputBlock) * 16)
316
+ : ((((outputBlock * tensors.kernelHeight + y) *
252
317
  tensors.kernelWidth +
253
318
  x) *
254
319
  inputBlocks +
255
320
  inputBlock) *
256
321
  16);
257
322
  // One vec4 contains the four output lanes for an input lane.
258
- // Direct and tiled kernels share this layout, keeping model
259
- // descriptors independent from the selected kernel.
323
+ // Both layouts preserve the vec4's output lanes and the
324
+ // convolution's reduction order; only the vec4 address changes.
260
325
  const packedIndex =
261
326
  packedBlock + inputLane * 4 + outputLane;
262
327
  if (
@@ -285,13 +350,14 @@ function packConvTensors(
285
350
 
286
351
  const weights = createMappedBuffer(
287
352
  device,
288
- `oidn/${id}/weights/${precision}`,
353
+ `oidn/${id}/weights/${precision}/${weightLayout}`,
289
354
  packedWeights,
290
355
  GPUBufferUsage.STORAGE
291
356
  );
292
357
  try {
293
358
  return {
294
359
  weights,
360
+ weightLayout,
295
361
  bias: createMappedBuffer(
296
362
  device,
297
363
  `oidn/${id}/bias`,
@@ -382,21 +448,58 @@ ${accumulator}
382
448
  `;
383
449
  }
384
450
 
385
- function tiledAccumulationCode(precision: NativeUNetPrecision) {
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
+ }
386
489
  if (precision === 'fp16') {
387
490
  return /* wgsl */ `
388
- var partial: array<vec4<f16>, ${TILED_CONV_ROWS_PER_THREAD}>;
389
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
491
+ var partial: array<vec4<f16>, ${rowsPerThread}>;
492
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
390
493
  partial[row] = vec4<f16>(0.0h);
391
494
  }
392
- for (var tileK = 0u; tileK < ${TILED_CONV_K_BLOCKS}u; tileK++) {
495
+ for (var tileK = 0u; tileK < ${tileKBlocks}u; tileK++) {
393
496
  let weightBase =
394
- (tileK * ${TILED_CONV_N_BLOCKS}u + localId.x) * 4u;
395
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
497
+ ${weightBase};
498
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
396
499
  let tileSpatial =
397
- localId.y * ${TILED_CONV_ROWS_PER_THREAD}u + row;
500
+ localId.y * ${rowsPerThread}u + row;
398
501
  let inputValue =
399
- inputTile[tileSpatial * ${TILED_CONV_K_BLOCKS}u + tileK];
502
+ inputTile[${inputIndex}];
400
503
  partial[row] = fma(
401
504
  weightTile[weightBase],
402
505
  vec4<f16>(inputValue.x),
@@ -419,20 +522,20 @@ function tiledAccumulationCode(precision: NativeUNetPrecision) {
419
522
  );
420
523
  }
421
524
  }
422
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
525
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
423
526
  acc[row] += vec4<f32>(partial[row]);
424
527
  }
425
528
  `;
426
529
  }
427
530
  return /* wgsl */ `
428
- for (var tileK = 0u; tileK < ${TILED_CONV_K_BLOCKS}u; tileK++) {
531
+ for (var tileK = 0u; tileK < ${tileKBlocks}u; tileK++) {
429
532
  let weightBase =
430
- (tileK * ${TILED_CONV_N_BLOCKS}u + localId.x) * 4u;
431
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
533
+ ${weightBase};
534
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
432
535
  let tileSpatial =
433
- localId.y * ${TILED_CONV_ROWS_PER_THREAD}u + row;
536
+ localId.y * ${rowsPerThread}u + row;
434
537
  let inputValue =
435
- inputTile[tileSpatial * ${TILED_CONV_K_BLOCKS}u + tileK];
538
+ inputTile[${inputIndex}];
436
539
  acc[row] = fma(
437
540
  weightTile[weightBase],
438
541
  vec4<f32>(inputValue.x),
@@ -690,29 +793,109 @@ fn main(
690
793
  `;
691
794
  }
692
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
+
693
869
  /**
694
- * Implicit-GEMM convolution for FP32. This follows the proven packed WebGPU
695
- * shape: one 8x8 workgroup computes 32 spatial rows by 8 vec4 output blocks,
696
- * with four output rows per thread and an eight-vec4 K tile.
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.
697
873
  */
698
874
  function createTiledConvShader(
699
875
  precision: NativeUNetPrecision,
700
876
  outputPrecision: NativeUNetPrecision,
701
877
  activation: Conv2DNodeSpec['activation'],
702
878
  inputBlocks: number,
703
- outputBlocks: number
879
+ outputBlocks: number,
880
+ gemm: Readonly<Required<NativeUNetGemmOptions>>
704
881
  ) {
882
+ const { workgroupX, workgroupY, rowsPerThread, tileM, tileNBlocks, tileKBlocks } = gemmTile(gemm);
883
+ const { addressMode, weightLayout } = gemm;
705
884
  const inputType = storageVecType(precision);
706
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;
707
890
  const stored = storeExpression(
708
891
  activationExpression('acc', activation),
709
892
  outputPrecision
710
893
  );
711
894
  const workgroupThreads =
712
- TILED_CONV_WORKGROUP * TILED_CONV_WORKGROUP;
713
- const inputTileValues = TILED_CONV_M * TILED_CONV_K_BLOCKS;
895
+ workgroupX * workgroupY;
896
+ const inputTileValues = tileM * tileKBlocks;
714
897
  const weightTileValues =
715
- TILED_CONV_K_BLOCKS * TILED_CONV_N_BLOCKS * 4;
898
+ tileKBlocks * tileNBlocks * 4;
716
899
  return /* wgsl */ `${shaderPreamble(precision)}
717
900
  struct Params {
718
901
  inputWidth: u32,
@@ -722,46 +905,53 @@ struct Params {
722
905
  inputBlocks: u32,
723
906
  outputBlocks: u32,
724
907
  }
725
- @group(0) @binding(0) var<storage, read> inputData: array<${inputType}>;
726
- @group(0) @binding(1) var<storage, read> weights: array<${inputType}>;
908
+ @group(0) @binding(0) var<storage, read> inputData: array<${inputStorageType}>;
909
+ @group(0) @binding(1) var<storage, read> weights: array<${weightStorageType}>;
727
910
  @group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
728
911
  @group(0) @binding(3) var<storage, read_write> outputData: array<${outputType}>;
729
912
  @group(0) @binding(4) var<uniform> params: Params;
730
913
 
731
- var<workgroup> inputTile: array<${inputType}, ${inputTileValues}>;
732
- var<workgroup> weightTile: array<${inputType}, ${weightTileValues}>;
914
+ var<workgroup> inputTile: array<${inputType}, ${gemmSharedSizes(gemm).input}>;
915
+ var<workgroup> weightTile: array<${inputType}, ${gemmSharedSizes(gemm).weights}>;
733
916
 
734
- @compute @workgroup_size(${TILED_CONV_WORKGROUP}, ${TILED_CONV_WORKGROUP}, 1)
917
+ @compute @workgroup_size(${workgroupX}, ${workgroupY}, 1)
735
918
  fn main(
736
919
  @builtin(local_invocation_id) localId: vec3<u32>,
737
920
  @builtin(workgroup_id) workgroupId: vec3<u32>
738
921
  ) {
739
922
  let spatialBase =
740
- workgroupId.x * ${TILED_CONV_M}u +
741
- localId.y * ${TILED_CONV_ROWS_PER_THREAD}u;
923
+ workgroupId.x * ${tileM}u +
924
+ localId.y * ${rowsPerThread}u;
742
925
  let outputBlock =
743
- workgroupId.y * ${TILED_CONV_N_BLOCKS}u + localId.x;
926
+ workgroupId.y * ${tileNBlocks}u + localId.x;
744
927
  let spatialCount = params.outputWidth * params.outputHeight;
745
- var acc: array<vec4<f32>, ${TILED_CONV_ROWS_PER_THREAD}>;
928
+ var acc: array<vec4<f32>, ${rowsPerThread}>;
746
929
  if (outputBlock < ${outputBlocks}u) {
747
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
930
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
748
931
  acc[row] = bias[outputBlock];
749
932
  }
750
933
  }
751
934
 
752
935
  let localLinear =
753
- localId.y * ${TILED_CONV_WORKGROUP}u + localId.x;
936
+ localId.y * ${workgroupX}u + localId.x;
937
+ ${addressMode !== 'analytic' ? tiledAddressCacheCode(inputBlocks, false, gemm) : ''}
754
938
  let totalK = ${inputBlocks * 9}u;
755
- for (var kBase = 0u; kBase < totalK; kBase += ${TILED_CONV_K_BLOCKS}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 */ `
756
946
  for (
757
947
  var loadIndex = localLinear;
758
948
  loadIndex < ${inputTileValues}u;
759
949
  loadIndex += ${workgroupThreads}u
760
950
  ) {
761
- let tileSpatial = loadIndex / ${TILED_CONV_K_BLOCKS}u;
762
- let tileK = loadIndex % ${TILED_CONV_K_BLOCKS}u;
951
+ let tileSpatial = loadIndex / ${tileKBlocks}u;
952
+ let tileK = loadIndex % ${tileKBlocks}u;
763
953
  let inputSpatialIndex =
764
- workgroupId.x * ${TILED_CONV_M}u + tileSpatial;
954
+ workgroupId.x * ${tileM}u + tileSpatial;
765
955
  let kIndex = kBase + tileK;
766
956
  var value = ${inputType}(0.0);
767
957
  if (inputSpatialIndex < spatialCount && kIndex < totalK) {
@@ -780,26 +970,32 @@ fn main(
780
970
  let inputIndex =
781
971
  (u32(inputY) * params.inputWidth + u32(inputX)) *
782
972
  ${inputBlocks}u + inputBlock;
783
- value = inputData[inputIndex];
973
+ value = ${gemmLoadExpression('inputData[inputIndex]', packedInput)};
784
974
  }
785
975
  }
786
- inputTile[loadIndex] = value;
976
+ inputTile[${gemmSharedInputIndex('loadIndex', gemm)}] = value;
787
977
  }
978
+ `}
788
979
 
789
980
  for (
790
981
  var loadIndex = localLinear;
791
982
  loadIndex < ${weightTileValues}u;
792
983
  loadIndex += ${workgroupThreads}u
793
984
  ) {
794
- let tileK = loadIndex / ${TILED_CONV_N_BLOCKS * 4}u;
795
- let outputRemainder = loadIndex % ${TILED_CONV_N_BLOCKS * 4}u;
985
+ let tileK = loadIndex / ${tileNBlocks * 4}u;
986
+ let outputRemainder = loadIndex % ${tileNBlocks * 4}u;
796
987
  let tileOutputBlock = outputRemainder / 4u;
797
988
  let outputLane = outputRemainder % 4u;
798
989
  let loadedOutputBlock =
799
- workgroupId.y * ${TILED_CONV_N_BLOCKS}u + tileOutputBlock;
990
+ workgroupId.y * ${tileNBlocks}u + tileOutputBlock;
800
991
  let kIndex = kBase + tileK;
801
992
  var value = ${inputType}(0.0);
802
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 */ `
803
999
  let inputBlock = kIndex % ${inputBlocks}u;
804
1000
  let kernelIndex = kIndex / ${inputBlocks}u;
805
1001
  let kernelY = kernelIndex / 3u;
@@ -807,18 +1003,20 @@ fn main(
807
1003
  let weightIndex =
808
1004
  ((((loadedOutputBlock * 3u + kernelY) * 3u + kernelX) *
809
1005
  ${inputBlocks}u + inputBlock) * 4u + outputLane);
810
- value = weights[weightIndex];
1006
+ `}
1007
+ value = ${gemmLoadExpression('weights[weightIndex]', packedWeights)};
811
1008
  }
812
- weightTile[loadIndex] = value;
1009
+ weightTile[${gemmSharedWeightIndex('loadIndex', gemm)}] = value;
813
1010
  }
814
1011
 
815
1012
  workgroupBarrier();
816
- ${tiledAccumulationCode(precision)}
1013
+ ${tiledAccumulationCode(precision, gemm)}
817
1014
  workgroupBarrier();
1015
+ ${addressMode !== 'analytic' ? tiledAddressAdvanceCode(inputBlocks, gemm) : ''}
818
1016
  }
819
1017
 
820
1018
  if (outputBlock < ${outputBlocks}u) {
821
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
1019
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
822
1020
  let spatialIndex = spatialBase + row;
823
1021
  if (spatialIndex < spatialCount) {
824
1022
  let outputIndex = spatialIndex * ${outputBlocks}u + outputBlock;
@@ -893,7 +1091,8 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
893
1091
 
894
1092
  function createMaxPoolShader(
895
1093
  precision: NativeUNetPrecision,
896
- outputBlocks: number
1094
+ outputBlocks: number,
1095
+ coalesced: boolean
897
1096
  ) {
898
1097
  const valueType = storageVecType(precision);
899
1098
  return /* wgsl */ `${shaderPreamble(precision)}
@@ -909,11 +1108,17 @@ struct Params {
909
1108
  @group(0) @binding(2) var<uniform> params: Params;
910
1109
 
911
1110
  @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
912
- fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
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;`}
913
1118
  if (
914
- gid.x >= params.outputWidth ||
1119
+ outputX >= params.outputWidth ||
915
1120
  gid.y >= params.outputHeight ||
916
- gid.z >= ${outputBlocks}u
1121
+ outputBlock >= ${outputBlocks}u
917
1122
  ) {
918
1123
  return;
919
1124
  }
@@ -922,15 +1127,15 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
922
1127
  let inputY = gid.y * 2u + py;
923
1128
  if (inputY >= params.inputHeight) { continue; }
924
1129
  for (var px = 0u; px < 2u; px++) {
925
- let inputX = gid.x * 2u + px;
1130
+ let inputX = outputX * 2u + px;
926
1131
  if (inputX >= params.inputWidth) { continue; }
927
1132
  let inputIndex =
928
- (inputY * params.inputWidth + inputX) * ${outputBlocks}u + gid.z;
1133
+ (inputY * params.inputWidth + inputX) * ${outputBlocks}u + outputBlock;
929
1134
  pooled = max(pooled, vec4<f32>(inputData[inputIndex]));
930
1135
  }
931
1136
  }
932
1137
  let outputIndex =
933
- (gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
1138
+ (gid.y * params.outputWidth + outputX) * ${outputBlocks}u + outputBlock;
934
1139
  outputData[outputIndex] = ${storeExpression('pooled', precision)};
935
1140
  }
936
1141
  `;
@@ -938,22 +1143,31 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
938
1143
 
939
1144
  function createTiledDecoderShader(
940
1145
  precision: NativeUNetPrecision,
1146
+ outputPrecision: NativeUNetPrecision,
941
1147
  activation: Conv2DNodeSpec['activation'],
942
1148
  sourceBlocks: readonly [number, number],
943
1149
  upsampledSource: 0 | 1,
944
- outputBlocks: number
1150
+ outputBlocks: number,
1151
+ gemm: Readonly<Required<NativeUNetGemmOptions>>
945
1152
  ) {
1153
+ const { workgroupX, workgroupY, rowsPerThread, tileM, tileNBlocks, tileKBlocks } = gemmTile(gemm);
1154
+ const { addressMode, weightLayout } = gemm;
946
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;
947
1161
  const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
948
1162
  const stored = storeExpression(
949
1163
  activationExpression('acc[row]', activation),
950
- precision
1164
+ outputPrecision
951
1165
  );
952
1166
  const workgroupThreads =
953
- TILED_CONV_WORKGROUP * TILED_CONV_WORKGROUP;
954
- const inputTileValues = TILED_CONV_M * TILED_CONV_K_BLOCKS;
1167
+ workgroupX * workgroupY;
1168
+ const inputTileValues = tileM * tileKBlocks;
955
1169
  const weightTileValues =
956
- TILED_CONV_K_BLOCKS * TILED_CONV_N_BLOCKS * 4;
1170
+ tileKBlocks * tileNBlocks * 4;
957
1171
  const sourceRead = (source: 0 | 1, blockExpression: string) => {
958
1172
  const sourceX =
959
1173
  source === upsampledSource
@@ -966,12 +1180,15 @@ function createTiledDecoderShader(
966
1180
  return /* wgsl */ `
967
1181
  {
968
1182
  let sourceBlock = ${blockExpression};
969
- let sourceX = ${sourceX};
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};
970
1187
  let sourceY = ${sourceY};
971
1188
  let sourceIndex =
972
1189
  (sourceY * params.source${source}Width + sourceX) *
973
- ${sourceBlocks[source]}u + sourceBlock;
974
- value = input${source}[sourceIndex];
1190
+ ${sourceBlocks[source]}u + sourceBlock;` }
1191
+ value = ${gemmLoadExpression(`input${source}[sourceIndex]`, packedInput)};
975
1192
  }`;
976
1193
  };
977
1194
  return /* wgsl */ `${shaderPreamble(precision)}
@@ -985,47 +1202,61 @@ struct Params {
985
1202
  source1Width: u32,
986
1203
  source1Height: u32,
987
1204
  }
988
- @group(0) @binding(0) var<storage, read> input0: array<${valueType}>;
989
- @group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
990
- @group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
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}>;
991
1208
  @group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
992
- @group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
1209
+ @group(0) @binding(4) var<storage, read_write> outputData: array<${outputType}>;
993
1210
  @group(0) @binding(5) var<uniform> params: Params;
994
1211
 
995
- var<workgroup> inputTile: array<${valueType}, ${inputTileValues}>;
996
- var<workgroup> weightTile: array<${valueType}, ${weightTileValues}>;
1212
+ var<workgroup> inputTile: array<${valueType}, ${gemmSharedSizes(gemm).input}>;
1213
+ var<workgroup> weightTile: array<${valueType}, ${gemmSharedSizes(gemm).weights}>;
997
1214
 
998
- @compute @workgroup_size(${TILED_CONV_WORKGROUP}, ${TILED_CONV_WORKGROUP}, 1)
1215
+ @compute @workgroup_size(${workgroupX}, ${workgroupY}, 1)
999
1216
  fn main(
1000
1217
  @builtin(local_invocation_id) localId: vec3<u32>,
1001
1218
  @builtin(workgroup_id) workgroupId: vec3<u32>
1002
1219
  ) {
1003
1220
  let spatialBase =
1004
- workgroupId.x * ${TILED_CONV_M}u +
1005
- localId.y * ${TILED_CONV_ROWS_PER_THREAD}u;
1221
+ workgroupId.x * ${tileM}u +
1222
+ localId.y * ${rowsPerThread}u;
1006
1223
  let outputBlock =
1007
- workgroupId.y * ${TILED_CONV_N_BLOCKS}u + localId.x;
1224
+ workgroupId.y * ${tileNBlocks}u + localId.x;
1008
1225
  let spatialCount = params.outputWidth * params.outputHeight;
1009
- var acc: array<vec4<f32>, ${TILED_CONV_ROWS_PER_THREAD}>;
1226
+ var acc: array<vec4<f32>, ${rowsPerThread}>;
1010
1227
  if (outputBlock < ${outputBlocks}u) {
1011
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
1228
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
1012
1229
  acc[row] = bias[outputBlock];
1013
1230
  }
1014
1231
  }
1015
1232
 
1016
1233
  let localLinear =
1017
- localId.y * ${TILED_CONV_WORKGROUP}u + localId.x;
1234
+ localId.y * ${workgroupX}u + localId.x;
1235
+ ${addressMode !== 'analytic' ? tiledAddressCacheCode(inputBlocks, true, gemm, { blocks: sourceBlocks, upsampled: upsampledSource }) : ''}
1018
1236
  let totalK = ${inputBlocks * 9}u;
1019
- for (var kBase = 0u; kBase < totalK; kBase += ${TILED_CONV_K_BLOCKS}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 */ `
1020
1251
  for (
1021
1252
  var loadIndex = localLinear;
1022
1253
  loadIndex < ${inputTileValues}u;
1023
1254
  loadIndex += ${workgroupThreads}u
1024
1255
  ) {
1025
- let tileSpatial = loadIndex / ${TILED_CONV_K_BLOCKS}u;
1026
- let tileK = loadIndex % ${TILED_CONV_K_BLOCKS}u;
1256
+ let tileSpatial = loadIndex / ${tileKBlocks}u;
1257
+ let tileK = loadIndex % ${tileKBlocks}u;
1027
1258
  let outputSpatialIndex =
1028
- workgroupId.x * ${TILED_CONV_M}u + tileSpatial;
1259
+ workgroupId.x * ${tileM}u + tileSpatial;
1029
1260
  let kIndex = kBase + tileK;
1030
1261
  var value = ${valueType}(0.0);
1031
1262
  if (outputSpatialIndex < spatialCount && kIndex < totalK) {
@@ -1048,23 +1279,29 @@ fn main(
1048
1279
  }
1049
1280
  }
1050
1281
  }
1051
- inputTile[loadIndex] = value;
1282
+ inputTile[${gemmSharedInputIndex('loadIndex', gemm)}] = value;
1052
1283
  }
1284
+ `}
1053
1285
 
1054
1286
  for (
1055
1287
  var loadIndex = localLinear;
1056
1288
  loadIndex < ${weightTileValues}u;
1057
1289
  loadIndex += ${workgroupThreads}u
1058
1290
  ) {
1059
- let tileK = loadIndex / ${TILED_CONV_N_BLOCKS * 4}u;
1060
- let outputRemainder = loadIndex % ${TILED_CONV_N_BLOCKS * 4}u;
1291
+ let tileK = loadIndex / ${tileNBlocks * 4}u;
1292
+ let outputRemainder = loadIndex % ${tileNBlocks * 4}u;
1061
1293
  let tileOutputBlock = outputRemainder / 4u;
1062
1294
  let outputLane = outputRemainder % 4u;
1063
1295
  let loadedOutputBlock =
1064
- workgroupId.y * ${TILED_CONV_N_BLOCKS}u + tileOutputBlock;
1296
+ workgroupId.y * ${tileNBlocks}u + tileOutputBlock;
1065
1297
  let kIndex = kBase + tileK;
1066
1298
  var value = ${valueType}(0.0);
1067
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 */ `
1068
1305
  let inputBlock = kIndex % ${inputBlocks}u;
1069
1306
  let kernelIndex = kIndex / ${inputBlocks}u;
1070
1307
  let kernelY = kernelIndex / 3u;
@@ -1072,18 +1309,20 @@ fn main(
1072
1309
  let weightIndex =
1073
1310
  ((((loadedOutputBlock * 3u + kernelY) * 3u + kernelX) *
1074
1311
  ${inputBlocks}u + inputBlock) * 4u + outputLane);
1075
- value = weights[weightIndex];
1312
+ `}
1313
+ value = ${gemmLoadExpression('weights[weightIndex]', packedWeights)};
1076
1314
  }
1077
- weightTile[loadIndex] = value;
1315
+ weightTile[${gemmSharedWeightIndex('loadIndex', gemm)}] = value;
1078
1316
  }
1079
1317
 
1080
1318
  workgroupBarrier();
1081
- ${tiledAccumulationCode(precision)}
1319
+ ${tiledAccumulationCode(precision, gemm)}
1082
1320
  workgroupBarrier();
1321
+ ${addressMode !== 'analytic' ? tiledAddressAdvanceCode(inputBlocks, gemm) : ''}
1083
1322
  }
1084
1323
 
1085
1324
  if (outputBlock < ${outputBlocks}u) {
1086
- for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
1325
+ for (var row = 0u; row < ${rowsPerThread}u; row++) {
1087
1326
  let spatialIndex = spatialBase + row;
1088
1327
  if (spatialIndex < spatialCount) {
1089
1328
  let outputIndex = spatialIndex * ${outputBlocks}u + outputBlock;
@@ -1097,12 +1336,14 @@ fn main(
1097
1336
 
1098
1337
  function createFusedDecoderShader(
1099
1338
  precision: NativeUNetPrecision,
1339
+ outputPrecision: NativeUNetPrecision,
1100
1340
  activation: Conv2DNodeSpec['activation'],
1101
1341
  sourceBlocks: readonly [number, number],
1102
1342
  upsampledSource: 0 | 1,
1103
1343
  outputBlocks: number
1104
1344
  ) {
1105
1345
  const valueType = storageVecType(precision);
1346
+ const outputType = storageVecType(outputPrecision);
1106
1347
  const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
1107
1348
  const sourceCode = (source: 0 | 1, blockOffset: number) => {
1108
1349
  const isUpsampled = source === upsampledSource;
@@ -1124,7 +1365,7 @@ function createFusedDecoderShader(
1124
1365
  };
1125
1366
  const stored = storeExpression(
1126
1367
  activationExpression('acc', activation),
1127
- precision
1368
+ outputPrecision
1128
1369
  );
1129
1370
  return /* wgsl */ `${shaderPreamble(precision)}
1130
1371
  struct Params {
@@ -1141,7 +1382,7 @@ struct Params {
1141
1382
  @group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
1142
1383
  @group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
1143
1384
  @group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
1144
- @group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
1385
+ @group(0) @binding(4) var<storage, read_write> outputData: array<${outputType}>;
1145
1386
  @group(0) @binding(5) var<uniform> params: Params;
1146
1387
 
1147
1388
  @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
@@ -1168,12 +1409,14 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
1168
1409
 
1169
1410
  function createSpatialDecoderShader(
1170
1411
  precision: NativeUNetPrecision,
1412
+ outputPrecision: NativeUNetPrecision,
1171
1413
  activation: Conv2DNodeSpec['activation'],
1172
1414
  sourceBlocks: readonly [number, number],
1173
1415
  upsampledSource: 0 | 1,
1174
1416
  outputBlocks: number
1175
1417
  ) {
1176
1418
  const valueType = storageVecType(precision);
1419
+ const outputType = storageVecType(outputPrecision);
1177
1420
  const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
1178
1421
  const patchValues =
1179
1422
  SPATIAL_CONV_PATCH * SPATIAL_CONV_PATCH * inputBlocks;
@@ -1181,7 +1424,7 @@ function createSpatialDecoderShader(
1181
1424
  SPATIAL_CONV_WORKGROUP * SPATIAL_CONV_WORKGROUP;
1182
1425
  const stored = storeExpression(
1183
1426
  activationExpression('acc', activation),
1184
- precision
1427
+ outputPrecision
1185
1428
  );
1186
1429
  const sourceRead = (source: 0 | 1, blockExpression: string) => {
1187
1430
  const sourceX =
@@ -1211,7 +1454,7 @@ struct Params {
1211
1454
  @group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
1212
1455
  @group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
1213
1456
  @group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
1214
- @group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
1457
+ @group(0) @binding(4) var<storage, read_write> outputData: array<${outputType}>;
1215
1458
  @group(0) @binding(5) var<uniform> params: Params;
1216
1459
 
1217
1460
  var<workgroup> inputPatch: array<${valueType}, ${patchValues}>;
@@ -1361,10 +1604,12 @@ export function resolveNativeUNetPrecision(
1361
1604
  export class NativeUNetExecutor {
1362
1605
  readonly precision: NativeUNetPrecision;
1363
1606
  readonly kernelSetting: NativeUNetKernelSetting;
1607
+ readonly gemm: Readonly<Required<NativeUNetGemmOptions>>;
1364
1608
  readonly maxSpatialInputBlocks: number;
1365
1609
  readonly subgroupsAvailable: boolean;
1366
1610
 
1367
1611
  private _model: UNetModelGraph;
1612
+ private _gemmByOutputBlocks = new Map<number, Readonly<Required<NativeUNetGemmOptions>>>();
1368
1613
  private _packedConvs = new Map<string, PackedConvBuffers>();
1369
1614
  private _pipelineCache: Map<string, GPUComputePipeline>;
1370
1615
  private _pipelinePromises: Map<string, Promise<GPUComputePipeline>>;
@@ -1391,6 +1636,63 @@ export class NativeUNetExecutor {
1391
1636
  options.precision ?? 'auto'
1392
1637
  );
1393
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
+ }
1394
1696
  this.subgroupsAvailable = _device.features.has(
1395
1697
  'subgroups' as GPUFeatureName
1396
1698
  );
@@ -1440,7 +1742,18 @@ export class NativeUNetExecutor {
1440
1742
  }
1441
1743
  try {
1442
1744
  for (const [id, tensors] of model.convTensors) {
1443
- const packed = packConvTensors(_device, id, tensors, this.precision);
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
+ );
1444
1757
  this._resources.track('gpu-buffer', packed.weights);
1445
1758
  this._resources.track('gpu-buffer', packed.bias);
1446
1759
  this._packedConvs.set(id, packed);
@@ -1513,11 +1826,36 @@ export class NativeUNetExecutor {
1513
1826
  const outputBlocks = blocksForChannels(
1514
1827
  this._model.convChannels.get(node.id)!.outputChannels
1515
1828
  );
1829
+ const gemm = this._gemmForOutput(outputBlocks);
1516
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
+ }
1517
1851
  const key =
1518
1852
  `conv-${kernel}/${this.precision}/` +
1519
1853
  `${outputPrecision}/${node.activation}/` +
1520
- `in${inputBlocks}/out${outputBlocks}`;
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
+ : '');
1521
1859
  return {
1522
1860
  key,
1523
1861
  kernel,
@@ -1527,7 +1865,8 @@ export class NativeUNetExecutor {
1527
1865
  outputPrecision,
1528
1866
  node.activation,
1529
1867
  inputBlocks,
1530
- outputBlocks
1868
+ outputBlocks,
1869
+ gemm
1531
1870
  )
1532
1871
  : kernel === 'spatial'
1533
1872
  ? createSpatialConvShader(
@@ -1558,11 +1897,12 @@ export class NativeUNetExecutor {
1558
1897
  const outputBlocks = blocksForChannels(
1559
1898
  this._model.channelsByValue.get(node.id)!
1560
1899
  );
1561
- const key = `max-pool/${this.precision}/out${outputBlocks}`;
1900
+ const coalesced = this._coalescedPool();
1901
+ const key = `max-pool/${this.precision}/out${outputBlocks}/${coalesced ? 'channels' : 'spatial'}`;
1562
1902
  return {
1563
1903
  key,
1564
1904
  kernel: 'direct',
1565
- code: createMaxPoolShader(this.precision, outputBlocks)
1905
+ code: createMaxPoolShader(this.precision, outputBlocks, coalesced)
1566
1906
  };
1567
1907
  }
1568
1908
  if (node.op === 'fusedConvReluMaxPool2d') {
@@ -1598,7 +1938,8 @@ export class NativeUNetExecutor {
1598
1938
  throw new Error(`Native fused decoder ${node.id} has no upsample input`);
1599
1939
  }
1600
1940
  const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
1601
- const selectedKernel = this._selectConvKernel(inputBlocks, false);
1941
+ const outputPrecision = isFinal ? 'fp32' : this.precision;
1942
+ const selectedKernel = this._selectConvKernel(inputBlocks, isFinal);
1602
1943
  // The subgroup broadcast path currently targets the common standalone
1603
1944
  // convolution layout; fused decoder reads use the direct kernel.
1604
1945
  const kernel = selectedKernel === 'subgroup'
@@ -1607,24 +1948,32 @@ export class NativeUNetExecutor {
1607
1948
  const outputBlocks = blocksForChannels(
1608
1949
  this._model.convChannels.get(node.conv.id)!.outputChannels
1609
1950
  );
1951
+ const gemm = this._gemmForOutput(outputBlocks);
1610
1952
  const key =
1611
- `decoder-${kernel}/${this.precision}/` +
1953
+ `decoder-${kernel}/${this.precision}/${outputPrecision}/` +
1612
1954
  `${node.conv.activation}/` +
1613
- `${sourceBlocks.join('+')}/out${outputBlocks}/up${upsampledSource}`;
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
+ : '');
1614
1960
  return {
1615
1961
  key,
1616
1962
  kernel,
1617
1963
  code: kernel === 'implicit-gemm'
1618
1964
  ? createTiledDecoderShader(
1619
1965
  this.precision,
1966
+ outputPrecision,
1620
1967
  node.conv.activation,
1621
1968
  sourceBlocks,
1622
1969
  upsampledSource,
1623
- outputBlocks
1970
+ outputBlocks,
1971
+ gemm
1624
1972
  )
1625
1973
  : kernel === 'spatial'
1626
1974
  ? createSpatialDecoderShader(
1627
1975
  this.precision,
1976
+ outputPrecision,
1628
1977
  node.conv.activation,
1629
1978
  sourceBlocks,
1630
1979
  upsampledSource,
@@ -1632,6 +1981,7 @@ export class NativeUNetExecutor {
1632
1981
  )
1633
1982
  : createFusedDecoderShader(
1634
1983
  this.precision,
1984
+ outputPrecision,
1635
1985
  node.conv.activation,
1636
1986
  sourceBlocks,
1637
1987
  upsampledSource,
@@ -1644,6 +1994,36 @@ export class NativeUNetExecutor {
1644
1994
  );
1645
1995
  }
1646
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
+
1647
2027
  private _selectConvKernel(
1648
2028
  inputBlocks: number,
1649
2029
  isFinal: boolean
@@ -1661,7 +2041,7 @@ export class NativeUNetExecutor {
1661
2041
  if (this.kernelSetting === 'subgroup') {
1662
2042
  return this.subgroupsAvailable ? 'subgroup' : 'direct';
1663
2043
  }
1664
- if (this.precision === 'fp32' && !isFinal) return 'implicit-gemm';
2044
+ if (!isFinal) return 'implicit-gemm';
1665
2045
  return 'direct';
1666
2046
  }
1667
2047
 
@@ -1939,6 +2319,13 @@ export class NativeUNetExecutor {
1939
2319
  return execution;
1940
2320
  }
1941
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
+
1942
2329
  /** Captures per-pass GPU timestamps for the next execute call when supported. */
1943
2330
  profileNextExecution() {
1944
2331
  if (!this._device.features.has('timestamp-query')) return false;
@@ -2056,13 +2443,25 @@ export class NativeUNetExecutor {
2056
2443
  );
2057
2444
  pass.setPipeline(execution.nodePipelines[index]);
2058
2445
  pass.setBindGroup(0, execution.nodeBindings[index]);
2059
- if (execution.nodeKernels[index] === 'implicit-gemm') {
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)));
2060
2459
  pass.dispatchWorkgroups(
2061
2460
  Math.ceil(
2062
- (outputShape.width * outputShape.height) / TILED_CONV_M
2461
+ (outputShape.width * outputShape.height) / tile.tileM
2063
2462
  ),
2064
2463
  Math.ceil(
2065
- blocksForChannels(outputShape.channels) / TILED_CONV_N_BLOCKS
2464
+ blocksForChannels(outputShape.channels) / tile.tileNBlocks
2066
2465
  ),
2067
2466
  1
2068
2467
  );
@@ -2135,7 +2534,7 @@ export class NativeUNetExecutor {
2135
2534
  return execution.valueBuffers.get(execution.plan.spec.output)!;
2136
2535
  }
2137
2536
 
2138
- /** Compatibility path for ImageData/HDR arrays without TensorFlow.js. */
2537
+ /** Executes interleaved CPU image data through the native GPU runtime. */
2139
2538
  async executeCPU(
2140
2539
  interleavedInput: Float32Array,
2141
2540
  width: number,