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.
Files changed (52) hide show
  1. package/LICENSE +21 -0
  2. package/README.md +2 -0
  3. package/dist/oidn.mjs +22909 -0
  4. package/dist/oidn.umd.js +5905 -0
  5. package/lib/UNet.d.ts +54 -0
  6. package/lib/UNet.js +467 -0
  7. package/lib/UNet.js.map +1 -0
  8. package/lib/WGPUComputePass.d.ts +53 -0
  9. package/lib/WGPUComputePass.js +220 -0
  10. package/lib/WGPUComputePass.js.map +1 -0
  11. package/lib/WGPUFullQuadPass.d.ts +51 -0
  12. package/lib/WGPUFullQuadPass.js +261 -0
  13. package/lib/WGPUFullQuadPass.js.map +1 -0
  14. package/lib/backend.d.ts +5 -0
  15. package/lib/backend.js +41 -0
  16. package/lib/backend.js.map +1 -0
  17. package/lib/hdr.d.ts +31 -0
  18. package/lib/hdr.js +340 -0
  19. package/lib/hdr.js.map +1 -0
  20. package/lib/helper.d.ts +1 -0
  21. package/lib/helper.js +27 -0
  22. package/lib/helper.js.map +1 -0
  23. package/lib/kernels.d.ts +1 -0
  24. package/lib/kernels.js +26 -0
  25. package/lib/kernels.js.map +1 -0
  26. package/lib/main.d.ts +20 -0
  27. package/lib/main.js +20 -0
  28. package/lib/main.js.map +1 -0
  29. package/lib/process.d.ts +39 -0
  30. package/lib/process.js +309 -0
  31. package/lib/process.js.map +1 -0
  32. package/lib/tza.d.ts +15 -0
  33. package/lib/tza.js +114 -0
  34. package/lib/tza.js.map +1 -0
  35. package/package.json +26 -0
  36. package/src/UNet.ts +708 -0
  37. package/src/WGPUComputePass.ts +318 -0
  38. package/src/WGPUFullQuadPass.ts +348 -0
  39. package/src/backend.ts +53 -0
  40. package/src/hdr.ts +398 -0
  41. package/src/helper.ts +35 -0
  42. package/src/kernels.ts +31 -0
  43. package/src/main.ts +42 -0
  44. package/src/process.ts +362 -0
  45. package/src/tza.ts +136 -0
  46. package/weights/.gitattributes +1 -0
  47. package/weights/LICENSE.txt +202 -0
  48. package/weights/README.md +7 -0
  49. package/weights/rt_hdr.tza +0 -0
  50. package/weights/rt_hdr_alb_nrm.tza +0 -0
  51. package/weights/rt_ldr.tza +0 -0
  52. 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
+ }