oidn-web 0.3.4 → 0.4.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 +49 -0
- package/README.md +140 -4
- package/benchmarks/compare.mjs +651 -0
- package/benchmarks/leak.mjs +255 -0
- package/benchmarks/results/before-spatial.json +391 -0
- package/benchmarks/results/before-spatial.md +47 -0
- package/benchmarks/results/int8-scan.json +2007 -0
- package/benchmarks/results/int8-scan.md +160 -0
- package/benchmarks/results/int8-w8a8-scan.json +2007 -0
- package/benchmarks/results/int8-w8a8-scan.md +160 -0
- package/benchmarks/results/int8-weight-channel.json +1413 -0
- package/benchmarks/results/int8-weight-channel.md +118 -0
- package/benchmarks/results/int8-weight-only.json +1437 -0
- package/benchmarks/results/int8-weight-only.md +118 -0
- package/benchmarks/results/kernel-webnn-final.json +1115 -0
- package/benchmarks/results/kernel-webnn-final.md +104 -0
- package/benchmarks/results/latest-optimized.json +375 -0
- package/benchmarks/results/latest-optimized.md +47 -0
- package/benchmarks/results/latest.json +391 -0
- package/benchmarks/results/latest.md +47 -0
- package/benchmarks/results/profile-baseline.json +331 -0
- package/benchmarks/results/profile-baseline.md +12 -0
- package/benchmarks/results/profile-conv2x.json +331 -0
- package/benchmarks/results/profile-conv2x.md +12 -0
- package/benchmarks/results/profile-fast-init.json +385 -0
- package/benchmarks/results/profile-fast-init.md +47 -0
- package/benchmarks/results/profile-fp16-fma.json +369 -0
- package/benchmarks/results/profile-fp16-fma.md +47 -0
- package/benchmarks/results/profile-fp16-tiled-decoder.json +369 -0
- package/benchmarks/results/profile-fp16-tiled-decoder.md +47 -0
- package/benchmarks/results/profile-fp16-tiled-encoder.json +369 -0
- package/benchmarks/results/profile-fp16-tiled-encoder.md +47 -0
- package/benchmarks/results/profile-fp16-unfused-pool.json +385 -0
- package/benchmarks/results/profile-fp16-unfused-pool.md +47 -0
- package/benchmarks/results/profile-input-major.json +347 -0
- package/benchmarks/results/profile-input-major.md +12 -0
- package/benchmarks/results/profile-k16.json +347 -0
- package/benchmarks/results/profile-k16.md +12 -0
- package/benchmarks/results/profile-k4.json +347 -0
- package/benchmarks/results/profile-k4.md +12 -0
- package/benchmarks/results/profile-pool-reuse.json +331 -0
- package/benchmarks/results/profile-pool-reuse.md +12 -0
- package/benchmarks/results/profile-precompiled.json +385 -0
- package/benchmarks/results/profile-precompiled.md +47 -0
- package/benchmarks/results/profile-static-channels.json +385 -0
- package/benchmarks/results/profile-static-channels.md +47 -0
- package/benchmarks/results/profile-static-io.json +385 -0
- package/benchmarks/results/profile-static-io.md +47 -0
- package/benchmarks/results/profile-tiled-conv.json +331 -0
- package/benchmarks/results/profile-tiled-conv.md +12 -0
- package/benchmarks/results/profile-tiled-decoder.json +347 -0
- package/benchmarks/results/profile-tiled-decoder.md +12 -0
- package/benchmarks/results/profile-tiled-matmul.json +331 -0
- package/benchmarks/results/profile-tiled-matmul.md +12 -0
- package/benchmarks/results/profile-unfused-decoder.json +379 -0
- package/benchmarks/results/profile-unfused-decoder.md +12 -0
- package/benchmarks/results/profile-unfused-pool.json +347 -0
- package/benchmarks/results/profile-unfused-pool.md +12 -0
- package/benchmarks/results/spatial-auto.json +575 -0
- package/benchmarks/results/spatial-auto.md +61 -0
- package/benchmarks/results/subgroup-smoke.json +1094 -0
- package/benchmarks/results/subgroup-smoke.md +104 -0
- package/benchmarks/results/webnn-smoke.json +739 -0
- package/benchmarks/results/webnn-smoke.md +76 -0
- package/dist/oidn.js +4189 -22603
- package/dist/oidn.umd.cjs +784 -5807
- package/lib/UNet.d.ts +66 -15
- package/lib/UNet.js +162 -257
- package/lib/UNet.js.map +1 -1
- package/lib/WGPUComputePass.d.ts +1 -1
- package/lib/WGPUComputePass.js +6 -4
- package/lib/WGPUComputePass.js.map +1 -1
- package/lib/backend.d.ts +8 -4
- package/lib/backend.js +36 -44
- package/lib/backend.js.map +1 -1
- package/lib/graphOptimizer.d.ts +54 -0
- package/lib/graphOptimizer.js +216 -0
- package/lib/graphOptimizer.js.map +1 -0
- package/lib/main.d.ts +33 -10
- package/lib/main.js +5 -0
- package/lib/main.js.map +1 -1
- package/lib/modelSpec.d.ts +80 -0
- package/lib/modelSpec.js +270 -0
- package/lib/modelSpec.js.map +1 -0
- package/lib/nativeUNet.d.ts +67 -0
- package/lib/nativeUNet.js +1735 -0
- package/lib/nativeUNet.js.map +1 -0
- package/lib/process.js +38 -35
- package/lib/process.js.map +1 -1
- package/lib/resourceTracker.d.ts +26 -0
- package/lib/resourceTracker.js +65 -0
- package/lib/resourceTracker.js.map +1 -0
- package/lib/tileScheduler.d.ts +33 -0
- package/lib/tileScheduler.js +86 -0
- package/lib/tileScheduler.js.map +1 -0
- package/lib/webnnUNet.d.ts +52 -0
- package/lib/webnnUNet.js +535 -0
- package/lib/webnnUNet.js.map +1 -0
- package/package.json +9 -5
- package/scripts/inspect-model.mjs +64 -0
- package/src/UNet.ts +236 -339
- package/src/WGPUComputePass.ts +6 -4
- package/src/backend.ts +42 -55
- package/src/graphOptimizer.ts +301 -0
- package/src/main.ts +71 -11
- package/src/modelSpec.ts +414 -0
- package/src/nativeUNet.ts +2256 -0
- package/src/process.ts +38 -36
- package/src/resourceTracker.ts +94 -0
- package/src/tileScheduler.ts +138 -0
- package/src/webnnUNet.ts +812 -0
- package/tests/modelSpec.test.mjs +128 -0
- package/tests/resourceLifecycle.test.mjs +383 -0
- package/tests/tileScheduler.test.mjs +90 -0
- package/lib/helper.d.ts +0 -4
- package/lib/helper.js +0 -33
- package/lib/helper.js.map +0 -1
- package/lib/kernels.d.ts +0 -1
- package/lib/kernels.js +0 -26
- package/lib/kernels.js.map +0 -1
- package/src/helper.ts +0 -43
- package/src/kernels.ts +0 -31
|
@@ -0,0 +1,1735 @@
|
|
|
1
|
+
import { Float16Array } from '@petamoriken/float16';
|
|
2
|
+
import { optimizeModelGraph, planModelExecution } from './graphOptimizer.js';
|
|
3
|
+
import { OIDNResourceTracker } from './resourceTracker.js';
|
|
4
|
+
const pipelineCachesByDevice = new WeakMap();
|
|
5
|
+
function sharedPipelineCache(device) {
|
|
6
|
+
let cache = pipelineCachesByDevice.get(device);
|
|
7
|
+
if (!cache) {
|
|
8
|
+
cache = { ready: new Map(), pending: new Map() };
|
|
9
|
+
pipelineCachesByDevice.set(device, cache);
|
|
10
|
+
}
|
|
11
|
+
return cache;
|
|
12
|
+
}
|
|
13
|
+
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
|
+
const TILED_CONV_K_BLOCKS = 8;
|
|
19
|
+
const SPATIAL_CONV_WORKGROUP = 8;
|
|
20
|
+
const SPATIAL_CONV_PATCH = SPATIAL_CONV_WORKGROUP + 2;
|
|
21
|
+
function roundUp(value, alignment) {
|
|
22
|
+
return Math.ceil(value / alignment) * alignment;
|
|
23
|
+
}
|
|
24
|
+
function blocksForChannels(channels) {
|
|
25
|
+
return Math.ceil(channels / 4);
|
|
26
|
+
}
|
|
27
|
+
function activationByteSize(shape, bytesPerScalar) {
|
|
28
|
+
return (shape.width *
|
|
29
|
+
shape.height *
|
|
30
|
+
blocksForChannels(shape.channels) *
|
|
31
|
+
4 *
|
|
32
|
+
bytesPerScalar);
|
|
33
|
+
}
|
|
34
|
+
let float16ToFloat32Lookup;
|
|
35
|
+
function halfBitsToNumber(bits) {
|
|
36
|
+
const sign = bits & 0x8000 ? -1 : 1;
|
|
37
|
+
const exponent = (bits >>> 10) & 0x1f;
|
|
38
|
+
const fraction = bits & 0x3ff;
|
|
39
|
+
if (exponent === 0) {
|
|
40
|
+
return sign * fraction * 2 ** -24;
|
|
41
|
+
}
|
|
42
|
+
if (exponent === 0x1f) {
|
|
43
|
+
return fraction === 0 ? sign * Infinity : NaN;
|
|
44
|
+
}
|
|
45
|
+
return sign * (1 + fraction / 1024) * 2 ** (exponent - 15);
|
|
46
|
+
}
|
|
47
|
+
function halfLookup() {
|
|
48
|
+
if (!float16ToFloat32Lookup) {
|
|
49
|
+
float16ToFloat32Lookup = new Float32Array(1 << 16);
|
|
50
|
+
for (let bits = 0; bits < float16ToFloat32Lookup.length; bits++) {
|
|
51
|
+
float16ToFloat32Lookup[bits] = halfBitsToNumber(bits);
|
|
52
|
+
}
|
|
53
|
+
}
|
|
54
|
+
return float16ToFloat32Lookup;
|
|
55
|
+
}
|
|
56
|
+
function tensorFloat32Values(tensor) {
|
|
57
|
+
if (tensor.desc.dataType === 'Float32') {
|
|
58
|
+
return new Float32Array(tensor.data.buffer, tensor.data.byteOffset, tensor.data.byteLength / 4);
|
|
59
|
+
}
|
|
60
|
+
const bits = new Uint16Array(tensor.data.buffer, tensor.data.byteOffset, tensor.data.byteLength / 2);
|
|
61
|
+
const values = new Float32Array(bits.length);
|
|
62
|
+
if (bits.length < 4096) {
|
|
63
|
+
for (let index = 0; index < bits.length; index++) {
|
|
64
|
+
values[index] = halfBitsToNumber(bits[index]);
|
|
65
|
+
}
|
|
66
|
+
}
|
|
67
|
+
else {
|
|
68
|
+
const lookup = halfLookup();
|
|
69
|
+
for (let index = 0; index < bits.length; index++) {
|
|
70
|
+
values[index] = lookup[bits[index]];
|
|
71
|
+
}
|
|
72
|
+
}
|
|
73
|
+
return values;
|
|
74
|
+
}
|
|
75
|
+
function createMappedBuffer(device, label, data, usage) {
|
|
76
|
+
const size = roundUp(data.byteLength, 4);
|
|
77
|
+
const buffer = device.createBuffer({
|
|
78
|
+
label,
|
|
79
|
+
size,
|
|
80
|
+
usage,
|
|
81
|
+
mappedAtCreation: true
|
|
82
|
+
});
|
|
83
|
+
new Uint8Array(buffer.getMappedRange()).set(new Uint8Array(data.buffer, data.byteOffset, data.byteLength));
|
|
84
|
+
buffer.unmap();
|
|
85
|
+
return buffer;
|
|
86
|
+
}
|
|
87
|
+
function createUniformBuffer(device, label, values) {
|
|
88
|
+
const data = new Uint32Array(roundUp(values.length, 4));
|
|
89
|
+
data.set(values);
|
|
90
|
+
return createMappedBuffer(device, label, data, GPUBufferUsage.UNIFORM);
|
|
91
|
+
}
|
|
92
|
+
function packConvTensors(device, id, tensors, precision) {
|
|
93
|
+
const inputBlocks = blocksForChannels(tensors.inputChannels);
|
|
94
|
+
const outputBlocks = blocksForChannels(tensors.outputChannels);
|
|
95
|
+
const packedWeightCount = outputBlocks *
|
|
96
|
+
tensors.kernelHeight *
|
|
97
|
+
tensors.kernelWidth *
|
|
98
|
+
inputBlocks *
|
|
99
|
+
4 *
|
|
100
|
+
4;
|
|
101
|
+
const canCopyHalfBits = precision === 'fp16' && tensors.weight.desc.dataType === 'Float16';
|
|
102
|
+
const packedWeights = canCopyHalfBits
|
|
103
|
+
? new Uint16Array(packedWeightCount)
|
|
104
|
+
: precision === 'fp16'
|
|
105
|
+
? new Float16Array(packedWeightCount)
|
|
106
|
+
: new Float32Array(packedWeightCount);
|
|
107
|
+
const sourceWeights = canCopyHalfBits
|
|
108
|
+
? new Uint16Array(tensors.weight.data.buffer, tensors.weight.data.byteOffset, tensors.weight.data.byteLength / 2)
|
|
109
|
+
: tensorFloat32Values(tensors.weight);
|
|
110
|
+
for (let outputBlock = 0; outputBlock < outputBlocks; outputBlock++) {
|
|
111
|
+
for (let y = 0; y < tensors.kernelHeight; y++) {
|
|
112
|
+
for (let x = 0; x < tensors.kernelWidth; x++) {
|
|
113
|
+
for (let inputBlock = 0; inputBlock < inputBlocks; inputBlock++) {
|
|
114
|
+
for (let outputLane = 0; outputLane < 4; outputLane++) {
|
|
115
|
+
const outputChannel = outputBlock * 4 + outputLane;
|
|
116
|
+
for (let inputLane = 0; inputLane < 4; inputLane++) {
|
|
117
|
+
const inputChannel = inputBlock * 4 + inputLane;
|
|
118
|
+
const packedBlock = ((((outputBlock * tensors.kernelHeight + y) *
|
|
119
|
+
tensors.kernelWidth +
|
|
120
|
+
x) *
|
|
121
|
+
inputBlocks +
|
|
122
|
+
inputBlock) *
|
|
123
|
+
16);
|
|
124
|
+
// One vec4 contains the four output lanes for an input lane.
|
|
125
|
+
// Direct and tiled kernels share this layout, keeping model
|
|
126
|
+
// descriptors independent from the selected kernel.
|
|
127
|
+
const packedIndex = packedBlock + inputLane * 4 + outputLane;
|
|
128
|
+
if (outputChannel < tensors.outputChannels &&
|
|
129
|
+
inputChannel < tensors.inputChannels) {
|
|
130
|
+
const sourceIndex = ((outputChannel * tensors.inputChannels + inputChannel) *
|
|
131
|
+
tensors.kernelHeight +
|
|
132
|
+
y) *
|
|
133
|
+
tensors.kernelWidth +
|
|
134
|
+
x;
|
|
135
|
+
packedWeights[packedIndex] = sourceWeights[sourceIndex];
|
|
136
|
+
}
|
|
137
|
+
}
|
|
138
|
+
}
|
|
139
|
+
}
|
|
140
|
+
}
|
|
141
|
+
}
|
|
142
|
+
}
|
|
143
|
+
// Bias stays f32 even for half activations because convolution accumulates
|
|
144
|
+
// in f32. Padded lanes are zero and never escape the final three channels.
|
|
145
|
+
const packedBias = new Float32Array(outputBlocks * 4);
|
|
146
|
+
packedBias.set(tensorFloat32Values(tensors.bias));
|
|
147
|
+
const weights = createMappedBuffer(device, `oidn/${id}/weights/${precision}`, packedWeights, GPUBufferUsage.STORAGE);
|
|
148
|
+
try {
|
|
149
|
+
return {
|
|
150
|
+
weights,
|
|
151
|
+
bias: createMappedBuffer(device, `oidn/${id}/bias`, packedBias, GPUBufferUsage.STORAGE)
|
|
152
|
+
};
|
|
153
|
+
}
|
|
154
|
+
catch (error) {
|
|
155
|
+
weights.destroy();
|
|
156
|
+
throw error;
|
|
157
|
+
}
|
|
158
|
+
}
|
|
159
|
+
function storageVecType(precision) {
|
|
160
|
+
return precision === 'fp16' ? 'vec4<f16>' : 'vec4<f32>';
|
|
161
|
+
}
|
|
162
|
+
function shaderPreamble(precision) {
|
|
163
|
+
return precision === 'fp16' ? 'enable f16;\n' : '';
|
|
164
|
+
}
|
|
165
|
+
function storeExpression(expression, outputPrecision) {
|
|
166
|
+
return outputPrecision === 'fp16'
|
|
167
|
+
? `vec4<f16>(${expression})`
|
|
168
|
+
: expression;
|
|
169
|
+
}
|
|
170
|
+
function activationExpression(expression, activation) {
|
|
171
|
+
return activation === 'relu'
|
|
172
|
+
? `max(${expression}, vec4<f32>(0.0))`
|
|
173
|
+
: expression;
|
|
174
|
+
}
|
|
175
|
+
function accumulationCode(inputExpression, weightBase, precision) {
|
|
176
|
+
if (precision === 'fp32') {
|
|
177
|
+
return /* wgsl */ `
|
|
178
|
+
let inputValue = vec4<f32>(${inputExpression});
|
|
179
|
+
let weightBase = ${weightBase};
|
|
180
|
+
acc = fma(vec4<f32>(weights[weightBase]), vec4<f32>(inputValue.x), acc);
|
|
181
|
+
acc = fma(vec4<f32>(weights[weightBase + 1u]), vec4<f32>(inputValue.y), acc);
|
|
182
|
+
acc = fma(vec4<f32>(weights[weightBase + 2u]), vec4<f32>(inputValue.z), acc);
|
|
183
|
+
acc = fma(vec4<f32>(weights[weightBase + 3u]), vec4<f32>(inputValue.w), acc);
|
|
184
|
+
`;
|
|
185
|
+
}
|
|
186
|
+
return /* wgsl */ `
|
|
187
|
+
let inputValue = vec4<f16>(${inputExpression});
|
|
188
|
+
let weightBase = ${weightBase};
|
|
189
|
+
var partial = vec4<f16>(0.0h);
|
|
190
|
+
partial = fma(weights[weightBase], vec4<f16>(inputValue.x), partial);
|
|
191
|
+
partial = fma(weights[weightBase + 1u], vec4<f16>(inputValue.y), partial);
|
|
192
|
+
partial = fma(weights[weightBase + 2u], vec4<f16>(inputValue.z), partial);
|
|
193
|
+
partial = fma(weights[weightBase + 3u], vec4<f16>(inputValue.w), partial);
|
|
194
|
+
acc += vec4<f32>(partial);
|
|
195
|
+
`;
|
|
196
|
+
}
|
|
197
|
+
function subgroupAccumulationCode(inputExpression, weightBase, precision) {
|
|
198
|
+
const inputType = precision === 'fp16' ? 'vec4<f16>' : 'vec4<f32>';
|
|
199
|
+
const accumulator = precision === 'fp16'
|
|
200
|
+
? `var partial = vec4<f16>(0.0h);
|
|
201
|
+
partial = fma(subgroupBroadcast(weights[weightBase], 0u), ${inputType}(inputValue.x), partial);
|
|
202
|
+
partial = fma(subgroupBroadcast(weights[weightBase + 1u], 0u), ${inputType}(inputValue.y), partial);
|
|
203
|
+
partial = fma(subgroupBroadcast(weights[weightBase + 2u], 0u), ${inputType}(inputValue.z), partial);
|
|
204
|
+
partial = fma(subgroupBroadcast(weights[weightBase + 3u], 0u), ${inputType}(inputValue.w), partial);
|
|
205
|
+
acc += vec4<f32>(partial);`
|
|
206
|
+
: `acc = fma(subgroupBroadcast(weights[weightBase], 0u), vec4<f32>(inputValue.x), acc);
|
|
207
|
+
acc = fma(subgroupBroadcast(weights[weightBase + 1u], 0u), vec4<f32>(inputValue.y), acc);
|
|
208
|
+
acc = fma(subgroupBroadcast(weights[weightBase + 2u], 0u), vec4<f32>(inputValue.z), acc);
|
|
209
|
+
acc = fma(subgroupBroadcast(weights[weightBase + 3u], 0u), vec4<f32>(inputValue.w), acc);`;
|
|
210
|
+
return /* wgsl */ `
|
|
211
|
+
let inputValue = ${inputType}(${inputExpression});
|
|
212
|
+
let weightBase = ${weightBase};
|
|
213
|
+
${accumulator}
|
|
214
|
+
`;
|
|
215
|
+
}
|
|
216
|
+
function tiledAccumulationCode(precision) {
|
|
217
|
+
if (precision === 'fp16') {
|
|
218
|
+
return /* wgsl */ `
|
|
219
|
+
var partial: array<vec4<f16>, ${TILED_CONV_ROWS_PER_THREAD}>;
|
|
220
|
+
for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
|
|
221
|
+
partial[row] = vec4<f16>(0.0h);
|
|
222
|
+
}
|
|
223
|
+
for (var tileK = 0u; tileK < ${TILED_CONV_K_BLOCKS}u; tileK++) {
|
|
224
|
+
let weightBase =
|
|
225
|
+
(tileK * ${TILED_CONV_N_BLOCKS}u + localId.x) * 4u;
|
|
226
|
+
for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
|
|
227
|
+
let tileSpatial =
|
|
228
|
+
localId.y * ${TILED_CONV_ROWS_PER_THREAD}u + row;
|
|
229
|
+
let inputValue =
|
|
230
|
+
inputTile[tileSpatial * ${TILED_CONV_K_BLOCKS}u + tileK];
|
|
231
|
+
partial[row] = fma(
|
|
232
|
+
weightTile[weightBase],
|
|
233
|
+
vec4<f16>(inputValue.x),
|
|
234
|
+
partial[row]
|
|
235
|
+
);
|
|
236
|
+
partial[row] = fma(
|
|
237
|
+
weightTile[weightBase + 1u],
|
|
238
|
+
vec4<f16>(inputValue.y),
|
|
239
|
+
partial[row]
|
|
240
|
+
);
|
|
241
|
+
partial[row] = fma(
|
|
242
|
+
weightTile[weightBase + 2u],
|
|
243
|
+
vec4<f16>(inputValue.z),
|
|
244
|
+
partial[row]
|
|
245
|
+
);
|
|
246
|
+
partial[row] = fma(
|
|
247
|
+
weightTile[weightBase + 3u],
|
|
248
|
+
vec4<f16>(inputValue.w),
|
|
249
|
+
partial[row]
|
|
250
|
+
);
|
|
251
|
+
}
|
|
252
|
+
}
|
|
253
|
+
for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
|
|
254
|
+
acc[row] += vec4<f32>(partial[row]);
|
|
255
|
+
}
|
|
256
|
+
`;
|
|
257
|
+
}
|
|
258
|
+
return /* wgsl */ `
|
|
259
|
+
for (var tileK = 0u; tileK < ${TILED_CONV_K_BLOCKS}u; tileK++) {
|
|
260
|
+
let weightBase =
|
|
261
|
+
(tileK * ${TILED_CONV_N_BLOCKS}u + localId.x) * 4u;
|
|
262
|
+
for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
|
|
263
|
+
let tileSpatial =
|
|
264
|
+
localId.y * ${TILED_CONV_ROWS_PER_THREAD}u + row;
|
|
265
|
+
let inputValue =
|
|
266
|
+
inputTile[tileSpatial * ${TILED_CONV_K_BLOCKS}u + tileK];
|
|
267
|
+
acc[row] = fma(
|
|
268
|
+
weightTile[weightBase],
|
|
269
|
+
vec4<f32>(inputValue.x),
|
|
270
|
+
acc[row]
|
|
271
|
+
);
|
|
272
|
+
acc[row] = fma(
|
|
273
|
+
weightTile[weightBase + 1u],
|
|
274
|
+
vec4<f32>(inputValue.y),
|
|
275
|
+
acc[row]
|
|
276
|
+
);
|
|
277
|
+
acc[row] = fma(
|
|
278
|
+
weightTile[weightBase + 2u],
|
|
279
|
+
vec4<f32>(inputValue.z),
|
|
280
|
+
acc[row]
|
|
281
|
+
);
|
|
282
|
+
acc[row] = fma(
|
|
283
|
+
weightTile[weightBase + 3u],
|
|
284
|
+
vec4<f32>(inputValue.w),
|
|
285
|
+
acc[row]
|
|
286
|
+
);
|
|
287
|
+
}
|
|
288
|
+
}
|
|
289
|
+
`;
|
|
290
|
+
}
|
|
291
|
+
function createConvShader(precision, outputPrecision, activation, inputBlocks, outputBlocks) {
|
|
292
|
+
const inputType = storageVecType(precision);
|
|
293
|
+
const weightType = storageVecType(precision);
|
|
294
|
+
const outputType = storageVecType(outputPrecision);
|
|
295
|
+
const stored = storeExpression(activationExpression('acc', activation), outputPrecision);
|
|
296
|
+
return /* wgsl */ `${shaderPreamble(precision)}
|
|
297
|
+
struct Params {
|
|
298
|
+
inputWidth: u32,
|
|
299
|
+
inputHeight: u32,
|
|
300
|
+
outputWidth: u32,
|
|
301
|
+
outputHeight: u32,
|
|
302
|
+
inputBlocks: u32,
|
|
303
|
+
outputBlocks: u32,
|
|
304
|
+
}
|
|
305
|
+
|
|
306
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${inputType}>;
|
|
307
|
+
@group(0) @binding(1) var<storage, read> weights: array<${weightType}>;
|
|
308
|
+
@group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
|
|
309
|
+
@group(0) @binding(3) var<storage, read_write> outputData: array<${outputType}>;
|
|
310
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
311
|
+
|
|
312
|
+
@compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
|
|
313
|
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
314
|
+
if (gid.x >= params.outputWidth || gid.y >= params.outputHeight || gid.z >= ${outputBlocks}u) {
|
|
315
|
+
return;
|
|
316
|
+
}
|
|
317
|
+
var acc = bias[gid.z];
|
|
318
|
+
for (var ky = 0u; ky < 3u; ky++) {
|
|
319
|
+
let inputY = i32(gid.y) + i32(ky) - 1;
|
|
320
|
+
if (inputY < 0 || inputY >= i32(params.inputHeight)) { continue; }
|
|
321
|
+
for (var kx = 0u; kx < 3u; kx++) {
|
|
322
|
+
let inputX = i32(gid.x) + i32(kx) - 1;
|
|
323
|
+
if (inputX < 0 || inputX >= i32(params.inputWidth)) { continue; }
|
|
324
|
+
let pixelBase = (u32(inputY) * params.inputWidth + u32(inputX)) * ${inputBlocks}u;
|
|
325
|
+
for (var inputBlock = 0u; inputBlock < ${inputBlocks}u; inputBlock++) {
|
|
326
|
+
${accumulationCode('inputData[pixelBase + inputBlock]', `((((gid.z * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`, precision)}
|
|
327
|
+
}
|
|
328
|
+
}
|
|
329
|
+
}
|
|
330
|
+
let outputIndex = (gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
|
|
331
|
+
outputData[outputIndex] = ${stored};
|
|
332
|
+
}
|
|
333
|
+
`;
|
|
334
|
+
}
|
|
335
|
+
/** Direct convolution with subgroup-wide weight broadcast. */
|
|
336
|
+
function createSubgroupConvShader(precision, outputPrecision, activation, inputBlocks, outputBlocks) {
|
|
337
|
+
const inputType = storageVecType(precision);
|
|
338
|
+
const weightType = storageVecType(precision);
|
|
339
|
+
const outputType = storageVecType(outputPrecision);
|
|
340
|
+
const stored = storeExpression(activationExpression('acc', activation), outputPrecision);
|
|
341
|
+
return /* wgsl */ `${shaderPreamble(precision)}
|
|
342
|
+
enable subgroups;
|
|
343
|
+
struct Params {
|
|
344
|
+
inputWidth: u32,
|
|
345
|
+
inputHeight: u32,
|
|
346
|
+
outputWidth: u32,
|
|
347
|
+
outputHeight: u32,
|
|
348
|
+
inputBlocks: u32,
|
|
349
|
+
outputBlocks: u32,
|
|
350
|
+
}
|
|
351
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${inputType}>;
|
|
352
|
+
@group(0) @binding(1) var<storage, read> weights: array<${weightType}>;
|
|
353
|
+
@group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
|
|
354
|
+
@group(0) @binding(3) var<storage, read_write> outputData: array<${outputType}>;
|
|
355
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
356
|
+
|
|
357
|
+
@compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
|
|
358
|
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
359
|
+
let outputInBounds =
|
|
360
|
+
gid.x < params.outputWidth && gid.y < params.outputHeight;
|
|
361
|
+
var acc = bias[gid.z];
|
|
362
|
+
for (var ky = 0u; ky < 3u; ky++) {
|
|
363
|
+
let inputY = i32(gid.y) + i32(ky) - 1;
|
|
364
|
+
let clampedY = u32(clamp(inputY, 0, i32(params.inputHeight) - 1));
|
|
365
|
+
for (var kx = 0u; kx < 3u; kx++) {
|
|
366
|
+
let inputX = i32(gid.x) + i32(kx) - 1;
|
|
367
|
+
let clampedX = u32(clamp(inputX, 0, i32(params.inputWidth) - 1));
|
|
368
|
+
let inputInBounds =
|
|
369
|
+
outputInBounds && inputX >= 0 && inputY >= 0 &&
|
|
370
|
+
inputX < i32(params.inputWidth) && inputY < i32(params.inputHeight);
|
|
371
|
+
let pixelBase =
|
|
372
|
+
(clampedY * params.inputWidth + clampedX) * ${inputBlocks}u;
|
|
373
|
+
for (var inputBlock = 0u; inputBlock < ${inputBlocks}u; inputBlock++) {
|
|
374
|
+
${subgroupAccumulationCode(`select(${inputType}(0.0), inputData[pixelBase + inputBlock], inputInBounds)`, `((((gid.z * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`, precision)}
|
|
375
|
+
}
|
|
376
|
+
}
|
|
377
|
+
}
|
|
378
|
+
if (outputInBounds) {
|
|
379
|
+
let outputIndex =
|
|
380
|
+
(gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
|
|
381
|
+
outputData[outputIndex] = ${stored};
|
|
382
|
+
}
|
|
383
|
+
}
|
|
384
|
+
`;
|
|
385
|
+
}
|
|
386
|
+
/**
|
|
387
|
+
* A 2D convolution tile which loads the complete 3x3 halo into workgroup
|
|
388
|
+
* memory once. Kernel choice only depends on the operation shape, precision,
|
|
389
|
+
* and device limits; it is deliberately independent of OIDN model names.
|
|
390
|
+
*/
|
|
391
|
+
function createSpatialConvShader(precision, outputPrecision, activation, inputBlocks, outputBlocks) {
|
|
392
|
+
const inputType = storageVecType(precision);
|
|
393
|
+
const outputType = storageVecType(outputPrecision);
|
|
394
|
+
const stored = storeExpression(activationExpression('acc', activation), outputPrecision);
|
|
395
|
+
const patchValues = SPATIAL_CONV_PATCH * SPATIAL_CONV_PATCH * inputBlocks;
|
|
396
|
+
const workgroupThreads = SPATIAL_CONV_WORKGROUP * SPATIAL_CONV_WORKGROUP;
|
|
397
|
+
return /* wgsl */ `${shaderPreamble(precision)}
|
|
398
|
+
struct Params {
|
|
399
|
+
inputWidth: u32,
|
|
400
|
+
inputHeight: u32,
|
|
401
|
+
outputWidth: u32,
|
|
402
|
+
outputHeight: u32,
|
|
403
|
+
inputBlocks: u32,
|
|
404
|
+
outputBlocks: u32,
|
|
405
|
+
}
|
|
406
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${inputType}>;
|
|
407
|
+
@group(0) @binding(1) var<storage, read> weights: array<${inputType}>;
|
|
408
|
+
@group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
|
|
409
|
+
@group(0) @binding(3) var<storage, read_write> outputData: array<${outputType}>;
|
|
410
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
411
|
+
|
|
412
|
+
var<workgroup> inputPatch: array<${inputType}, ${patchValues}>;
|
|
413
|
+
|
|
414
|
+
@compute @workgroup_size(${SPATIAL_CONV_WORKGROUP}, ${SPATIAL_CONV_WORKGROUP}, 1)
|
|
415
|
+
fn main(
|
|
416
|
+
@builtin(local_invocation_id) localId: vec3<u32>,
|
|
417
|
+
@builtin(workgroup_id) workgroupId: vec3<u32>
|
|
418
|
+
) {
|
|
419
|
+
let localLinear =
|
|
420
|
+
localId.y * ${SPATIAL_CONV_WORKGROUP}u + localId.x;
|
|
421
|
+
for (
|
|
422
|
+
var loadIndex = localLinear;
|
|
423
|
+
loadIndex < ${patchValues}u;
|
|
424
|
+
loadIndex += ${workgroupThreads}u
|
|
425
|
+
) {
|
|
426
|
+
let patchPixel = loadIndex / ${inputBlocks}u;
|
|
427
|
+
let inputBlock = loadIndex % ${inputBlocks}u;
|
|
428
|
+
let patchX = patchPixel % ${SPATIAL_CONV_PATCH}u;
|
|
429
|
+
let patchY = patchPixel / ${SPATIAL_CONV_PATCH}u;
|
|
430
|
+
let inputX =
|
|
431
|
+
i32(workgroupId.x * ${SPATIAL_CONV_WORKGROUP}u + patchX) - 1;
|
|
432
|
+
let inputY =
|
|
433
|
+
i32(workgroupId.y * ${SPATIAL_CONV_WORKGROUP}u + patchY) - 1;
|
|
434
|
+
var value = ${inputType}(0.0);
|
|
435
|
+
if (
|
|
436
|
+
inputX >= 0 && inputX < i32(params.inputWidth) &&
|
|
437
|
+
inputY >= 0 && inputY < i32(params.inputHeight)
|
|
438
|
+
) {
|
|
439
|
+
let inputIndex =
|
|
440
|
+
(u32(inputY) * params.inputWidth + u32(inputX)) *
|
|
441
|
+
${inputBlocks}u + inputBlock;
|
|
442
|
+
value = inputData[inputIndex];
|
|
443
|
+
}
|
|
444
|
+
inputPatch[loadIndex] = value;
|
|
445
|
+
}
|
|
446
|
+
workgroupBarrier();
|
|
447
|
+
|
|
448
|
+
let outputX =
|
|
449
|
+
workgroupId.x * ${SPATIAL_CONV_WORKGROUP}u + localId.x;
|
|
450
|
+
let outputY =
|
|
451
|
+
workgroupId.y * ${SPATIAL_CONV_WORKGROUP}u + localId.y;
|
|
452
|
+
let outputBlock = workgroupId.z;
|
|
453
|
+
if (
|
|
454
|
+
outputX >= params.outputWidth || outputY >= params.outputHeight ||
|
|
455
|
+
outputBlock >= ${outputBlocks}u
|
|
456
|
+
) {
|
|
457
|
+
return;
|
|
458
|
+
}
|
|
459
|
+
|
|
460
|
+
var acc = bias[outputBlock];
|
|
461
|
+
for (var ky = 0u; ky < 3u; ky++) {
|
|
462
|
+
for (var kx = 0u; kx < 3u; kx++) {
|
|
463
|
+
let patchBase =
|
|
464
|
+
((localId.y + ky) * ${SPATIAL_CONV_PATCH}u + localId.x + kx) *
|
|
465
|
+
${inputBlocks}u;
|
|
466
|
+
for (var inputBlock = 0u; inputBlock < ${inputBlocks}u; inputBlock++) {
|
|
467
|
+
${accumulationCode('inputPatch[patchBase + inputBlock]', `((((outputBlock * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`, precision)}
|
|
468
|
+
}
|
|
469
|
+
}
|
|
470
|
+
}
|
|
471
|
+
let outputIndex =
|
|
472
|
+
(outputY * params.outputWidth + outputX) * ${outputBlocks}u + outputBlock;
|
|
473
|
+
outputData[outputIndex] = ${stored};
|
|
474
|
+
}
|
|
475
|
+
`;
|
|
476
|
+
}
|
|
477
|
+
/**
|
|
478
|
+
* Implicit-GEMM convolution for FP32. This follows the proven packed WebGPU
|
|
479
|
+
* shape: one 8x8 workgroup computes 32 spatial rows by 8 vec4 output blocks,
|
|
480
|
+
* with four output rows per thread and an eight-vec4 K tile.
|
|
481
|
+
*/
|
|
482
|
+
function createTiledConvShader(precision, outputPrecision, activation, inputBlocks, outputBlocks) {
|
|
483
|
+
const inputType = storageVecType(precision);
|
|
484
|
+
const outputType = storageVecType(outputPrecision);
|
|
485
|
+
const stored = storeExpression(activationExpression('acc', activation), outputPrecision);
|
|
486
|
+
const workgroupThreads = TILED_CONV_WORKGROUP * TILED_CONV_WORKGROUP;
|
|
487
|
+
const inputTileValues = TILED_CONV_M * TILED_CONV_K_BLOCKS;
|
|
488
|
+
const weightTileValues = TILED_CONV_K_BLOCKS * TILED_CONV_N_BLOCKS * 4;
|
|
489
|
+
return /* wgsl */ `${shaderPreamble(precision)}
|
|
490
|
+
struct Params {
|
|
491
|
+
inputWidth: u32,
|
|
492
|
+
inputHeight: u32,
|
|
493
|
+
outputWidth: u32,
|
|
494
|
+
outputHeight: u32,
|
|
495
|
+
inputBlocks: u32,
|
|
496
|
+
outputBlocks: u32,
|
|
497
|
+
}
|
|
498
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${inputType}>;
|
|
499
|
+
@group(0) @binding(1) var<storage, read> weights: array<${inputType}>;
|
|
500
|
+
@group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
|
|
501
|
+
@group(0) @binding(3) var<storage, read_write> outputData: array<${outputType}>;
|
|
502
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
503
|
+
|
|
504
|
+
var<workgroup> inputTile: array<${inputType}, ${inputTileValues}>;
|
|
505
|
+
var<workgroup> weightTile: array<${inputType}, ${weightTileValues}>;
|
|
506
|
+
|
|
507
|
+
@compute @workgroup_size(${TILED_CONV_WORKGROUP}, ${TILED_CONV_WORKGROUP}, 1)
|
|
508
|
+
fn main(
|
|
509
|
+
@builtin(local_invocation_id) localId: vec3<u32>,
|
|
510
|
+
@builtin(workgroup_id) workgroupId: vec3<u32>
|
|
511
|
+
) {
|
|
512
|
+
let spatialBase =
|
|
513
|
+
workgroupId.x * ${TILED_CONV_M}u +
|
|
514
|
+
localId.y * ${TILED_CONV_ROWS_PER_THREAD}u;
|
|
515
|
+
let outputBlock =
|
|
516
|
+
workgroupId.y * ${TILED_CONV_N_BLOCKS}u + localId.x;
|
|
517
|
+
let spatialCount = params.outputWidth * params.outputHeight;
|
|
518
|
+
var acc: array<vec4<f32>, ${TILED_CONV_ROWS_PER_THREAD}>;
|
|
519
|
+
if (outputBlock < ${outputBlocks}u) {
|
|
520
|
+
for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
|
|
521
|
+
acc[row] = bias[outputBlock];
|
|
522
|
+
}
|
|
523
|
+
}
|
|
524
|
+
|
|
525
|
+
let localLinear =
|
|
526
|
+
localId.y * ${TILED_CONV_WORKGROUP}u + localId.x;
|
|
527
|
+
let totalK = ${inputBlocks * 9}u;
|
|
528
|
+
for (var kBase = 0u; kBase < totalK; kBase += ${TILED_CONV_K_BLOCKS}u) {
|
|
529
|
+
for (
|
|
530
|
+
var loadIndex = localLinear;
|
|
531
|
+
loadIndex < ${inputTileValues}u;
|
|
532
|
+
loadIndex += ${workgroupThreads}u
|
|
533
|
+
) {
|
|
534
|
+
let tileSpatial = loadIndex / ${TILED_CONV_K_BLOCKS}u;
|
|
535
|
+
let tileK = loadIndex % ${TILED_CONV_K_BLOCKS}u;
|
|
536
|
+
let inputSpatialIndex =
|
|
537
|
+
workgroupId.x * ${TILED_CONV_M}u + tileSpatial;
|
|
538
|
+
let kIndex = kBase + tileK;
|
|
539
|
+
var value = ${inputType}(0.0);
|
|
540
|
+
if (inputSpatialIndex < spatialCount && kIndex < totalK) {
|
|
541
|
+
let outputY = inputSpatialIndex / params.outputWidth;
|
|
542
|
+
let outputX = inputSpatialIndex % params.outputWidth;
|
|
543
|
+
let inputBlock = kIndex % ${inputBlocks}u;
|
|
544
|
+
let kernelIndex = kIndex / ${inputBlocks}u;
|
|
545
|
+
let kernelY = kernelIndex / 3u;
|
|
546
|
+
let kernelX = kernelIndex % 3u;
|
|
547
|
+
let inputY = i32(outputY) + i32(kernelY) - 1;
|
|
548
|
+
let inputX = i32(outputX) + i32(kernelX) - 1;
|
|
549
|
+
if (
|
|
550
|
+
inputY >= 0 && inputY < i32(params.inputHeight) &&
|
|
551
|
+
inputX >= 0 && inputX < i32(params.inputWidth)
|
|
552
|
+
) {
|
|
553
|
+
let inputIndex =
|
|
554
|
+
(u32(inputY) * params.inputWidth + u32(inputX)) *
|
|
555
|
+
${inputBlocks}u + inputBlock;
|
|
556
|
+
value = inputData[inputIndex];
|
|
557
|
+
}
|
|
558
|
+
}
|
|
559
|
+
inputTile[loadIndex] = value;
|
|
560
|
+
}
|
|
561
|
+
|
|
562
|
+
for (
|
|
563
|
+
var loadIndex = localLinear;
|
|
564
|
+
loadIndex < ${weightTileValues}u;
|
|
565
|
+
loadIndex += ${workgroupThreads}u
|
|
566
|
+
) {
|
|
567
|
+
let tileK = loadIndex / ${TILED_CONV_N_BLOCKS * 4}u;
|
|
568
|
+
let outputRemainder = loadIndex % ${TILED_CONV_N_BLOCKS * 4}u;
|
|
569
|
+
let tileOutputBlock = outputRemainder / 4u;
|
|
570
|
+
let outputLane = outputRemainder % 4u;
|
|
571
|
+
let loadedOutputBlock =
|
|
572
|
+
workgroupId.y * ${TILED_CONV_N_BLOCKS}u + tileOutputBlock;
|
|
573
|
+
let kIndex = kBase + tileK;
|
|
574
|
+
var value = ${inputType}(0.0);
|
|
575
|
+
if (loadedOutputBlock < ${outputBlocks}u && kIndex < totalK) {
|
|
576
|
+
let inputBlock = kIndex % ${inputBlocks}u;
|
|
577
|
+
let kernelIndex = kIndex / ${inputBlocks}u;
|
|
578
|
+
let kernelY = kernelIndex / 3u;
|
|
579
|
+
let kernelX = kernelIndex % 3u;
|
|
580
|
+
let weightIndex =
|
|
581
|
+
((((loadedOutputBlock * 3u + kernelY) * 3u + kernelX) *
|
|
582
|
+
${inputBlocks}u + inputBlock) * 4u + outputLane);
|
|
583
|
+
value = weights[weightIndex];
|
|
584
|
+
}
|
|
585
|
+
weightTile[loadIndex] = value;
|
|
586
|
+
}
|
|
587
|
+
|
|
588
|
+
workgroupBarrier();
|
|
589
|
+
${tiledAccumulationCode(precision)}
|
|
590
|
+
workgroupBarrier();
|
|
591
|
+
}
|
|
592
|
+
|
|
593
|
+
if (outputBlock < ${outputBlocks}u) {
|
|
594
|
+
for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
|
|
595
|
+
let spatialIndex = spatialBase + row;
|
|
596
|
+
if (spatialIndex < spatialCount) {
|
|
597
|
+
let outputIndex = spatialIndex * ${outputBlocks}u + outputBlock;
|
|
598
|
+
outputData[outputIndex] = ${stored.replaceAll('acc', 'acc[row]')};
|
|
599
|
+
}
|
|
600
|
+
}
|
|
601
|
+
}
|
|
602
|
+
}
|
|
603
|
+
`;
|
|
604
|
+
}
|
|
605
|
+
function createFusedConvPoolShader(precision, activation, inputBlocks, outputBlocks) {
|
|
606
|
+
const valueType = storageVecType(precision);
|
|
607
|
+
const activated = activationExpression('acc', activation);
|
|
608
|
+
return /* wgsl */ `${shaderPreamble(precision)}
|
|
609
|
+
struct Params {
|
|
610
|
+
inputWidth: u32,
|
|
611
|
+
inputHeight: u32,
|
|
612
|
+
outputWidth: u32,
|
|
613
|
+
outputHeight: u32,
|
|
614
|
+
inputBlocks: u32,
|
|
615
|
+
outputBlocks: u32,
|
|
616
|
+
}
|
|
617
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${valueType}>;
|
|
618
|
+
@group(0) @binding(1) var<storage, read> weights: array<${valueType}>;
|
|
619
|
+
@group(0) @binding(2) var<storage, read> bias: array<vec4<f32>>;
|
|
620
|
+
@group(0) @binding(3) var<storage, read_write> outputData: array<${valueType}>;
|
|
621
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
622
|
+
|
|
623
|
+
@compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
|
|
624
|
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
625
|
+
if (gid.x >= params.outputWidth || gid.y >= params.outputHeight || gid.z >= ${outputBlocks}u) {
|
|
626
|
+
return;
|
|
627
|
+
}
|
|
628
|
+
var pooled = vec4<f32>(-3.402823466e+38);
|
|
629
|
+
for (var py = 0u; py < 2u; py++) {
|
|
630
|
+
let centerY = gid.y * 2u + py;
|
|
631
|
+
if (centerY >= params.inputHeight) { continue; }
|
|
632
|
+
for (var px = 0u; px < 2u; px++) {
|
|
633
|
+
let centerX = gid.x * 2u + px;
|
|
634
|
+
if (centerX >= params.inputWidth) { continue; }
|
|
635
|
+
var acc = bias[gid.z];
|
|
636
|
+
for (var ky = 0u; ky < 3u; ky++) {
|
|
637
|
+
let inputY = i32(centerY) + i32(ky) - 1;
|
|
638
|
+
if (inputY < 0 || inputY >= i32(params.inputHeight)) { continue; }
|
|
639
|
+
for (var kx = 0u; kx < 3u; kx++) {
|
|
640
|
+
let inputX = i32(centerX) + i32(kx) - 1;
|
|
641
|
+
if (inputX < 0 || inputX >= i32(params.inputWidth)) { continue; }
|
|
642
|
+
let pixelBase = (u32(inputY) * params.inputWidth + u32(inputX)) * ${inputBlocks}u;
|
|
643
|
+
for (var inputBlock = 0u; inputBlock < ${inputBlocks}u; inputBlock++) {
|
|
644
|
+
${accumulationCode('inputData[pixelBase + inputBlock]', `((((gid.z * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`, precision)}
|
|
645
|
+
}
|
|
646
|
+
}
|
|
647
|
+
}
|
|
648
|
+
pooled = max(pooled, ${activated});
|
|
649
|
+
}
|
|
650
|
+
}
|
|
651
|
+
let outputIndex = (gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
|
|
652
|
+
outputData[outputIndex] = ${storeExpression('pooled', precision)};
|
|
653
|
+
}
|
|
654
|
+
`;
|
|
655
|
+
}
|
|
656
|
+
function createMaxPoolShader(precision, outputBlocks) {
|
|
657
|
+
const valueType = storageVecType(precision);
|
|
658
|
+
return /* wgsl */ `${shaderPreamble(precision)}
|
|
659
|
+
struct Params {
|
|
660
|
+
inputWidth: u32,
|
|
661
|
+
inputHeight: u32,
|
|
662
|
+
outputWidth: u32,
|
|
663
|
+
outputHeight: u32,
|
|
664
|
+
outputBlocks: u32,
|
|
665
|
+
}
|
|
666
|
+
@group(0) @binding(0) var<storage, read> inputData: array<${valueType}>;
|
|
667
|
+
@group(0) @binding(1) var<storage, read_write> outputData: array<${valueType}>;
|
|
668
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
669
|
+
|
|
670
|
+
@compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
|
|
671
|
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
672
|
+
if (
|
|
673
|
+
gid.x >= params.outputWidth ||
|
|
674
|
+
gid.y >= params.outputHeight ||
|
|
675
|
+
gid.z >= ${outputBlocks}u
|
|
676
|
+
) {
|
|
677
|
+
return;
|
|
678
|
+
}
|
|
679
|
+
var pooled = vec4<f32>(-3.402823466e+38);
|
|
680
|
+
for (var py = 0u; py < 2u; py++) {
|
|
681
|
+
let inputY = gid.y * 2u + py;
|
|
682
|
+
if (inputY >= params.inputHeight) { continue; }
|
|
683
|
+
for (var px = 0u; px < 2u; px++) {
|
|
684
|
+
let inputX = gid.x * 2u + px;
|
|
685
|
+
if (inputX >= params.inputWidth) { continue; }
|
|
686
|
+
let inputIndex =
|
|
687
|
+
(inputY * params.inputWidth + inputX) * ${outputBlocks}u + gid.z;
|
|
688
|
+
pooled = max(pooled, vec4<f32>(inputData[inputIndex]));
|
|
689
|
+
}
|
|
690
|
+
}
|
|
691
|
+
let outputIndex =
|
|
692
|
+
(gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
|
|
693
|
+
outputData[outputIndex] = ${storeExpression('pooled', precision)};
|
|
694
|
+
}
|
|
695
|
+
`;
|
|
696
|
+
}
|
|
697
|
+
function createTiledDecoderShader(precision, activation, sourceBlocks, upsampledSource, outputBlocks) {
|
|
698
|
+
const valueType = storageVecType(precision);
|
|
699
|
+
const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
|
|
700
|
+
const stored = storeExpression(activationExpression('acc[row]', activation), precision);
|
|
701
|
+
const workgroupThreads = TILED_CONV_WORKGROUP * TILED_CONV_WORKGROUP;
|
|
702
|
+
const inputTileValues = TILED_CONV_M * TILED_CONV_K_BLOCKS;
|
|
703
|
+
const weightTileValues = TILED_CONV_K_BLOCKS * TILED_CONV_N_BLOCKS * 4;
|
|
704
|
+
const sourceRead = (source, blockExpression) => {
|
|
705
|
+
const sourceX = source === upsampledSource
|
|
706
|
+
? 'u32(inputX) / 2u'
|
|
707
|
+
: 'u32(inputX)';
|
|
708
|
+
const sourceY = source === upsampledSource
|
|
709
|
+
? 'u32(inputY) / 2u'
|
|
710
|
+
: 'u32(inputY)';
|
|
711
|
+
return /* wgsl */ `
|
|
712
|
+
{
|
|
713
|
+
let sourceBlock = ${blockExpression};
|
|
714
|
+
let sourceX = ${sourceX};
|
|
715
|
+
let sourceY = ${sourceY};
|
|
716
|
+
let sourceIndex =
|
|
717
|
+
(sourceY * params.source${source}Width + sourceX) *
|
|
718
|
+
${sourceBlocks[source]}u + sourceBlock;
|
|
719
|
+
value = input${source}[sourceIndex];
|
|
720
|
+
}`;
|
|
721
|
+
};
|
|
722
|
+
return /* wgsl */ `${shaderPreamble(precision)}
|
|
723
|
+
struct Params {
|
|
724
|
+
outputWidth: u32,
|
|
725
|
+
outputHeight: u32,
|
|
726
|
+
outputBlocks: u32,
|
|
727
|
+
inputBlocks: u32,
|
|
728
|
+
source0Width: u32,
|
|
729
|
+
source0Height: u32,
|
|
730
|
+
source1Width: u32,
|
|
731
|
+
source1Height: u32,
|
|
732
|
+
}
|
|
733
|
+
@group(0) @binding(0) var<storage, read> input0: array<${valueType}>;
|
|
734
|
+
@group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
|
|
735
|
+
@group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
|
|
736
|
+
@group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
|
|
737
|
+
@group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
|
|
738
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
739
|
+
|
|
740
|
+
var<workgroup> inputTile: array<${valueType}, ${inputTileValues}>;
|
|
741
|
+
var<workgroup> weightTile: array<${valueType}, ${weightTileValues}>;
|
|
742
|
+
|
|
743
|
+
@compute @workgroup_size(${TILED_CONV_WORKGROUP}, ${TILED_CONV_WORKGROUP}, 1)
|
|
744
|
+
fn main(
|
|
745
|
+
@builtin(local_invocation_id) localId: vec3<u32>,
|
|
746
|
+
@builtin(workgroup_id) workgroupId: vec3<u32>
|
|
747
|
+
) {
|
|
748
|
+
let spatialBase =
|
|
749
|
+
workgroupId.x * ${TILED_CONV_M}u +
|
|
750
|
+
localId.y * ${TILED_CONV_ROWS_PER_THREAD}u;
|
|
751
|
+
let outputBlock =
|
|
752
|
+
workgroupId.y * ${TILED_CONV_N_BLOCKS}u + localId.x;
|
|
753
|
+
let spatialCount = params.outputWidth * params.outputHeight;
|
|
754
|
+
var acc: array<vec4<f32>, ${TILED_CONV_ROWS_PER_THREAD}>;
|
|
755
|
+
if (outputBlock < ${outputBlocks}u) {
|
|
756
|
+
for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
|
|
757
|
+
acc[row] = bias[outputBlock];
|
|
758
|
+
}
|
|
759
|
+
}
|
|
760
|
+
|
|
761
|
+
let localLinear =
|
|
762
|
+
localId.y * ${TILED_CONV_WORKGROUP}u + localId.x;
|
|
763
|
+
let totalK = ${inputBlocks * 9}u;
|
|
764
|
+
for (var kBase = 0u; kBase < totalK; kBase += ${TILED_CONV_K_BLOCKS}u) {
|
|
765
|
+
for (
|
|
766
|
+
var loadIndex = localLinear;
|
|
767
|
+
loadIndex < ${inputTileValues}u;
|
|
768
|
+
loadIndex += ${workgroupThreads}u
|
|
769
|
+
) {
|
|
770
|
+
let tileSpatial = loadIndex / ${TILED_CONV_K_BLOCKS}u;
|
|
771
|
+
let tileK = loadIndex % ${TILED_CONV_K_BLOCKS}u;
|
|
772
|
+
let outputSpatialIndex =
|
|
773
|
+
workgroupId.x * ${TILED_CONV_M}u + tileSpatial;
|
|
774
|
+
let kIndex = kBase + tileK;
|
|
775
|
+
var value = ${valueType}(0.0);
|
|
776
|
+
if (outputSpatialIndex < spatialCount && kIndex < totalK) {
|
|
777
|
+
let outputY = outputSpatialIndex / params.outputWidth;
|
|
778
|
+
let outputX = outputSpatialIndex % params.outputWidth;
|
|
779
|
+
let inputBlock = kIndex % ${inputBlocks}u;
|
|
780
|
+
let kernelIndex = kIndex / ${inputBlocks}u;
|
|
781
|
+
let kernelY = kernelIndex / 3u;
|
|
782
|
+
let kernelX = kernelIndex % 3u;
|
|
783
|
+
let inputY = i32(outputY) + i32(kernelY) - 1;
|
|
784
|
+
let inputX = i32(outputX) + i32(kernelX) - 1;
|
|
785
|
+
if (
|
|
786
|
+
inputY >= 0 && inputY < i32(params.outputHeight) &&
|
|
787
|
+
inputX >= 0 && inputX < i32(params.outputWidth)
|
|
788
|
+
) {
|
|
789
|
+
if (inputBlock < ${sourceBlocks[0]}u) {
|
|
790
|
+
${sourceRead(0, 'inputBlock')}
|
|
791
|
+
} else {
|
|
792
|
+
${sourceRead(1, `inputBlock - ${sourceBlocks[0]}u`)}
|
|
793
|
+
}
|
|
794
|
+
}
|
|
795
|
+
}
|
|
796
|
+
inputTile[loadIndex] = value;
|
|
797
|
+
}
|
|
798
|
+
|
|
799
|
+
for (
|
|
800
|
+
var loadIndex = localLinear;
|
|
801
|
+
loadIndex < ${weightTileValues}u;
|
|
802
|
+
loadIndex += ${workgroupThreads}u
|
|
803
|
+
) {
|
|
804
|
+
let tileK = loadIndex / ${TILED_CONV_N_BLOCKS * 4}u;
|
|
805
|
+
let outputRemainder = loadIndex % ${TILED_CONV_N_BLOCKS * 4}u;
|
|
806
|
+
let tileOutputBlock = outputRemainder / 4u;
|
|
807
|
+
let outputLane = outputRemainder % 4u;
|
|
808
|
+
let loadedOutputBlock =
|
|
809
|
+
workgroupId.y * ${TILED_CONV_N_BLOCKS}u + tileOutputBlock;
|
|
810
|
+
let kIndex = kBase + tileK;
|
|
811
|
+
var value = ${valueType}(0.0);
|
|
812
|
+
if (loadedOutputBlock < ${outputBlocks}u && kIndex < totalK) {
|
|
813
|
+
let inputBlock = kIndex % ${inputBlocks}u;
|
|
814
|
+
let kernelIndex = kIndex / ${inputBlocks}u;
|
|
815
|
+
let kernelY = kernelIndex / 3u;
|
|
816
|
+
let kernelX = kernelIndex % 3u;
|
|
817
|
+
let weightIndex =
|
|
818
|
+
((((loadedOutputBlock * 3u + kernelY) * 3u + kernelX) *
|
|
819
|
+
${inputBlocks}u + inputBlock) * 4u + outputLane);
|
|
820
|
+
value = weights[weightIndex];
|
|
821
|
+
}
|
|
822
|
+
weightTile[loadIndex] = value;
|
|
823
|
+
}
|
|
824
|
+
|
|
825
|
+
workgroupBarrier();
|
|
826
|
+
${tiledAccumulationCode(precision)}
|
|
827
|
+
workgroupBarrier();
|
|
828
|
+
}
|
|
829
|
+
|
|
830
|
+
if (outputBlock < ${outputBlocks}u) {
|
|
831
|
+
for (var row = 0u; row < ${TILED_CONV_ROWS_PER_THREAD}u; row++) {
|
|
832
|
+
let spatialIndex = spatialBase + row;
|
|
833
|
+
if (spatialIndex < spatialCount) {
|
|
834
|
+
let outputIndex = spatialIndex * ${outputBlocks}u + outputBlock;
|
|
835
|
+
outputData[outputIndex] = ${stored};
|
|
836
|
+
}
|
|
837
|
+
}
|
|
838
|
+
}
|
|
839
|
+
}
|
|
840
|
+
`;
|
|
841
|
+
}
|
|
842
|
+
function createFusedDecoderShader(precision, activation, sourceBlocks, upsampledSource, outputBlocks) {
|
|
843
|
+
const valueType = storageVecType(precision);
|
|
844
|
+
const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
|
|
845
|
+
const sourceCode = (source, blockOffset) => {
|
|
846
|
+
const isUpsampled = source === upsampledSource;
|
|
847
|
+
return /* wgsl */ `
|
|
848
|
+
{
|
|
849
|
+
let sourceX = ${isUpsampled ? 'u32(inputX) / 2u' : 'u32(inputX)'};
|
|
850
|
+
let sourceY = ${isUpsampled ? 'u32(inputY) / 2u' : 'u32(inputY)'};
|
|
851
|
+
let sourcePixelBase = (sourceY * params.source${source}Width + sourceX) * ${sourceBlocks[source]}u;
|
|
852
|
+
for (var sourceBlock = 0u; sourceBlock < ${sourceBlocks[source]}u; sourceBlock++) {
|
|
853
|
+
let inputBlock = ${blockOffset}u + sourceBlock;
|
|
854
|
+
${accumulationCode(`input${source}[sourcePixelBase + sourceBlock]`, `((((gid.z * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`, precision)}
|
|
855
|
+
}
|
|
856
|
+
}
|
|
857
|
+
`;
|
|
858
|
+
};
|
|
859
|
+
const stored = storeExpression(activationExpression('acc', activation), precision);
|
|
860
|
+
return /* wgsl */ `${shaderPreamble(precision)}
|
|
861
|
+
struct Params {
|
|
862
|
+
outputWidth: u32,
|
|
863
|
+
outputHeight: u32,
|
|
864
|
+
outputBlocks: u32,
|
|
865
|
+
inputBlocks: u32,
|
|
866
|
+
source0Width: u32,
|
|
867
|
+
source0Height: u32,
|
|
868
|
+
source1Width: u32,
|
|
869
|
+
source1Height: u32,
|
|
870
|
+
}
|
|
871
|
+
@group(0) @binding(0) var<storage, read> input0: array<${valueType}>;
|
|
872
|
+
@group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
|
|
873
|
+
@group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
|
|
874
|
+
@group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
|
|
875
|
+
@group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
|
|
876
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
877
|
+
|
|
878
|
+
@compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
|
|
879
|
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
880
|
+
if (gid.x >= params.outputWidth || gid.y >= params.outputHeight || gid.z >= ${outputBlocks}u) {
|
|
881
|
+
return;
|
|
882
|
+
}
|
|
883
|
+
var acc = bias[gid.z];
|
|
884
|
+
for (var ky = 0u; ky < 3u; ky++) {
|
|
885
|
+
let inputY = i32(gid.y) + i32(ky) - 1;
|
|
886
|
+
if (inputY < 0 || inputY >= i32(params.outputHeight)) { continue; }
|
|
887
|
+
for (var kx = 0u; kx < 3u; kx++) {
|
|
888
|
+
let inputX = i32(gid.x) + i32(kx) - 1;
|
|
889
|
+
if (inputX < 0 || inputX >= i32(params.outputWidth)) { continue; }
|
|
890
|
+
${sourceCode(0, 0)}
|
|
891
|
+
${sourceCode(1, sourceBlocks[0])}
|
|
892
|
+
}
|
|
893
|
+
}
|
|
894
|
+
let outputIndex = (gid.y * params.outputWidth + gid.x) * ${outputBlocks}u + gid.z;
|
|
895
|
+
outputData[outputIndex] = ${stored};
|
|
896
|
+
}
|
|
897
|
+
`;
|
|
898
|
+
}
|
|
899
|
+
function createSpatialDecoderShader(precision, activation, sourceBlocks, upsampledSource, outputBlocks) {
|
|
900
|
+
const valueType = storageVecType(precision);
|
|
901
|
+
const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
|
|
902
|
+
const patchValues = SPATIAL_CONV_PATCH * SPATIAL_CONV_PATCH * inputBlocks;
|
|
903
|
+
const workgroupThreads = SPATIAL_CONV_WORKGROUP * SPATIAL_CONV_WORKGROUP;
|
|
904
|
+
const stored = storeExpression(activationExpression('acc', activation), precision);
|
|
905
|
+
const sourceRead = (source, blockExpression) => {
|
|
906
|
+
const sourceX = source === upsampledSource ? 'u32(inputX) / 2u' : 'u32(inputX)';
|
|
907
|
+
const sourceY = source === upsampledSource ? 'u32(inputY) / 2u' : 'u32(inputY)';
|
|
908
|
+
return /* wgsl */ `
|
|
909
|
+
let sourceBlock = ${blockExpression};
|
|
910
|
+
let sourceIndex =
|
|
911
|
+
(${sourceY} * params.source${source}Width + ${sourceX}) *
|
|
912
|
+
${sourceBlocks[source]}u + sourceBlock;
|
|
913
|
+
value = input${source}[sourceIndex];`;
|
|
914
|
+
};
|
|
915
|
+
return /* wgsl */ `${shaderPreamble(precision)}
|
|
916
|
+
struct Params {
|
|
917
|
+
outputWidth: u32,
|
|
918
|
+
outputHeight: u32,
|
|
919
|
+
outputBlocks: u32,
|
|
920
|
+
inputBlocks: u32,
|
|
921
|
+
source0Width: u32,
|
|
922
|
+
source0Height: u32,
|
|
923
|
+
source1Width: u32,
|
|
924
|
+
source1Height: u32,
|
|
925
|
+
}
|
|
926
|
+
@group(0) @binding(0) var<storage, read> input0: array<${valueType}>;
|
|
927
|
+
@group(0) @binding(1) var<storage, read> input1: array<${valueType}>;
|
|
928
|
+
@group(0) @binding(2) var<storage, read> weights: array<${valueType}>;
|
|
929
|
+
@group(0) @binding(3) var<storage, read> bias: array<vec4<f32>>;
|
|
930
|
+
@group(0) @binding(4) var<storage, read_write> outputData: array<${valueType}>;
|
|
931
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
932
|
+
|
|
933
|
+
var<workgroup> inputPatch: array<${valueType}, ${patchValues}>;
|
|
934
|
+
|
|
935
|
+
@compute @workgroup_size(${SPATIAL_CONV_WORKGROUP}, ${SPATIAL_CONV_WORKGROUP}, 1)
|
|
936
|
+
fn main(
|
|
937
|
+
@builtin(local_invocation_id) localId: vec3<u32>,
|
|
938
|
+
@builtin(workgroup_id) workgroupId: vec3<u32>
|
|
939
|
+
) {
|
|
940
|
+
let localLinear =
|
|
941
|
+
localId.y * ${SPATIAL_CONV_WORKGROUP}u + localId.x;
|
|
942
|
+
for (
|
|
943
|
+
var loadIndex = localLinear;
|
|
944
|
+
loadIndex < ${patchValues}u;
|
|
945
|
+
loadIndex += ${workgroupThreads}u
|
|
946
|
+
) {
|
|
947
|
+
let patchPixel = loadIndex / ${inputBlocks}u;
|
|
948
|
+
let inputBlock = loadIndex % ${inputBlocks}u;
|
|
949
|
+
let patchX = patchPixel % ${SPATIAL_CONV_PATCH}u;
|
|
950
|
+
let patchY = patchPixel / ${SPATIAL_CONV_PATCH}u;
|
|
951
|
+
let inputX =
|
|
952
|
+
i32(workgroupId.x * ${SPATIAL_CONV_WORKGROUP}u + patchX) - 1;
|
|
953
|
+
let inputY =
|
|
954
|
+
i32(workgroupId.y * ${SPATIAL_CONV_WORKGROUP}u + patchY) - 1;
|
|
955
|
+
var value = ${valueType}(0.0);
|
|
956
|
+
if (
|
|
957
|
+
inputX >= 0 && inputX < i32(params.outputWidth) &&
|
|
958
|
+
inputY >= 0 && inputY < i32(params.outputHeight)
|
|
959
|
+
) {
|
|
960
|
+
if (inputBlock < ${sourceBlocks[0]}u) {
|
|
961
|
+
${sourceRead(0, 'inputBlock')}
|
|
962
|
+
} else {
|
|
963
|
+
${sourceRead(1, `inputBlock - ${sourceBlocks[0]}u`)}
|
|
964
|
+
}
|
|
965
|
+
}
|
|
966
|
+
inputPatch[loadIndex] = value;
|
|
967
|
+
}
|
|
968
|
+
workgroupBarrier();
|
|
969
|
+
|
|
970
|
+
let outputX =
|
|
971
|
+
workgroupId.x * ${SPATIAL_CONV_WORKGROUP}u + localId.x;
|
|
972
|
+
let outputY =
|
|
973
|
+
workgroupId.y * ${SPATIAL_CONV_WORKGROUP}u + localId.y;
|
|
974
|
+
let outputBlock = workgroupId.z;
|
|
975
|
+
if (
|
|
976
|
+
outputX >= params.outputWidth || outputY >= params.outputHeight ||
|
|
977
|
+
outputBlock >= ${outputBlocks}u
|
|
978
|
+
) {
|
|
979
|
+
return;
|
|
980
|
+
}
|
|
981
|
+
|
|
982
|
+
var acc = bias[outputBlock];
|
|
983
|
+
for (var ky = 0u; ky < 3u; ky++) {
|
|
984
|
+
for (var kx = 0u; kx < 3u; kx++) {
|
|
985
|
+
let patchBase =
|
|
986
|
+
((localId.y + ky) * ${SPATIAL_CONV_PATCH}u + localId.x + kx) *
|
|
987
|
+
${inputBlocks}u;
|
|
988
|
+
for (var inputBlock = 0u; inputBlock < ${inputBlocks}u; inputBlock++) {
|
|
989
|
+
${accumulationCode('inputPatch[patchBase + inputBlock]', `((((outputBlock * 3u + ky) * 3u + kx) * ${inputBlocks}u + inputBlock) * 4u)`, precision)}
|
|
990
|
+
}
|
|
991
|
+
}
|
|
992
|
+
}
|
|
993
|
+
let outputIndex =
|
|
994
|
+
(outputY * params.outputWidth + outputX) * ${outputBlocks}u + outputBlock;
|
|
995
|
+
outputData[outputIndex] = ${stored};
|
|
996
|
+
}
|
|
997
|
+
`;
|
|
998
|
+
}
|
|
999
|
+
function createInputPackShader(precision, sourceCount) {
|
|
1000
|
+
const outputType = storageVecType(precision);
|
|
1001
|
+
const inputBindings = Array.from({ length: sourceCount }, (_, index) => `@group(0) @binding(${index}) var<storage, read> input${index}: array<vec4<f32>>;`).join('\n');
|
|
1002
|
+
const readBranches = Array.from({ length: sourceCount }, (_, index) => {
|
|
1003
|
+
const firstChannel = index * 3;
|
|
1004
|
+
return `if (channel < ${firstChannel + 3}u) { return input${index}[pixel][channel - ${firstChannel}u]; }`;
|
|
1005
|
+
}).join('\n ');
|
|
1006
|
+
const outputBinding = sourceCount;
|
|
1007
|
+
const paramsBinding = sourceCount + 1;
|
|
1008
|
+
return /* wgsl */ `${shaderPreamble(precision)}
|
|
1009
|
+
struct Params {
|
|
1010
|
+
width: u32,
|
|
1011
|
+
height: u32,
|
|
1012
|
+
outputBlocks: u32,
|
|
1013
|
+
inputChannels: u32,
|
|
1014
|
+
}
|
|
1015
|
+
${inputBindings}
|
|
1016
|
+
@group(0) @binding(${outputBinding}) var<storage, read_write> outputData: array<${outputType}>;
|
|
1017
|
+
@group(0) @binding(${paramsBinding}) var<uniform> params: Params;
|
|
1018
|
+
|
|
1019
|
+
fn readChannel(pixel: u32, channel: u32) -> f32 {
|
|
1020
|
+
${readBranches}
|
|
1021
|
+
return 0.0;
|
|
1022
|
+
}
|
|
1023
|
+
|
|
1024
|
+
@compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
|
|
1025
|
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
|
1026
|
+
if (gid.x >= params.width || gid.y >= params.height || gid.z >= params.outputBlocks) {
|
|
1027
|
+
return;
|
|
1028
|
+
}
|
|
1029
|
+
let pixel = gid.y * params.width + gid.x;
|
|
1030
|
+
let firstChannel = gid.z * 4u;
|
|
1031
|
+
let value = vec4<f32>(
|
|
1032
|
+
readChannel(pixel, firstChannel),
|
|
1033
|
+
readChannel(pixel, firstChannel + 1u),
|
|
1034
|
+
readChannel(pixel, firstChannel + 2u),
|
|
1035
|
+
readChannel(pixel, firstChannel + 3u)
|
|
1036
|
+
);
|
|
1037
|
+
outputData[pixel * params.outputBlocks + gid.z] = ${storeExpression('value', precision)};
|
|
1038
|
+
}
|
|
1039
|
+
`;
|
|
1040
|
+
}
|
|
1041
|
+
function executableInputs(node) {
|
|
1042
|
+
if (node.op === 'concat')
|
|
1043
|
+
return node.inputs;
|
|
1044
|
+
if (node.op === 'fusedUpsampleConcatConv2d') {
|
|
1045
|
+
return node.inputs.map((input) => input.value);
|
|
1046
|
+
}
|
|
1047
|
+
return [node.input];
|
|
1048
|
+
}
|
|
1049
|
+
export function resolveNativeUNetPrecision(device, requested = 'auto') {
|
|
1050
|
+
const hasShaderF16 = device.features.has('shader-f16');
|
|
1051
|
+
if (requested === 'fp16' && !hasShaderF16) {
|
|
1052
|
+
throw new Error('OIDN FP16 was requested but the GPUDevice does not have shader-f16 enabled');
|
|
1053
|
+
}
|
|
1054
|
+
if (requested === 'auto')
|
|
1055
|
+
return hasShaderF16 ? 'fp16' : 'fp32';
|
|
1056
|
+
return requested;
|
|
1057
|
+
}
|
|
1058
|
+
/** Native, model-driven OIDN U-Net executor. */
|
|
1059
|
+
export class NativeUNetExecutor {
|
|
1060
|
+
_device;
|
|
1061
|
+
precision;
|
|
1062
|
+
kernelSetting;
|
|
1063
|
+
maxSpatialInputBlocks;
|
|
1064
|
+
subgroupsAvailable;
|
|
1065
|
+
_model;
|
|
1066
|
+
_packedConvs = new Map();
|
|
1067
|
+
_pipelineCache;
|
|
1068
|
+
_pipelinePromises;
|
|
1069
|
+
_executionCache = new Map();
|
|
1070
|
+
_retiredExecutions = new Set();
|
|
1071
|
+
_clock = 0;
|
|
1072
|
+
_shapeCacheSize;
|
|
1073
|
+
_profileNextExecution = false;
|
|
1074
|
+
_lastExecutionProfile;
|
|
1075
|
+
_profileOperations = 0;
|
|
1076
|
+
_resources = new OIDNResourceTracker();
|
|
1077
|
+
_disposed = false;
|
|
1078
|
+
constructor(_device, model, options = {}) {
|
|
1079
|
+
this._device = _device;
|
|
1080
|
+
const pipelineCache = sharedPipelineCache(_device);
|
|
1081
|
+
this._pipelineCache = pipelineCache.ready;
|
|
1082
|
+
this._pipelinePromises = pipelineCache.pending;
|
|
1083
|
+
this.precision = resolveNativeUNetPrecision(_device, options.precision ?? 'auto');
|
|
1084
|
+
this.kernelSetting = options.kernel ?? 'auto';
|
|
1085
|
+
this.subgroupsAvailable = _device.features.has('subgroups');
|
|
1086
|
+
this.maxSpatialInputBlocks =
|
|
1087
|
+
this.precision === 'fp16' &&
|
|
1088
|
+
_device.limits.maxComputeInvocationsPerWorkgroup >=
|
|
1089
|
+
SPATIAL_CONV_WORKGROUP * SPATIAL_CONV_WORKGROUP &&
|
|
1090
|
+
_device.limits.maxComputeWorkgroupSizeX >= SPATIAL_CONV_WORKGROUP &&
|
|
1091
|
+
_device.limits.maxComputeWorkgroupSizeY >= SPATIAL_CONV_WORKGROUP
|
|
1092
|
+
? Math.floor(_device.limits.maxComputeWorkgroupStorageSize /
|
|
1093
|
+
(SPATIAL_CONV_PATCH * SPATIAL_CONV_PATCH * 4 * 2))
|
|
1094
|
+
: 0;
|
|
1095
|
+
this._shapeCacheSize = Math.max(1, options.shapeCacheSize ?? 2);
|
|
1096
|
+
if (model.inputChannels % 3 !== 0 ||
|
|
1097
|
+
model.inputChannels < 3 ||
|
|
1098
|
+
model.inputChannels > 9) {
|
|
1099
|
+
throw new Error(`Native OIDN expects 3, 6, or 9 input channels, got ${model.inputChannels}`);
|
|
1100
|
+
}
|
|
1101
|
+
this._model = {
|
|
1102
|
+
spec: model.spec,
|
|
1103
|
+
inputChannels: model.inputChannels,
|
|
1104
|
+
outputChannels: model.outputChannels,
|
|
1105
|
+
channelsByValue: new Map(model.channelsByValue),
|
|
1106
|
+
convChannels: new Map(model.convChannels)
|
|
1107
|
+
};
|
|
1108
|
+
for (const node of optimizeModelGraph(this._model, {
|
|
1109
|
+
fuseConvPool: false
|
|
1110
|
+
}).nodes) {
|
|
1111
|
+
if (node.op !== 'conv2d' &&
|
|
1112
|
+
node.op !== 'maxPool2d' &&
|
|
1113
|
+
node.op !== 'fusedConvReluMaxPool2d' &&
|
|
1114
|
+
node.op !== 'fusedUpsampleConcatConv2d') {
|
|
1115
|
+
throw new Error(`Native OIDN descriptor ${model.spec.id} leaves unsupported ` +
|
|
1116
|
+
`${node.op} node ${node.id} after graph optimization`);
|
|
1117
|
+
}
|
|
1118
|
+
}
|
|
1119
|
+
try {
|
|
1120
|
+
for (const [id, tensors] of model.convTensors) {
|
|
1121
|
+
const packed = packConvTensors(_device, id, tensors, this.precision);
|
|
1122
|
+
this._resources.track('gpu-buffer', packed.weights);
|
|
1123
|
+
this._resources.track('gpu-buffer', packed.bias);
|
|
1124
|
+
this._packedConvs.set(id, packed);
|
|
1125
|
+
}
|
|
1126
|
+
}
|
|
1127
|
+
catch (error) {
|
|
1128
|
+
for (const packed of this._packedConvs.values()) {
|
|
1129
|
+
this._releaseBuffer(packed.weights);
|
|
1130
|
+
this._releaseBuffer(packed.bias);
|
|
1131
|
+
}
|
|
1132
|
+
this._packedConvs.clear();
|
|
1133
|
+
throw error;
|
|
1134
|
+
}
|
|
1135
|
+
}
|
|
1136
|
+
_pipeline(key, code) {
|
|
1137
|
+
let pipeline = this._pipelineCache.get(key);
|
|
1138
|
+
if (!pipeline) {
|
|
1139
|
+
pipeline = this._device.createComputePipeline({
|
|
1140
|
+
label: `oidn/${key}`,
|
|
1141
|
+
layout: 'auto',
|
|
1142
|
+
compute: {
|
|
1143
|
+
module: this._device.createShaderModule({
|
|
1144
|
+
label: `oidn/${key}`,
|
|
1145
|
+
code
|
|
1146
|
+
}),
|
|
1147
|
+
entryPoint: 'main'
|
|
1148
|
+
}
|
|
1149
|
+
});
|
|
1150
|
+
this._pipelineCache.set(key, pipeline);
|
|
1151
|
+
}
|
|
1152
|
+
return pipeline;
|
|
1153
|
+
}
|
|
1154
|
+
_pipelineAsync(key, code) {
|
|
1155
|
+
const ready = this._pipelineCache.get(key);
|
|
1156
|
+
if (ready)
|
|
1157
|
+
return Promise.resolve(ready);
|
|
1158
|
+
const pending = this._pipelinePromises.get(key);
|
|
1159
|
+
if (pending)
|
|
1160
|
+
return pending;
|
|
1161
|
+
const promise = this._device.createComputePipelineAsync({
|
|
1162
|
+
label: `oidn/${key}`,
|
|
1163
|
+
layout: 'auto',
|
|
1164
|
+
compute: {
|
|
1165
|
+
module: this._device.createShaderModule({
|
|
1166
|
+
label: `oidn/${key}`,
|
|
1167
|
+
code
|
|
1168
|
+
}),
|
|
1169
|
+
entryPoint: 'main'
|
|
1170
|
+
}
|
|
1171
|
+
}).then((pipeline) => {
|
|
1172
|
+
this._pipelineCache.set(key, pipeline);
|
|
1173
|
+
this._pipelinePromises.delete(key);
|
|
1174
|
+
return pipeline;
|
|
1175
|
+
}, (error) => {
|
|
1176
|
+
this._pipelinePromises.delete(key);
|
|
1177
|
+
throw error;
|
|
1178
|
+
});
|
|
1179
|
+
this._pipelinePromises.set(key, promise);
|
|
1180
|
+
return promise;
|
|
1181
|
+
}
|
|
1182
|
+
_nodePipelineSpec(node, isFinal) {
|
|
1183
|
+
if (node.op === 'conv2d') {
|
|
1184
|
+
const outputPrecision = isFinal ? 'fp32' : this.precision;
|
|
1185
|
+
const inputBlocks = blocksForChannels(this._model.convChannels.get(node.id).inputChannels);
|
|
1186
|
+
const outputBlocks = blocksForChannels(this._model.convChannels.get(node.id).outputChannels);
|
|
1187
|
+
const kernel = this._selectConvKernel(inputBlocks, isFinal);
|
|
1188
|
+
const key = `conv-${kernel}/${this.precision}/` +
|
|
1189
|
+
`${outputPrecision}/${node.activation}/` +
|
|
1190
|
+
`in${inputBlocks}/out${outputBlocks}`;
|
|
1191
|
+
return {
|
|
1192
|
+
key,
|
|
1193
|
+
kernel,
|
|
1194
|
+
code: kernel === 'implicit-gemm'
|
|
1195
|
+
? createTiledConvShader(this.precision, outputPrecision, node.activation, inputBlocks, outputBlocks)
|
|
1196
|
+
: kernel === 'spatial'
|
|
1197
|
+
? createSpatialConvShader(this.precision, outputPrecision, node.activation, inputBlocks, outputBlocks)
|
|
1198
|
+
: kernel === 'subgroup'
|
|
1199
|
+
? createSubgroupConvShader(this.precision, outputPrecision, node.activation, inputBlocks, outputBlocks)
|
|
1200
|
+
: createConvShader(this.precision, outputPrecision, node.activation, inputBlocks, outputBlocks)
|
|
1201
|
+
};
|
|
1202
|
+
}
|
|
1203
|
+
if (node.op === 'maxPool2d') {
|
|
1204
|
+
const outputBlocks = blocksForChannels(this._model.channelsByValue.get(node.id));
|
|
1205
|
+
const key = `max-pool/${this.precision}/out${outputBlocks}`;
|
|
1206
|
+
return {
|
|
1207
|
+
key,
|
|
1208
|
+
kernel: 'direct',
|
|
1209
|
+
code: createMaxPoolShader(this.precision, outputBlocks)
|
|
1210
|
+
};
|
|
1211
|
+
}
|
|
1212
|
+
if (node.op === 'fusedConvReluMaxPool2d') {
|
|
1213
|
+
const inputBlocks = blocksForChannels(this._model.convChannels.get(node.conv.id).inputChannels);
|
|
1214
|
+
const outputBlocks = blocksForChannels(this._model.convChannels.get(node.conv.id).outputChannels);
|
|
1215
|
+
const key = `conv-pool/${this.precision}/${node.conv.activation}/` +
|
|
1216
|
+
`in${inputBlocks}/out${outputBlocks}`;
|
|
1217
|
+
return {
|
|
1218
|
+
key,
|
|
1219
|
+
kernel: 'direct',
|
|
1220
|
+
code: createFusedConvPoolShader(this.precision, node.conv.activation, inputBlocks, outputBlocks)
|
|
1221
|
+
};
|
|
1222
|
+
}
|
|
1223
|
+
if (node.op === 'fusedUpsampleConcatConv2d') {
|
|
1224
|
+
if (node.inputs.length !== 2) {
|
|
1225
|
+
throw new Error(`Native fused decoder ${node.id} requires two inputs`);
|
|
1226
|
+
}
|
|
1227
|
+
const sourceBlocks = node.inputs.map((input) => blocksForChannels(this._model.channelsByValue.get(input.value)));
|
|
1228
|
+
const upsampledSource = node.inputs.findIndex((input) => input.upsample);
|
|
1229
|
+
if (upsampledSource !== 0 && upsampledSource !== 1) {
|
|
1230
|
+
throw new Error(`Native fused decoder ${node.id} has no upsample input`);
|
|
1231
|
+
}
|
|
1232
|
+
const inputBlocks = sourceBlocks[0] + sourceBlocks[1];
|
|
1233
|
+
const selectedKernel = this._selectConvKernel(inputBlocks, false);
|
|
1234
|
+
// The subgroup broadcast path currently targets the common standalone
|
|
1235
|
+
// convolution layout; fused decoder reads use the direct kernel.
|
|
1236
|
+
const kernel = selectedKernel === 'subgroup'
|
|
1237
|
+
? 'direct'
|
|
1238
|
+
: selectedKernel;
|
|
1239
|
+
const outputBlocks = blocksForChannels(this._model.convChannels.get(node.conv.id).outputChannels);
|
|
1240
|
+
const key = `decoder-${kernel}/${this.precision}/` +
|
|
1241
|
+
`${node.conv.activation}/` +
|
|
1242
|
+
`${sourceBlocks.join('+')}/out${outputBlocks}/up${upsampledSource}`;
|
|
1243
|
+
return {
|
|
1244
|
+
key,
|
|
1245
|
+
kernel,
|
|
1246
|
+
code: kernel === 'implicit-gemm'
|
|
1247
|
+
? createTiledDecoderShader(this.precision, node.conv.activation, sourceBlocks, upsampledSource, outputBlocks)
|
|
1248
|
+
: kernel === 'spatial'
|
|
1249
|
+
? createSpatialDecoderShader(this.precision, node.conv.activation, sourceBlocks, upsampledSource, outputBlocks)
|
|
1250
|
+
: createFusedDecoderShader(this.precision, node.conv.activation, sourceBlocks, upsampledSource, outputBlocks)
|
|
1251
|
+
};
|
|
1252
|
+
}
|
|
1253
|
+
throw new Error(`Native OIDN does not implement unfused ${node.op} node ${node.id}`);
|
|
1254
|
+
}
|
|
1255
|
+
_selectConvKernel(inputBlocks, isFinal) {
|
|
1256
|
+
const spatialFits = this.precision === 'fp16' &&
|
|
1257
|
+
inputBlocks <= this.maxSpatialInputBlocks;
|
|
1258
|
+
if (this.kernelSetting === 'direct')
|
|
1259
|
+
return 'direct';
|
|
1260
|
+
if (this.kernelSetting === 'spatial') {
|
|
1261
|
+
return spatialFits ? 'spatial' : 'direct';
|
|
1262
|
+
}
|
|
1263
|
+
if (this.kernelSetting === 'implicit-gemm') {
|
|
1264
|
+
return isFinal ? 'direct' : 'implicit-gemm';
|
|
1265
|
+
}
|
|
1266
|
+
if (this.kernelSetting === 'subgroup') {
|
|
1267
|
+
return this.subgroupsAvailable ? 'subgroup' : 'direct';
|
|
1268
|
+
}
|
|
1269
|
+
if (this.precision === 'fp32' && !isFinal)
|
|
1270
|
+
return 'implicit-gemm';
|
|
1271
|
+
return 'direct';
|
|
1272
|
+
}
|
|
1273
|
+
_nodePipeline(node, isFinal) {
|
|
1274
|
+
const { key, code } = this._nodePipelineSpec(node, isFinal);
|
|
1275
|
+
return this._pipeline(key, code);
|
|
1276
|
+
}
|
|
1277
|
+
/** Compiles all shape-independent kernels before the model reports ready. */
|
|
1278
|
+
async prepare() {
|
|
1279
|
+
if (this._disposed)
|
|
1280
|
+
throw new Error('Native OIDN executor is disposed');
|
|
1281
|
+
const graph = optimizeModelGraph(this._model, { fuseConvPool: false });
|
|
1282
|
+
const sourceCount = this._model.inputChannels / 3;
|
|
1283
|
+
const specs = [
|
|
1284
|
+
{
|
|
1285
|
+
key: `pack/${this.precision}/${sourceCount}`,
|
|
1286
|
+
code: createInputPackShader(this.precision, sourceCount)
|
|
1287
|
+
},
|
|
1288
|
+
...graph.nodes.map((node) => this._nodePipelineSpec(node, node.id === graph.spec.output))
|
|
1289
|
+
];
|
|
1290
|
+
await Promise.all(specs.map(({ key, code }) => this._pipelineAsync(key, code)));
|
|
1291
|
+
}
|
|
1292
|
+
_createExecution(width, height) {
|
|
1293
|
+
const plan = planModelExecution(this._model, width, height, {
|
|
1294
|
+
fuseConvPool: false
|
|
1295
|
+
});
|
|
1296
|
+
const valueBuffers = new Map();
|
|
1297
|
+
const slots = [];
|
|
1298
|
+
const lastUses = new Map();
|
|
1299
|
+
plan.nodes.forEach((node, index) => {
|
|
1300
|
+
for (const input of executableInputs(node))
|
|
1301
|
+
lastUses.set(input, index);
|
|
1302
|
+
});
|
|
1303
|
+
lastUses.set(plan.spec.output, plan.nodes.length);
|
|
1304
|
+
const createdBuffers = [];
|
|
1305
|
+
const own = (buffer) => {
|
|
1306
|
+
createdBuffers.push(buffer);
|
|
1307
|
+
return this._resources.track('gpu-buffer', buffer);
|
|
1308
|
+
};
|
|
1309
|
+
try {
|
|
1310
|
+
const allocate = (value, shape, bytesPerScalar, index) => {
|
|
1311
|
+
for (const slot of slots) {
|
|
1312
|
+
if (slot.activeValue &&
|
|
1313
|
+
(lastUses.get(slot.activeValue) ?? -1) < index) {
|
|
1314
|
+
slot.activeValue = undefined;
|
|
1315
|
+
}
|
|
1316
|
+
}
|
|
1317
|
+
const requiredSize = activationByteSize(shape, bytesPerScalar);
|
|
1318
|
+
let slot = slots
|
|
1319
|
+
.filter((candidate) => !candidate.activeValue && candidate.capacity >= requiredSize)
|
|
1320
|
+
.sort((a, b) => a.capacity - b.capacity)[0];
|
|
1321
|
+
if (!slot) {
|
|
1322
|
+
const buffer = own(this._device.createBuffer({
|
|
1323
|
+
label: `oidn/activation/${width}x${height}/${slots.length}`,
|
|
1324
|
+
size: roundUp(requiredSize, 4),
|
|
1325
|
+
usage: GPUBufferUsage.STORAGE |
|
|
1326
|
+
GPUBufferUsage.COPY_SRC |
|
|
1327
|
+
GPUBufferUsage.COPY_DST
|
|
1328
|
+
}));
|
|
1329
|
+
slot = { buffer, capacity: requiredSize };
|
|
1330
|
+
slots.push(slot);
|
|
1331
|
+
}
|
|
1332
|
+
slot.activeValue = value;
|
|
1333
|
+
valueBuffers.set(value, slot.buffer);
|
|
1334
|
+
};
|
|
1335
|
+
allocate(plan.spec.input, plan.inputShape, this.precision === 'fp16' ? 2 : 4, -1);
|
|
1336
|
+
plan.plannedNodes.forEach(({ node, outputShape }, index) => {
|
|
1337
|
+
const isFinal = node.id === plan.spec.output;
|
|
1338
|
+
allocate(node.id, outputShape, isFinal || this.precision === 'fp32' ? 4 : 2, index);
|
|
1339
|
+
});
|
|
1340
|
+
const inputSourceCount = this._model.inputChannels / 3;
|
|
1341
|
+
const inputKey = `pack/${this.precision}/${inputSourceCount}`;
|
|
1342
|
+
const inputPipeline = this._pipeline(inputKey, createInputPackShader(this.precision, inputSourceCount));
|
|
1343
|
+
const nodePipelines = [];
|
|
1344
|
+
const nodeKernels = [];
|
|
1345
|
+
const nodeBindings = [];
|
|
1346
|
+
const ownedBuffers = [];
|
|
1347
|
+
plan.plannedNodes.forEach(({ node, outputShape }, index) => {
|
|
1348
|
+
const isFinal = node.id === plan.spec.output;
|
|
1349
|
+
const pipelineSpec = this._nodePipelineSpec(node, isFinal);
|
|
1350
|
+
const pipeline = this._pipeline(pipelineSpec.key, pipelineSpec.code);
|
|
1351
|
+
nodePipelines.push(pipeline);
|
|
1352
|
+
nodeKernels.push(pipelineSpec.kernel ?? 'direct');
|
|
1353
|
+
const output = valueBuffers.get(node.id);
|
|
1354
|
+
let entries;
|
|
1355
|
+
let uniformValues;
|
|
1356
|
+
let convId;
|
|
1357
|
+
if (node.op === 'conv2d') {
|
|
1358
|
+
const inputShape = plan.valueShapes.get(node.input);
|
|
1359
|
+
convId = node.id;
|
|
1360
|
+
uniformValues = [
|
|
1361
|
+
inputShape.width,
|
|
1362
|
+
inputShape.height,
|
|
1363
|
+
outputShape.width,
|
|
1364
|
+
outputShape.height,
|
|
1365
|
+
blocksForChannels(inputShape.channels),
|
|
1366
|
+
blocksForChannels(outputShape.channels)
|
|
1367
|
+
];
|
|
1368
|
+
const packed = this._packedConvs.get(convId);
|
|
1369
|
+
entries = [
|
|
1370
|
+
{ binding: 0, resource: { buffer: valueBuffers.get(node.input) } },
|
|
1371
|
+
{ binding: 1, resource: { buffer: packed.weights } },
|
|
1372
|
+
{ binding: 2, resource: { buffer: packed.bias } },
|
|
1373
|
+
{ binding: 3, resource: { buffer: output } }
|
|
1374
|
+
];
|
|
1375
|
+
}
|
|
1376
|
+
else if (node.op === 'maxPool2d') {
|
|
1377
|
+
const inputShape = plan.valueShapes.get(node.input);
|
|
1378
|
+
uniformValues = [
|
|
1379
|
+
inputShape.width,
|
|
1380
|
+
inputShape.height,
|
|
1381
|
+
outputShape.width,
|
|
1382
|
+
outputShape.height,
|
|
1383
|
+
blocksForChannels(outputShape.channels)
|
|
1384
|
+
];
|
|
1385
|
+
entries = [
|
|
1386
|
+
{ binding: 0, resource: { buffer: valueBuffers.get(node.input) } },
|
|
1387
|
+
{ binding: 1, resource: { buffer: output } }
|
|
1388
|
+
];
|
|
1389
|
+
}
|
|
1390
|
+
else if (node.op === 'fusedConvReluMaxPool2d') {
|
|
1391
|
+
const inputShape = plan.valueShapes.get(node.input);
|
|
1392
|
+
convId = node.conv.id;
|
|
1393
|
+
uniformValues = [
|
|
1394
|
+
inputShape.width,
|
|
1395
|
+
inputShape.height,
|
|
1396
|
+
outputShape.width,
|
|
1397
|
+
outputShape.height,
|
|
1398
|
+
blocksForChannels(inputShape.channels),
|
|
1399
|
+
blocksForChannels(outputShape.channels)
|
|
1400
|
+
];
|
|
1401
|
+
const packed = this._packedConvs.get(convId);
|
|
1402
|
+
entries = [
|
|
1403
|
+
{ binding: 0, resource: { buffer: valueBuffers.get(node.input) } },
|
|
1404
|
+
{ binding: 1, resource: { buffer: packed.weights } },
|
|
1405
|
+
{ binding: 2, resource: { buffer: packed.bias } },
|
|
1406
|
+
{ binding: 3, resource: { buffer: output } }
|
|
1407
|
+
];
|
|
1408
|
+
}
|
|
1409
|
+
else if (node.op === 'fusedUpsampleConcatConv2d') {
|
|
1410
|
+
convId = node.conv.id;
|
|
1411
|
+
const firstShape = plan.valueShapes.get(node.inputs[0].value);
|
|
1412
|
+
const secondShape = plan.valueShapes.get(node.inputs[1].value);
|
|
1413
|
+
uniformValues = [
|
|
1414
|
+
outputShape.width,
|
|
1415
|
+
outputShape.height,
|
|
1416
|
+
blocksForChannels(outputShape.channels),
|
|
1417
|
+
blocksForChannels(this._model.convChannels.get(convId).inputChannels),
|
|
1418
|
+
firstShape.width,
|
|
1419
|
+
firstShape.height,
|
|
1420
|
+
secondShape.width,
|
|
1421
|
+
secondShape.height
|
|
1422
|
+
];
|
|
1423
|
+
const packed = this._packedConvs.get(convId);
|
|
1424
|
+
entries = [
|
|
1425
|
+
{
|
|
1426
|
+
binding: 0,
|
|
1427
|
+
resource: { buffer: valueBuffers.get(node.inputs[0].value) }
|
|
1428
|
+
},
|
|
1429
|
+
{
|
|
1430
|
+
binding: 1,
|
|
1431
|
+
resource: { buffer: valueBuffers.get(node.inputs[1].value) }
|
|
1432
|
+
},
|
|
1433
|
+
{ binding: 2, resource: { buffer: packed.weights } },
|
|
1434
|
+
{ binding: 3, resource: { buffer: packed.bias } },
|
|
1435
|
+
{ binding: 4, resource: { buffer: output } }
|
|
1436
|
+
];
|
|
1437
|
+
}
|
|
1438
|
+
else {
|
|
1439
|
+
throw new Error(`Unexpected native node ${node.op}`);
|
|
1440
|
+
}
|
|
1441
|
+
const uniform = own(createUniformBuffer(this._device, `oidn/${node.id}/params/${width}x${height}`, uniformValues));
|
|
1442
|
+
ownedBuffers.push(uniform);
|
|
1443
|
+
entries.push({ binding: entries.length, resource: { buffer: uniform } });
|
|
1444
|
+
nodeBindings.push(this._device.createBindGroup({
|
|
1445
|
+
label: `oidn/${node.id}/bindings`,
|
|
1446
|
+
layout: pipeline.getBindGroupLayout(0),
|
|
1447
|
+
entries
|
|
1448
|
+
}));
|
|
1449
|
+
});
|
|
1450
|
+
const inputUniform = own(createUniformBuffer(this._device, `oidn/input/params/${width}x${height}`, [
|
|
1451
|
+
width,
|
|
1452
|
+
height,
|
|
1453
|
+
blocksForChannels(this._model.inputChannels),
|
|
1454
|
+
this._model.inputChannels
|
|
1455
|
+
]));
|
|
1456
|
+
ownedBuffers.push(inputUniform);
|
|
1457
|
+
return {
|
|
1458
|
+
plan,
|
|
1459
|
+
valueBuffers,
|
|
1460
|
+
slots,
|
|
1461
|
+
nodeBindings,
|
|
1462
|
+
nodePipelines,
|
|
1463
|
+
nodeKernels,
|
|
1464
|
+
inputPipeline,
|
|
1465
|
+
inputUniform,
|
|
1466
|
+
ownedBuffers,
|
|
1467
|
+
lastUsed: ++this._clock
|
|
1468
|
+
};
|
|
1469
|
+
}
|
|
1470
|
+
catch (error) {
|
|
1471
|
+
for (const buffer of createdBuffers)
|
|
1472
|
+
this._releaseBuffer(buffer);
|
|
1473
|
+
throw error;
|
|
1474
|
+
}
|
|
1475
|
+
}
|
|
1476
|
+
_execution(width, height) {
|
|
1477
|
+
if (this._disposed)
|
|
1478
|
+
throw new Error('Native OIDN executor is disposed');
|
|
1479
|
+
const key = `${width}x${height}`;
|
|
1480
|
+
let execution = this._executionCache.get(key);
|
|
1481
|
+
if (!execution) {
|
|
1482
|
+
execution = this._createExecution(width, height);
|
|
1483
|
+
this._executionCache.set(key, execution);
|
|
1484
|
+
if (this._executionCache.size > this._shapeCacheSize) {
|
|
1485
|
+
const oldest = [...this._executionCache.entries()]
|
|
1486
|
+
.filter(([candidateKey]) => candidateKey !== key)
|
|
1487
|
+
.sort((a, b) => a[1].lastUsed - b[1].lastUsed)[0];
|
|
1488
|
+
if (oldest) {
|
|
1489
|
+
this._executionCache.delete(oldest[0]);
|
|
1490
|
+
// Commands using an evicted plan may still be submitted. Defer actual
|
|
1491
|
+
// destruction until all work currently on the shared queue completes.
|
|
1492
|
+
this._retiredExecutions.add(oldest[1]);
|
|
1493
|
+
void this._device.queue.onSubmittedWorkDone()
|
|
1494
|
+
.catch(() => undefined)
|
|
1495
|
+
.then(() => {
|
|
1496
|
+
this._retiredExecutions.delete(oldest[1]);
|
|
1497
|
+
this._destroyExecution(oldest[1]);
|
|
1498
|
+
});
|
|
1499
|
+
}
|
|
1500
|
+
}
|
|
1501
|
+
}
|
|
1502
|
+
execution.lastUsed = ++this._clock;
|
|
1503
|
+
return execution;
|
|
1504
|
+
}
|
|
1505
|
+
/** Captures per-pass GPU timestamps for the next execute call when supported. */
|
|
1506
|
+
profileNextExecution() {
|
|
1507
|
+
if (!this._device.features.has('timestamp-query'))
|
|
1508
|
+
return false;
|
|
1509
|
+
this._profileNextExecution = true;
|
|
1510
|
+
return true;
|
|
1511
|
+
}
|
|
1512
|
+
getLastExecutionProfile() {
|
|
1513
|
+
return this._lastExecutionProfile;
|
|
1514
|
+
}
|
|
1515
|
+
execute(inputBuffers, width, height) {
|
|
1516
|
+
const sourceCount = this._model.inputChannels / 3;
|
|
1517
|
+
if (inputBuffers.length !== sourceCount) {
|
|
1518
|
+
throw new Error(`Native OIDN expected ${sourceCount} input buffers, got ${inputBuffers.length}`);
|
|
1519
|
+
}
|
|
1520
|
+
const execution = this._execution(width, height);
|
|
1521
|
+
const profileLabels = [
|
|
1522
|
+
'input-pack',
|
|
1523
|
+
...execution.plan.nodes.map((node) => node.id)
|
|
1524
|
+
];
|
|
1525
|
+
const shouldProfile = this._profileNextExecution &&
|
|
1526
|
+
this._device.features.has('timestamp-query');
|
|
1527
|
+
this._profileNextExecution = false;
|
|
1528
|
+
const queryCount = profileLabels.length * 2;
|
|
1529
|
+
const querySet = shouldProfile
|
|
1530
|
+
? this._resources.track('gpu-query-set', this._device.createQuerySet({ type: 'timestamp', count: queryCount }))
|
|
1531
|
+
: undefined;
|
|
1532
|
+
const queryBufferSize = queryCount * 8;
|
|
1533
|
+
let queryResolveBuffer;
|
|
1534
|
+
let queryReadbackBuffer;
|
|
1535
|
+
try {
|
|
1536
|
+
queryResolveBuffer = shouldProfile
|
|
1537
|
+
? this._resources.track('gpu-buffer', this._device.createBuffer({
|
|
1538
|
+
label: `oidn/profile/resolve/${width}x${height}`,
|
|
1539
|
+
size: queryBufferSize,
|
|
1540
|
+
usage: GPUBufferUsage.QUERY_RESOLVE | GPUBufferUsage.COPY_SRC
|
|
1541
|
+
}))
|
|
1542
|
+
: undefined;
|
|
1543
|
+
queryReadbackBuffer = shouldProfile
|
|
1544
|
+
? this._resources.track('gpu-buffer', this._device.createBuffer({
|
|
1545
|
+
label: `oidn/profile/readback/${width}x${height}`,
|
|
1546
|
+
size: queryBufferSize,
|
|
1547
|
+
usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.MAP_READ
|
|
1548
|
+
}))
|
|
1549
|
+
: undefined;
|
|
1550
|
+
}
|
|
1551
|
+
catch (error) {
|
|
1552
|
+
this._releaseBuffer(queryReadbackBuffer);
|
|
1553
|
+
this._releaseBuffer(queryResolveBuffer);
|
|
1554
|
+
this._releaseQuerySet(querySet);
|
|
1555
|
+
throw error;
|
|
1556
|
+
}
|
|
1557
|
+
const passDescriptor = (label, index) => ({
|
|
1558
|
+
label,
|
|
1559
|
+
...(querySet
|
|
1560
|
+
? {
|
|
1561
|
+
timestampWrites: {
|
|
1562
|
+
querySet,
|
|
1563
|
+
beginningOfPassWriteIndex: index * 2,
|
|
1564
|
+
endOfPassWriteIndex: index * 2 + 1
|
|
1565
|
+
}
|
|
1566
|
+
}
|
|
1567
|
+
: {})
|
|
1568
|
+
});
|
|
1569
|
+
try {
|
|
1570
|
+
const encoder = this._device.createCommandEncoder({
|
|
1571
|
+
label: `oidn/native/${width}x${height}`
|
|
1572
|
+
});
|
|
1573
|
+
const inputEntries = inputBuffers.map((buffer, binding) => ({ binding, resource: { buffer } }));
|
|
1574
|
+
inputEntries.push({
|
|
1575
|
+
binding: sourceCount,
|
|
1576
|
+
resource: { buffer: execution.valueBuffers.get(execution.plan.spec.input) }
|
|
1577
|
+
});
|
|
1578
|
+
inputEntries.push({
|
|
1579
|
+
binding: sourceCount + 1,
|
|
1580
|
+
resource: { buffer: execution.inputUniform }
|
|
1581
|
+
});
|
|
1582
|
+
const inputBindings = this._device.createBindGroup({
|
|
1583
|
+
label: 'oidn/input/bindings',
|
|
1584
|
+
layout: execution.inputPipeline.getBindGroupLayout(0),
|
|
1585
|
+
entries: inputEntries
|
|
1586
|
+
});
|
|
1587
|
+
{
|
|
1588
|
+
const pass = encoder.beginComputePass(passDescriptor('oidn/input-pack', 0));
|
|
1589
|
+
pass.setPipeline(execution.inputPipeline);
|
|
1590
|
+
pass.setBindGroup(0, inputBindings);
|
|
1591
|
+
pass.dispatchWorkgroups(Math.ceil(width / WORKGROUP_SIZE), Math.ceil(height / WORKGROUP_SIZE), blocksForChannels(this._model.inputChannels));
|
|
1592
|
+
pass.end();
|
|
1593
|
+
}
|
|
1594
|
+
execution.plan.plannedNodes.forEach(({ node, outputShape }, index) => {
|
|
1595
|
+
const pass = encoder.beginComputePass(passDescriptor(`oidn/${execution.plan.nodes[index].id}`, index + 1));
|
|
1596
|
+
pass.setPipeline(execution.nodePipelines[index]);
|
|
1597
|
+
pass.setBindGroup(0, execution.nodeBindings[index]);
|
|
1598
|
+
if (execution.nodeKernels[index] === 'implicit-gemm') {
|
|
1599
|
+
pass.dispatchWorkgroups(Math.ceil((outputShape.width * outputShape.height) / TILED_CONV_M), Math.ceil(blocksForChannels(outputShape.channels) / TILED_CONV_N_BLOCKS), 1);
|
|
1600
|
+
}
|
|
1601
|
+
else {
|
|
1602
|
+
pass.dispatchWorkgroups(Math.ceil(outputShape.width / WORKGROUP_SIZE), Math.ceil(outputShape.height / WORKGROUP_SIZE), blocksForChannels(outputShape.channels));
|
|
1603
|
+
}
|
|
1604
|
+
pass.end();
|
|
1605
|
+
});
|
|
1606
|
+
if (querySet) {
|
|
1607
|
+
encoder.resolveQuerySet(querySet, 0, queryCount, queryResolveBuffer, 0);
|
|
1608
|
+
encoder.copyBufferToBuffer(queryResolveBuffer, 0, queryReadbackBuffer, 0, queryBufferSize);
|
|
1609
|
+
}
|
|
1610
|
+
this._device.queue.submit([encoder.finish()]);
|
|
1611
|
+
}
|
|
1612
|
+
catch (error) {
|
|
1613
|
+
this._releaseBuffer(queryReadbackBuffer);
|
|
1614
|
+
this._releaseBuffer(queryResolveBuffer);
|
|
1615
|
+
this._releaseQuerySet(querySet);
|
|
1616
|
+
throw error;
|
|
1617
|
+
}
|
|
1618
|
+
if (querySet) {
|
|
1619
|
+
this._profileOperations++;
|
|
1620
|
+
this._lastExecutionProfile = (async () => {
|
|
1621
|
+
try {
|
|
1622
|
+
await queryReadbackBuffer.mapAsync(GPUMapMode.READ);
|
|
1623
|
+
const timestamps = new BigUint64Array(queryReadbackBuffer.getMappedRange());
|
|
1624
|
+
const layers = profileLabels.map((id, index) => ({
|
|
1625
|
+
id,
|
|
1626
|
+
durationMs: Number(timestamps[index * 2 + 1] - timestamps[index * 2]) /
|
|
1627
|
+
1_000_000
|
|
1628
|
+
}));
|
|
1629
|
+
return {
|
|
1630
|
+
totalMs: layers.reduce((sum, layer) => sum + layer.durationMs, 0),
|
|
1631
|
+
layers
|
|
1632
|
+
};
|
|
1633
|
+
}
|
|
1634
|
+
finally {
|
|
1635
|
+
if (queryReadbackBuffer.mapState === 'mapped') {
|
|
1636
|
+
queryReadbackBuffer.unmap();
|
|
1637
|
+
}
|
|
1638
|
+
this._releaseQuerySet(querySet);
|
|
1639
|
+
this._releaseBuffer(queryResolveBuffer);
|
|
1640
|
+
this._releaseBuffer(queryReadbackBuffer);
|
|
1641
|
+
this._profileOperations--;
|
|
1642
|
+
}
|
|
1643
|
+
})();
|
|
1644
|
+
}
|
|
1645
|
+
return execution.valueBuffers.get(execution.plan.spec.output);
|
|
1646
|
+
}
|
|
1647
|
+
/** Compatibility path for ImageData/HDR arrays without TensorFlow.js. */
|
|
1648
|
+
async executeCPU(interleavedInput, width, height) {
|
|
1649
|
+
const expectedLength = width * height * this._model.inputChannels;
|
|
1650
|
+
if (interleavedInput.length !== expectedLength) {
|
|
1651
|
+
throw new Error(`Native OIDN CPU input has ${interleavedInput.length} values, expected ${expectedLength}`);
|
|
1652
|
+
}
|
|
1653
|
+
const execution = this._execution(width, height);
|
|
1654
|
+
const sourceCount = this._model.inputChannels / 3;
|
|
1655
|
+
const pixelCount = width * height;
|
|
1656
|
+
if (!execution.cpuInputBuffers) {
|
|
1657
|
+
execution.cpuInputBuffers = Array.from({ length: sourceCount }, (_, index) => {
|
|
1658
|
+
const buffer = this._device.createBuffer({
|
|
1659
|
+
label: `oidn/cpu-input/${width}x${height}/${index}`,
|
|
1660
|
+
size: pixelCount * 16,
|
|
1661
|
+
usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST
|
|
1662
|
+
});
|
|
1663
|
+
this._resources.track('gpu-buffer', buffer);
|
|
1664
|
+
execution.ownedBuffers.push(buffer);
|
|
1665
|
+
return buffer;
|
|
1666
|
+
});
|
|
1667
|
+
execution.cpuReadbackBuffer = this._device.createBuffer({
|
|
1668
|
+
label: `oidn/cpu-readback/${width}x${height}`,
|
|
1669
|
+
size: pixelCount * 16,
|
|
1670
|
+
usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.MAP_READ
|
|
1671
|
+
});
|
|
1672
|
+
this._resources.track('gpu-buffer', execution.cpuReadbackBuffer);
|
|
1673
|
+
execution.ownedBuffers.push(execution.cpuReadbackBuffer);
|
|
1674
|
+
}
|
|
1675
|
+
for (let source = 0; source < sourceCount; source++) {
|
|
1676
|
+
const upload = new Float32Array(pixelCount * 4);
|
|
1677
|
+
for (let pixel = 0; pixel < pixelCount; pixel++) {
|
|
1678
|
+
const inputOffset = pixel * this._model.inputChannels + source * 3;
|
|
1679
|
+
const outputOffset = pixel * 4;
|
|
1680
|
+
upload[outputOffset] = interleavedInput[inputOffset];
|
|
1681
|
+
upload[outputOffset + 1] = interleavedInput[inputOffset + 1];
|
|
1682
|
+
upload[outputOffset + 2] = interleavedInput[inputOffset + 2];
|
|
1683
|
+
}
|
|
1684
|
+
this._device.queue.writeBuffer(execution.cpuInputBuffers[source], 0, upload);
|
|
1685
|
+
}
|
|
1686
|
+
const output = this.execute(execution.cpuInputBuffers, width, height);
|
|
1687
|
+
const encoder = this._device.createCommandEncoder({
|
|
1688
|
+
label: `oidn/cpu-readback/${width}x${height}`
|
|
1689
|
+
});
|
|
1690
|
+
encoder.copyBufferToBuffer(output, 0, execution.cpuReadbackBuffer, 0, pixelCount * 16);
|
|
1691
|
+
this._device.queue.submit([encoder.finish()]);
|
|
1692
|
+
await execution.cpuReadbackBuffer.mapAsync(GPUMapMode.READ);
|
|
1693
|
+
const rgba = new Float32Array(execution.cpuReadbackBuffer.getMappedRange());
|
|
1694
|
+
const rgb = new Float32Array(pixelCount * 3);
|
|
1695
|
+
for (let pixel = 0; pixel < pixelCount; pixel++) {
|
|
1696
|
+
rgb[pixel * 3] = rgba[pixel * 4];
|
|
1697
|
+
rgb[pixel * 3 + 1] = rgba[pixel * 4 + 1];
|
|
1698
|
+
rgb[pixel * 3 + 2] = rgba[pixel * 4 + 2];
|
|
1699
|
+
}
|
|
1700
|
+
execution.cpuReadbackBuffer.unmap();
|
|
1701
|
+
return rgb;
|
|
1702
|
+
}
|
|
1703
|
+
_releaseBuffer(buffer) {
|
|
1704
|
+
this._resources.release('gpu-buffer', buffer, () => buffer.destroy());
|
|
1705
|
+
}
|
|
1706
|
+
_releaseQuerySet(querySet) {
|
|
1707
|
+
this._resources.release('gpu-query-set', querySet, () => querySet.destroy());
|
|
1708
|
+
}
|
|
1709
|
+
_destroyExecution(execution) {
|
|
1710
|
+
execution.slots.forEach((slot) => this._releaseBuffer(slot.buffer));
|
|
1711
|
+
execution.ownedBuffers.forEach((buffer) => this._releaseBuffer(buffer));
|
|
1712
|
+
}
|
|
1713
|
+
getResourceInfo() {
|
|
1714
|
+
return this._resources.snapshot(this._retiredExecutions.size + this._profileOperations);
|
|
1715
|
+
}
|
|
1716
|
+
dispose() {
|
|
1717
|
+
if (this._disposed)
|
|
1718
|
+
return;
|
|
1719
|
+
this._disposed = true;
|
|
1720
|
+
for (const packed of this._packedConvs.values()) {
|
|
1721
|
+
this._releaseBuffer(packed.weights);
|
|
1722
|
+
this._releaseBuffer(packed.bias);
|
|
1723
|
+
}
|
|
1724
|
+
this._packedConvs.clear();
|
|
1725
|
+
for (const execution of this._executionCache.values()) {
|
|
1726
|
+
this._destroyExecution(execution);
|
|
1727
|
+
}
|
|
1728
|
+
this._executionCache.clear();
|
|
1729
|
+
for (const execution of this._retiredExecutions) {
|
|
1730
|
+
this._destroyExecution(execution);
|
|
1731
|
+
}
|
|
1732
|
+
this._retiredExecutions.clear();
|
|
1733
|
+
}
|
|
1734
|
+
}
|
|
1735
|
+
//# sourceMappingURL=nativeUNet.js.map
|