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