-
Notifications
You must be signed in to change notification settings - Fork 17
CUDA pipeline performance series: vectorized VBx + threaded pdist, fbank session pool, folded segmentation export #15
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 3 commits
1f4a076
90200c1
7687da7
a82f09d
0c4cdbb
0e7a12a
d8f00a8
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -26,30 +26,83 @@ pub fn cluster(embeddings: &ArrayView2<f32>, config: AhcConfig) -> Vec<usize> { | |
| } | ||
|
|
||
| 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<f32>) -> Vec<f32> { | ||
| // 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::<f32>() | ||
| .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<f32> = 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<i}(n-1-r) = i*(n-1) - i*(i-1)/2 | ||
| let seg_start = |i: usize| i * (observations - 1) - i * i.saturating_sub(1) / 2; | ||
|
|
||
| // hand each block its contiguous slice | ||
| let mut blocks: Vec<(usize, usize, &mut [f32])> = 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; | ||
| } | ||
| } | ||
|
|
||
| std::thread::scope(|scope| { | ||
| for (bi, bi_end, slice) in blocks { | ||
| let emb = &embeddings; | ||
| let norms = &sq_norms; | ||
| scope.spawn(move || { | ||
| 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) | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
When clustering a large recording, this loop starts every block worker before joining any of them, so all Gram matrices remain live alongside the full condensed-distance vector. For the documented 21k-observation workload, this approximately doubles pairwise-distance storage and can terminate clustering through allocation failure or the OOM killer. Knowledge Base Used: Clustering algorithms Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time! |
||
| 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; | ||
| } | ||
| } | ||
| }); | ||
|
coderabbitai[bot] marked this conversation as resolved.
|
||
| } | ||
| }); | ||
| condensed | ||
| } | ||
|
|
||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -58,6 +58,44 @@ impl EmbeddingModel { | |
| &mut self, | ||
| audios: &[&[f32]], | ||
| ) -> Result<Vec<Array2<f32>>, 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() { | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. When a Core ML split-backend model processes multiple chunks, the mode-independent CPU pool satisfies this branch before native fbank batching is considered. Fbank extraction therefore runs through CPU ONNX sessions instead of the configured native accelerator, causing a substantial performance regression on the optimized Core ML path. Knowledge Base Used: Speaker embedding inference Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time! |
||
| 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<Option<Vec<Array2<f32>>>> = 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<Vec<Array2<f32>>, 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()); | ||
| } | ||
|
Comment on lines
+61
to
+97
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🚀 Performance & Scalability | 🟠 Major | ⚡ Quick win Keep CoreML batch fbank on CoreML. At Line 63, this branch selects the CPU ORT pool before 🤖 Prompt for AI Agents |
||
|
|
||
| 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<Array2<f32>, ort::Error> { | ||
| let mut buf = ndarray::Array3::<f32>::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::<f32>()?; | ||
| array2_from_shape_vec( | ||
| shape[1] as usize, | ||
| shape[2] as usize, | ||
| data.to_vec(), | ||
| "pool chunk fbank output", | ||
| ) | ||
| } | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
🎯 Functional Correctness | 🟠 Major | ⚡ Quick win
🔎 Supported by static analysis
🏁 Script executed:
Repository: avencera/speakrs
Length of output: 6984
🌐 Web query:
onnxsim simplify function default check_n test_input_shapes input_data onnx_simplifier.py💡 Result:
In the onnxsim library, the simplify function is the primary entry point for model simplification [1][2]. The parameters check_n, test_input_shapes, and input_data are used to control the validation process of the simplified model [3][4]. 1. check_n (int, default 0): This parameter determines the number of times the simplified model is validated against the original model using random inputs [3][1]. When set to a value greater than 0, the library generates random inputs (or uses custom data if provided), runs both the original and simplified models, and compares their outputs to ensure numerical equivalence [1]. 2. test_input_shapes (dict or None, default None): When a model has dynamic input shapes, this parameter allows the user to specify fixed shapes to be used specifically for generating the random test inputs during the validation phase [3][4]. 3. input_data (dict or None, default None): This parameter allows users to provide custom input data (as a dictionary of numpy arrays) for the validation process [1]. If provided, this data is used instead of randomly generated data when check_n is greater than 0 [1]. These parameters are part of the simplify function signature, which is defined in the onnx_simplifier.py module [3][5]. Note that while the command-line interface provides similar functionality, these specific parameters are used when invoking the simplify function directly via the Python API [5][2].
Citations:
Reject simplification unless outputs match.
onnxsim.simplify(model)defaults tocheck_n=0. This validates the graph but does not compare outputs with the original model. A behavior-changing graph can therefore pass Lines 77-78 and overwrite the export at Line 79. Usecheck_n > 0with inputs for all three exports, or run an explicit original-versus-simplified inference check before saving.🤖 Prompt for AI Agents
Source: MCP tools