Skip to content

Commit fb17a3c

Browse files
hi-ogawaOpenCode
andauthored
feat: add native and web benchmark harness (#68)
Co-authored-by: Hiroshi Ogawa <4232207+hi-ogawa@users.noreply.github.com> Co-authored-by: OpenCode <noreply@opencode.ai>
1 parent b45f838 commit fb17a3c

18 files changed

Lines changed: 645 additions & 19 deletions

File tree

Cargo.lock

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

README.md

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -70,6 +70,36 @@ pnpm build-model htdemucs
7070

7171
Run `pnpm build-model --all` to build the standard model and all fine-tuned specialists.
7272

73+
## Benchmark
74+
75+
Compare native ONNX Runtime inference with Chromium's ONNX Runtime WASM backend using a deterministic 30-second workload:
76+
77+
```bash
78+
mkdir -p data/benchmark
79+
ffmpeg -f lavfi -i "sine=frequency=440:sample_rate=44100:duration=30" \
80+
-filter_complex "[0:a]asplit=2[left][right];[left][right]join=inputs=2:channel_layout=stereo" \
81+
-c:a pcm_f32le -y data/benchmark/input-30s.wav
82+
cargo build --release -p demucs-cli
83+
pnpm build-wasm
84+
for threads in 0 1 2 4 8 16; do
85+
label=$threads
86+
test "$threads" = 0 && label=default
87+
for run in 0 1 2 3; do
88+
rm -rf "data/benchmark/native-threads-$label-run-$run"
89+
target/release/demucs separate \
90+
--models data/onnx-lean \
91+
--threads "$threads" \
92+
--timings-json "data/benchmark/native-threads-$label-run-$run.json" \
93+
data/benchmark/input-30s.wav \
94+
"data/benchmark/native-threads-$label-run-$run"
95+
done
96+
done
97+
pnpm -C packages/app benchmark
98+
pnpm tsx tools/benchmark-summary.ts
99+
```
100+
101+
See [`docs/benchmark.md`](docs/benchmark.md) for the workload, timing boundaries, and generated results.
102+
73103
## Web App
74104

75105
Build the WASM binding and start the fully client-side app:

crates/cli/Cargo.toml

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,8 @@ clap = { version = "4", features = ["derive"] }
1313
console = "0.16"
1414
demucs-core = { path = "../core" }
1515
indicatif = "0.18"
16+
serde = { version = "1", features = ["derive"] }
17+
serde_json = "1"
1618
# default features with tls-native swapped for tls-rustls, so the binary downloader
1719
# doesn't need system OpenSSL headers
1820
ort = { version = "2.0.0-rc.10", default-features = false, features = [

crates/cli/src/main.rs

Lines changed: 73 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@ use clap::{Args, Parser, Subcommand};
44
use console::style;
55
use demucs_core as core;
66
use indicatif::{MultiProgress, ProgressBar, ProgressStyle};
7+
use serde::Serialize;
78
use std::path::PathBuf;
89
use std::time::{Duration, Instant};
910

@@ -69,6 +70,14 @@ struct SeparateArgs {
6970
)]
7071
shifts: u32,
7172

73+
/// ONNX Runtime intra-op threads; 0 uses the runtime default
74+
#[arg(long, default_value_t = 4, value_name = "N", hide_default_value = true)]
75+
threads: usize,
76+
77+
/// Write machine-readable phase timings to this JSON file
78+
#[arg(long, value_name = "FILE")]
79+
timings_json: Option<PathBuf>,
80+
7281
/// Input WAV file
7382
#[arg(value_name = "INPUT.WAV")]
7483
input: String,
@@ -112,9 +121,14 @@ fn separate(args: SeparateArgs) -> Result<()> {
112121

113122
eprintln!("prepared audio in {}", format_duration(prepare_elapsed));
114123
let mut progress = CliProgress::new();
115-
let outputs = ort_driver::run_all(&args.models_dir, &members, wav, opts, |event| {
116-
progress.update(event)
117-
})?;
124+
let outputs = ort_driver::run_all(
125+
&args.models_dir,
126+
&members,
127+
wav,
128+
opts,
129+
args.threads,
130+
|event| progress.update(event),
131+
)?;
118132
progress.finish();
119133

120134
let named: Vec<(String, [Vec<f32>; core::CHANNELS])> = match outputs {
@@ -150,9 +164,50 @@ fn separate(args: SeparateArgs) -> Result<()> {
150164
format_duration(progress.finalize_elapsed),
151165
format_duration(write_elapsed),
152166
);
167+
if let Some(path) = args.timings_json {
168+
let timings = Timings {
169+
prepare_ms: millis(prepare_elapsed),
170+
load_ms: millis(progress.load_elapsed),
171+
inference_ms: millis(progress.inference_elapsed),
172+
chunks: progress.chunks,
173+
finalize_ms: millis(progress.finalize_elapsed),
174+
write_ms: millis(write_elapsed),
175+
total_ms: millis(total_elapsed),
176+
};
177+
std::fs::write(&path, serde_json::to_vec_pretty(&timings)?)
178+
.with_context(|| format!("write {}", path.display()))?;
179+
}
153180
Ok(())
154181
}
155182

183+
#[derive(Serialize)]
184+
#[serde(rename_all = "camelCase")]
185+
struct Timings {
186+
prepare_ms: f64,
187+
load_ms: f64,
188+
inference_ms: f64,
189+
chunks: Vec<ChunkTiming>,
190+
finalize_ms: f64,
191+
write_ms: f64,
192+
total_ms: f64,
193+
}
194+
195+
#[derive(Serialize)]
196+
#[serde(rename_all = "camelCase")]
197+
struct ChunkTiming {
198+
member: usize,
199+
shift: usize,
200+
chunk: usize,
201+
prepare_input_ms: f64,
202+
ort_run_ms: f64,
203+
output_copy_ms: f64,
204+
process_output_ms: f64,
205+
}
206+
207+
fn millis(duration: Duration) -> f64 {
208+
duration.as_secs_f64() * 1000.0
209+
}
210+
156211
// Interactive layouts stay model-oriented while the overall row remains stable:
157212
// htdemucs.onnx
158213
// Load done | 5.2s
@@ -165,6 +220,7 @@ struct CliProgress {
165220
phase: Option<ProgressBar>,
166221
load_elapsed: Duration,
167222
inference_elapsed: Duration,
223+
chunks: Vec<ChunkTiming>,
168224
finalize_elapsed: Duration,
169225
loaded: usize,
170226
eta: Option<Duration>,
@@ -181,6 +237,7 @@ impl CliProgress {
181237
phase: None,
182238
load_elapsed: Duration::ZERO,
183239
inference_elapsed: Duration::ZERO,
240+
chunks: Vec::new(),
184241
finalize_elapsed: Duration::ZERO,
185242
loaded: 0,
186243
eta: None,
@@ -257,9 +314,22 @@ impl CliProgress {
257314
shift,
258315
shifts,
259316
member_done,
317+
member,
260318
elapsed,
319+
prepare_elapsed,
320+
run_elapsed,
321+
process_elapsed,
261322
} => {
262323
self.inference_elapsed += elapsed;
324+
self.chunks.push(ChunkTiming {
325+
member,
326+
shift,
327+
chunk: member_done,
328+
prepare_input_ms: millis(prepare_elapsed),
329+
ort_run_ms: millis(run_elapsed),
330+
output_copy_ms: 0.0,
331+
process_output_ms: millis(process_elapsed),
332+
});
263333
self.overall.set_length(total as u64);
264334
self.overall.set_position(done as u64);
265335
if let Some(phase) = &self.phase {

crates/cli/src/ort_driver.rs

Lines changed: 20 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -22,10 +22,14 @@ pub enum Progress<'a> {
2222
done: usize,
2323
total: usize,
2424
members: usize,
25+
member: usize,
2526
shift: usize,
2627
shifts: usize,
2728
member_done: usize,
2829
elapsed: Duration,
30+
prepare_elapsed: Duration,
31+
run_elapsed: Duration,
32+
process_elapsed: Duration,
2933
},
3034
MemberFinished {
3135
chunks: usize,
@@ -57,6 +61,7 @@ pub fn run_all(
5761
members: &[core::vocab::Member],
5862
wav: [Vec<f32>; core::CHANNELS],
5963
opts: core::Options,
64+
threads: usize,
6065
mut on_progress: impl FnMut(Progress<'_>),
6166
) -> Result<core::Outputs> {
6267
let mut separation = core::Separation::new(wav, opts)?;
@@ -90,10 +95,11 @@ pub fn run_all(
9095
});
9196
let path = models_dir.join(file);
9297
let load_started = Instant::now();
93-
let mut session = ort::session::Session::builder()
94-
.map_err(ort_err)?
95-
.with_intra_threads(4)
96-
.map_err(ort_err)?
98+
let mut builder = ort::session::Session::builder().map_err(ort_err)?;
99+
if threads > 0 {
100+
builder = builder.with_intra_threads(threads).map_err(ort_err)?;
101+
}
102+
let mut session = builder
97103
.commit_from_file(&path)
98104
.map_err(ort_err)
99105
.with_context(|| format!("load {}", path.display()))?;
@@ -107,31 +113,41 @@ pub fn run_all(
107113
let mut chunk_processor = separation.plan.create_chunk_processor(shift);
108114
for &chunk in &shift.chunks {
109115
let chunk_started = Instant::now();
116+
let prepare_started = Instant::now();
110117
chunk_processor.prepare_input(chunk, &mut input)?;
111118
let value = ort::value::TensorRef::from_array_view((
112119
[1usize, core::CHANNELS, core::SEGMENT],
113120
input.as_slice(),
114121
))
115122
.map_err(ort_err)?;
123+
let prepare_elapsed = prepare_started.elapsed();
124+
let run_started = Instant::now();
116125
let run = session
117126
.run(ort::inputs!["input" => value])
118127
.map_err(ort_err)?;
128+
let run_elapsed = run_started.elapsed();
129+
let process_started = Instant::now();
119130
let (shape, data) = run["output"].try_extract_tensor::<f32>().map_err(ort_err)?;
120131
let dims: Vec<usize> = shape.iter().map(|&d| d as usize).collect();
121132
if dims != [1, core::NUM_SOURCES, core::CHANNELS, core::SEGMENT] {
122133
bail!("unexpected output shape {dims:?}");
123134
}
124135
chunk_processor.process_output(chunk, data)?;
136+
let process_elapsed = process_started.elapsed();
125137
done += 1;
126138
member_done += 1;
127139
on_progress(Progress::Inference {
128140
done,
129141
total,
130142
members: members.len(),
143+
member: member_index + 1,
131144
shift: shift_index + 1,
132145
shifts: member_plan.shifts.len(),
133146
member_done,
134147
elapsed: chunk_started.elapsed(),
148+
prepare_elapsed,
149+
run_elapsed,
150+
process_elapsed,
135151
});
136152
}
137153
shift_merger.add(chunk_processor.finish());

0 commit comments

Comments
 (0)