@@ -4,6 +4,7 @@ use clap::{Args, Parser, Subcommand};
44use console:: style;
55use demucs_core as core;
66use indicatif:: { MultiProgress , ProgressBar , ProgressStyle } ;
7+ use serde:: Serialize ;
78use std:: path:: PathBuf ;
89use 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 {
0 commit comments