oidn-web 0.3.5 → 0.5.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (65) hide show
  1. package/CHANGELOG.md +94 -0
  2. package/README.md +208 -8
  3. package/dist/oidn.js +4699 -22516
  4. package/dist/oidn.umd.cjs +989 -5796
  5. package/lib/UNet.d.ts +111 -26
  6. package/lib/UNet.js +310 -329
  7. package/lib/UNet.js.map +1 -1
  8. package/lib/WGPUComputePass.d.ts +1 -1
  9. package/lib/WGPUComputePass.js +6 -4
  10. package/lib/WGPUComputePass.js.map +1 -1
  11. package/lib/backend.d.ts +1 -4
  12. package/lib/backend.js +28 -44
  13. package/lib/backend.js.map +1 -1
  14. package/lib/finalRgbShader.d.ts +13 -0
  15. package/lib/finalRgbShader.js +160 -0
  16. package/lib/finalRgbShader.js.map +1 -0
  17. package/lib/graphOptimizer.d.ts +54 -0
  18. package/lib/graphOptimizer.js +215 -0
  19. package/lib/graphOptimizer.js.map +1 -0
  20. package/lib/hdrTransfer.d.ts +14 -0
  21. package/lib/hdrTransfer.js +61 -0
  22. package/lib/hdrTransfer.js.map +1 -0
  23. package/lib/main.d.ts +43 -11
  24. package/lib/main.js +9 -5
  25. package/lib/main.js.map +1 -1
  26. package/lib/modelSpec.d.ts +80 -0
  27. package/lib/modelSpec.js +270 -0
  28. package/lib/modelSpec.js.map +1 -0
  29. package/lib/nativeUNet.d.ts +103 -0
  30. package/lib/nativeUNet.js +2064 -0
  31. package/lib/nativeUNet.js.map +1 -0
  32. package/lib/process.d.ts +5 -11
  33. package/lib/process.js +38 -49
  34. package/lib/process.js.map +1 -1
  35. package/lib/resourceTracker.d.ts +26 -0
  36. package/lib/resourceTracker.js +65 -0
  37. package/lib/resourceTracker.js.map +1 -0
  38. package/lib/tileScheduler.d.ts +61 -0
  39. package/lib/tileScheduler.js +199 -0
  40. package/lib/tileScheduler.js.map +1 -0
  41. package/lib/webnnUNet.d.ts +52 -0
  42. package/lib/webnnUNet.js +535 -0
  43. package/lib/webnnUNet.js.map +1 -0
  44. package/package.json +16 -5
  45. package/src/UNet.ts +463 -437
  46. package/src/WGPUComputePass.ts +6 -4
  47. package/src/backend.ts +33 -59
  48. package/src/finalRgbShader.ts +186 -0
  49. package/src/graphOptimizer.ts +300 -0
  50. package/src/hdrTransfer.ts +88 -0
  51. package/src/main.ts +95 -20
  52. package/src/modelSpec.ts +414 -0
  53. package/src/nativeUNet.ts +2655 -0
  54. package/src/process.ts +46 -71
  55. package/src/resourceTracker.ts +94 -0
  56. package/src/tileScheduler.ts +330 -0
  57. package/src/webnnUNet.ts +812 -0
  58. package/lib/helper.d.ts +0 -4
  59. package/lib/helper.js +0 -33
  60. package/lib/helper.js.map +0 -1
  61. package/lib/kernels.d.ts +0 -1
  62. package/lib/kernels.js +0 -26
  63. package/lib/kernels.js.map +0 -1
  64. package/src/helper.ts +0 -43
  65. package/src/kernels.ts +0 -31
@@ -0,0 +1,812 @@
1
+ import { Float16Array } from '@petamoriken/float16';
2
+ import type {
3
+ Conv2DNodeSpec,
4
+ ValidatedUNetModel
5
+ } from './modelSpec';
6
+ import type { HostTensor } from './tza';
7
+ import type {
8
+ NativeUNetPrecision,
9
+ NativeUNetPrecisionSetting
10
+ } from './nativeUNet';
11
+ import {
12
+ OIDNResourceTracker,
13
+ type OIDNResourceSnapshot
14
+ } from './resourceTracker.js';
15
+
16
+ type MLDataType = 'float16' | 'float32';
17
+ type MLOperandLike = object;
18
+ type MLGraphLike = { destroy?: () => void; devices?: readonly string[] };
19
+ type MLTensorLike = { destroy: () => void };
20
+
21
+ interface MLContextLike {
22
+ createExportableTensor(
23
+ descriptor: Record<string, unknown>,
24
+ device: GPUDevice
25
+ ): Promise<MLTensorLike>;
26
+ dispatch(
27
+ graph: MLGraphLike,
28
+ inputs: Record<string, MLTensorLike>,
29
+ outputs: Record<string, MLTensorLike>
30
+ ): void;
31
+ exportToGPU(tensor: MLTensorLike): Promise<GPUBuffer>;
32
+ opSupportLimits?: () => Record<string, any>;
33
+ readTensor(tensor: MLTensorLike): Promise<ArrayBuffer>;
34
+ writeTensor(tensor: MLTensorLike, data: ArrayBufferView): void;
35
+ destroy?: () => void;
36
+ }
37
+
38
+ interface MLGraphBuilderLike {
39
+ input(name: string, descriptor: Record<string, unknown>): MLOperandLike;
40
+ constant(
41
+ descriptor: Record<string, unknown>,
42
+ data: ArrayBufferView
43
+ ): MLOperandLike;
44
+ conv2d(
45
+ input: MLOperandLike,
46
+ filter: MLOperandLike,
47
+ options: Record<string, unknown>
48
+ ): MLOperandLike;
49
+ relu(input: MLOperandLike): MLOperandLike;
50
+ maxPool2d(
51
+ input: MLOperandLike,
52
+ options: Record<string, unknown>
53
+ ): MLOperandLike;
54
+ resample2d(
55
+ input: MLOperandLike,
56
+ options: Record<string, unknown>
57
+ ): MLOperandLike;
58
+ concat(inputs: readonly MLOperandLike[], axis: number): MLOperandLike;
59
+ build(outputs: Record<string, MLOperandLike>): Promise<MLGraphLike>;
60
+ }
61
+
62
+ interface WebNNShapeExecution {
63
+ graph: MLGraphLike;
64
+ inputTensor: MLTensorLike;
65
+ outputTensor: MLTensorLike;
66
+ outputBuffer: GPUBuffer;
67
+ inputUniform: GPUBuffer;
68
+ outputUniform: GPUBuffer;
69
+ width: number;
70
+ height: number;
71
+ lastUsed: number;
72
+ }
73
+
74
+ interface WebNNInteropPipelines {
75
+ input: GPUComputePipeline;
76
+ output: GPUComputePipeline;
77
+ }
78
+
79
+ const interopPipelinesByDevice = new WeakMap<
80
+ GPUDevice,
81
+ Map<number, WebNNInteropPipelines>
82
+ >();
83
+
84
+ function interopPipelines(device: GPUDevice, sourceCount: number) {
85
+ let bySourceCount = interopPipelinesByDevice.get(device);
86
+ if (!bySourceCount) {
87
+ bySourceCount = new Map();
88
+ interopPipelinesByDevice.set(device, bySourceCount);
89
+ }
90
+ let pipelines = bySourceCount.get(sourceCount);
91
+ if (!pipelines) {
92
+ const inputModule = device.createShaderModule({
93
+ label: `oidn/webnn/input-pack/${sourceCount}`,
94
+ code: createInputPackShader(sourceCount)
95
+ });
96
+ const outputModule = device.createShaderModule({
97
+ label: 'oidn/webnn/output-unpack',
98
+ code: createOutputUnpackShader()
99
+ });
100
+ pipelines = {
101
+ input: device.createComputePipeline({
102
+ label: `oidn/webnn/input-pack/${sourceCount}`,
103
+ layout: 'auto',
104
+ compute: { module: inputModule, entryPoint: 'main' }
105
+ }),
106
+ output: device.createComputePipeline({
107
+ label: 'oidn/webnn/output-unpack',
108
+ layout: 'auto',
109
+ compute: { module: outputModule, entryPoint: 'main' }
110
+ })
111
+ };
112
+ bySourceCount.set(sourceCount, pipelines);
113
+ }
114
+ return pipelines;
115
+ }
116
+
117
+ export interface WebNNUNetOptions {
118
+ precision?: NativeUNetPrecisionSetting;
119
+ shapeCacheSize?: number;
120
+ }
121
+
122
+ export interface WebNNRuntimeSupport {
123
+ available: boolean;
124
+ reason?: string;
125
+ fp16Conv: boolean;
126
+ gpuInterop: boolean;
127
+ }
128
+
129
+ const WORKGROUP_SIZE = 8;
130
+
131
+ function roundUp(value: number, alignment: number) {
132
+ return Math.ceil(value / alignment) * alignment;
133
+ }
134
+
135
+ function createMappedBuffer(
136
+ device: GPUDevice,
137
+ label: string,
138
+ data: ArrayBufferView,
139
+ usage: GPUBufferUsageFlags
140
+ ) {
141
+ const buffer = device.createBuffer({
142
+ label,
143
+ size: roundUp(data.byteLength, 4),
144
+ usage,
145
+ mappedAtCreation: true
146
+ });
147
+ new Uint8Array(buffer.getMappedRange()).set(
148
+ new Uint8Array(data.buffer, data.byteOffset, data.byteLength)
149
+ );
150
+ buffer.unmap();
151
+ return buffer;
152
+ }
153
+
154
+ function uniformBuffer(
155
+ device: GPUDevice,
156
+ label: string,
157
+ values: readonly number[]
158
+ ) {
159
+ const data = new Uint32Array(roundUp(values.length, 4));
160
+ data.set(values);
161
+ return createMappedBuffer(device, label, data, GPUBufferUsage.UNIFORM);
162
+ }
163
+
164
+ function hasDataType(
165
+ limits: Record<string, any>,
166
+ op: string,
167
+ operand: string,
168
+ dataType: MLDataType
169
+ ) {
170
+ return Boolean(
171
+ limits?.[op]?.[operand]?.dataTypes?.includes?.(dataType)
172
+ );
173
+ }
174
+
175
+ function tensorBytes(tensor: HostTensor, precision: NativeUNetPrecision) {
176
+ if (
177
+ precision === 'fp16' &&
178
+ tensor.desc.dataType === 'Float16'
179
+ ) {
180
+ return new Uint8Array(
181
+ tensor.data.buffer,
182
+ tensor.data.byteOffset,
183
+ tensor.data.byteLength
184
+ );
185
+ }
186
+ if (
187
+ precision === 'fp32' &&
188
+ tensor.desc.dataType === 'Float32'
189
+ ) {
190
+ return new Uint8Array(
191
+ tensor.data.buffer,
192
+ tensor.data.byteOffset,
193
+ tensor.data.byteLength
194
+ );
195
+ }
196
+
197
+ const source = tensor.desc.dataType === 'Float32'
198
+ ? new Float32Array(
199
+ tensor.data.buffer,
200
+ tensor.data.byteOffset,
201
+ tensor.data.byteLength / 4
202
+ )
203
+ : new Float16Array(
204
+ tensor.data.buffer,
205
+ tensor.data.byteOffset,
206
+ tensor.data.byteLength / 2
207
+ );
208
+ const converted = precision === 'fp16'
209
+ ? new Float16Array(source)
210
+ : new Float32Array(source);
211
+ return new Uint8Array(
212
+ converted.buffer,
213
+ converted.byteOffset,
214
+ converted.byteLength
215
+ );
216
+ }
217
+
218
+ function createInputPackShader(sourceCount: number) {
219
+ const sources = Array.from(
220
+ { length: sourceCount },
221
+ (_, index) =>
222
+ `@group(0) @binding(${index}) var<storage, read> input${index}: array<vec4<f32>>;`
223
+ ).join('\n');
224
+ const branches = Array.from({ length: sourceCount }, (_, index) => {
225
+ const firstChannel = index * 3;
226
+ return `if (channel < ${firstChannel + 3}u) {
227
+ return input${index}[pixel][channel - ${firstChannel}u];
228
+ }`;
229
+ }).join('\n ');
230
+ return /* wgsl */ `enable f16;
231
+ struct Params { width: u32, height: u32, channels: u32, padding: u32 }
232
+ ${sources}
233
+ @group(0) @binding(${sourceCount}) var<storage, read_write> outputData: array<f16>;
234
+ @group(0) @binding(${sourceCount + 1}) var<uniform> params: Params;
235
+
236
+ fn readChannel(pixel: u32, channel: u32) -> f32 {
237
+ ${branches}
238
+ return 0.0;
239
+ }
240
+
241
+ @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
242
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
243
+ if (gid.x >= params.width || gid.y >= params.height || gid.z >= params.channels) {
244
+ return;
245
+ }
246
+ let pixel = gid.y * params.width + gid.x;
247
+ let outputIndex = (gid.z * params.height + gid.y) * params.width + gid.x;
248
+ outputData[outputIndex] = f16(readChannel(pixel, gid.z));
249
+ }
250
+ `;
251
+ }
252
+
253
+ function createOutputUnpackShader() {
254
+ return /* wgsl */ `enable f16;
255
+ struct Params { width: u32, height: u32, padding0: u32, padding1: u32 }
256
+ @group(0) @binding(0) var<storage, read> inputData: array<f16>;
257
+ @group(0) @binding(1) var<storage, read_write> outputData: array<vec4<f32>>;
258
+ @group(0) @binding(2) var<uniform> params: Params;
259
+
260
+ @compute @workgroup_size(${WORKGROUP_SIZE}, ${WORKGROUP_SIZE}, 1)
261
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
262
+ if (gid.x >= params.width || gid.y >= params.height) { return; }
263
+ let pixel = gid.y * params.width + gid.x;
264
+ let plane = params.width * params.height;
265
+ outputData[pixel] = vec4<f32>(
266
+ f32(inputData[pixel]),
267
+ f32(inputData[plane + pixel]),
268
+ f32(inputData[plane * 2u + pixel]),
269
+ 0.0
270
+ );
271
+ }
272
+ `;
273
+ }
274
+
275
+ function resolveWebNNPrecision(
276
+ device: GPUDevice,
277
+ requested: NativeUNetPrecisionSetting
278
+ ): NativeUNetPrecision {
279
+ if (requested === 'fp32') {
280
+ throw new Error(
281
+ 'OIDN WebNN GPU interop currently requires FP16 exportable tensors'
282
+ );
283
+ }
284
+ if (!device.features.has('shader-f16')) {
285
+ throw new Error('OIDN WebNN requires shader-f16 on the shared GPUDevice');
286
+ }
287
+ return 'fp16';
288
+ }
289
+
290
+ function activation(
291
+ builder: MLGraphBuilderLike,
292
+ operand: MLOperandLike,
293
+ kind: Conv2DNodeSpec['activation']
294
+ ) {
295
+ return kind === 'relu' ? builder.relu(operand) : operand;
296
+ }
297
+
298
+ /** Experimental model-driven WebNN executor with FP16 WebGPU interop. */
299
+ export class WebNNUNetExecutor {
300
+ readonly precision: NativeUNetPrecision;
301
+ readonly support: WebNNRuntimeSupport;
302
+
303
+ private _context!: MLContextLike;
304
+ private _builderConstructor!: new (
305
+ context: MLContextLike
306
+ ) => MLGraphBuilderLike;
307
+ private _shapeCache = new Map<string, WebNNShapeExecution>();
308
+ private _shapePromises = new Map<string, Promise<WebNNShapeExecution>>();
309
+ private _retiredExecutions = new Set<WebNNShapeExecution>();
310
+ private _pendingCreationCount = 0;
311
+ private _shapeCacheSize: number;
312
+ private _clock = 0;
313
+ private _inputPipeline: GPUComputePipeline;
314
+ private _outputPipeline: GPUComputePipeline;
315
+ private _resources = new OIDNResourceTracker();
316
+ private _disposed = false;
317
+
318
+ constructor(
319
+ private _device: GPUDevice,
320
+ private _model: ValidatedUNetModel,
321
+ options: WebNNUNetOptions = {}
322
+ ) {
323
+ this.precision = resolveWebNNPrecision(
324
+ _device,
325
+ options.precision ?? 'auto'
326
+ );
327
+ this._shapeCacheSize = Math.max(1, options.shapeCacheSize ?? 2);
328
+ this.support = {
329
+ available: false,
330
+ fp16Conv: false,
331
+ gpuInterop: false
332
+ };
333
+
334
+ const pipelines = interopPipelines(
335
+ _device,
336
+ _model.inputChannels / 3
337
+ );
338
+ this._inputPipeline = pipelines.input;
339
+ this._outputPipeline = pipelines.output;
340
+ }
341
+
342
+ async prepare() {
343
+ if (this._disposed) throw new Error('OIDN WebNN executor is disposed');
344
+ const webNN = (globalThis.navigator as any)?.ml;
345
+ const Builder = (globalThis as any).MLGraphBuilder;
346
+ if (!webNN?.createContext || typeof Builder !== 'function') {
347
+ this.support.reason = 'WebNN is not exposed by this browser';
348
+ throw new Error(this.support.reason);
349
+ }
350
+ this._builderConstructor = Builder;
351
+ try {
352
+ try {
353
+ // Chromium's experimental implementation only enables WebGPU tensor
354
+ // interop for an explicitly GPU-backed context.
355
+ this._context = await webNN.createContext({
356
+ deviceType: 'gpu',
357
+ powerPreference: 'high-performance'
358
+ });
359
+ } catch {
360
+ this._context = await webNN.createContext({ deviceType: 'gpu' });
361
+ }
362
+ this._resources.track('ml-context', this._context);
363
+ if (this._disposed) throw new Error('OIDN WebNN executor is disposed');
364
+
365
+ if (
366
+ typeof this._context.createExportableTensor !== 'function' ||
367
+ typeof this._context.exportToGPU !== 'function'
368
+ ) {
369
+ this.support.reason = 'WebNN WebGPU tensor interop is unavailable';
370
+ throw new Error(this.support.reason);
371
+ }
372
+ const limits = this._context.opSupportLimits?.() ?? {};
373
+ this.support.fp16Conv =
374
+ hasDataType(limits, 'conv2d', 'input', 'float16') &&
375
+ hasDataType(limits, 'conv2d', 'filter', 'float16') &&
376
+ hasDataType(limits, 'conv2d', 'output', 'float16');
377
+ if (!this.support.fp16Conv) {
378
+ this.support.reason = 'WebNN does not support FP16 conv2d';
379
+ throw new Error(this.support.reason);
380
+ }
381
+
382
+ let probe: MLTensorLike | undefined;
383
+ let probeBuffer: GPUBuffer | undefined;
384
+ try {
385
+ probe = this._resources.track(
386
+ 'ml-tensor',
387
+ await this._context.createExportableTensor(
388
+ { dataType: 'float16', shape: [4] },
389
+ this._device
390
+ )
391
+ );
392
+ probeBuffer = this._resources.track(
393
+ 'gpu-buffer',
394
+ await this._context.exportToGPU(probe)
395
+ );
396
+ this.support.gpuInterop = true;
397
+ } catch (error) {
398
+ this.support.reason =
399
+ `WebNN FP16 WebGPU interop failed: ${String(error)}`;
400
+ throw new Error(this.support.reason);
401
+ } finally {
402
+ this._releaseBuffer(probeBuffer);
403
+ this._releaseTensor(probe);
404
+ }
405
+ this.support.available = true;
406
+ } catch (error) {
407
+ this._releaseContext();
408
+ throw error;
409
+ }
410
+ }
411
+
412
+ private _constant(
413
+ builder: MLGraphBuilderLike,
414
+ tensor: HostTensor
415
+ ) {
416
+ return builder.constant(
417
+ {
418
+ dataType: 'float16',
419
+ shape: [...tensor.desc.dims]
420
+ },
421
+ tensorBytes(tensor, this.precision)
422
+ );
423
+ }
424
+
425
+ private async _createExecution(width: number, height: number) {
426
+ if (this._disposed) throw new Error('OIDN WebNN executor is disposed');
427
+ const builder = new this._builderConstructor(this._context);
428
+ const values = new Map<string, MLOperandLike>();
429
+ const shapes = new Map<string, [number, number, number]>();
430
+ values.set(
431
+ this._model.spec.input,
432
+ builder.input('input', {
433
+ dataType: 'float16',
434
+ shape: [1, this._model.inputChannels, height, width]
435
+ })
436
+ );
437
+ shapes.set(this._model.spec.input, [this._model.inputChannels, height, width]);
438
+
439
+ for (const node of this._model.spec.nodes) {
440
+ let result: MLOperandLike;
441
+ let shape: [number, number, number];
442
+ if (node.op === 'conv2d') {
443
+ const inputShape = shapes.get(node.input)!;
444
+ const tensors = this._model.convTensors.get(node.id)!;
445
+ const convolution = builder.conv2d(
446
+ values.get(node.input)!,
447
+ this._constant(builder, tensors.weight),
448
+ {
449
+ bias: this._constant(builder, tensors.bias),
450
+ padding: [1, 1, 1, 1],
451
+ inputLayout: 'nchw',
452
+ filterLayout: 'oihw'
453
+ }
454
+ );
455
+ result = activation(builder, convolution, node.activation);
456
+ shape = [tensors.outputChannels, inputShape[1], inputShape[2]];
457
+ } else if (node.op === 'maxPool2d') {
458
+ const inputShape = shapes.get(node.input)!;
459
+ result = builder.maxPool2d(values.get(node.input)!, {
460
+ windowDimensions: [2, 2],
461
+ strides: [2, 2],
462
+ padding: [0, inputShape[1] % 2, 0, inputShape[2] % 2],
463
+ layout: 'nchw'
464
+ });
465
+ shape = [
466
+ inputShape[0],
467
+ Math.ceil(inputShape[1] / 2),
468
+ Math.ceil(inputShape[2] / 2)
469
+ ];
470
+ } else if (node.op === 'upsample2d') {
471
+ const inputShape = shapes.get(node.input)!;
472
+ result = builder.resample2d(values.get(node.input)!, {
473
+ mode: 'nearest-neighbor',
474
+ axes: [2, 3],
475
+ scales: [2, 2]
476
+ });
477
+ shape = [inputShape[0], inputShape[1] * 2, inputShape[2] * 2];
478
+ } else {
479
+ const inputShapes = node.inputs.map((input) => shapes.get(input)!);
480
+ if (
481
+ inputShapes.some(
482
+ (candidate) =>
483
+ candidate[1] !== inputShapes[0][1] ||
484
+ candidate[2] !== inputShapes[0][2]
485
+ )
486
+ ) {
487
+ throw new Error(
488
+ `WebNN concat ${node.id} has mismatched spatial shapes`
489
+ );
490
+ }
491
+ result = builder.concat(
492
+ node.inputs.map((input) => values.get(input)!),
493
+ 1
494
+ );
495
+ shape = [
496
+ inputShapes.reduce((sum, candidate) => sum + candidate[0], 0),
497
+ inputShapes[0][1],
498
+ inputShapes[0][2]
499
+ ];
500
+ }
501
+ values.set(node.id, result);
502
+ shapes.set(node.id, shape);
503
+ }
504
+
505
+ let graph: MLGraphLike | undefined;
506
+ let inputTensor: MLTensorLike | undefined;
507
+ let outputTensor: MLTensorLike | undefined;
508
+ let outputBuffer: GPUBuffer | undefined;
509
+ let inputUniform: GPUBuffer | undefined;
510
+ let outputUniform: GPUBuffer | undefined;
511
+ try {
512
+ graph = this._resources.track(
513
+ 'ml-graph',
514
+ await builder.build({
515
+ output: values.get(this._model.spec.output)!
516
+ })
517
+ );
518
+ inputTensor = this._resources.track(
519
+ 'ml-tensor',
520
+ await this._context.createExportableTensor(
521
+ {
522
+ dataType: 'float16',
523
+ shape: [1, this._model.inputChannels, height, width],
524
+ writable: true
525
+ },
526
+ this._device
527
+ )
528
+ );
529
+ outputTensor = this._resources.track(
530
+ 'ml-tensor',
531
+ await this._context.createExportableTensor(
532
+ {
533
+ dataType: 'float16',
534
+ shape: [1, this._model.outputChannels, height, width],
535
+ readable: true
536
+ },
537
+ this._device
538
+ )
539
+ );
540
+ outputBuffer = this._resources.track(
541
+ 'gpu-buffer',
542
+ this._device.createBuffer({
543
+ label: `oidn/webnn/output/${width}x${height}`,
544
+ size: width * height * 4 * 4,
545
+ usage:
546
+ GPUBufferUsage.STORAGE |
547
+ GPUBufferUsage.COPY_SRC |
548
+ GPUBufferUsage.COPY_DST
549
+ })
550
+ );
551
+ inputUniform = this._resources.track(
552
+ 'gpu-buffer',
553
+ uniformBuffer(
554
+ this._device,
555
+ `oidn/webnn/input/${width}x${height}`,
556
+ [width, height, this._model.inputChannels]
557
+ )
558
+ );
559
+ outputUniform = this._resources.track(
560
+ 'gpu-buffer',
561
+ uniformBuffer(
562
+ this._device,
563
+ `oidn/webnn/output/${width}x${height}`,
564
+ [width, height]
565
+ )
566
+ );
567
+ const execution: WebNNShapeExecution = {
568
+ graph,
569
+ inputTensor,
570
+ outputTensor,
571
+ outputBuffer,
572
+ inputUniform,
573
+ outputUniform,
574
+ width,
575
+ height,
576
+ lastUsed: ++this._clock
577
+ };
578
+ return execution;
579
+ } catch (error) {
580
+ this._releaseBuffer(outputUniform);
581
+ this._releaseBuffer(inputUniform);
582
+ this._releaseBuffer(outputBuffer);
583
+ this._releaseTensor(outputTensor);
584
+ this._releaseTensor(inputTensor);
585
+ this._releaseGraph(graph);
586
+ throw error;
587
+ }
588
+ }
589
+
590
+ private async _execution(width: number, height: number) {
591
+ const key = `${width}x${height}`;
592
+ let execution = this._shapeCache.get(key);
593
+ if (!execution) {
594
+ let pending = this._shapePromises.get(key);
595
+ if (!pending) {
596
+ pending = (async () => {
597
+ this._pendingCreationCount++;
598
+ try {
599
+ return await this._createExecution(width, height);
600
+ } finally {
601
+ this._pendingCreationCount--;
602
+ }
603
+ })();
604
+ this._shapePromises.set(key, pending);
605
+ }
606
+ try {
607
+ execution = await pending;
608
+ if (this._disposed) {
609
+ this._destroyExecution(execution);
610
+ throw new Error('OIDN WebNN executor is disposed');
611
+ }
612
+ this._shapeCache.set(key, execution);
613
+ } finally {
614
+ if (this._shapePromises.get(key) === pending) {
615
+ this._shapePromises.delete(key);
616
+ }
617
+ }
618
+ if (this._shapeCache.size > this._shapeCacheSize) {
619
+ const oldest = [...this._shapeCache.entries()]
620
+ .filter(([candidate]) => candidate !== key)
621
+ .sort((left, right) => left[1].lastUsed - right[1].lastUsed)[0];
622
+ if (oldest) {
623
+ this._shapeCache.delete(oldest[0]);
624
+ this._retireExecution(oldest[1]);
625
+ }
626
+ }
627
+ }
628
+ execution.lastUsed = ++this._clock;
629
+ return execution;
630
+ }
631
+
632
+ /** Compiles common tile shapes while the host still reports model loading. */
633
+ async prewarm(shapes: readonly { width: number; height: number }[]) {
634
+ for (const shape of shapes) {
635
+ await this._execution(shape.width, shape.height);
636
+ }
637
+ }
638
+
639
+ async execute(
640
+ inputBuffers: readonly GPUBuffer[],
641
+ width: number,
642
+ height: number
643
+ ) {
644
+ const sourceCount = this._model.inputChannels / 3;
645
+ if (inputBuffers.length !== sourceCount) {
646
+ throw new Error(
647
+ `OIDN WebNN expected ${sourceCount} input buffers, got ${inputBuffers.length}`
648
+ );
649
+ }
650
+ const execution = await this._execution(width, height);
651
+ const inputGPUBuffer = this._resources.track(
652
+ 'gpu-buffer',
653
+ await this._context.exportToGPU(execution.inputTensor)
654
+ );
655
+ try {
656
+ const inputEntries: GPUBindGroupEntry[] = inputBuffers.map(
657
+ (buffer, binding) => ({ binding, resource: { buffer } })
658
+ );
659
+ inputEntries.push({
660
+ binding: sourceCount,
661
+ resource: { buffer: inputGPUBuffer }
662
+ });
663
+ inputEntries.push({
664
+ binding: sourceCount + 1,
665
+ resource: { buffer: execution.inputUniform }
666
+ });
667
+ const inputBindGroup = this._device.createBindGroup({
668
+ label: 'oidn/webnn/input-bindings',
669
+ layout: this._inputPipeline.getBindGroupLayout(0),
670
+ entries: inputEntries
671
+ });
672
+ const inputEncoder = this._device.createCommandEncoder({
673
+ label: 'oidn/webnn/input-pack'
674
+ });
675
+ const inputPass = inputEncoder.beginComputePass();
676
+ inputPass.setPipeline(this._inputPipeline);
677
+ inputPass.setBindGroup(0, inputBindGroup);
678
+ inputPass.dispatchWorkgroups(
679
+ Math.ceil(width / WORKGROUP_SIZE),
680
+ Math.ceil(height / WORKGROUP_SIZE),
681
+ this._model.inputChannels
682
+ );
683
+ inputPass.end();
684
+ this._device.queue.submit([inputEncoder.finish()]);
685
+ } finally {
686
+ this._releaseBuffer(inputGPUBuffer);
687
+ }
688
+
689
+ this._context.dispatch(
690
+ execution.graph,
691
+ { input: execution.inputTensor },
692
+ { output: execution.outputTensor }
693
+ );
694
+ const outputGPUBuffer = this._resources.track(
695
+ 'gpu-buffer',
696
+ await this._context.exportToGPU(execution.outputTensor)
697
+ );
698
+ try {
699
+ const outputBindGroup = this._device.createBindGroup({
700
+ label: 'oidn/webnn/output-bindings',
701
+ layout: this._outputPipeline.getBindGroupLayout(0),
702
+ entries: [
703
+ { binding: 0, resource: { buffer: outputGPUBuffer } },
704
+ { binding: 1, resource: { buffer: execution.outputBuffer } },
705
+ { binding: 2, resource: { buffer: execution.outputUniform } }
706
+ ]
707
+ });
708
+ const outputEncoder = this._device.createCommandEncoder({
709
+ label: 'oidn/webnn/output-unpack'
710
+ });
711
+ const outputPass = outputEncoder.beginComputePass();
712
+ outputPass.setPipeline(this._outputPipeline);
713
+ outputPass.setBindGroup(0, outputBindGroup);
714
+ outputPass.dispatchWorkgroups(
715
+ Math.ceil(width / WORKGROUP_SIZE),
716
+ Math.ceil(height / WORKGROUP_SIZE)
717
+ );
718
+ outputPass.end();
719
+ this._device.queue.submit([outputEncoder.finish()]);
720
+ } finally {
721
+ this._releaseBuffer(outputGPUBuffer);
722
+ }
723
+ return execution.outputBuffer;
724
+ }
725
+
726
+ async executeCPU(input: Float32Array, width: number, height: number) {
727
+ const execution = await this._execution(width, height);
728
+ const plane = width * height;
729
+ const packed = new Float16Array(plane * this._model.inputChannels);
730
+ for (let pixel = 0; pixel < plane; pixel++) {
731
+ for (let channel = 0; channel < this._model.inputChannels; channel++) {
732
+ packed[channel * plane + pixel] =
733
+ input[pixel * this._model.inputChannels + channel];
734
+ }
735
+ }
736
+ this._context.writeTensor(execution.inputTensor, packed);
737
+ this._context.dispatch(
738
+ execution.graph,
739
+ { input: execution.inputTensor },
740
+ { output: execution.outputTensor }
741
+ );
742
+ const result = new Float16Array(
743
+ await this._context.readTensor(execution.outputTensor)
744
+ );
745
+ const unpacked = new Float32Array(plane * this._model.outputChannels);
746
+ for (let pixel = 0; pixel < plane; pixel++) {
747
+ for (let channel = 0; channel < this._model.outputChannels; channel++) {
748
+ unpacked[pixel * this._model.outputChannels + channel] =
749
+ result[channel * plane + pixel];
750
+ }
751
+ }
752
+ return unpacked;
753
+ }
754
+
755
+ private _destroyExecution(execution: WebNNShapeExecution) {
756
+ this._releaseGraph(execution.graph);
757
+ this._releaseTensor(execution.inputTensor);
758
+ this._releaseTensor(execution.outputTensor);
759
+ this._releaseBuffer(execution.outputBuffer);
760
+ this._releaseBuffer(execution.inputUniform);
761
+ this._releaseBuffer(execution.outputUniform);
762
+ }
763
+
764
+ private _retireExecution(execution: WebNNShapeExecution) {
765
+ this._retiredExecutions.add(execution);
766
+ void this._device.queue.onSubmittedWorkDone().catch(() => undefined).then(() => {
767
+ this._retiredExecutions.delete(execution);
768
+ this._destroyExecution(execution);
769
+ });
770
+ }
771
+
772
+ private _releaseBuffer(buffer: GPUBuffer | undefined) {
773
+ this._resources.release('gpu-buffer', buffer, () => buffer!.destroy());
774
+ }
775
+
776
+ private _releaseTensor(tensor: MLTensorLike | undefined) {
777
+ this._resources.release('ml-tensor', tensor, () => tensor!.destroy());
778
+ }
779
+
780
+ private _releaseGraph(graph: MLGraphLike | undefined) {
781
+ this._resources.release('ml-graph', graph, () => graph!.destroy?.());
782
+ }
783
+
784
+ private _releaseContext() {
785
+ this._resources.release(
786
+ 'ml-context',
787
+ this._context,
788
+ () => this._context.destroy?.()
789
+ );
790
+ }
791
+
792
+ getResourceInfo(): OIDNResourceSnapshot {
793
+ return this._resources.snapshot(
794
+ this._pendingCreationCount + this._retiredExecutions.size
795
+ );
796
+ }
797
+
798
+ dispose() {
799
+ if (this._disposed) return;
800
+ this._disposed = true;
801
+ for (const execution of this._shapeCache.values()) {
802
+ this._destroyExecution(execution);
803
+ }
804
+ this._shapeCache.clear();
805
+ for (const execution of this._retiredExecutions) {
806
+ this._destroyExecution(execution);
807
+ }
808
+ this._retiredExecutions.clear();
809
+ this._shapePromises.clear();
810
+ this._releaseContext();
811
+ }
812
+ }