diff --git a/Cargo.lock b/Cargo.lock index 95787055..a87a6e0e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -888,6 +888,12 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" +[[package]] +name = "hermit-abi" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" + [[package]] name = "hf-hub" version = "0.5.0" @@ -1376,7 +1382,10 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a06de3016e9fae57a36fd14dba131fccf49f74b40b7fbdb472f96e361ec71a08" dependencies = [ "autocfg", + "num_cpus", + "once_cell", "rawpointer", + "thread-tree", ] [[package]] @@ -1540,6 +1549,16 @@ dependencies = [ "autocfg", ] +[[package]] +name = "num_cpus" +version = "1.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91df4bbde75afed763b708b7eee1e8e7651e02d97f6d5dd763e89367e957b23b" +dependencies = [ + "hermit-abi", + "libc", +] + [[package]] name = "objc2" version = "0.6.4" @@ -2582,6 +2601,15 @@ dependencies = [ "syn", ] +[[package]] +name = "thread-tree" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ffbd370cb847953a25954d9f63e14824a36113f8c72eecf6eccef5dc4b45d630" +dependencies = [ + "crossbeam-channel", +] + [[package]] name = "thread_local" version = "1.1.9" diff --git a/Cargo.toml b/Cargo.toml index 46c765f5..1ae14f97 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -40,7 +40,7 @@ hf-hub = { version = "0.5", optional = true, default-features = false } # math kodama = "0.3.0" -ndarray = "0.17.2" +ndarray = { version = "0.17.2", features = ["matrixmultiply-threading"] } ndarray-linalg-mkl = { package = "ndarray-linalg", version = "0.18.1", features = ["intel-mkl-static"], optional = true } ndarray-linalg-static = { package = "ndarray-linalg", version = "0.18.1", features = ["openblas-static"], optional = true } ndarray-linalg-system = { package = "ndarray-linalg", version = "0.18.1", features = ["openblas-system"], optional = true } diff --git a/scripts/export_models.py b/scripts/export_models.py index e6d01acd..1d1e361e 100644 --- a/scripts/export_models.py +++ b/scripts/export_models.py @@ -6,6 +6,7 @@ # "numpy", # "onnx", # "onnxscript", +# "onnxsim", # ] # /// """Download and export ONNX models + PLDA params for speakrs. @@ -66,6 +67,18 @@ def main() -> None: print("Done!") +def fold_onnx_graph(path: str) -> None: + """Constant-fold a graph in place; a folding that changes outputs is a hard error.""" + import onnx + from onnxsim import simplify + + model = onnx.load(path) + simplified, ok = simplify(model) + if not ok: + raise RuntimeError(f"onnxsim could not validate the simplified graph for {path}") + onnx.save(simplified, path) + + def export_segmentation(pipeline: Any, models_dir: str) -> None: print("Exporting segmentation model...") seg_model = pipeline._segmentation.model @@ -102,6 +115,14 @@ def export_segmentation(pipeline: Any, models_dir: str) -> None: dynamo=False, ) + # Constant-fold the exported graphs: SincNet synthesizes its filterbank from frozen + # parameters every forward (Sin/Cos/If subgraph). On ORT's CUDA EP those ops fall back + # to CPU with Memcpy nodes inserted, costing 2x per batch-32 in serving tests. Folding + # is bit-exact (max_abs_diff 0.0, argmax mismatch 0) and shrinks the graph 179->40 nodes. + fold_onnx_graph(os.path.join(models_dir, "segmentation-3.0.onnx")) + fold_onnx_graph(os.path.join(models_dir, "segmentation-3.0-b32.onnx")) + fold_onnx_graph(os.path.join(models_dir, "segmentation-3.0-b64.onnx")) + sz = os.path.getsize(os.path.join(models_dir, "segmentation-3.0.onnx")) / 1e6 print(f" segmentation-3.0.onnx ({sz:.1f} MB)") bsz = os.path.getsize(os.path.join(models_dir, "segmentation-3.0-b32.onnx")) / 1e6 diff --git a/src/clustering/ahc.rs b/src/clustering/ahc.rs index 007ffb37..ca79eb6c 100644 --- a/src/clustering/ahc.rs +++ b/src/clustering/ahc.rs @@ -26,33 +26,121 @@ pub fn cluster(embeddings: &ArrayView2, config: AhcConfig) -> Vec { } let normalized = l2_normalize_rows(embeddings); + let t0 = std::time::Instant::now(); let mut condensed = condensed_euclidean(&normalized); + let pdist_ms = t0.elapsed().as_millis(); + let t1 = std::time::Instant::now(); let dendrogram = linkage(&mut condensed, observations, Method::Centroid); - flat_clusters(observations, dendrogram.steps(), config.threshold) + let linkage_ms = t1.elapsed().as_millis(); + let t2 = std::time::Instant::now(); + let labels = flat_clusters(observations, dendrogram.steps(), config.threshold); + tracing::debug!( + observations, + pdist_ms, + linkage_ms, + flat_ms = t2.elapsed().as_millis(), + "AHC stage timing" + ); + labels } fn condensed_euclidean(embeddings: &Array2) -> Vec { + condensed_euclidean_with_workers(embeddings, pdist_worker_count()) +} + +fn condensed_euclidean_with_workers(embeddings: &Array2, workers: usize) -> Vec { + // diar-native patch: blocked Gram-matrix formulation with scoped threads. + // The original per-pair scalar loop cost 64.5 s at N=21k; dist^2 = |a|^2 + |b|^2 - 2ab + // via matmul blocks is ~20x faster and each block writes a disjoint contiguous + // range of the condensed vector, so blocks parallelize without locks. let observations = embeddings.nrows(); - let mut condensed = Vec::with_capacity(observations * (observations - 1) / 2); - for row in 0..observations.saturating_sub(1) { - for col in row + 1..observations { - let lhs = embeddings.row(row); - let rhs = embeddings.row(col); - let distance = lhs - .iter() - .zip(rhs.iter()) - .map(|(left, right)| { - let delta = left - right; - delta * delta - }) - .sum::() - .sqrt(); - condensed.push(distance); + if observations < 2 { + return Vec::new(); + } + let total = observations * (observations - 1) / 2; + let mut condensed = vec![0f32; total]; + let sq_norms: Vec = embeddings + .rows() + .into_iter() + .map(|row| row.dot(&row)) + .collect(); + + const BLOCK: usize = 1024; + // start offset of row i's segment in the condensed vector: + // sum_{r = Vec::new(); + { + let mut rest: &mut [f32] = &mut condensed; + let mut consumed = 0usize; + let mut bi = 0usize; + while bi < observations.saturating_sub(1) { + let bi_end = (bi + BLOCK).min(observations - 1); + let end_offset = seg_start(bi_end); + let (head, tail) = rest.split_at_mut(end_offset - consumed); + blocks.push((bi, bi_end, head)); + consumed = end_offset; + rest = tail; + bi = bi_end; } } + + // Bounded worker pool: one thread per block would scale with meeting length + // (n=21k => 21 threads, each driving its own multi-threaded BLAS `dot`), which + // oversubscribes small/shared hosts and keeps every block's Gram matrix alive at + // once. Workers pull blocks from a shared queue instead, so peak concurrency and + // peak scratch memory scale with core count, not with n. + let workers = workers.min(blocks.len()).max(1); + let queue = std::sync::Mutex::new(blocks); + std::thread::scope(|scope| { + for _ in 0..workers { + let queue = &queue; + let emb = &embeddings; + let norms = &sq_norms; + scope.spawn(move || { + loop { + let next = queue.lock().expect("pdist queue poisoned").pop(); + let Some((bi, bi_end, slice)) = next else { + break; + }; + let a = emb.slice(ndarray::s![bi..bi_end, ..]); + let b = emb.slice(ndarray::s![bi.., ..]); + let gram = a.dot(&b.t()); // (bi_end-bi) x (observations-bi) + let mut offset = 0usize; + for (local, i) in (bi..bi_end).enumerate() { + for j in (i + 1)..observations { + let dot = gram[[local, j - bi]]; + let d2 = (norms[i] + norms[j] - 2.0 * dot).max(0.0); + slice[offset] = d2.sqrt(); + offset += 1; + } + } + } + }); + } + }); condensed } +/// Number of concurrent workers used for the blocked pairwise-distance computation. +/// +/// Defaults to `available_parallelism()` capped at 8 (each worker also drives a +/// multi-threaded BLAS `dot`, so a higher cap oversubscribes rather than helps). +/// Override with `SPEAKRS_AHC_THREADS`; values are clamped to at least 1. +fn pdist_worker_count() -> usize { + std::env::var("SPEAKRS_AHC_THREADS") + .ok() + .and_then(|v| v.parse::().ok()) + .filter(|v| *v > 0) + .unwrap_or_else(|| { + std::thread::available_parallelism() + .map(|c| c.get().min(8)) + .unwrap_or(1) + }) +} + fn flat_clusters(observations: usize, steps: &[Step], threshold: f32) -> Vec { if observations == 0 { return Vec::new(); @@ -168,6 +256,39 @@ mod tests { .join(name) } + #[test] + fn condensed_euclidean_is_bit_identical_across_worker_counts() { + // 2600 rows => 3 blocks at BLOCK=1024, so worker counts below/at/above the + // block count all get exercised. + let rows = 2600; + let cols = 16; + let data: Vec = (0..rows * cols) + .map(|i| ((i * 37 % 101) as f32 / 101.0) - 0.5) + .collect(); + let embeddings = Array2::from_shape_vec((rows, cols), data).unwrap(); + + let reference = condensed_euclidean_with_workers(&embeddings, 1); + for workers in [2, 3, 8, 64] { + let got = condensed_euclidean_with_workers(&embeddings, workers); + assert_eq!(got.len(), reference.len()); + assert!( + got.iter() + .zip(reference.iter()) + .all(|(a, b)| a.to_bits() == b.to_bits()), + "worker count {workers} changed pdist output" + ); + } + } + + #[test] + fn pdist_worker_count_is_bounded() { + let workers = pdist_worker_count(); + assert!(workers >= 1); + if std::env::var_os("SPEAKRS_AHC_THREADS").is_none() { + assert!(workers <= 8, "default worker count should stay bounded"); + } + } + #[test] fn separates_two_clusters() { let embeddings = array![[1.0, 0.0], [0.95, 0.05], [-1.0, 0.0], [-0.95, -0.05],]; diff --git a/src/clustering/vbx.rs b/src/clustering/vbx.rs index 100b53af..1933960a 100644 --- a/src/clustering/vbx.rs +++ b/src/clustering/vbx.rs @@ -76,65 +76,59 @@ pub fn vbx( // m-step: compute speaker models // invL[k,d] = 1.0 / (1 + Fa/Fb * N_k * Phi[d]) // alpha[k,d] = Fa/Fb * invL[k,d] * sum_t(gamma[t,k] * rho[t,d]) + // diar-native patch: vectorized — the original per-element loops cost + // O(N*K*D) scalar work per iteration (305 s at N=21k, K=1.9k, D=128). let n_k: Array1 = gamma.sum_axis(Axis(0)); - let mut inv_l = Array2::zeros((n_speakers, dim)); - let mut alpha = Array2::zeros((n_speakers, dim)); - - for speaker_idx in 0..n_speakers { - for dim_idx in 0..dim { - inv_l[[speaker_idx, dim_idx]] = - 1.0 / (1.0 + fa_over_fb * n_k[speaker_idx] * phi_f64[dim_idx]); - } - - // gamma.T @ rho for this speaker - let mut f_k = Array1::::zeros(dim); - for sample_idx in 0..n_samples { - f_k.scaled_add(gamma[[sample_idx, speaker_idx]], &rho.row(sample_idx)); - } - - for dim_idx in 0..dim { - alpha[[speaker_idx, dim_idx]] = - fa_over_fb * inv_l[[speaker_idx, dim_idx]] * f_k[dim_idx]; - } + let mut inv_l = Array2::::zeros((n_speakers, dim)); + for (speaker_idx, mut row) in inv_l.rows_mut().into_iter().enumerate() { + let scale = fa_over_fb * n_k[speaker_idx]; + row.assign(&phi_f64.mapv(|p| 1.0 / (1.0 + scale * p))); } + // f = gamma.T @ rho (K x D), alpha = Fa/Fb * invL ⊙ f + let f = gamma.t().dot(&rho); + let mut alpha = &inv_l * &f; + alpha.mapv_inplace(|v| v * fa_over_fb); + // e-step // log_p_[t,k] = Fa * (rho[t] . alpha[k] - 0.5 * (invL[k] + alpha[k]^2) . Phi + G[t]) - let mut log_p = Array2::::zeros((n_samples, n_speakers)); - for sample_idx in 0..n_samples { - for speaker_idx in 0..n_speakers { - let rho_dot_alpha: f64 = rho.row(sample_idx).dot(&alpha.row(speaker_idx)); - let penalty: f64 = (0..dim) - .map(|dim_idx| { - (inv_l[[speaker_idx, dim_idx]] - + alpha[[speaker_idx, dim_idx]] * alpha[[speaker_idx, dim_idx]]) - * phi_f64[dim_idx] - }) - .sum(); - log_p[[sample_idx, speaker_idx]] = - fa * (rho_dot_alpha - 0.5 * penalty + frame_constants[sample_idx]); - } + // penalty depends only on k — compute once per iteration, not per sample. + let penalty: Array1 = (0..n_speakers) + .map(|speaker_idx| { + inv_l + .row(speaker_idx) + .iter() + .zip(alpha.row(speaker_idx).iter()) + .zip(phi_f64.iter()) + .map(|((&il, &a), &p)| (il + a * a) * p) + .sum() + }) + .collect(); + + let mut log_p = rho.dot(&alpha.t()); // N x K + for (sample_idx, mut row) in log_p.rows_mut().into_iter().enumerate() { + let g = frame_constants[sample_idx]; + row.zip_mut_with(&penalty, |value, &pen| { + *value = fa * (*value - 0.5 * pen + g); + }); } - // GMM-style update with pi priors + // GMM-style update with pi priors (single fused pass per row) let lpi: Array1 = pi.mapv(|p| (p + 1e-8).ln()); - // log_p_x[sample_idx] = logsumexp(log_p[sample_idx] + lpi) let mut log_p_x = Array1::::zeros(n_samples); - for sample_idx in 0..n_samples { - scratch.assign(&log_p.row(sample_idx)); + for ((log_p_row, mut gamma_row), log_p_x_slot) in log_p + .rows() + .into_iter() + .zip(gamma.rows_mut()) + .zip(log_p_x.iter_mut()) + { + scratch.assign(&log_p_row); scratch += &lpi; - log_p_x[sample_idx] = logsumexp_f64(&scratch.view()); - } - - // gamma[sample_idx,speaker_idx] = exp(log_p[sample_idx,speaker_idx] + lpi[speaker_idx] - log_p_x[sample_idx]) - for sample_idx in 0..n_samples { - for speaker_idx in 0..n_speakers { - gamma[[sample_idx, speaker_idx]] = - (log_p[[sample_idx, speaker_idx]] + lpi[speaker_idx] - log_p_x[sample_idx]) - .exp(); - } + let lse = logsumexp_f64(&scratch.view()); + *log_p_x_slot = lse; + gamma_row.zip_mut_with(&scratch, |g, &s| *g = (s - lse).exp()); } // update pi diff --git a/src/inference/embedding.rs b/src/inference/embedding.rs index 2fd95127..ba5dd87c 100644 --- a/src/inference/embedding.rs +++ b/src/inference/embedding.rs @@ -74,6 +74,7 @@ struct OrtEmbeddingState { session: Session, primary_batched_session: Option, split_fbank_session: Option, + split_fbank_pool: Vec, split_fbank_batched_session: Option, split_tail_session: Option, split_tail_batched_session: Option, diff --git a/src/inference/embedding/fbank.rs b/src/inference/embedding/fbank.rs index 59c66fdd..5db2f5d8 100644 --- a/src/inference/embedding/fbank.rs +++ b/src/inference/embedding/fbank.rs @@ -58,6 +58,44 @@ impl EmbeddingModel { &mut self, audios: &[&[f32]], ) -> Result>, ort::Error> { + // Fan per-chunk fbank out across a pool of CPU sessions. + // fbank is otherwise ~76% of CUDA E2E wall time (intra-op threads don't scale it). + if audios.len() > 1 && !self.ort.split_fbank_pool.is_empty() { + let window_samples = self.meta.window_samples; + let pool = &mut self.ort.split_fbank_pool; + let per_worker = audios.len().div_ceil(pool.len()); + let mut collected: Vec>>> = Vec::new(); + collected.resize_with(pool.len(), || None); + std::thread::scope(|scope| { + let mut handles = Vec::new(); + for (worker_idx, session) in pool.iter_mut().enumerate() { + let start = worker_idx * per_worker; + if start >= audios.len() { + break; + } + let end = (start + per_worker).min(audios.len()); + let slice = &audios[start..end]; + handles.push(( + worker_idx, + scope.spawn(move || -> Result>, ort::Error> { + slice + .iter() + .map(|audio| fbank_via_session(session, audio, window_samples)) + .collect() + }), + )); + } + for (worker_idx, handle) in handles { + let vals = handle + .join() + .map_err(|_| ort::Error::new("fbank pool worker panicked"))??; + collected[worker_idx] = Some(vals); + } + Ok::<(), ort::Error>(()) + })?; + return Ok(collected.into_iter().flatten().flatten().collect()); + } + let has_batched = self.has_batched_fbank(); if !has_batched { tracing::debug!( @@ -171,3 +209,25 @@ impl EmbeddingModel { Ok(true) } } + +// Session-local fbank for the parallel pool path (no shared buffers). +fn fbank_via_session( + session: &mut ort::session::Session, + audio: &[f32], + window_samples: usize, +) -> Result, ort::Error> { + let mut buf = ndarray::Array3::::zeros((1, 1, window_samples)); + let copy_len = audio.len().min(window_samples); + buf.slice_mut(s![0, 0, ..copy_len]) + .assign(&ndarray::ArrayView1::from(&audio[..copy_len])); + let waveform_tensor = TensorRef::from_array_view(buf.view())?; + let outputs = session.run(ort::inputs!["waveform" => waveform_tensor])?; + let output = first_output(outputs.values(), "pool chunk fbank output")?; + let (shape, data) = output.try_extract_tensor::()?; + array2_from_shape_vec( + shape[1] as usize, + shape[2] as usize, + data.to_vec(), + "pool chunk fbank output", + ) +} diff --git a/src/inference/embedding/load/sessions.rs b/src/inference/embedding/load/sessions.rs index ff5ae7f8..42d00d19 100644 --- a/src/inference/embedding/load/sessions.rs +++ b/src/inference/embedding/load/sessions.rs @@ -26,6 +26,7 @@ pub(super) struct LoadedOrtSessions { session: Session, primary_batched_session: Option, split_fbank_session: Option, + split_fbank_pool: Vec, split_fbank_batched_session: Option, split_tail_session: Option, split_tail_batched_session: Option, @@ -54,6 +55,20 @@ pub(super) struct LoadedSessions { coreml: LoadedCoreMlState, } +/// Fallback sizing for the CPU fbank pool when `RuntimeConfig::fbank_pool` is `None`: +/// the `SPEAKRS_FBANK_POOL` override if it parses, else one session per four cores +/// (clamped to `1..=8`). Callers that set `fbank_pool` explicitly never reach the environment. +fn auto_fbank_pool_size() -> usize { + std::env::var("SPEAKRS_FBANK_POOL") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or_else(|| { + std::thread::available_parallelism() + .map(|c| (c.get() / 4).clamp(1, 8)) + .unwrap_or(1) + }) +} + impl LoadedSessions { pub(super) fn load( model_path: &Path, @@ -67,8 +82,6 @@ impl LoadedSessions { let split_primary_tail_batched_path = split_tail_model_path(model_path, PRIMARY_BATCH_SIZE); #[cfg(feature = "coreml")] let native_chunk_compute_units = config.chunk_emb_compute_units.to_ml_compute_units(); - #[cfg(not(feature = "coreml"))] - let _ = config; let use_split_backend = EmbeddingModel::split_backend_available(model_path); #[cfg(feature = "coreml")] @@ -236,10 +249,25 @@ impl LoadedSessions { ); } + // Pool of extra CPU fbank sessions for parallel per-chunk fbank + // (single-session fbank measured at ~76% of CUDA E2E wall on many-core hosts). + // CoreML modes have a native batched fbank path that the CPU pool would shadow, + // so the pool is skipped entirely there (also avoids loading unused CPU sessions). + let split_fbank_pool: Vec = if use_split_backend && !mode.is_coreml() { + let pool_size = config.fbank_pool.unwrap_or_else(auto_fbank_pool_size); + tracing::debug!(fbank_pool = pool_size, "fbank session pool"); + (0..pool_size) + .map(|_| EmbeddingModel::build_fbank_session(&split_fbank_path, ExecutionMode::Cpu)) + .collect::, _>>()? + } else { + Vec::new() + }; + let ort = LoadedOrtSessions { session, primary_batched_session, split_fbank_session, + split_fbank_pool, split_fbank_batched_session, split_tail_session, split_tail_batched_session, @@ -288,6 +316,7 @@ impl LoadedSessions { session: self.ort.session, primary_batched_session: self.ort.primary_batched_session, split_fbank_session: self.ort.split_fbank_session, + split_fbank_pool: self.ort.split_fbank_pool, split_fbank_batched_session: self.ort.split_fbank_batched_session, split_tail_session: self.ort.split_tail_session, split_tail_batched_session: self.ort.split_tail_batched_session, diff --git a/src/inference/embedding/session.rs b/src/inference/embedding/session.rs index 751dd24c..c0d1d136 100644 --- a/src/inference/embedding/session.rs +++ b/src/inference/embedding/session.rs @@ -62,9 +62,16 @@ impl EmbeddingModel { model_path: &Path, mode: ExecutionMode, ) -> Result { - let threads = std::thread::available_parallelism() - .map(|count| count.get().min(4)) - .unwrap_or(1); + // The default cap of 4 intra-op threads leaves fbank as ~76% of + // CUDA E2E wall time on many-core hosts; allow override via SPEAKRS_FBANK_THREADS. + let threads = std::env::var("SPEAKRS_FBANK_THREADS") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or_else(|| { + std::thread::available_parallelism() + .map(|count| count.get().min(4)) + .unwrap_or(1) + }); let builder = Session::builder()? .with_independent_thread_pool()? .with_intra_threads(threads)? diff --git a/src/pipeline/config.rs b/src/pipeline/config.rs index b4579608..6d8e84ae 100644 --- a/src/pipeline/config.rs +++ b/src/pipeline/config.rs @@ -77,6 +77,16 @@ impl PipelineConfig { pub struct RuntimeConfig { /// Number of chunk embedding workers pub chunk_emb_workers: usize, + /// Size of the CPU fbank session pool used for parallel per-chunk fbank. + /// + /// `None` (the default) auto-sizes: the `SPEAKRS_FBANK_POOL` environment override when it + /// parses, otherwise one session per four cores clamped to `1..=8`. `Some(0)` disables the + /// pool and falls back to the single fbank session. + /// + /// Setting this explicitly lets an embedder size the pool without touching the environment, + /// which matters because `setenv` is not thread-safe: a host that loads models lazily or + /// concurrently cannot safely use the environment override. + pub fbank_pool: Option, /// CoreML compute units for chunk embedding (CoreML modes only) #[cfg(feature = "coreml")] #[cfg_attr(docsrs, doc(cfg(feature = "coreml")))] @@ -87,6 +97,7 @@ impl Default for RuntimeConfig { fn default() -> Self { Self { chunk_emb_workers: 1, + fbank_pool: None, #[cfg(feature = "coreml")] chunk_emb_compute_units: CoreMlComputeUnits::All, }