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.
- package/CHANGELOG.md +45 -0
- package/README.md +83 -19
- package/dist/oidn.js +3239 -2642
- package/dist/oidn.umd.cjs +512 -296
- package/lib/UNet.d.ts +54 -20
- package/lib/UNet.js +194 -118
- package/lib/UNet.js.map +1 -1
- package/lib/backend.d.ts +1 -8
- package/lib/backend.js +1 -9
- package/lib/backend.js.map +1 -1
- package/lib/finalRgbShader.d.ts +13 -0
- package/lib/finalRgbShader.js +160 -0
- package/lib/finalRgbShader.js.map +1 -0
- package/lib/graphOptimizer.js +1 -2
- package/lib/graphOptimizer.js.map +1 -1
- package/lib/hdrTransfer.d.ts +14 -0
- package/lib/hdrTransfer.js +61 -0
- package/lib/hdrTransfer.js.map +1 -0
- package/lib/main.d.ts +16 -7
- package/lib/main.js +4 -5
- package/lib/main.js.map +1 -1
- package/lib/nativeUNet.d.ts +39 -3
- package/lib/nativeUNet.js +449 -120
- package/lib/nativeUNet.js.map +1 -1
- package/lib/process.d.ts +5 -11
- package/lib/process.js +35 -49
- package/lib/process.js.map +1 -1
- package/lib/tileScheduler.d.ts +32 -4
- package/lib/tileScheduler.js +133 -20
- package/lib/tileScheduler.js.map +1 -1
- package/package.json +9 -2
- package/src/UNet.ts +287 -158
- package/src/backend.ts +1 -14
- package/src/finalRgbShader.ts +186 -0
- package/src/graphOptimizer.ts +1 -2
- package/src/hdrTransfer.ts +88 -0
- package/src/main.ts +28 -13
- package/src/nativeUNet.ts +515 -116
- package/src/process.ts +43 -70
- package/src/tileScheduler.ts +216 -24
- package/benchmarks/compare.mjs +0 -651
- package/benchmarks/leak.mjs +0 -255
- package/benchmarks/results/before-spatial.json +0 -391
- package/benchmarks/results/before-spatial.md +0 -47
- package/benchmarks/results/int8-scan.json +0 -2007
- package/benchmarks/results/int8-scan.md +0 -160
- package/benchmarks/results/int8-w8a8-scan.json +0 -2007
- package/benchmarks/results/int8-w8a8-scan.md +0 -160
- package/benchmarks/results/int8-weight-channel.json +0 -1413
- package/benchmarks/results/int8-weight-channel.md +0 -118
- package/benchmarks/results/int8-weight-only.json +0 -1437
- package/benchmarks/results/int8-weight-only.md +0 -118
- package/benchmarks/results/kernel-webnn-final.json +0 -1115
- package/benchmarks/results/kernel-webnn-final.md +0 -104
- package/benchmarks/results/latest-optimized.json +0 -375
- package/benchmarks/results/latest-optimized.md +0 -47
- package/benchmarks/results/latest.json +0 -391
- package/benchmarks/results/latest.md +0 -47
- package/benchmarks/results/profile-baseline.json +0 -331
- package/benchmarks/results/profile-baseline.md +0 -12
- package/benchmarks/results/profile-conv2x.json +0 -331
- package/benchmarks/results/profile-conv2x.md +0 -12
- package/benchmarks/results/profile-fast-init.json +0 -385
- package/benchmarks/results/profile-fast-init.md +0 -47
- package/benchmarks/results/profile-fp16-fma.json +0 -369
- package/benchmarks/results/profile-fp16-fma.md +0 -47
- package/benchmarks/results/profile-fp16-tiled-decoder.json +0 -369
- package/benchmarks/results/profile-fp16-tiled-decoder.md +0 -47
- package/benchmarks/results/profile-fp16-tiled-encoder.json +0 -369
- package/benchmarks/results/profile-fp16-tiled-encoder.md +0 -47
- package/benchmarks/results/profile-fp16-unfused-pool.json +0 -385
- package/benchmarks/results/profile-fp16-unfused-pool.md +0 -47
- package/benchmarks/results/profile-input-major.json +0 -347
- package/benchmarks/results/profile-input-major.md +0 -12
- package/benchmarks/results/profile-k16.json +0 -347
- package/benchmarks/results/profile-k16.md +0 -12
- package/benchmarks/results/profile-k4.json +0 -347
- package/benchmarks/results/profile-k4.md +0 -12
- package/benchmarks/results/profile-pool-reuse.json +0 -331
- package/benchmarks/results/profile-pool-reuse.md +0 -12
- package/benchmarks/results/profile-precompiled.json +0 -385
- package/benchmarks/results/profile-precompiled.md +0 -47
- package/benchmarks/results/profile-static-channels.json +0 -385
- package/benchmarks/results/profile-static-channels.md +0 -47
- package/benchmarks/results/profile-static-io.json +0 -385
- package/benchmarks/results/profile-static-io.md +0 -47
- package/benchmarks/results/profile-tiled-conv.json +0 -331
- package/benchmarks/results/profile-tiled-conv.md +0 -12
- package/benchmarks/results/profile-tiled-decoder.json +0 -347
- package/benchmarks/results/profile-tiled-decoder.md +0 -12
- package/benchmarks/results/profile-tiled-matmul.json +0 -331
- package/benchmarks/results/profile-tiled-matmul.md +0 -12
- package/benchmarks/results/profile-unfused-decoder.json +0 -379
- package/benchmarks/results/profile-unfused-decoder.md +0 -12
- package/benchmarks/results/profile-unfused-pool.json +0 -347
- package/benchmarks/results/profile-unfused-pool.md +0 -12
- package/benchmarks/results/spatial-auto.json +0 -575
- package/benchmarks/results/spatial-auto.md +0 -61
- package/benchmarks/results/subgroup-smoke.json +0 -1094
- package/benchmarks/results/subgroup-smoke.md +0 -104
- package/benchmarks/results/webnn-smoke.json +0 -739
- package/benchmarks/results/webnn-smoke.md +0 -76
- package/scripts/inspect-model.mjs +0 -64
- package/tests/modelSpec.test.mjs +0 -128
- package/tests/resourceLifecycle.test.mjs +0 -383
- 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
|
|
34
|
-
*
|
|
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
|
-
((((
|
|
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
|
-
//
|
|
259
|
-
//
|
|
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
|
|
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>, ${
|
|
389
|
-
for (var row = 0u; 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 < ${
|
|
495
|
+
for (var tileK = 0u; tileK < ${tileKBlocks}u; tileK++) {
|
|
393
496
|
let weightBase =
|
|
394
|
-
|
|
395
|
-
for (var row = 0u; row < ${
|
|
497
|
+
${weightBase};
|
|
498
|
+
for (var row = 0u; row < ${rowsPerThread}u; row++) {
|
|
396
499
|
let tileSpatial =
|
|
397
|
-
localId.y * ${
|
|
500
|
+
localId.y * ${rowsPerThread}u + row;
|
|
398
501
|
let inputValue =
|
|
399
|
-
inputTile[
|
|
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 < ${
|
|
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 < ${
|
|
531
|
+
for (var tileK = 0u; tileK < ${tileKBlocks}u; tileK++) {
|
|
429
532
|
let weightBase =
|
|
430
|
-
|
|
431
|
-
for (var row = 0u; row < ${
|
|
533
|
+
${weightBase};
|
|
534
|
+
for (var row = 0u; row < ${rowsPerThread}u; row++) {
|
|
432
535
|
let tileSpatial =
|
|
433
|
-
localId.y * ${
|
|
536
|
+
localId.y * ${rowsPerThread}u + row;
|
|
434
537
|
let inputValue =
|
|
435
|
-
inputTile[
|
|
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
|
|
695
|
-
* shape:
|
|
696
|
-
*
|
|
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
|
-
|
|
713
|
-
const inputTileValues =
|
|
895
|
+
workgroupX * workgroupY;
|
|
896
|
+
const inputTileValues = tileM * tileKBlocks;
|
|
714
897
|
const weightTileValues =
|
|
715
|
-
|
|
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<${
|
|
726
|
-
@group(0) @binding(1) var<storage, read> weights: array<${
|
|
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}, ${
|
|
732
|
-
var<workgroup> weightTile: array<${inputType}, ${
|
|
914
|
+
var<workgroup> inputTile: array<${inputType}, ${gemmSharedSizes(gemm).input}>;
|
|
915
|
+
var<workgroup> weightTile: array<${inputType}, ${gemmSharedSizes(gemm).weights}>;
|
|
733
916
|
|
|
734
|
-
@compute @workgroup_size(${
|
|
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 * ${
|
|
741
|
-
localId.y * ${
|
|
923
|
+
workgroupId.x * ${tileM}u +
|
|
924
|
+
localId.y * ${rowsPerThread}u;
|
|
742
925
|
let outputBlock =
|
|
743
|
-
workgroupId.y * ${
|
|
926
|
+
workgroupId.y * ${tileNBlocks}u + localId.x;
|
|
744
927
|
let spatialCount = params.outputWidth * params.outputHeight;
|
|
745
|
-
var acc: array<vec4<f32>, ${
|
|
928
|
+
var acc: array<vec4<f32>, ${rowsPerThread}>;
|
|
746
929
|
if (outputBlock < ${outputBlocks}u) {
|
|
747
|
-
for (var row = 0u; 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 * ${
|
|
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 += ${
|
|
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 / ${
|
|
762
|
-
let tileK = loadIndex % ${
|
|
951
|
+
let tileSpatial = loadIndex / ${tileKBlocks}u;
|
|
952
|
+
let tileK = loadIndex % ${tileKBlocks}u;
|
|
763
953
|
let inputSpatialIndex =
|
|
764
|
-
workgroupId.x * ${
|
|
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 / ${
|
|
795
|
-
let outputRemainder = loadIndex % ${
|
|
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 * ${
|
|
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
|
-
|
|
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 < ${
|
|
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(
|
|
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
|
-
|
|
1119
|
+
outputX >= params.outputWidth ||
|
|
915
1120
|
gid.y >= params.outputHeight ||
|
|
916
|
-
|
|
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 =
|
|
1130
|
+
let inputX = outputX * 2u + px;
|
|
926
1131
|
if (inputX >= params.inputWidth) { continue; }
|
|
927
1132
|
let inputIndex =
|
|
928
|
-
(inputY * params.inputWidth + inputX) * ${outputBlocks}u +
|
|
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 +
|
|
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
|
-
|
|
1164
|
+
outputPrecision
|
|
951
1165
|
);
|
|
952
1166
|
const workgroupThreads =
|
|
953
|
-
|
|
954
|
-
const inputTileValues =
|
|
1167
|
+
workgroupX * workgroupY;
|
|
1168
|
+
const inputTileValues = tileM * tileKBlocks;
|
|
955
1169
|
const weightTileValues =
|
|
956
|
-
|
|
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
|
|
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<${
|
|
989
|
-
@group(0) @binding(1) var<storage, read> input1: array<${
|
|
990
|
-
@group(0) @binding(2) var<storage, read> weights: array<${
|
|
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<${
|
|
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}, ${
|
|
996
|
-
var<workgroup> weightTile: array<${valueType}, ${
|
|
1212
|
+
var<workgroup> inputTile: array<${valueType}, ${gemmSharedSizes(gemm).input}>;
|
|
1213
|
+
var<workgroup> weightTile: array<${valueType}, ${gemmSharedSizes(gemm).weights}>;
|
|
997
1214
|
|
|
998
|
-
@compute @workgroup_size(${
|
|
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 * ${
|
|
1005
|
-
localId.y * ${
|
|
1221
|
+
workgroupId.x * ${tileM}u +
|
|
1222
|
+
localId.y * ${rowsPerThread}u;
|
|
1006
1223
|
let outputBlock =
|
|
1007
|
-
workgroupId.y * ${
|
|
1224
|
+
workgroupId.y * ${tileNBlocks}u + localId.x;
|
|
1008
1225
|
let spatialCount = params.outputWidth * params.outputHeight;
|
|
1009
|
-
var acc: array<vec4<f32>, ${
|
|
1226
|
+
var acc: array<vec4<f32>, ${rowsPerThread}>;
|
|
1010
1227
|
if (outputBlock < ${outputBlocks}u) {
|
|
1011
|
-
for (var row = 0u; 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 * ${
|
|
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 += ${
|
|
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 / ${
|
|
1026
|
-
let tileK = loadIndex % ${
|
|
1256
|
+
let tileSpatial = loadIndex / ${tileKBlocks}u;
|
|
1257
|
+
let tileK = loadIndex % ${tileKBlocks}u;
|
|
1027
1258
|
let outputSpatialIndex =
|
|
1028
|
-
workgroupId.x * ${
|
|
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 / ${
|
|
1060
|
-
let outputRemainder = loadIndex % ${
|
|
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 * ${
|
|
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
|
-
|
|
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 < ${
|
|
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
|
-
|
|
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<${
|
|
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
|
-
|
|
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<${
|
|
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
|
-
|
|
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
|
|
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
|
|
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 (
|
|
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 (
|
|
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) /
|
|
2461
|
+
(outputShape.width * outputShape.height) / tile.tileM
|
|
2063
2462
|
),
|
|
2064
2463
|
Math.ceil(
|
|
2065
|
-
blocksForChannels(outputShape.channels) /
|
|
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
|
-
/**
|
|
2537
|
+
/** Executes interleaved CPU image data through the native GPU runtime. */
|
|
2139
2538
|
async executeCPU(
|
|
2140
2539
|
interleavedInput: Float32Array,
|
|
2141
2540
|
width: number,
|