oidn-web 0.1.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/LICENSE +21 -0
- package/README.md +2 -0
- package/dist/oidn.mjs +22909 -0
- package/dist/oidn.umd.js +5905 -0
- package/lib/UNet.d.ts +54 -0
- package/lib/UNet.js +467 -0
- package/lib/UNet.js.map +1 -0
- package/lib/WGPUComputePass.d.ts +53 -0
- package/lib/WGPUComputePass.js +220 -0
- package/lib/WGPUComputePass.js.map +1 -0
- package/lib/WGPUFullQuadPass.d.ts +51 -0
- package/lib/WGPUFullQuadPass.js +261 -0
- package/lib/WGPUFullQuadPass.js.map +1 -0
- package/lib/backend.d.ts +5 -0
- package/lib/backend.js +41 -0
- package/lib/backend.js.map +1 -0
- package/lib/hdr.d.ts +31 -0
- package/lib/hdr.js +340 -0
- package/lib/hdr.js.map +1 -0
- package/lib/helper.d.ts +1 -0
- package/lib/helper.js +27 -0
- package/lib/helper.js.map +1 -0
- package/lib/kernels.d.ts +1 -0
- package/lib/kernels.js +26 -0
- package/lib/kernels.js.map +1 -0
- package/lib/main.d.ts +20 -0
- package/lib/main.js +20 -0
- package/lib/main.js.map +1 -0
- package/lib/process.d.ts +39 -0
- package/lib/process.js +309 -0
- package/lib/process.js.map +1 -0
- package/lib/tza.d.ts +15 -0
- package/lib/tza.js +114 -0
- package/lib/tza.js.map +1 -0
- package/package.json +26 -0
- package/src/UNet.ts +708 -0
- package/src/WGPUComputePass.ts +318 -0
- package/src/WGPUFullQuadPass.ts +348 -0
- package/src/backend.ts +53 -0
- package/src/hdr.ts +398 -0
- package/src/helper.ts +35 -0
- package/src/kernels.ts +31 -0
- package/src/main.ts +42 -0
- package/src/process.ts +362 -0
- package/src/tza.ts +136 -0
- package/weights/.gitattributes +1 -0
- package/weights/LICENSE.txt +202 -0
- package/weights/README.md +7 -0
- package/weights/rt_hdr.tza +0 -0
- package/weights/rt_hdr_alb_nrm.tza +0 -0
- package/weights/rt_ldr.tza +0 -0
- package/weights/rt_ldr_alb_nrm.tza +0 -0
|
@@ -0,0 +1,318 @@
|
|
|
1
|
+
function isStorageParamsEqual(
|
|
2
|
+
params: WGPUComputePassOutput,
|
|
3
|
+
other: WGPUComputePassOutput
|
|
4
|
+
) {
|
|
5
|
+
return params.channels === other.channels;
|
|
6
|
+
}
|
|
7
|
+
|
|
8
|
+
export const WORKGROUP_SIZE = 8;
|
|
9
|
+
|
|
10
|
+
export interface Uniform {
|
|
11
|
+
label: string;
|
|
12
|
+
type: string;
|
|
13
|
+
data: Float32Array | Int32Array | Uint32Array;
|
|
14
|
+
}
|
|
15
|
+
export interface WGPUComputePassInput {
|
|
16
|
+
buffer: GPUBuffer;
|
|
17
|
+
channels: number;
|
|
18
|
+
}
|
|
19
|
+
|
|
20
|
+
export interface WGPUComputePassOutput {
|
|
21
|
+
channels: number;
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
export class WGPUComputePass<I extends string, O extends string> {
|
|
25
|
+
private _label;
|
|
26
|
+
|
|
27
|
+
private _device;
|
|
28
|
+
private _outputBuffers: Record<
|
|
29
|
+
string,
|
|
30
|
+
{
|
|
31
|
+
buffer: GPUBuffer;
|
|
32
|
+
params: WGPUComputePassOutput;
|
|
33
|
+
}
|
|
34
|
+
> = {};
|
|
35
|
+
|
|
36
|
+
private _pipeline!: GPUComputePipeline;
|
|
37
|
+
private _bindGroups: GPUBindGroup[] = [];
|
|
38
|
+
private _needsUpdatePipeline = true;
|
|
39
|
+
|
|
40
|
+
private _inputs: string[] = [];
|
|
41
|
+
private _outputs: string[] = [];
|
|
42
|
+
private _uniforms: Uniform[] = [];
|
|
43
|
+
private _uniformBuffers: Record<string, GPUBuffer> = {};
|
|
44
|
+
|
|
45
|
+
private _width = 10;
|
|
46
|
+
private _height = 10;
|
|
47
|
+
|
|
48
|
+
private _execWidth?: number;
|
|
49
|
+
private _execHeight?: number;
|
|
50
|
+
|
|
51
|
+
private _csCode = '';
|
|
52
|
+
private _csMain;
|
|
53
|
+
private _csDefine;
|
|
54
|
+
|
|
55
|
+
constructor(
|
|
56
|
+
label: string,
|
|
57
|
+
device: GPUDevice,
|
|
58
|
+
opts: {
|
|
59
|
+
inputs: I[];
|
|
60
|
+
outputs: O[];
|
|
61
|
+
csMain: string;
|
|
62
|
+
csDefine?: string;
|
|
63
|
+
uniforms: Uniform[];
|
|
64
|
+
}
|
|
65
|
+
) {
|
|
66
|
+
this._label = label;
|
|
67
|
+
this._device = device;
|
|
68
|
+
this._csMain = opts.csMain;
|
|
69
|
+
this._csDefine = opts.csDefine;
|
|
70
|
+
this._inputs = opts.inputs;
|
|
71
|
+
this._outputs = opts.outputs;
|
|
72
|
+
this._uniforms = opts.uniforms;
|
|
73
|
+
|
|
74
|
+
opts.uniforms.forEach((uniform) => {
|
|
75
|
+
this._uniformBuffers[uniform.label] = device.createBuffer({
|
|
76
|
+
label: this._label,
|
|
77
|
+
size: uniform.data.byteLength,
|
|
78
|
+
usage: GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST
|
|
79
|
+
});
|
|
80
|
+
this._device.queue.writeBuffer(
|
|
81
|
+
this._uniformBuffers[uniform.label],
|
|
82
|
+
0,
|
|
83
|
+
uniform.data
|
|
84
|
+
);
|
|
85
|
+
});
|
|
86
|
+
}
|
|
87
|
+
|
|
88
|
+
setSize(width: number, height: number) {
|
|
89
|
+
width = Math.ceil(width);
|
|
90
|
+
height = Math.ceil(height);
|
|
91
|
+
const sizeChanged = width !== this._width || height !== this._height;
|
|
92
|
+
this._width = width;
|
|
93
|
+
this._height = height;
|
|
94
|
+
if (sizeChanged) {
|
|
95
|
+
this._resizeOutputBuffers();
|
|
96
|
+
this._needsUpdatePipeline = true;
|
|
97
|
+
}
|
|
98
|
+
}
|
|
99
|
+
|
|
100
|
+
setExecuteSize(width: number, height: number) {
|
|
101
|
+
width = Math.ceil(width);
|
|
102
|
+
height = Math.ceil(height);
|
|
103
|
+
this._execWidth = width;
|
|
104
|
+
this._execHeight = height;
|
|
105
|
+
}
|
|
106
|
+
|
|
107
|
+
setOutputParams(outputParams: Record<O, WGPUComputePassOutput>) {
|
|
108
|
+
this._updateOutputBuffers(outputParams);
|
|
109
|
+
this._needsUpdatePipeline = true;
|
|
110
|
+
}
|
|
111
|
+
|
|
112
|
+
setUniform(label: string, data: Float32Array | Int32Array | Uint32Array) {
|
|
113
|
+
const buffer = this._uniformBuffers[label];
|
|
114
|
+
this._device.queue.writeBuffer(buffer, 0, data);
|
|
115
|
+
}
|
|
116
|
+
|
|
117
|
+
getOutputBuffer(name: O) {
|
|
118
|
+
return this._outputBuffers[name].buffer;
|
|
119
|
+
}
|
|
120
|
+
|
|
121
|
+
dispose() {
|
|
122
|
+
Object.keys(this._uniformBuffers).forEach((key) => {
|
|
123
|
+
(this._uniformBuffers as any)[key].destroy();
|
|
124
|
+
});
|
|
125
|
+
Object.keys(this._outputBuffers).forEach((key) => {
|
|
126
|
+
(this._outputBuffers as any)[key].texture.destroy();
|
|
127
|
+
});
|
|
128
|
+
}
|
|
129
|
+
|
|
130
|
+
private _createBuffer(params: WGPUComputePassOutput) {
|
|
131
|
+
// const byteLength = this._width * this._height * params.channels * 4;
|
|
132
|
+
// Buffer data needs to be aligned with 8, 16
|
|
133
|
+
const byteLength = this._width * this._height * 4 * 4;
|
|
134
|
+
return this._device.createBuffer({
|
|
135
|
+
label: this._label,
|
|
136
|
+
// webgpu needs buffer at least 80 bytes.
|
|
137
|
+
size: Math.max(byteLength, 80),
|
|
138
|
+
usage:
|
|
139
|
+
GPUBufferUsage.STORAGE |
|
|
140
|
+
GPUBufferUsage.COPY_DST |
|
|
141
|
+
GPUBufferUsage.COPY_SRC
|
|
142
|
+
});
|
|
143
|
+
}
|
|
144
|
+
|
|
145
|
+
private _resizeOutputBuffers() {
|
|
146
|
+
const outputBuffers = this._outputBuffers;
|
|
147
|
+
for (const key in outputBuffers) {
|
|
148
|
+
const { buffer, params } = outputBuffers[key];
|
|
149
|
+
buffer.destroy();
|
|
150
|
+
outputBuffers[key].buffer = this._createBuffer(params);
|
|
151
|
+
}
|
|
152
|
+
}
|
|
153
|
+
|
|
154
|
+
private _updateOutputBuffers(
|
|
155
|
+
outputParams: Record<string, WGPUComputePassOutput>
|
|
156
|
+
) {
|
|
157
|
+
const outputBuffers = this._outputBuffers;
|
|
158
|
+
for (const key in outputParams) {
|
|
159
|
+
const params = outputParams[key];
|
|
160
|
+
if (
|
|
161
|
+
!isStorageParamsEqual(
|
|
162
|
+
params,
|
|
163
|
+
outputBuffers[key]?.params || ({} as WGPUComputePassOutput)
|
|
164
|
+
)
|
|
165
|
+
) {
|
|
166
|
+
outputBuffers[key]?.buffer.destroy();
|
|
167
|
+
const buffer = this._createBuffer(params);
|
|
168
|
+
outputBuffers[key] = {
|
|
169
|
+
buffer,
|
|
170
|
+
params
|
|
171
|
+
};
|
|
172
|
+
}
|
|
173
|
+
}
|
|
174
|
+
}
|
|
175
|
+
|
|
176
|
+
private _updatePipeline(inputParams: Record<string, WGPUComputePassInput>) {
|
|
177
|
+
if (!this._needsUpdatePipeline) {
|
|
178
|
+
return;
|
|
179
|
+
}
|
|
180
|
+
this._needsUpdatePipeline = false;
|
|
181
|
+
const device = this._device;
|
|
182
|
+
const csCode = this._getFullCs(inputParams);
|
|
183
|
+
if (csCode === this._csCode) {
|
|
184
|
+
return;
|
|
185
|
+
}
|
|
186
|
+
this._csCode = csCode;
|
|
187
|
+
this._pipeline = device.createComputePipeline({
|
|
188
|
+
label: this._label,
|
|
189
|
+
layout: 'auto',
|
|
190
|
+
compute: {
|
|
191
|
+
module: device.createShaderModule({
|
|
192
|
+
label: this._label,
|
|
193
|
+
code: csCode
|
|
194
|
+
}),
|
|
195
|
+
entryPoint: 'main'
|
|
196
|
+
}
|
|
197
|
+
});
|
|
198
|
+
this._updateBindGroups();
|
|
199
|
+
}
|
|
200
|
+
|
|
201
|
+
private _getFullCs(inputParams: Record<string, WGPUComputePassInput>) {
|
|
202
|
+
const inputs = this._inputs;
|
|
203
|
+
const hasInputs = inputs.length > 0;
|
|
204
|
+
const cs = `
|
|
205
|
+
${inputs
|
|
206
|
+
.sort()
|
|
207
|
+
.map(
|
|
208
|
+
(bufferName, idx) =>
|
|
209
|
+
// TODO more channels option.
|
|
210
|
+
`@group(0) @binding(${idx}) var<storage, read> in_${bufferName}: array<vec${inputParams[bufferName].channels}f>;`
|
|
211
|
+
)
|
|
212
|
+
.join('\n')}
|
|
213
|
+
${this._uniforms
|
|
214
|
+
.map(
|
|
215
|
+
(uniform, idx) =>
|
|
216
|
+
`@group(${hasInputs ? 1 : 0}) @binding(${idx}) var<uniform> ${
|
|
217
|
+
uniform.label
|
|
218
|
+
}: ${uniform.type};`
|
|
219
|
+
)
|
|
220
|
+
.join('\n')}
|
|
221
|
+
|
|
222
|
+
${this._outputs
|
|
223
|
+
.map(
|
|
224
|
+
(name, idx) =>
|
|
225
|
+
`@group(${
|
|
226
|
+
hasInputs ? 2 : 1
|
|
227
|
+
}) @binding(${idx}) var<storage, read_write> out_${name}: array<vec${
|
|
228
|
+
this._outputBuffers[name].params.channels
|
|
229
|
+
}f>;`
|
|
230
|
+
)
|
|
231
|
+
.join('\n')}
|
|
232
|
+
${this._csDefine ?? ''}
|
|
233
|
+
@compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
|
|
234
|
+
fn main(@builtin(global_invocation_id) globalId: vec3u) {
|
|
235
|
+
${this._csMain}
|
|
236
|
+
}
|
|
237
|
+
`;
|
|
238
|
+
|
|
239
|
+
return cs;
|
|
240
|
+
}
|
|
241
|
+
|
|
242
|
+
private _updateBindGroups() {
|
|
243
|
+
const bindGroups: GPUBindGroup[] = [];
|
|
244
|
+
const device = this._device;
|
|
245
|
+
|
|
246
|
+
//TODO
|
|
247
|
+
const uniformBindGroupIndex = this._inputs.length > 0 ? 1 : 0;
|
|
248
|
+
if (this._uniforms.length > 0) {
|
|
249
|
+
bindGroups[uniformBindGroupIndex] = device.createBindGroup({
|
|
250
|
+
label: this._label,
|
|
251
|
+
layout: this._pipeline.getBindGroupLayout(uniformBindGroupIndex),
|
|
252
|
+
entries: this._uniforms.map(
|
|
253
|
+
(uniform, idx) =>
|
|
254
|
+
({
|
|
255
|
+
binding: idx,
|
|
256
|
+
resource: {
|
|
257
|
+
buffer: this._uniformBuffers[uniform.label]
|
|
258
|
+
}
|
|
259
|
+
} as GPUBindGroupEntry)
|
|
260
|
+
)
|
|
261
|
+
});
|
|
262
|
+
}
|
|
263
|
+
|
|
264
|
+
this._bindGroups = bindGroups;
|
|
265
|
+
}
|
|
266
|
+
|
|
267
|
+
createPass(
|
|
268
|
+
commandEncoder: GPUCommandEncoder,
|
|
269
|
+
inputBuffers: Record<I, WGPUComputePassInput>
|
|
270
|
+
) {
|
|
271
|
+
this._updatePipeline(inputBuffers);
|
|
272
|
+
|
|
273
|
+
const hasInputs = this._inputs.length > 0;
|
|
274
|
+
// TODO createBindGroup every time?
|
|
275
|
+
if (hasInputs) {
|
|
276
|
+
this._bindGroups[0] = this._device.createBindGroup({
|
|
277
|
+
label: this._label,
|
|
278
|
+
layout: this._pipeline.getBindGroupLayout(0),
|
|
279
|
+
entries: this._inputs.map((bufferName, idx) => ({
|
|
280
|
+
binding: idx,
|
|
281
|
+
// TODO
|
|
282
|
+
resource: {
|
|
283
|
+
buffer: inputBuffers[bufferName as I].buffer
|
|
284
|
+
}
|
|
285
|
+
}))
|
|
286
|
+
});
|
|
287
|
+
}
|
|
288
|
+
|
|
289
|
+
// Outputs
|
|
290
|
+
this._bindGroups[hasInputs ? 2 : 1] = this._device.createBindGroup({
|
|
291
|
+
label: this._label,
|
|
292
|
+
layout: this._pipeline.getBindGroupLayout(hasInputs ? 2 : 1),
|
|
293
|
+
entries: this._outputs.map((bufferName, idx) => ({
|
|
294
|
+
binding: idx,
|
|
295
|
+
resource: {
|
|
296
|
+
buffer: this._outputBuffers[bufferName].buffer
|
|
297
|
+
}
|
|
298
|
+
}))
|
|
299
|
+
});
|
|
300
|
+
|
|
301
|
+
// Begin the render pass
|
|
302
|
+
const computePass = commandEncoder.beginComputePass();
|
|
303
|
+
|
|
304
|
+
// Draw a full quad
|
|
305
|
+
computePass.setPipeline(this._pipeline);
|
|
306
|
+
// Bind groups
|
|
307
|
+
this._bindGroups.forEach((bindGroup, idx) => {
|
|
308
|
+
computePass.setBindGroup(idx, bindGroup);
|
|
309
|
+
});
|
|
310
|
+
computePass.dispatchWorkgroups(
|
|
311
|
+
Math.ceil((this._execWidth ?? this._width) / WORKGROUP_SIZE),
|
|
312
|
+
Math.ceil((this._execHeight ?? this._height) / WORKGROUP_SIZE),
|
|
313
|
+
1
|
|
314
|
+
);
|
|
315
|
+
// End the render pass
|
|
316
|
+
computePass.end();
|
|
317
|
+
}
|
|
318
|
+
}
|
|
@@ -0,0 +1,348 @@
|
|
|
1
|
+
export const fullScreenQuadVertexShaderWGSL = /*wgsl */ `
|
|
2
|
+
@vertex
|
|
3
|
+
fn main(
|
|
4
|
+
@builtin(vertex_index) vertexIndex: u32
|
|
5
|
+
) -> @builtin(position) vec4f {
|
|
6
|
+
const pos = array(
|
|
7
|
+
vec2(-1.0, -1.0), vec2(1.0, -1.0), vec2(-1.0, 1.0),
|
|
8
|
+
vec2(-1.0, 1.0), vec2(1.0, -1.0), vec2(1.0, 1.0),
|
|
9
|
+
);
|
|
10
|
+
|
|
11
|
+
return vec4<f32>(pos[vertexIndex], 0.0, 1.0);
|
|
12
|
+
}
|
|
13
|
+
`;
|
|
14
|
+
|
|
15
|
+
function isTextureParamsEqual(
|
|
16
|
+
params: WGPUFullQuadPassOutput,
|
|
17
|
+
other: WGPUFullQuadPassOutput
|
|
18
|
+
) {
|
|
19
|
+
return params.format === other.format;
|
|
20
|
+
}
|
|
21
|
+
|
|
22
|
+
export interface Uniform {
|
|
23
|
+
label: string;
|
|
24
|
+
type: string;
|
|
25
|
+
data: Float32Array | Int32Array | Uint32Array;
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
export interface WGPUFullQuadPassOutput {
|
|
29
|
+
format: GPUTextureFormat;
|
|
30
|
+
}
|
|
31
|
+
export class WGPUFullQuadPass<I extends string, O extends string> {
|
|
32
|
+
private _label;
|
|
33
|
+
|
|
34
|
+
private _device;
|
|
35
|
+
private _outputTextures: Record<
|
|
36
|
+
string,
|
|
37
|
+
{
|
|
38
|
+
texture: GPUTexture;
|
|
39
|
+
params: WGPUFullQuadPassOutput;
|
|
40
|
+
}
|
|
41
|
+
> = {};
|
|
42
|
+
|
|
43
|
+
private _pipeline!: GPURenderPipeline;
|
|
44
|
+
private _bindGroups: GPUBindGroup[] = [];
|
|
45
|
+
private _needsUpdatePipeline = true;
|
|
46
|
+
/**
|
|
47
|
+
* When render to canvas
|
|
48
|
+
*/
|
|
49
|
+
private _renderToScreen?: {
|
|
50
|
+
screenTexture: GPUTexture;
|
|
51
|
+
presentationFormat: GPUTextureFormat;
|
|
52
|
+
};
|
|
53
|
+
private _inputs: string[] = [];
|
|
54
|
+
private _outputs: string[] = [];
|
|
55
|
+
private _uniforms: Uniform[] = [];
|
|
56
|
+
private _uniformBuffers: Record<string, GPUBuffer> = {};
|
|
57
|
+
|
|
58
|
+
private _width = 10;
|
|
59
|
+
private _height = 10;
|
|
60
|
+
|
|
61
|
+
private _fsCode = '';
|
|
62
|
+
private _fsMain;
|
|
63
|
+
private _fsDefine;
|
|
64
|
+
|
|
65
|
+
constructor(
|
|
66
|
+
label: string,
|
|
67
|
+
device: GPUDevice,
|
|
68
|
+
opts: {
|
|
69
|
+
inputs: I[];
|
|
70
|
+
outputs: O[];
|
|
71
|
+
fsMain: string;
|
|
72
|
+
fsDefine?: string;
|
|
73
|
+
uniforms: Uniform[];
|
|
74
|
+
}
|
|
75
|
+
) {
|
|
76
|
+
this._label = label;
|
|
77
|
+
this._device = device;
|
|
78
|
+
this._fsMain = opts.fsMain;
|
|
79
|
+
this._fsDefine = opts.fsDefine;
|
|
80
|
+
this._inputs = opts.inputs;
|
|
81
|
+
this._outputs = opts.outputs;
|
|
82
|
+
this._uniforms = opts.uniforms;
|
|
83
|
+
|
|
84
|
+
opts.uniforms.forEach((uniform) => {
|
|
85
|
+
this._uniformBuffers[uniform.label] = device.createBuffer({
|
|
86
|
+
label: this._label,
|
|
87
|
+
size: uniform.data.byteLength,
|
|
88
|
+
usage: GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST
|
|
89
|
+
});
|
|
90
|
+
this._device.queue.writeBuffer(
|
|
91
|
+
this._uniformBuffers[uniform.label],
|
|
92
|
+
0,
|
|
93
|
+
uniform.data
|
|
94
|
+
);
|
|
95
|
+
});
|
|
96
|
+
}
|
|
97
|
+
|
|
98
|
+
setSize(width: number, height: number) {
|
|
99
|
+
width = Math.ceil(width);
|
|
100
|
+
height = Math.ceil(height);
|
|
101
|
+
const sizeChanged = width !== this._width || height !== this._height;
|
|
102
|
+
this._width = width;
|
|
103
|
+
this._height = height;
|
|
104
|
+
if (sizeChanged) {
|
|
105
|
+
this._resizeOutputTextures();
|
|
106
|
+
this._needsUpdatePipeline = true;
|
|
107
|
+
}
|
|
108
|
+
}
|
|
109
|
+
|
|
110
|
+
setOutputParams(outputParams: Record<O, WGPUFullQuadPassOutput>) {
|
|
111
|
+
this._renderToScreen = undefined;
|
|
112
|
+
this._updateOutputTextures(outputParams);
|
|
113
|
+
this._needsUpdatePipeline = true;
|
|
114
|
+
}
|
|
115
|
+
|
|
116
|
+
setRenderToScreen(
|
|
117
|
+
screenTexture: GPUTexture,
|
|
118
|
+
presentationFormat: GPUTextureFormat
|
|
119
|
+
) {
|
|
120
|
+
this._renderToScreen = {
|
|
121
|
+
screenTexture,
|
|
122
|
+
presentationFormat
|
|
123
|
+
};
|
|
124
|
+
}
|
|
125
|
+
|
|
126
|
+
setUniform(label: string, data: Float32Array | Int32Array | Uint32Array) {
|
|
127
|
+
const buffer = this._uniformBuffers[label];
|
|
128
|
+
this._device.queue.writeBuffer(buffer, 0, data);
|
|
129
|
+
}
|
|
130
|
+
|
|
131
|
+
getOutputTexture(name: O) {
|
|
132
|
+
return this._outputTextures[name].texture;
|
|
133
|
+
}
|
|
134
|
+
|
|
135
|
+
dispose() {
|
|
136
|
+
Object.keys(this._uniformBuffers).forEach((key) => {
|
|
137
|
+
(this._uniformBuffers as any)[key].destroy();
|
|
138
|
+
});
|
|
139
|
+
Object.keys(this._outputTextures).forEach((key) => {
|
|
140
|
+
(this._outputTextures as any)[key].texture.destroy();
|
|
141
|
+
});
|
|
142
|
+
}
|
|
143
|
+
|
|
144
|
+
private _createTexture(params: WGPUFullQuadPassOutput) {
|
|
145
|
+
return this._device.createTexture({
|
|
146
|
+
label: this._label,
|
|
147
|
+
size: {
|
|
148
|
+
width: this._width,
|
|
149
|
+
height: this._height,
|
|
150
|
+
depthOrArrayLayers: 1
|
|
151
|
+
},
|
|
152
|
+
format: params.format,
|
|
153
|
+
usage: GPUTextureUsage.RENDER_ATTACHMENT | GPUTextureUsage.TEXTURE_BINDING
|
|
154
|
+
});
|
|
155
|
+
}
|
|
156
|
+
|
|
157
|
+
private _resizeOutputTextures() {
|
|
158
|
+
const outputTextures = this._outputTextures;
|
|
159
|
+
for (const key in outputTextures) {
|
|
160
|
+
const { texture, params } = outputTextures[key];
|
|
161
|
+
texture.destroy();
|
|
162
|
+
outputTextures[key].texture = this._createTexture(params);
|
|
163
|
+
}
|
|
164
|
+
}
|
|
165
|
+
|
|
166
|
+
private _updateOutputTextures(
|
|
167
|
+
outputParams: Record<string, WGPUFullQuadPassOutput>
|
|
168
|
+
) {
|
|
169
|
+
const outputTextures = this._outputTextures;
|
|
170
|
+
for (const key in outputParams) {
|
|
171
|
+
const params = outputParams[key];
|
|
172
|
+
if (
|
|
173
|
+
!isTextureParamsEqual(
|
|
174
|
+
params,
|
|
175
|
+
outputTextures[key]?.params || ({} as WGPUFullQuadPassOutput)
|
|
176
|
+
)
|
|
177
|
+
) {
|
|
178
|
+
outputTextures[key]?.texture.destroy();
|
|
179
|
+
const texture = this._createTexture(params);
|
|
180
|
+
outputTextures[key] = {
|
|
181
|
+
texture,
|
|
182
|
+
params
|
|
183
|
+
};
|
|
184
|
+
}
|
|
185
|
+
}
|
|
186
|
+
}
|
|
187
|
+
|
|
188
|
+
private _updatePipeline() {
|
|
189
|
+
if (!this._needsUpdatePipeline) {
|
|
190
|
+
return;
|
|
191
|
+
}
|
|
192
|
+
this._needsUpdatePipeline = false;
|
|
193
|
+
const device = this._device;
|
|
194
|
+
const fsCode = this._getFullFs();
|
|
195
|
+
if (fsCode === this._fsCode) {
|
|
196
|
+
return;
|
|
197
|
+
}
|
|
198
|
+
const { screenTexture, presentationFormat } = this._renderToScreen || {};
|
|
199
|
+
this._fsCode = fsCode;
|
|
200
|
+
this._pipeline = device.createRenderPipeline({
|
|
201
|
+
label: this._label,
|
|
202
|
+
layout: 'auto',
|
|
203
|
+
vertex: {
|
|
204
|
+
module: device.createShaderModule({
|
|
205
|
+
label: this._label,
|
|
206
|
+
code: fullScreenQuadVertexShaderWGSL
|
|
207
|
+
}),
|
|
208
|
+
entryPoint: 'main'
|
|
209
|
+
},
|
|
210
|
+
fragment: {
|
|
211
|
+
module: device.createShaderModule({
|
|
212
|
+
label: this._label,
|
|
213
|
+
code: fsCode
|
|
214
|
+
}),
|
|
215
|
+
entryPoint: 'main',
|
|
216
|
+
targets: screenTexture
|
|
217
|
+
? [
|
|
218
|
+
{
|
|
219
|
+
format: presentationFormat!
|
|
220
|
+
}
|
|
221
|
+
]
|
|
222
|
+
: this._outputs.map((key) => ({
|
|
223
|
+
format: this._outputTextures[key].params.format
|
|
224
|
+
}))
|
|
225
|
+
},
|
|
226
|
+
primitive: {
|
|
227
|
+
topology: 'triangle-list'
|
|
228
|
+
}
|
|
229
|
+
});
|
|
230
|
+
this._updateBindGroups();
|
|
231
|
+
}
|
|
232
|
+
|
|
233
|
+
private _getFullFs() {
|
|
234
|
+
const inputs = this._inputs;
|
|
235
|
+
const hasInputs = inputs.length > 0;
|
|
236
|
+
const fs = `
|
|
237
|
+
${inputs
|
|
238
|
+
.sort()
|
|
239
|
+
.map(
|
|
240
|
+
(textureName, idx) =>
|
|
241
|
+
`@group(0) @binding(${idx}) var ${textureName}: texture_2d<f32>;`
|
|
242
|
+
)
|
|
243
|
+
.join('\n')}
|
|
244
|
+
${this._uniforms
|
|
245
|
+
.map(
|
|
246
|
+
(uniform, idx) =>
|
|
247
|
+
`@group(${hasInputs ? 1 : 0}) @binding(${idx}) var<uniform> ${
|
|
248
|
+
uniform.label
|
|
249
|
+
}: ${uniform.type};`
|
|
250
|
+
)
|
|
251
|
+
.join('\n')}
|
|
252
|
+
|
|
253
|
+
struct FSOutput {
|
|
254
|
+
${this._outputs
|
|
255
|
+
.map((name, idx) => `@location(${idx}) ${name}: vec4f,`)
|
|
256
|
+
.join('\n')}
|
|
257
|
+
}
|
|
258
|
+
${this._fsDefine ?? ''}
|
|
259
|
+
@fragment
|
|
260
|
+
fn main(
|
|
261
|
+
@builtin(position) coord: vec4f
|
|
262
|
+
) -> FSOutput {
|
|
263
|
+
var uv = vec2i(floor(coord.xy));
|
|
264
|
+
var output: FSOutput;
|
|
265
|
+
${this._fsMain}
|
|
266
|
+
return output;
|
|
267
|
+
}
|
|
268
|
+
`;
|
|
269
|
+
|
|
270
|
+
return fs;
|
|
271
|
+
}
|
|
272
|
+
|
|
273
|
+
private _updateBindGroups() {
|
|
274
|
+
const bindGroups: GPUBindGroup[] = [];
|
|
275
|
+
const device = this._device;
|
|
276
|
+
|
|
277
|
+
//TODO
|
|
278
|
+
const uniformBindGroupIndex = this._inputs.length > 0 ? 1 : 0;
|
|
279
|
+
if (this._uniforms.length > 0) {
|
|
280
|
+
bindGroups[uniformBindGroupIndex] = device.createBindGroup({
|
|
281
|
+
label: this._label,
|
|
282
|
+
layout: this._pipeline.getBindGroupLayout(uniformBindGroupIndex),
|
|
283
|
+
entries: this._uniforms.map(
|
|
284
|
+
(uniform, idx) =>
|
|
285
|
+
({
|
|
286
|
+
binding: idx,
|
|
287
|
+
resource: {
|
|
288
|
+
buffer: this._uniformBuffers[uniform.label]
|
|
289
|
+
}
|
|
290
|
+
} satisfies GPUBindGroupEntry)
|
|
291
|
+
)
|
|
292
|
+
});
|
|
293
|
+
}
|
|
294
|
+
|
|
295
|
+
this._bindGroups = bindGroups;
|
|
296
|
+
}
|
|
297
|
+
|
|
298
|
+
createPass(
|
|
299
|
+
commandEncoder: GPUCommandEncoder,
|
|
300
|
+
inputTextures: Record<I, GPUTexture>
|
|
301
|
+
) {
|
|
302
|
+
this._updatePipeline();
|
|
303
|
+
|
|
304
|
+
// TODO createBindGrou every time?
|
|
305
|
+
if (this._inputs.length > 0) {
|
|
306
|
+
this._bindGroups[0] = this._device.createBindGroup({
|
|
307
|
+
label: this._label,
|
|
308
|
+
layout: this._pipeline.getBindGroupLayout(0),
|
|
309
|
+
entries: this._inputs.map((textureName, idx) => ({
|
|
310
|
+
binding: idx,
|
|
311
|
+
// TODO
|
|
312
|
+
resource: inputTextures[textureName as I].createView()
|
|
313
|
+
}))
|
|
314
|
+
});
|
|
315
|
+
}
|
|
316
|
+
// Begin the render pass
|
|
317
|
+
const renderPass = commandEncoder.beginRenderPass({
|
|
318
|
+
colorAttachments: this._renderToScreen
|
|
319
|
+
? [
|
|
320
|
+
{
|
|
321
|
+
view: this._renderToScreen.screenTexture.createView(),
|
|
322
|
+
clearValue: { r: 0, g: 0, b: 0, a: 0 },
|
|
323
|
+
storeOp: 'store' as GPUStoreOp,
|
|
324
|
+
loadOp: 'clear' as GPULoadOp
|
|
325
|
+
}
|
|
326
|
+
]
|
|
327
|
+
: this._outputs.map(
|
|
328
|
+
(textureName) =>
|
|
329
|
+
({
|
|
330
|
+
view: this._outputTextures[textureName].texture.createView(),
|
|
331
|
+
clearValue: { r: 0, g: 0, b: 0, a: 0 },
|
|
332
|
+
loadOp: 'clear' as GPULoadOp,
|
|
333
|
+
storeOp: 'store' as GPUStoreOp
|
|
334
|
+
} satisfies GPURenderPassColorAttachment)
|
|
335
|
+
)
|
|
336
|
+
});
|
|
337
|
+
|
|
338
|
+
// Draw a full quad
|
|
339
|
+
renderPass.setPipeline(this._pipeline);
|
|
340
|
+
// Bind groups
|
|
341
|
+
this._bindGroups.forEach((bindGroup, idx) => {
|
|
342
|
+
renderPass.setBindGroup(idx, bindGroup);
|
|
343
|
+
});
|
|
344
|
+
renderPass.draw(6, 1, 0, 0);
|
|
345
|
+
// End the render pass
|
|
346
|
+
renderPass.end();
|
|
347
|
+
}
|
|
348
|
+
}
|
package/src/backend.ts
ADDED
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
import { WebGPUBackend } from '@tensorflow/tfjs-backend-webgpu/dist/base';
|
|
2
|
+
import { ENGINE } from '@tensorflow/tfjs-core/dist/engine';
|
|
3
|
+
|
|
4
|
+
import './kernels';
|
|
5
|
+
|
|
6
|
+
export async function initWebGPUBackend() {
|
|
7
|
+
try {
|
|
8
|
+
const gpuDescriptor: GPURequestAdapterOptions = {
|
|
9
|
+
powerPreference: 'high-performance'
|
|
10
|
+
};
|
|
11
|
+
|
|
12
|
+
const adapter = (await navigator.gpu.requestAdapter(gpuDescriptor))!;
|
|
13
|
+
const deviceDescriptor: GPUDeviceDescriptor = {};
|
|
14
|
+
|
|
15
|
+
const requiredFeatures = [];
|
|
16
|
+
if (adapter.features.has('timestamp-query')) {
|
|
17
|
+
requiredFeatures.push('timestamp-query');
|
|
18
|
+
}
|
|
19
|
+
if (adapter.features.has('bgra8unorm-storage')) {
|
|
20
|
+
requiredFeatures.push(['bgra8unorm-storage']);
|
|
21
|
+
}
|
|
22
|
+
deviceDescriptor.requiredFeatures =
|
|
23
|
+
requiredFeatures as Iterable<GPUFeatureName>;
|
|
24
|
+
|
|
25
|
+
const adapterLimits = adapter.limits;
|
|
26
|
+
deviceDescriptor.requiredLimits = {
|
|
27
|
+
maxComputeWorkgroupStorageSize:
|
|
28
|
+
adapterLimits.maxComputeWorkgroupStorageSize,
|
|
29
|
+
maxComputeWorkgroupsPerDimension:
|
|
30
|
+
adapterLimits.maxComputeWorkgroupsPerDimension,
|
|
31
|
+
maxStorageBufferBindingSize: adapterLimits.maxStorageBufferBindingSize,
|
|
32
|
+
maxBufferSize: adapterLimits.maxBufferSize,
|
|
33
|
+
maxComputeWorkgroupSizeX: adapterLimits.maxComputeWorkgroupSizeX,
|
|
34
|
+
maxComputeInvocationsPerWorkgroup:
|
|
35
|
+
adapterLimits.maxComputeInvocationsPerWorkgroup
|
|
36
|
+
};
|
|
37
|
+
const device = await adapter.requestDevice(deviceDescriptor);
|
|
38
|
+
const adapterInfo = await adapter.requestAdapterInfo();
|
|
39
|
+
|
|
40
|
+
return initWebGPUBackendWithDevice(device, adapterInfo);
|
|
41
|
+
} catch (e) {}
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
export async function initWebGPUBackendWithDevice(
|
|
45
|
+
device: GPUDevice,
|
|
46
|
+
adapter: GPUAdapterInfo
|
|
47
|
+
) {
|
|
48
|
+
const backend = new WebGPUBackend(device, adapter);
|
|
49
|
+
ENGINE.registerBackend('webgpu-oidn', () => backend);
|
|
50
|
+
await ENGINE.setBackend('webgpu-oidn');
|
|
51
|
+
|
|
52
|
+
return backend;
|
|
53
|
+
}
|