oidn-web 0.2.2 → 0.3.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.
@@ -13,7 +13,8 @@ export interface Uniform {
13
13
  data: Float32Array | Int32Array | Uint32Array;
14
14
  }
15
15
  export interface WGPUComputePassInput {
16
- buffer: GPUBuffer;
16
+ buffer?: GPUBuffer;
17
+ texture?: GPUTexture;
17
18
  channels: number;
18
19
  }
19
20
 
@@ -22,6 +23,8 @@ export interface WGPUComputePassOutput {
22
23
  }
23
24
 
24
25
  export class WGPUComputePass<I extends string, O extends string> {
26
+ autoUpdateOutputBuffer = true;
27
+
25
28
  private _label;
26
29
 
27
30
  private _device;
@@ -36,6 +39,7 @@ export class WGPUComputePass<I extends string, O extends string> {
36
39
  private _pipeline!: GPUComputePipeline;
37
40
  private _bindGroups: GPUBindGroup[] = [];
38
41
  private _needsUpdatePipeline = true;
42
+ private _needsResizeBuffer = true;
39
43
 
40
44
  private _inputs: string[] = [];
41
45
  private _outputs: string[] = [];
@@ -52,6 +56,12 @@ export class WGPUComputePass<I extends string, O extends string> {
52
56
  private _csMain;
53
57
  private _csDefine;
54
58
 
59
+ private _groupOffsets = {
60
+ inputs: 0,
61
+ uniforms: 1,
62
+ outputs: 2
63
+ };
64
+
55
65
  constructor(
56
66
  label: string,
57
67
  device: GPUDevice,
@@ -61,6 +71,7 @@ export class WGPUComputePass<I extends string, O extends string> {
61
71
  csMain: string;
62
72
  csDefine?: string;
63
73
  uniforms: Uniform[];
74
+ autoUpdateOutputBuffer?: boolean;
64
75
  }
65
76
  ) {
66
77
  this._label = label;
@@ -70,6 +81,7 @@ export class WGPUComputePass<I extends string, O extends string> {
70
81
  this._inputs = opts.inputs;
71
82
  this._outputs = opts.outputs;
72
83
  this._uniforms = opts.uniforms;
84
+ this.autoUpdateOutputBuffer = opts.autoUpdateOutputBuffer ?? true;
73
85
 
74
86
  opts.uniforms.forEach((uniform) => {
75
87
  this._uniformBuffers[uniform.label] = device.createBuffer({
@@ -85,6 +97,12 @@ export class WGPUComputePass<I extends string, O extends string> {
85
97
  });
86
98
  }
87
99
 
100
+ setCSCode({ csDefine, csMain }: { csDefine: string; csMain: string }) {
101
+ this._csDefine = csDefine;
102
+ this._csMain = csMain;
103
+ this._needsUpdatePipeline = true;
104
+ }
105
+
88
106
  setSize(width: number, height: number) {
89
107
  width = Math.ceil(width);
90
108
  height = Math.ceil(height);
@@ -92,7 +110,7 @@ export class WGPUComputePass<I extends string, O extends string> {
92
110
  this._width = width;
93
111
  this._height = height;
94
112
  if (sizeChanged) {
95
- this._resizeOutputBuffers();
113
+ this._needsResizeBuffer = true;
96
114
  this._needsUpdatePipeline = true;
97
115
  }
98
116
  }
@@ -105,16 +123,32 @@ export class WGPUComputePass<I extends string, O extends string> {
105
123
  }
106
124
 
107
125
  setOutputParams(outputParams: Record<O, WGPUComputePassOutput>) {
108
- this._updateOutputBuffers(outputParams);
126
+ if (this.autoUpdateOutputBuffer) {
127
+ this._updateOutputBuffers(outputParams);
128
+ }
109
129
  this._needsUpdatePipeline = true;
110
130
  }
111
131
 
132
+ setOutputBuffers(outputBuffers: Record<O, GPUBuffer>) {
133
+ this._outputBuffers = Object.keys(outputBuffers).reduce((obj, key) => {
134
+ obj[key] = {
135
+ buffer: outputBuffers[key as O],
136
+ params: { channels: 4 }
137
+ };
138
+ return obj;
139
+ }, {} as WGPUComputePass<I, O>['_outputBuffers']);
140
+ }
141
+
112
142
  setUniform(label: string, data: Float32Array | Int32Array | Uint32Array) {
113
143
  const buffer = this._uniformBuffers[label];
114
144
  this._device.queue.writeBuffer(buffer, 0, data);
115
145
  }
116
146
 
117
- getOutputBuffer(name: O) {
147
+ getOutput(name: O) {
148
+ if (this._needsResizeBuffer && this.autoUpdateOutputBuffer) {
149
+ this._resizeOutputBuffers();
150
+ this._needsResizeBuffer = false;
151
+ }
118
152
  return this._outputBuffers[name].buffer;
119
153
  }
120
154
 
@@ -123,7 +157,7 @@ export class WGPUComputePass<I extends string, O extends string> {
123
157
  (this._uniformBuffers as any)[key].destroy();
124
158
  });
125
159
  Object.keys(this._outputBuffers).forEach((key) => {
126
- (this._outputBuffers as any)[key].texture.destroy();
160
+ (this._outputBuffers as any)[key].buffer.destroy();
127
161
  });
128
162
  }
129
163
 
@@ -173,13 +207,16 @@ export class WGPUComputePass<I extends string, O extends string> {
173
207
  }
174
208
  }
175
209
 
176
- private _updatePipeline(inputParams: Record<string, WGPUComputePassInput>) {
210
+ private _updatePipeline(
211
+ inputParams: Record<string, WGPUComputePassInput>,
212
+ inputTypes: Record<string, 'texture' | 'buffer'>
213
+ ) {
177
214
  if (!this._needsUpdatePipeline) {
178
215
  return;
179
216
  }
180
217
  this._needsUpdatePipeline = false;
181
218
  const device = this._device;
182
- const csCode = this._getFullCs(inputParams);
219
+ const csCode = this._getFullCs(inputParams, inputTypes);
183
220
  if (csCode === this._csCode) {
184
221
  return;
185
222
  }
@@ -198,35 +235,49 @@ export class WGPUComputePass<I extends string, O extends string> {
198
235
  this._updateBindGroups();
199
236
  }
200
237
 
201
- private _getFullCs(inputParams: Record<string, WGPUComputePassInput>) {
238
+ private _getFullCs(
239
+ inputParams: Record<string, WGPUComputePassInput>,
240
+ inputTypes: Record<string, 'texture' | 'buffer'>
241
+ ) {
202
242
  const inputs = this._inputs;
203
- const hasInputs = inputs.length > 0;
243
+ const uniforms = this._uniforms;
244
+ let offset = 0;
245
+ const groupOffsets = (this._groupOffsets = {
246
+ inputs: 0,
247
+ uniforms: 0,
248
+ outputs: 0
249
+ });
250
+ if (inputs.length > 0) {
251
+ offset++;
252
+ }
253
+ if (uniforms.length > 0) {
254
+ groupOffsets.uniforms = offset;
255
+ offset++;
256
+ }
257
+ groupOffsets.outputs = offset;
204
258
  const cs = `
205
259
  ${inputs
206
260
  .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
- )
261
+ .map((inputName, idx) => {
262
+ const bindingPrefix = `@group(${groupOffsets.inputs}) @binding(${idx}) `;
263
+ const varName = `in_${inputName}`;
264
+ // TODO more channels option.
265
+ return inputTypes[inputName] === 'texture'
266
+ ? `${bindingPrefix} var ${varName}: texture_2d<f32>;`
267
+ : `${bindingPrefix} var<storage, read> ${varName}: array<vec${inputParams[inputName].channels}f>;`;
268
+ })
212
269
  .join('\n')}
213
270
  ${this._uniforms
214
271
  .map(
215
272
  (uniform, idx) =>
216
- `@group(${hasInputs ? 1 : 0}) @binding(${idx}) var<uniform> ${
217
- uniform.label
218
- }: ${uniform.type};`
273
+ `@group(${groupOffsets.uniforms}) @binding(${idx}) var<uniform> ${uniform.label}: ${uniform.type};`
219
274
  )
220
275
  .join('\n')}
221
276
 
222
277
  ${this._outputs
223
278
  .map(
224
279
  (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>;`
280
+ `@group(${groupOffsets.outputs}) @binding(${idx}) var<storage, read_write> out_${name}: array<vec${this._outputBuffers[name].params.channels}f>;`
230
281
  )
231
282
  .join('\n')}
232
283
  ${this._csDefine ?? ''}
@@ -242,13 +293,12 @@ ${this._csMain}
242
293
  private _updateBindGroups() {
243
294
  const bindGroups: GPUBindGroup[] = [];
244
295
  const device = this._device;
296
+ const groupOffsets = this._groupOffsets;
245
297
 
246
- //TODO
247
- const uniformBindGroupIndex = this._inputs.length > 0 ? 1 : 0;
248
298
  if (this._uniforms.length > 0) {
249
- bindGroups[uniformBindGroupIndex] = device.createBindGroup({
299
+ bindGroups[groupOffsets.uniforms] = device.createBindGroup({
250
300
  label: this._label,
251
- layout: this._pipeline.getBindGroupLayout(uniformBindGroupIndex),
301
+ layout: this._pipeline.getBindGroupLayout(groupOffsets.uniforms),
252
302
  entries: this._uniforms.map(
253
303
  (uniform, idx) =>
254
304
  ({
@@ -266,30 +316,42 @@ ${this._csMain}
266
316
 
267
317
  createPass(
268
318
  commandEncoder: GPUCommandEncoder,
269
- inputBuffers: Record<I, WGPUComputePassInput>
319
+ inputs: Record<I, WGPUComputePassInput>
270
320
  ) {
271
- this._updatePipeline(inputBuffers);
321
+ if (this._needsResizeBuffer && this.autoUpdateOutputBuffer) {
322
+ this._resizeOutputBuffers();
323
+ this._needsResizeBuffer = false;
324
+ }
325
+
326
+ const inputTypes = this._inputs.reduce((obj, inputName) => {
327
+ obj[inputName] = inputs[inputName as I].buffer ? 'buffer' : 'texture';
328
+ return obj;
329
+ }, {} as Record<string, 'texture' | 'buffer'>);
272
330
 
273
- const hasInputs = this._inputs.length > 0;
331
+ this._updatePipeline(inputs, inputTypes);
332
+
333
+ const groupOffsets = this._groupOffsets;
274
334
  // TODO createBindGroup every time?
275
- if (hasInputs) {
276
- this._bindGroups[0] = this._device.createBindGroup({
335
+ if (this._inputs.length > 0) {
336
+ this._bindGroups[groupOffsets.inputs] = this._device.createBindGroup({
277
337
  label: this._label,
278
- layout: this._pipeline.getBindGroupLayout(0),
279
- entries: this._inputs.map((bufferName, idx) => ({
338
+ layout: this._pipeline.getBindGroupLayout(groupOffsets.inputs),
339
+ entries: this._inputs.map((inputName, idx) => ({
280
340
  binding: idx,
281
341
  // TODO
282
- resource: {
283
- buffer: inputBuffers[bufferName as I].buffer
284
- }
342
+ resource: inputs[inputName as I].buffer
343
+ ? {
344
+ buffer: inputs[inputName as I].buffer!
345
+ }
346
+ : inputs[inputName as I].texture!.createView()
285
347
  }))
286
348
  });
287
349
  }
288
350
 
289
351
  // Outputs
290
- this._bindGroups[hasInputs ? 2 : 1] = this._device.createBindGroup({
352
+ this._bindGroups[groupOffsets.outputs] = this._device.createBindGroup({
291
353
  label: this._label,
292
- layout: this._pipeline.getBindGroupLayout(hasInputs ? 2 : 1),
354
+ layout: this._pipeline.getBindGroupLayout(groupOffsets.outputs),
293
355
  entries: this._outputs.map((bufferName, idx) => ({
294
356
  binding: idx,
295
357
  resource: {
package/src/backend.ts CHANGED
@@ -45,9 +45,15 @@ export async function initWebGPUBackendWithDevice(
45
45
  device: GPUDevice,
46
46
  adapter: GPUAdapterInfo
47
47
  ) {
48
- const backend = new WebGPUBackend(device, adapter);
48
+ // TODO multiple device and adapter in one backend
49
+ let backend = ENGINE.findBackend('webgpu-oidn');
50
+ if (backend != null) {
51
+ return backend as WebGPUBackend;
52
+ }
53
+
54
+ backend = new WebGPUBackend(device, adapter);
49
55
  ENGINE.registerBackend('webgpu-oidn', () => backend);
50
56
  await ENGINE.setBackend('webgpu-oidn');
51
57
 
52
- return backend;
58
+ return backend as WebGPUBackend;
53
59
  }
package/src/helper.ts CHANGED
@@ -33,3 +33,11 @@ export function profileAndLogKernelCode(execute: () => void, disabled = true) {
33
33
  console.log(code);
34
34
  });
35
35
  }
36
+
37
+ export function memory() {
38
+ return tfjs.memory();
39
+ }
40
+
41
+ export function tidy(f: () => void) {
42
+ return tfjs.tidy(f);
43
+ }