Skip to content

Commit a906452

Browse files
feat: add Node WASM CLI
1 parent 7bdb73a commit a906452

4 files changed

Lines changed: 290 additions & 0 deletions

File tree

README.md

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,18 @@ pnpm cli-separate --name htdemucs_ft --two-stems bass --two-stems-mix minus data
4141

4242
This creates `bass.wav` and `no_bass.wav`.
4343

44+
Run the same separation flow through the Rust/WASM driver and
45+
`onnxruntime-web`'s Node WASM runtime with:
46+
47+
```bash
48+
pnpm build-wasm
49+
pnpm wasm-separate --models data/onnx-lean data/input/song.wav data/output/song
50+
```
51+
52+
The WASM CLI accepts the same `--name`, `--two-stems`, `--two-stems-mix`, and
53+
`--shifts` options as the native CLI. Its WAV decoder currently requires 44.1
54+
kHz PCM or float input.
55+
4456
| Option | Explanation |
4557
| ------------------------------ | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- |
4658
| `--name htdemucs\|htdemucs_ft` | Chooses the standard general-purpose model or the fine-tuned source-specialist models. |

package.json

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
"verify-parity": "uv run python tools/model-export/verify_parity.py",
1111
"model-release": "uv run --no-sync tools/model_release.py",
1212
"build-wasm": "wasm-pack build crates/wasm --target web --release",
13+
"wasm-separate": "node tools/wasm-cli.mjs separate",
1314
"build": "pnpm build-wasm && pnpm -C packages/app build",
1415
"build-cf": "(command -v cargo >/dev/null || curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --profile minimal) && PATH=\"$HOME/.cargo/bin:$PATH\" pnpm build",
1516
"cli-separate": "cargo run --release -p demucs-cli -- separate --models data/onnx-lean",
@@ -19,6 +20,7 @@
1920
"test-e2e": "pnpm -C packages/app test-e2e"
2021
},
2122
"devDependencies": {
23+
"onnxruntime-web": "^1.27.0",
2224
"vite-plus": "^0.2.4",
2325
"wasm-pack": "0.15.0",
2426
"wrangler": "4.110.0"

pnpm-lock.yaml

Lines changed: 3 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

tools/wasm-cli.mjs

Lines changed: 273 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,273 @@
1+
import { mkdir, readFile, writeFile } from "node:fs/promises";
2+
import { basename, join } from "node:path";
3+
import { parseArgs } from "node:util";
4+
import * as ort from "onnxruntime-web";
5+
import init, {
6+
separate as separateWasm,
7+
} from "../crates/wasm/pkg/demucs_wasm.js";
8+
9+
const SAMPLE_RATE = 44_100;
10+
const SEGMENT = 343_980;
11+
const INPUT_LENGTH = 2 * SEGMENT;
12+
const OUTPUT_LENGTH = 4 * 2 * SEGMENT;
13+
const SOURCES = ["drums", "bass", "other", "vocals"];
14+
15+
function usage() {
16+
console.error(`Usage: pnpm wasm-separate [OPTIONS] <INPUT.WAV> <OUT_DIR>
17+
18+
Options:
19+
--models <DIR> Directory containing ONNX models (required)
20+
--name <MODEL> htdemucs or htdemucs_ft (default: htdemucs)
21+
--two-stems <SOURCE> drums, bass, other, or vocals
22+
--two-stems-mix <METHOD> add or minus
23+
--shifts <N> Number of processing passes (default: 1)`);
24+
}
25+
26+
function parseCli() {
27+
const command = process.argv[2];
28+
if (command !== "separate") {
29+
usage();
30+
process.exit(2);
31+
}
32+
let parsed;
33+
try {
34+
parsed = parseArgs({
35+
args: process.argv.slice(3),
36+
allowPositionals: true,
37+
options: {
38+
models: { type: "string" },
39+
name: { type: "string", default: "htdemucs" },
40+
"two-stems": { type: "string" },
41+
"two-stems-mix": { type: "string" },
42+
shifts: { type: "string", default: "1" },
43+
help: { type: "boolean", short: "h" },
44+
},
45+
});
46+
} catch (error) {
47+
console.error(String(error));
48+
usage();
49+
process.exit(2);
50+
}
51+
if (parsed.values.help) {
52+
usage();
53+
process.exit(0);
54+
}
55+
const { models, name, shifts: shiftsText } = parsed.values;
56+
const [input, outDir] = parsed.positionals;
57+
const twoStems = parsed.values["two-stems"];
58+
const method = parsed.values["two-stems-mix"];
59+
const shifts = Number(shiftsText);
60+
if (!models || !input || !outDir || parsed.positionals.length !== 2) {
61+
usage();
62+
process.exit(2);
63+
}
64+
if (!["htdemucs", "htdemucs_ft"].includes(name)) {
65+
throw new Error(`unknown model ${name}`);
66+
}
67+
if (twoStems && !SOURCES.includes(twoStems)) {
68+
throw new Error(`unknown source ${twoStems}`);
69+
}
70+
if (method && !["add", "minus"].includes(method)) {
71+
throw new Error(`unknown two-stems mix ${method}`);
72+
}
73+
if (!Number.isSafeInteger(shifts) || shifts < 1) {
74+
throw new Error("shifts must be an integer >= 1");
75+
}
76+
return { models, name, twoStems, method, shifts, input, outDir };
77+
}
78+
79+
function decodeWav(bytes) {
80+
const view = new DataView(bytes.buffer, bytes.byteOffset, bytes.byteLength);
81+
if (readFourCc(view, 0) !== "RIFF" || readFourCc(view, 8) !== "WAVE") {
82+
throw new Error("expected a RIFF/WAVE file");
83+
}
84+
let format;
85+
let dataOffset;
86+
let dataLength;
87+
for (let offset = 12; offset + 8 <= view.byteLength; ) {
88+
const id = readFourCc(view, offset);
89+
const length = view.getUint32(offset + 4, true);
90+
if (id === "fmt ") {
91+
format = {
92+
encoding: view.getUint16(offset + 8, true),
93+
channels: view.getUint16(offset + 10, true),
94+
sampleRate: view.getUint32(offset + 12, true),
95+
blockAlign: view.getUint16(offset + 20, true),
96+
bits: view.getUint16(offset + 22, true),
97+
};
98+
} else if (id === "data") {
99+
dataOffset = offset + 8;
100+
dataLength = length;
101+
}
102+
offset += 8 + length + (length & 1);
103+
}
104+
if (!format || dataOffset === undefined || dataLength === undefined) {
105+
throw new Error("WAV is missing fmt or data chunk");
106+
}
107+
if (format.sampleRate !== SAMPLE_RATE) {
108+
throw new Error(
109+
`expected ${SAMPLE_RATE}Hz WAV, got ${format.sampleRate}Hz`,
110+
);
111+
}
112+
if (format.channels < 1) {
113+
throw new Error("WAV has no channels");
114+
}
115+
const frames = Math.floor(dataLength / format.blockAlign);
116+
const left = new Float32Array(frames);
117+
const right = new Float32Array(frames);
118+
for (let frame = 0; frame < frames; frame++) {
119+
const offset = dataOffset + frame * format.blockAlign;
120+
left[frame] = readSample(view, offset, format);
121+
right[frame] =
122+
format.channels === 1
123+
? left[frame]
124+
: readSample(view, offset + format.bits / 8, format);
125+
}
126+
return { left, right };
127+
}
128+
129+
function readFourCc(view, offset) {
130+
return String.fromCharCode(
131+
view.getUint8(offset),
132+
view.getUint8(offset + 1),
133+
view.getUint8(offset + 2),
134+
view.getUint8(offset + 3),
135+
);
136+
}
137+
138+
function readSample(view, offset, format) {
139+
if (format.encoding === 3 && format.bits === 32) {
140+
return view.getFloat32(offset, true);
141+
}
142+
if (format.encoding !== 1) {
143+
throw new Error(`unsupported WAV encoding ${format.encoding}`);
144+
}
145+
switch (format.bits) {
146+
case 8:
147+
return (view.getUint8(offset) - 128) / 128;
148+
case 16:
149+
return view.getInt16(offset, true) / 32_768;
150+
case 24: {
151+
let value =
152+
view.getUint8(offset) |
153+
(view.getUint8(offset + 1) << 8) |
154+
(view.getUint8(offset + 2) << 16);
155+
if (value & 0x80_0000) {
156+
value |= 0xff00_0000;
157+
}
158+
return value / 8_388_608;
159+
}
160+
case 32:
161+
return view.getInt32(offset, true) / 2_147_483_648;
162+
default:
163+
throw new Error(`unsupported PCM bit depth ${format.bits}`);
164+
}
165+
}
166+
167+
function encodeWav(left, right) {
168+
const buffer = new ArrayBuffer(44 + left.length * 8);
169+
const view = new DataView(buffer);
170+
writeFourCc(view, 0, "RIFF");
171+
view.setUint32(4, buffer.byteLength - 8, true);
172+
writeFourCc(view, 8, "WAVE");
173+
writeFourCc(view, 12, "fmt ");
174+
view.setUint32(16, 16, true);
175+
view.setUint16(20, 3, true);
176+
view.setUint16(22, 2, true);
177+
view.setUint32(24, SAMPLE_RATE, true);
178+
view.setUint32(28, SAMPLE_RATE * 8, true);
179+
view.setUint16(32, 8, true);
180+
view.setUint16(34, 32, true);
181+
writeFourCc(view, 36, "data");
182+
view.setUint32(40, left.length * 8, true);
183+
for (let i = 0; i < left.length; i++) {
184+
view.setFloat32(44 + i * 8, left[i], true);
185+
view.setFloat32(48 + i * 8, right[i], true);
186+
}
187+
return new Uint8Array(buffer);
188+
}
189+
190+
function writeFourCc(view, offset, value) {
191+
for (let i = 0; i < 4; i++) {
192+
view.setUint8(offset + i, value.charCodeAt(i));
193+
}
194+
}
195+
196+
async function main() {
197+
const args = parseCli();
198+
const inputBytes = await readFile(args.input);
199+
const { left, right } = decodeWav(inputBytes);
200+
console.error(
201+
`input: ${left.length} samples (${(left.length / SAMPLE_RATE).toFixed(2)}s) | model ${args.name} | shifts ${args.shifts}`,
202+
);
203+
204+
const wasmBytes = await readFile(
205+
new URL("../crates/wasm/pkg/demucs_wasm_bg.wasm", import.meta.url),
206+
);
207+
const wasm = await init({ module_or_path: wasmBytes });
208+
let dft;
209+
const host = {
210+
event(type, ...event) {
211+
if (type === "model-loading") {
212+
console.error(`loading ${event[3]}`);
213+
}
214+
if (type === "model-complete") {
215+
console.error("model complete");
216+
}
217+
},
218+
async initialize() {
219+
dft = await readFile(join(args.models, "dft.bin"));
220+
},
221+
async loadModel(model, source) {
222+
const file = source ? `${model}_${source}.onnx` : `${model}.onnx`;
223+
const modelBytes = await readFile(join(args.models, file));
224+
return ort.InferenceSession.create(modelBytes, {
225+
executionProviders: ["wasm"],
226+
externalData: [{ data: dft, path: "dft.bin" }],
227+
});
228+
},
229+
async runModel(session, inputPtr, outputPtr) {
230+
const input = new Float32Array(
231+
wasm.memory.buffer,
232+
inputPtr,
233+
INPUT_LENGTH,
234+
);
235+
const result = await session.run({
236+
input: new ort.Tensor("float32", input, [1, 2, SEGMENT]),
237+
});
238+
const output = new Float32Array(
239+
wasm.memory.buffer,
240+
outputPtr,
241+
OUTPUT_LENGTH,
242+
);
243+
output.set(result.output.data);
244+
},
245+
async releaseModel(session) {
246+
await session.release();
247+
},
248+
};
249+
250+
const tracks = await separateWasm(
251+
args.name,
252+
args.twoStems,
253+
args.method,
254+
args.shifts,
255+
left,
256+
right,
257+
host,
258+
);
259+
const names = args.twoStems
260+
? [args.twoStems, `no_${args.twoStems}`]
261+
: SOURCES;
262+
await mkdir(args.outDir, { recursive: true });
263+
for (const [index, name] of names.entries()) {
264+
const path = join(args.outDir, `${name}.wav`);
265+
await writeFile(path, encodeWav(tracks[index * 2], tracks[index * 2 + 1]));
266+
console.error(`wrote ${path}`);
267+
}
268+
}
269+
270+
main().catch((error) => {
271+
console.error(`${basename(process.argv[1])}: ${error.message ?? error}`);
272+
process.exitCode = 1;
273+
});

0 commit comments

Comments
 (0)