oidn-web 0.3.5 → 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.
Files changed (122) hide show
  1. package/CHANGELOG.md +49 -0
  2. package/README.md +140 -4
  3. package/benchmarks/compare.mjs +651 -0
  4. package/benchmarks/leak.mjs +255 -0
  5. package/benchmarks/results/before-spatial.json +391 -0
  6. package/benchmarks/results/before-spatial.md +47 -0
  7. package/benchmarks/results/int8-scan.json +2007 -0
  8. package/benchmarks/results/int8-scan.md +160 -0
  9. package/benchmarks/results/int8-w8a8-scan.json +2007 -0
  10. package/benchmarks/results/int8-w8a8-scan.md +160 -0
  11. package/benchmarks/results/int8-weight-channel.json +1413 -0
  12. package/benchmarks/results/int8-weight-channel.md +118 -0
  13. package/benchmarks/results/int8-weight-only.json +1437 -0
  14. package/benchmarks/results/int8-weight-only.md +118 -0
  15. package/benchmarks/results/kernel-webnn-final.json +1115 -0
  16. package/benchmarks/results/kernel-webnn-final.md +104 -0
  17. package/benchmarks/results/latest-optimized.json +375 -0
  18. package/benchmarks/results/latest-optimized.md +47 -0
  19. package/benchmarks/results/latest.json +391 -0
  20. package/benchmarks/results/latest.md +47 -0
  21. package/benchmarks/results/profile-baseline.json +331 -0
  22. package/benchmarks/results/profile-baseline.md +12 -0
  23. package/benchmarks/results/profile-conv2x.json +331 -0
  24. package/benchmarks/results/profile-conv2x.md +12 -0
  25. package/benchmarks/results/profile-fast-init.json +385 -0
  26. package/benchmarks/results/profile-fast-init.md +47 -0
  27. package/benchmarks/results/profile-fp16-fma.json +369 -0
  28. package/benchmarks/results/profile-fp16-fma.md +47 -0
  29. package/benchmarks/results/profile-fp16-tiled-decoder.json +369 -0
  30. package/benchmarks/results/profile-fp16-tiled-decoder.md +47 -0
  31. package/benchmarks/results/profile-fp16-tiled-encoder.json +369 -0
  32. package/benchmarks/results/profile-fp16-tiled-encoder.md +47 -0
  33. package/benchmarks/results/profile-fp16-unfused-pool.json +385 -0
  34. package/benchmarks/results/profile-fp16-unfused-pool.md +47 -0
  35. package/benchmarks/results/profile-input-major.json +347 -0
  36. package/benchmarks/results/profile-input-major.md +12 -0
  37. package/benchmarks/results/profile-k16.json +347 -0
  38. package/benchmarks/results/profile-k16.md +12 -0
  39. package/benchmarks/results/profile-k4.json +347 -0
  40. package/benchmarks/results/profile-k4.md +12 -0
  41. package/benchmarks/results/profile-pool-reuse.json +331 -0
  42. package/benchmarks/results/profile-pool-reuse.md +12 -0
  43. package/benchmarks/results/profile-precompiled.json +385 -0
  44. package/benchmarks/results/profile-precompiled.md +47 -0
  45. package/benchmarks/results/profile-static-channels.json +385 -0
  46. package/benchmarks/results/profile-static-channels.md +47 -0
  47. package/benchmarks/results/profile-static-io.json +385 -0
  48. package/benchmarks/results/profile-static-io.md +47 -0
  49. package/benchmarks/results/profile-tiled-conv.json +331 -0
  50. package/benchmarks/results/profile-tiled-conv.md +12 -0
  51. package/benchmarks/results/profile-tiled-decoder.json +347 -0
  52. package/benchmarks/results/profile-tiled-decoder.md +12 -0
  53. package/benchmarks/results/profile-tiled-matmul.json +331 -0
  54. package/benchmarks/results/profile-tiled-matmul.md +12 -0
  55. package/benchmarks/results/profile-unfused-decoder.json +379 -0
  56. package/benchmarks/results/profile-unfused-decoder.md +12 -0
  57. package/benchmarks/results/profile-unfused-pool.json +347 -0
  58. package/benchmarks/results/profile-unfused-pool.md +12 -0
  59. package/benchmarks/results/spatial-auto.json +575 -0
  60. package/benchmarks/results/spatial-auto.md +61 -0
  61. package/benchmarks/results/subgroup-smoke.json +1094 -0
  62. package/benchmarks/results/subgroup-smoke.md +104 -0
  63. package/benchmarks/results/webnn-smoke.json +739 -0
  64. package/benchmarks/results/webnn-smoke.md +76 -0
  65. package/dist/oidn.js +4166 -22580
  66. package/dist/oidn.umd.cjs +776 -5799
  67. package/lib/UNet.d.ts +66 -15
  68. package/lib/UNet.js +162 -257
  69. package/lib/UNet.js.map +1 -1
  70. package/lib/WGPUComputePass.d.ts +1 -1
  71. package/lib/WGPUComputePass.js +6 -4
  72. package/lib/WGPUComputePass.js.map +1 -1
  73. package/lib/backend.d.ts +8 -4
  74. package/lib/backend.js +36 -44
  75. package/lib/backend.js.map +1 -1
  76. package/lib/graphOptimizer.d.ts +54 -0
  77. package/lib/graphOptimizer.js +216 -0
  78. package/lib/graphOptimizer.js.map +1 -0
  79. package/lib/main.d.ts +33 -10
  80. package/lib/main.js +5 -0
  81. package/lib/main.js.map +1 -1
  82. package/lib/modelSpec.d.ts +80 -0
  83. package/lib/modelSpec.js +270 -0
  84. package/lib/modelSpec.js.map +1 -0
  85. package/lib/nativeUNet.d.ts +67 -0
  86. package/lib/nativeUNet.js +1735 -0
  87. package/lib/nativeUNet.js.map +1 -0
  88. package/lib/process.js +3 -0
  89. package/lib/process.js.map +1 -1
  90. package/lib/resourceTracker.d.ts +26 -0
  91. package/lib/resourceTracker.js +65 -0
  92. package/lib/resourceTracker.js.map +1 -0
  93. package/lib/tileScheduler.d.ts +33 -0
  94. package/lib/tileScheduler.js +86 -0
  95. package/lib/tileScheduler.js.map +1 -0
  96. package/lib/webnnUNet.d.ts +52 -0
  97. package/lib/webnnUNet.js +535 -0
  98. package/lib/webnnUNet.js.map +1 -0
  99. package/package.json +9 -5
  100. package/scripts/inspect-model.mjs +64 -0
  101. package/src/UNet.ts +236 -339
  102. package/src/WGPUComputePass.ts +6 -4
  103. package/src/backend.ts +42 -55
  104. package/src/graphOptimizer.ts +301 -0
  105. package/src/main.ts +71 -11
  106. package/src/modelSpec.ts +414 -0
  107. package/src/nativeUNet.ts +2256 -0
  108. package/src/process.ts +3 -1
  109. package/src/resourceTracker.ts +94 -0
  110. package/src/tileScheduler.ts +138 -0
  111. package/src/webnnUNet.ts +812 -0
  112. package/tests/modelSpec.test.mjs +128 -0
  113. package/tests/resourceLifecycle.test.mjs +383 -0
  114. package/tests/tileScheduler.test.mjs +90 -0
  115. package/lib/helper.d.ts +0 -4
  116. package/lib/helper.js +0 -33
  117. package/lib/helper.js.map +0 -1
  118. package/lib/kernels.d.ts +0 -1
  119. package/lib/kernels.js +0 -26
  120. package/lib/kernels.js.map +0 -1
  121. package/src/helper.ts +0 -43
  122. package/src/kernels.ts +0 -31
@@ -0,0 +1,128 @@
1
+ import test from 'node:test';
2
+ import assert from 'node:assert/strict';
3
+
4
+ import {
5
+ detectUNetModelSpec,
6
+ OIDN_UNET_LARGE_SPEC,
7
+ OIDN_UNET_SMALL_SPEC,
8
+ validateUNetModel
9
+ } from '../lib/modelSpec.js';
10
+ import { optimizeModelGraph, planModelExecution } from '../lib/graphOptimizer.js';
11
+ import { HostTensor, TensorDesc } from '../lib/tza.js';
12
+
13
+ function hostTensor(dims, layout = 'x', dataType = 'Float16') {
14
+ const desc = new TensorDesc();
15
+ desc.dims = [...dims];
16
+ desc.paddedDims = [...dims];
17
+ desc.layout = layout;
18
+ desc.dataType = dataType;
19
+ return new HostTensor(desc, new Uint8Array(desc.getByteSize()));
20
+ }
21
+
22
+ function makeModelTensors(spec, inputChannels = 3) {
23
+ const tensors = new Map();
24
+ const channels = new Map([[spec.input, inputChannels]]);
25
+
26
+ for (const node of spec.nodes) {
27
+ if (node.op === 'conv2d') {
28
+ const inChannels = channels.get(node.input);
29
+ assert.notEqual(inChannels, undefined);
30
+ const outChannels = node.id === spec.output ? 3 : 4;
31
+ tensors.set(
32
+ node.weight,
33
+ hostTensor([outChannels, inChannels, 3, 3], 'oihw')
34
+ );
35
+ tensors.set(node.bias, hostTensor([outChannels]));
36
+ channels.set(node.id, outChannels);
37
+ } else if (node.op === 'concat') {
38
+ channels.set(
39
+ node.id,
40
+ node.inputs.reduce((sum, input) => sum + channels.get(input), 0)
41
+ );
42
+ } else {
43
+ channels.set(node.id, channels.get(node.input));
44
+ }
45
+ }
46
+ return tensors;
47
+ }
48
+
49
+ test('detects and validates built-in small and large descriptors', () => {
50
+ for (const spec of [OIDN_UNET_SMALL_SPEC, OIDN_UNET_LARGE_SPEC]) {
51
+ const tensors = makeModelTensors(spec, 9);
52
+ assert.equal(detectUNetModelSpec(tensors), spec);
53
+ const validated = validateUNetModel(tensors);
54
+ assert.equal(validated.spec, spec);
55
+ assert.equal(validated.inputChannels, 9);
56
+ assert.equal(validated.outputChannels, 3);
57
+ assert.equal(validated.tensorDataType, 'Float16');
58
+ }
59
+ });
60
+
61
+ test('rejects an unknown topology instead of silently using a stale graph', () => {
62
+ const tensors = makeModelTensors(OIDN_UNET_SMALL_SPEC);
63
+ tensors.set('future_block.weight', hostTensor([4, 4, 3, 3], 'oihw'));
64
+ assert.throws(
65
+ () => detectUNetModelSpec(tensors),
66
+ /Unsupported OIDN model topology/
67
+ );
68
+ });
69
+
70
+ test('validates tensor byte length and graph channel flow', () => {
71
+ const tensors = makeModelTensors(OIDN_UNET_SMALL_SPEC);
72
+ tensors.get('enc_conv0.weight').data = new Uint8Array(2);
73
+ assert.throws(
74
+ () => validateUNetModel(tensors, OIDN_UNET_SMALL_SPEC),
75
+ /bytes, expected/
76
+ );
77
+
78
+ const wrongChannels = makeModelTensors(OIDN_UNET_SMALL_SPEC);
79
+ wrongChannels.set('enc_conv1.weight', hostTensor([4, 7, 3, 3], 'oihw'));
80
+ assert.throws(
81
+ () => validateUNetModel(wrongChannels, OIDN_UNET_SMALL_SPEC),
82
+ /expects 7 input channels/
83
+ );
84
+ });
85
+
86
+ test('fuses topology patterns without hard-coding a model family', () => {
87
+ const validated = validateUNetModel(
88
+ makeModelTensors(OIDN_UNET_SMALL_SPEC),
89
+ OIDN_UNET_SMALL_SPEC
90
+ );
91
+ const graph = optimizeModelGraph(validated);
92
+
93
+ assert.deepEqual(graph.fusions, {
94
+ convPool: 4,
95
+ upsampleConcatConv: 4
96
+ });
97
+ assert.equal(
98
+ graph.nodes.filter((node) => node.op === 'fusedConvReluMaxPool2d').length,
99
+ 4
100
+ );
101
+ assert.equal(
102
+ graph.nodes.filter((node) => node.op === 'fusedUpsampleConcatConv2d')
103
+ .length,
104
+ 4
105
+ );
106
+ assert.equal(graph.nodes.length, 16);
107
+ });
108
+
109
+ test('plans fused graph shapes and lifetimes', () => {
110
+ const validated = validateUNetModel(
111
+ makeModelTensors(OIDN_UNET_LARGE_SPEC, 9),
112
+ OIDN_UNET_LARGE_SPEC
113
+ );
114
+ const plan = planModelExecution(validated, 384, 256);
115
+
116
+ assert.deepEqual(plan.inputShape, { width: 384, height: 256, channels: 9 });
117
+ assert.deepEqual(plan.valueShapes.get('pool4'), {
118
+ width: 24,
119
+ height: 16,
120
+ channels: 4
121
+ });
122
+ assert.deepEqual(plan.valueShapes.get(OIDN_UNET_LARGE_SPEC.output), {
123
+ width: 384,
124
+ height: 256,
125
+ channels: 3
126
+ });
127
+ assert.equal(plan.plannedNodes.at(-1).lastUse, plan.nodes.length);
128
+ });
@@ -0,0 +1,383 @@
1
+ import test from 'node:test';
2
+ import assert from 'node:assert/strict';
3
+
4
+ import { validateUNetModel } from '../lib/modelSpec.js';
5
+ import { OIDNResourceTracker } from '../lib/resourceTracker.js';
6
+ import { HostTensor, TensorDesc } from '../lib/tza.js';
7
+ import { NativeUNetExecutor } from '../lib/nativeUNet.js';
8
+ import { WebNNUNetExecutor } from '../lib/webnnUNet.js';
9
+
10
+ const TEST_SPEC = {
11
+ schemaVersion: 1,
12
+ id: 'resource-test-unet',
13
+ family: 'resource-test',
14
+ input: 'input',
15
+ output: 'output',
16
+ receptiveField: 3,
17
+ nodes: [
18
+ {
19
+ op: 'conv2d',
20
+ id: 'output',
21
+ input: 'input',
22
+ weight: 'output.weight',
23
+ bias: 'output.bias',
24
+ activation: 'identity',
25
+ padding: 'same'
26
+ }
27
+ ]
28
+ };
29
+
30
+ function hostTensor(dims, layout = 'x') {
31
+ const desc = new TensorDesc();
32
+ desc.dims = [...dims];
33
+ desc.paddedDims = [...dims];
34
+ desc.layout = layout;
35
+ desc.dataType = 'Float16';
36
+ return new HostTensor(desc, new Uint8Array(desc.getByteSize()));
37
+ }
38
+
39
+ function testModel() {
40
+ return validateUNetModel(
41
+ new Map([
42
+ ['output.weight', hostTensor([3, 3, 3, 3], 'oihw')],
43
+ ['output.bias', hostTensor([3])]
44
+ ]),
45
+ TEST_SPEC
46
+ );
47
+ }
48
+
49
+ class FakeResource {
50
+ destroyCalls = 0;
51
+
52
+ destroy() {
53
+ this.destroyCalls++;
54
+ }
55
+ }
56
+
57
+ class FakeBuffer extends FakeResource {
58
+ mapState = 'unmapped';
59
+
60
+ constructor(size = 16, mappedAtCreation = false) {
61
+ super();
62
+ this.data = new ArrayBuffer(size);
63
+ this.mapState = mappedAtCreation ? 'mapped' : 'unmapped';
64
+ }
65
+
66
+ getMappedRange() {
67
+ return this.data;
68
+ }
69
+
70
+ unmap() {
71
+ this.mapState = 'unmapped';
72
+ }
73
+ }
74
+
75
+ function fakeDevice() {
76
+ const device = {
77
+ features: new Set(['shader-f16']),
78
+ limits: {
79
+ maxComputeInvocationsPerWorkgroup: 256,
80
+ maxComputeWorkgroupSizeX: 256,
81
+ maxComputeWorkgroupSizeY: 256,
82
+ maxComputeWorkgroupStorageSize: 32768
83
+ },
84
+ failBindGroup: false,
85
+ failBufferAt: undefined,
86
+ createBufferCalls: 0,
87
+ buffers: [],
88
+ queue: {
89
+ submit() {},
90
+ onSubmittedWorkDone: () => Promise.resolve()
91
+ },
92
+ createShaderModule: () => ({}),
93
+ createComputePipeline: () => ({ getBindGroupLayout: () => ({}) }),
94
+ createComputePipelineAsync: async () => ({
95
+ getBindGroupLayout: () => ({})
96
+ }),
97
+ createBuffer(descriptor) {
98
+ this.createBufferCalls++;
99
+ if (this.createBufferCalls === this.failBufferAt) {
100
+ throw new Error('injected buffer allocation failure');
101
+ }
102
+ const buffer = new FakeBuffer(
103
+ descriptor.size,
104
+ descriptor.mappedAtCreation
105
+ );
106
+ this.buffers.push(buffer);
107
+ return buffer;
108
+ },
109
+ createBindGroup() {
110
+ if (this.failBindGroup) throw new Error('injected bind group failure');
111
+ return {};
112
+ },
113
+ createCommandEncoder: () => ({
114
+ beginComputePass: () => ({
115
+ setPipeline() {},
116
+ setBindGroup() {},
117
+ dispatchWorkgroups() {},
118
+ end() {}
119
+ }),
120
+ finish: () => ({})
121
+ })
122
+ };
123
+ return device;
124
+ }
125
+
126
+ function deferred() {
127
+ let resolve;
128
+ let reject;
129
+ const promise = new Promise((resolvePromise, rejectPromise) => {
130
+ resolve = resolvePromise;
131
+ reject = rejectPromise;
132
+ });
133
+ return { promise, resolve, reject };
134
+ }
135
+
136
+ async function withFakeWebNN(run, { build } = {}) {
137
+ const previous = {
138
+ navigator: Object.getOwnPropertyDescriptor(globalThis, 'navigator'),
139
+ builder: Object.getOwnPropertyDescriptor(globalThis, 'MLGraphBuilder'),
140
+ usage: Object.getOwnPropertyDescriptor(globalThis, 'GPUBufferUsage')
141
+ };
142
+ const device = fakeDevice();
143
+ const context = new FakeResource();
144
+ Object.assign(context, {
145
+ async createExportableTensor() {
146
+ return new FakeResource();
147
+ },
148
+ async exportToGPU() {
149
+ return device.createBuffer({ size: 16 });
150
+ },
151
+ dispatch() {},
152
+ opSupportLimits() {
153
+ const operand = { dataTypes: ['float16'] };
154
+ return { conv2d: { input: operand, filter: operand, output: operand } };
155
+ },
156
+ async readTensor() {
157
+ return new ArrayBuffer(0);
158
+ },
159
+ writeTensor() {}
160
+ });
161
+
162
+ class FakeBuilder {
163
+ input() {
164
+ return {};
165
+ }
166
+ constant() {
167
+ return {};
168
+ }
169
+ conv2d() {
170
+ return {};
171
+ }
172
+ relu(value) {
173
+ return value;
174
+ }
175
+ maxPool2d() {
176
+ return {};
177
+ }
178
+ resample2d() {
179
+ return {};
180
+ }
181
+ concat() {
182
+ return {};
183
+ }
184
+ async build(outputs) {
185
+ return build ? build(outputs) : new FakeResource();
186
+ }
187
+ }
188
+
189
+ Object.defineProperties(globalThis, {
190
+ navigator: {
191
+ configurable: true,
192
+ value: { ml: { createContext: async () => context } }
193
+ },
194
+ MLGraphBuilder: { configurable: true, value: FakeBuilder },
195
+ GPUBufferUsage: {
196
+ configurable: true,
197
+ value: {
198
+ STORAGE: 1,
199
+ COPY_SRC: 2,
200
+ COPY_DST: 4,
201
+ UNIFORM: 8,
202
+ QUERY_RESOLVE: 16,
203
+ MAP_READ: 32
204
+ }
205
+ }
206
+ });
207
+
208
+ try {
209
+ await run({ device, context });
210
+ } finally {
211
+ for (const [key, descriptor] of Object.entries(previous)) {
212
+ const name =
213
+ key === 'builder'
214
+ ? 'MLGraphBuilder'
215
+ : key === 'usage'
216
+ ? 'GPUBufferUsage'
217
+ : key;
218
+ if (descriptor) Object.defineProperty(globalThis, name, descriptor);
219
+ else delete globalThis[name];
220
+ }
221
+ }
222
+ }
223
+
224
+ test('resource tracker releases each owned handle exactly once', () => {
225
+ const tracker = new OIDNResourceTracker();
226
+ const buffer = new FakeResource();
227
+ tracker.track('gpu-buffer', buffer);
228
+ tracker.track('gpu-buffer', buffer);
229
+
230
+ assert.deepEqual(tracker.snapshot(), {
231
+ live: 1,
232
+ created: 1,
233
+ destroyed: 0,
234
+ peakLive: 1,
235
+ pending: 0,
236
+ byKind: {
237
+ 'gpu-buffer': { created: 1, destroyed: 0, live: 1, peakLive: 1 }
238
+ }
239
+ });
240
+ tracker.release('gpu-buffer', buffer, () => buffer.destroy());
241
+ tracker.release('gpu-buffer', buffer, () => buffer.destroy());
242
+ assert.equal(buffer.destroyCalls, 1);
243
+ assert.equal(tracker.snapshot().live, 0);
244
+ assert.equal(tracker.snapshot().created, tracker.snapshot().destroyed);
245
+ });
246
+
247
+ test('native WGSL shape eviction and dispose release every owned GPU buffer', async () => {
248
+ await withFakeWebNN(async ({ device }) => {
249
+ const executor = new NativeUNetExecutor(device, testModel(), {
250
+ precision: 'fp16',
251
+ shapeCacheSize: 2,
252
+ kernel: 'direct'
253
+ });
254
+ const input = new FakeBuffer();
255
+ executor.execute([input], 16, 16);
256
+ executor.execute([input], 32, 16);
257
+ executor.execute([input], 48, 16);
258
+ await new Promise((resolve) => setImmediate(resolve));
259
+
260
+ const cached = executor.getResourceInfo();
261
+ assert.equal(cached.pending, 0);
262
+ assert.equal(cached.byKind['gpu-buffer'].live, 10);
263
+ assert.equal(cached.byKind['gpu-buffer'].destroyed, 4);
264
+
265
+ executor.dispose();
266
+ executor.dispose();
267
+ const disposed = executor.getResourceInfo();
268
+ assert.equal(disposed.pending, 0);
269
+ assert.equal(disposed.live, 0);
270
+ assert.equal(disposed.created, disposed.destroyed);
271
+ assert.ok(device.buffers.every((buffer) => buffer.destroyCalls === 1));
272
+ assert.equal(input.destroyCalls, 0);
273
+ });
274
+ });
275
+
276
+ test('native WGSL rolls back a partially allocated shape on failure', async () => {
277
+ await withFakeWebNN(async ({ device }) => {
278
+ const executor = new NativeUNetExecutor(device, testModel(), {
279
+ precision: 'fp16',
280
+ kernel: 'direct'
281
+ });
282
+ const baseline = executor.getResourceInfo().live;
283
+ device.failBufferAt = device.createBufferCalls + 3;
284
+
285
+ assert.throws(
286
+ () => executor.execute([new FakeBuffer()], 16, 16),
287
+ /injected buffer allocation failure/
288
+ );
289
+ assert.equal(executor.getResourceInfo().live, baseline);
290
+ executor.dispose();
291
+ assert.equal(executor.getResourceInfo().live, 0);
292
+ assert.ok(device.buffers.every((buffer) => buffer.destroyCalls === 1));
293
+ });
294
+ });
295
+
296
+ test('WebNN shape cache stays bounded and dispose returns to zero live resources', async () => {
297
+ await withFakeWebNN(async ({ device, context }) => {
298
+ const executor = new WebNNUNetExecutor(device, testModel(), {
299
+ precision: 'fp16',
300
+ shapeCacheSize: 2
301
+ });
302
+ await executor.prepare();
303
+ await executor.prewarm([
304
+ { width: 16, height: 16 },
305
+ { width: 32, height: 16 },
306
+ { width: 48, height: 16 }
307
+ ]);
308
+ await Promise.resolve();
309
+
310
+ const cached = executor.getResourceInfo();
311
+ assert.equal(cached.pending, 0);
312
+ assert.equal(cached.byKind['ml-context'].live, 1);
313
+ assert.equal(cached.byKind['ml-graph'].live, 2);
314
+ assert.equal(cached.byKind['ml-tensor'].live, 4);
315
+ assert.equal(cached.byKind['gpu-buffer'].live, 6);
316
+
317
+ executor.dispose();
318
+ executor.dispose();
319
+ const disposed = executor.getResourceInfo();
320
+ assert.equal(disposed.pending, 0);
321
+ assert.equal(disposed.live, 0);
322
+ assert.equal(disposed.created, disposed.destroyed);
323
+ assert.equal(context.destroyCalls, 1);
324
+ assert.ok(device.buffers.every((buffer) => buffer.destroyCalls === 1));
325
+ });
326
+ });
327
+
328
+ test('WebNN cleans a graph that finishes after dispose', async () => {
329
+ const build = deferred();
330
+ await withFakeWebNN(
331
+ async ({ device }) => {
332
+ const executor = new WebNNUNetExecutor(device, testModel(), {
333
+ precision: 'fp16'
334
+ });
335
+ await executor.prepare();
336
+ const prewarm = executor.prewarm([{ width: 16, height: 16 }]);
337
+ await Promise.resolve();
338
+ assert.equal(executor.getResourceInfo().pending, 1);
339
+
340
+ executor.dispose();
341
+ build.resolve(new FakeResource());
342
+ await assert.rejects(prewarm, /disposed/);
343
+
344
+ const resources = executor.getResourceInfo();
345
+ assert.equal(resources.pending, 0);
346
+ assert.equal(resources.live, 0);
347
+ assert.equal(resources.created, resources.destroyed);
348
+ },
349
+ { build: () => build.promise }
350
+ );
351
+ });
352
+
353
+ test('WebNN releases exported GPU buffers when command setup throws', async () => {
354
+ await withFakeWebNN(async ({ device }) => {
355
+ const executor = new WebNNUNetExecutor(device, testModel(), {
356
+ precision: 'fp16'
357
+ });
358
+ await executor.prepare();
359
+ await executor.prewarm([{ width: 16, height: 16 }]);
360
+ const baseline = executor.getResourceInfo().live;
361
+ device.failBindGroup = true;
362
+
363
+ await assert.rejects(
364
+ executor.execute([new FakeBuffer()], 16, 16),
365
+ /injected bind group failure/
366
+ );
367
+ assert.equal(executor.getResourceInfo().live, baseline);
368
+ executor.dispose();
369
+ assert.equal(executor.getResourceInfo().live, 0);
370
+ });
371
+ });
372
+
373
+ test('WebNN prepare failures destroy the partially initialized context', async () => {
374
+ await withFakeWebNN(async ({ device, context }) => {
375
+ context.opSupportLimits = () => ({});
376
+ const executor = new WebNNUNetExecutor(device, testModel(), {
377
+ precision: 'fp16'
378
+ });
379
+ await assert.rejects(executor.prepare(), /does not support FP16/);
380
+ assert.equal(context.destroyCalls, 1);
381
+ assert.equal(executor.getResourceInfo().live, 0);
382
+ });
383
+ });
@@ -0,0 +1,90 @@
1
+ import test from 'node:test';
2
+ import assert from 'node:assert/strict';
3
+
4
+ import {
5
+ DynamicTileController,
6
+ fitTileDimension,
7
+ waitForSubmittedGPUWork
8
+ } from '../lib/tileScheduler.js';
9
+
10
+ test('starts conservatively and respects the hard maximum', () => {
11
+ assert.equal(new DynamicTileController(512).tileSize, 384);
12
+ assert.equal(new DynamicTileController(320).tileSize, 320);
13
+ assert.equal(fitTileDimension(512, 384), 384);
14
+ assert.equal(fitTileDimension(300, 384), 304);
15
+ });
16
+
17
+ test('reduces slow tiles and grows fast tiles within configured bounds', () => {
18
+ const controller = new DynamicTileController(512);
19
+
20
+ assert.equal(controller.observe([30, 32, 34]), true);
21
+ assert.equal(controller.tileSize, 256);
22
+ assert.equal(controller.observe([40]), false);
23
+ assert.equal(controller.tileSize, 256);
24
+
25
+ assert.equal(controller.observe([4, 6, 8]), true);
26
+ assert.equal(controller.tileSize, 384);
27
+ assert.equal(controller.observe([5]), true);
28
+ assert.equal(controller.tileSize, 512);
29
+ assert.equal(controller.observe([5]), false);
30
+ });
31
+
32
+ test('uses the median and ignores invalid timings', () => {
33
+ const controller = new DynamicTileController(512);
34
+
35
+ assert.equal(controller.observe([1, 30, Number.NaN]), false);
36
+ assert.equal(controller.tileSize, 384);
37
+ assert.equal(controller.observe([Number.NaN, Number.POSITIVE_INFINITY]), false);
38
+ });
39
+
40
+ test('can restore fixed-size behavior', () => {
41
+ const controller = new DynamicTileController(500, false);
42
+
43
+ assert.equal(controller.tileSize, 496);
44
+ assert.equal(controller.observe([100]), false);
45
+ assert.equal(controller.tileSize, 496);
46
+ });
47
+
48
+ test('supports custom adaptive limits and timing targets', () => {
49
+ const controller = new DynamicTileController(768, {
50
+ minTileSize: 128,
51
+ initialTileSize: 512,
52
+ targetTileTimeMs: 24,
53
+ adjustmentStep: 64
54
+ });
55
+
56
+ assert.equal(controller.tileSize, 512);
57
+ controller.observe([40]);
58
+ assert.equal(controller.tileSize, 448);
59
+ controller.observe([8]);
60
+ assert.equal(controller.tileSize, 512);
61
+ });
62
+
63
+ test('waits for submitted GPU work before continuing', async () => {
64
+ let release;
65
+ let completed = false;
66
+ const pendingGPUWork = new Promise((resolve) => {
67
+ release = resolve;
68
+ });
69
+ const queue = {
70
+ onSubmittedWorkDone: () => pendingGPUWork
71
+ };
72
+
73
+ const wait = waitForSubmittedGPUWork(queue).then(() => {
74
+ completed = true;
75
+ });
76
+ await Promise.resolve();
77
+ assert.equal(completed, false);
78
+
79
+ release();
80
+ await wait;
81
+ assert.equal(completed, true);
82
+ });
83
+
84
+ test('does not strand scheduling when the GPU wait rejects', async () => {
85
+ await assert.doesNotReject(() =>
86
+ waitForSubmittedGPUWork({
87
+ onSubmittedWorkDone: () => Promise.reject(new Error('device lost'))
88
+ })
89
+ );
90
+ });
package/lib/helper.d.ts DELETED
@@ -1,4 +0,0 @@
1
- import * as tfjs from '@tensorflow/tfjs-core';
2
- export declare function profileAndLogKernelCode(execute: () => void, disabled?: boolean): void;
3
- export declare function memory(): tfjs.MemoryInfo;
4
- export declare function tidy(f: () => void): void;
package/lib/helper.js DELETED
@@ -1,33 +0,0 @@
1
- import * as tfjs from '@tensorflow/tfjs-core';
2
- export function profileAndLogKernelCode(execute, disabled = true) {
3
- if (disabled) {
4
- execute();
5
- return;
6
- }
7
- tfjs
8
- .profile(() => {
9
- execute();
10
- })
11
- .then((res) => {
12
- const kernelNames = Array.from(new Set(res.kernels.map((k) => k.name).filter((name) => !name.endsWith('_op'))));
13
- function nameToConfig(name) {
14
- return `${name[0].toLowerCase()}${name.slice(1)}Config`;
15
- }
16
- const importCode = kernelNames.map((name) => `import { ${nameToConfig(name)} } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/${name}';`);
17
- const configCode = kernelNames.map((name) => `${nameToConfig(name)},`);
18
- const code = `
19
- ${importCode.join('\n')}
20
- const kernelConfigs: KernelConfig[] = [
21
- ${configCode.join('\n')}
22
- ]
23
- `;
24
- console.log(code);
25
- });
26
- }
27
- export function memory() {
28
- return tfjs.memory();
29
- }
30
- export function tidy(f) {
31
- return tfjs.tidy(f);
32
- }
33
- //# sourceMappingURL=helper.js.map
package/lib/helper.js.map DELETED
@@ -1 +0,0 @@
1
- {"version":3,"file":"helper.js","sourceRoot":"","sources":["../src/helper.ts"],"names":[],"mappings":"AAAA,OAAO,KAAK,IAAI,MAAM,uBAAuB,CAAC;AAC9C,MAAM,UAAU,uBAAuB,CAAC,OAAmB,EAAE,QAAQ,GAAG,IAAI;IAC1E,IAAI,QAAQ,EAAE,CAAC;QACb,OAAO,EAAE,CAAC;QACV,OAAO;IACT,CAAC;IACD,IAAI;SACD,OAAO,CAAC,GAAG,EAAE;QACZ,OAAO,EAAE,CAAC;IACZ,CAAC,CAAC;SACD,IAAI,CAAC,CAAC,GAAG,EAAE,EAAE;QACZ,MAAM,WAAW,GAAG,KAAK,CAAC,IAAI,CAC5B,IAAI,GAAG,CACL,GAAG,CAAC,OAAO,CAAC,GAAG,CAAC,CAAC,CAAC,EAAE,EAAE,CAAC,CAAC,CAAC,IAAI,CAAC,CAAC,MAAM,CAAC,CAAC,IAAI,EAAE,EAAE,CAAC,CAAC,IAAI,CAAC,QAAQ,CAAC,KAAK,CAAC,CAAC,CACvE,CACF,CAAC;QACF,SAAS,YAAY,CAAC,IAAY;YAChC,OAAO,GAAG,IAAI,CAAC,CAAC,CAAC,CAAC,WAAW,EAAE,GAAG,IAAI,CAAC,KAAK,CAAC,CAAC,CAAC,QAAQ,CAAC;QAC1D,CAAC;QACD,MAAM,UAAU,GAAG,WAAW,CAAC,GAAG,CAChC,CAAC,IAAI,EAAE,EAAE,CACP,YAAY,YAAY,CACtB,IAAI,CACL,yDAAyD,IAAI,IAAI,CACrE,CAAC;QACF,MAAM,UAAU,GAAG,WAAW,CAAC,GAAG,CAAC,CAAC,IAAI,EAAE,EAAE,CAAC,GAAG,YAAY,CAAC,IAAI,CAAC,GAAG,CAAC,CAAC;QACvE,MAAM,IAAI,GAAG;IACf,UAAU,CAAC,IAAI,CAAC,IAAI,CAAC;;IAErB,UAAU,CAAC,IAAI,CAAC,IAAI,CAAC;;GAEtB,CAAC;QACE,OAAO,CAAC,GAAG,CAAC,IAAI,CAAC,CAAC;IACpB,CAAC,CAAC,CAAC;AACP,CAAC;AAED,MAAM,UAAU,MAAM;IACpB,OAAO,IAAI,CAAC,MAAM,EAAE,CAAC;AACvB,CAAC;AAED,MAAM,UAAU,IAAI,CAAC,CAAa;IAChC,OAAO,IAAI,CAAC,IAAI,CAAC,CAAC,CAAC,CAAC;AACtB,CAAC"}
package/lib/kernels.d.ts DELETED
@@ -1 +0,0 @@
1
- export {};
package/lib/kernels.js DELETED
@@ -1,26 +0,0 @@
1
- import { registerKernel } from '@tensorflow/tfjs-core/dist/kernel_registry';
2
- import { mirrorPadConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/MirrorPad';
3
- import { padV2Config } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/PadV2';
4
- import { sliceConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/Slice';
5
- import { fusedConv2DConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/FusedConv2D';
6
- import { maxPoolConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/MaxPool';
7
- import { resizeNearestNeighborConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/ResizeNearestNeighbor';
8
- import { concatConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/Concat';
9
- import { identityConfig } from '@tensorflow/tfjs-backend-webgpu/dist/kernels/Identity';
10
- const kernelConfigs = [
11
- mirrorPadConfig,
12
- padV2Config,
13
- sliceConfig,
14
- fusedConv2DConfig,
15
- maxPoolConfig,
16
- resizeNearestNeighborConfig,
17
- concatConfig,
18
- identityConfig
19
- ];
20
- for (const kernelConfig of kernelConfigs) {
21
- registerKernel({
22
- ...kernelConfig,
23
- backendName: 'webgpu-oidn'
24
- });
25
- }
26
- //# sourceMappingURL=kernels.js.map
@@ -1 +0,0 @@
1
- {"version":3,"file":"kernels.js","sourceRoot":"","sources":["../src/kernels.ts"],"names":[],"mappings":"AAAA,OAAO,EAEL,cAAc,EACf,MAAM,4CAA4C,CAAC;AAEpD,OAAO,EAAE,eAAe,EAAE,MAAM,wDAAwD,CAAC;AACzF,OAAO,EAAE,WAAW,EAAE,MAAM,oDAAoD,CAAC;AACjF,OAAO,EAAE,WAAW,EAAE,MAAM,oDAAoD,CAAC;AACjF,OAAO,EAAE,iBAAiB,EAAE,MAAM,0DAA0D,CAAC;AAC7F,OAAO,EAAE,aAAa,EAAE,MAAM,sDAAsD,CAAC;AACrF,OAAO,EAAE,2BAA2B,EAAE,MAAM,oEAAoE,CAAC;AACjH,OAAO,EAAE,YAAY,EAAE,MAAM,qDAAqD,CAAC;AACnF,OAAO,EAAE,cAAc,EAAE,MAAM,uDAAuD,CAAC;AAEvF,MAAM,aAAa,GAAmB;IACpC,eAAe;IACf,WAAW;IACX,WAAW;IACX,iBAAiB;IACjB,aAAa;IACb,2BAA2B;IAC3B,YAAY;IACZ,cAAc;CACf,CAAC;AAEF,KAAK,MAAM,YAAY,IAAI,aAAa,EAAE,CAAC;IACzC,cAAc,CAAC;QACb,GAAG,YAAY;QACf,WAAW,EAAE,aAAa;KAC3B,CAAC,CAAC;AACL,CAAC"}