oidn-web 0.3.4 → 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.
- package/CHANGELOG.md +49 -0
- package/README.md +140 -4
- package/benchmarks/compare.mjs +651 -0
- package/benchmarks/leak.mjs +255 -0
- package/benchmarks/results/before-spatial.json +391 -0
- package/benchmarks/results/before-spatial.md +47 -0
- package/benchmarks/results/int8-scan.json +2007 -0
- package/benchmarks/results/int8-scan.md +160 -0
- package/benchmarks/results/int8-w8a8-scan.json +2007 -0
- package/benchmarks/results/int8-w8a8-scan.md +160 -0
- package/benchmarks/results/int8-weight-channel.json +1413 -0
- package/benchmarks/results/int8-weight-channel.md +118 -0
- package/benchmarks/results/int8-weight-only.json +1437 -0
- package/benchmarks/results/int8-weight-only.md +118 -0
- package/benchmarks/results/kernel-webnn-final.json +1115 -0
- package/benchmarks/results/kernel-webnn-final.md +104 -0
- package/benchmarks/results/latest-optimized.json +375 -0
- package/benchmarks/results/latest-optimized.md +47 -0
- package/benchmarks/results/latest.json +391 -0
- package/benchmarks/results/latest.md +47 -0
- package/benchmarks/results/profile-baseline.json +331 -0
- package/benchmarks/results/profile-baseline.md +12 -0
- package/benchmarks/results/profile-conv2x.json +331 -0
- package/benchmarks/results/profile-conv2x.md +12 -0
- package/benchmarks/results/profile-fast-init.json +385 -0
- package/benchmarks/results/profile-fast-init.md +47 -0
- package/benchmarks/results/profile-fp16-fma.json +369 -0
- package/benchmarks/results/profile-fp16-fma.md +47 -0
- package/benchmarks/results/profile-fp16-tiled-decoder.json +369 -0
- package/benchmarks/results/profile-fp16-tiled-decoder.md +47 -0
- package/benchmarks/results/profile-fp16-tiled-encoder.json +369 -0
- package/benchmarks/results/profile-fp16-tiled-encoder.md +47 -0
- package/benchmarks/results/profile-fp16-unfused-pool.json +385 -0
- package/benchmarks/results/profile-fp16-unfused-pool.md +47 -0
- package/benchmarks/results/profile-input-major.json +347 -0
- package/benchmarks/results/profile-input-major.md +12 -0
- package/benchmarks/results/profile-k16.json +347 -0
- package/benchmarks/results/profile-k16.md +12 -0
- package/benchmarks/results/profile-k4.json +347 -0
- package/benchmarks/results/profile-k4.md +12 -0
- package/benchmarks/results/profile-pool-reuse.json +331 -0
- package/benchmarks/results/profile-pool-reuse.md +12 -0
- package/benchmarks/results/profile-precompiled.json +385 -0
- package/benchmarks/results/profile-precompiled.md +47 -0
- package/benchmarks/results/profile-static-channels.json +385 -0
- package/benchmarks/results/profile-static-channels.md +47 -0
- package/benchmarks/results/profile-static-io.json +385 -0
- package/benchmarks/results/profile-static-io.md +47 -0
- package/benchmarks/results/profile-tiled-conv.json +331 -0
- package/benchmarks/results/profile-tiled-conv.md +12 -0
- package/benchmarks/results/profile-tiled-decoder.json +347 -0
- package/benchmarks/results/profile-tiled-decoder.md +12 -0
- package/benchmarks/results/profile-tiled-matmul.json +331 -0
- package/benchmarks/results/profile-tiled-matmul.md +12 -0
- package/benchmarks/results/profile-unfused-decoder.json +379 -0
- package/benchmarks/results/profile-unfused-decoder.md +12 -0
- package/benchmarks/results/profile-unfused-pool.json +347 -0
- package/benchmarks/results/profile-unfused-pool.md +12 -0
- package/benchmarks/results/spatial-auto.json +575 -0
- package/benchmarks/results/spatial-auto.md +61 -0
- package/benchmarks/results/subgroup-smoke.json +1094 -0
- package/benchmarks/results/subgroup-smoke.md +104 -0
- package/benchmarks/results/webnn-smoke.json +739 -0
- package/benchmarks/results/webnn-smoke.md +76 -0
- package/dist/oidn.js +4189 -22603
- package/dist/oidn.umd.cjs +784 -5807
- package/lib/UNet.d.ts +66 -15
- package/lib/UNet.js +162 -257
- package/lib/UNet.js.map +1 -1
- package/lib/WGPUComputePass.d.ts +1 -1
- package/lib/WGPUComputePass.js +6 -4
- package/lib/WGPUComputePass.js.map +1 -1
- package/lib/backend.d.ts +8 -4
- package/lib/backend.js +36 -44
- package/lib/backend.js.map +1 -1
- package/lib/graphOptimizer.d.ts +54 -0
- package/lib/graphOptimizer.js +216 -0
- package/lib/graphOptimizer.js.map +1 -0
- package/lib/main.d.ts +33 -10
- package/lib/main.js +5 -0
- package/lib/main.js.map +1 -1
- package/lib/modelSpec.d.ts +80 -0
- package/lib/modelSpec.js +270 -0
- package/lib/modelSpec.js.map +1 -0
- package/lib/nativeUNet.d.ts +67 -0
- package/lib/nativeUNet.js +1735 -0
- package/lib/nativeUNet.js.map +1 -0
- package/lib/process.js +38 -35
- package/lib/process.js.map +1 -1
- package/lib/resourceTracker.d.ts +26 -0
- package/lib/resourceTracker.js +65 -0
- package/lib/resourceTracker.js.map +1 -0
- package/lib/tileScheduler.d.ts +33 -0
- package/lib/tileScheduler.js +86 -0
- package/lib/tileScheduler.js.map +1 -0
- package/lib/webnnUNet.d.ts +52 -0
- package/lib/webnnUNet.js +535 -0
- package/lib/webnnUNet.js.map +1 -0
- package/package.json +9 -5
- package/scripts/inspect-model.mjs +64 -0
- package/src/UNet.ts +236 -339
- package/src/WGPUComputePass.ts +6 -4
- package/src/backend.ts +42 -55
- package/src/graphOptimizer.ts +301 -0
- package/src/main.ts +71 -11
- package/src/modelSpec.ts +414 -0
- package/src/nativeUNet.ts +2256 -0
- package/src/process.ts +38 -36
- package/src/resourceTracker.ts +94 -0
- package/src/tileScheduler.ts +138 -0
- package/src/webnnUNet.ts +812 -0
- package/tests/modelSpec.test.mjs +128 -0
- package/tests/resourceLifecycle.test.mjs +383 -0
- package/tests/tileScheduler.test.mjs +90 -0
- package/lib/helper.d.ts +0 -4
- package/lib/helper.js +0 -33
- package/lib/helper.js.map +0 -1
- package/lib/kernels.d.ts +0 -1
- package/lib/kernels.js +0 -26
- package/lib/kernels.js.map +0 -1
- package/src/helper.ts +0 -43
- 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
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
|
package/lib/kernels.js.map
DELETED
|
@@ -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"}
|