Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/workflows/codecov.yml
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@ jobs:

- name: Codecov
# You may pin to the exact commit or the version.
uses: codecov/codecov-action@v5.1.2
uses: codecov/codecov-action@v4.6.0
with:
# Repository upload token - get it from codecov.io. Required only for private repositories
token: ${{ secrets.CODECOV_TOKEN }}
Expand Down
17 changes: 11 additions & 6 deletions src/cli.rs
Original file line number Diff line number Diff line change
Expand Up @@ -196,8 +196,8 @@ pub enum Commands {
#[arg(short)]
output: Option<String>,

/// Calculate sparse distances with k nearest-neighbours
#[arg(long, group = "query")]
/// Calculate sparse distances with k nearest-neighbours (ref-vs-ref or ref-vs-query)
#[arg(long)]
knn: Option<usize>,

/// Sample names to analyse
Expand All @@ -216,9 +216,14 @@ pub enum Commands {
#[arg(long, value_parser = valid_cpus, default_value_t = 1)]
threads: usize,

/// File listing sample and completeness estimate 0.0-1.0 (tab separated)
/// File listing reference sample completeness estimates 0.0-1.0 (tab separated: genome_id<tab>completeness)
#[arg(long)]
ref_completeness_file: Option<String>,

/// File listing query sample completeness estimates 0.0-1.0 (tab separated: genome_id<tab>completeness).
/// Only used in cross-query mode (when a query database is provided).
#[arg(long)]
completeness_file: Option<String>,
query_completeness_file: Option<String>,

/// minimum completeness for a sample to be corrected but the completeness correction
#[arg(long, default_value_t = 0.64)]
Expand Down Expand Up @@ -441,9 +446,9 @@ pub enum InvertedCommands {
#[arg(long, value_parser = valid_cpus, default_value_t = 1)]
threads: usize,

/// Completeness file
/// File listing sample completeness estimates 0.0-1.0 (tab separated: genome_id<tab>completeness)
#[arg(long)]
completeness_file: Option<String>,
ref_completeness_file: Option<String>,

/// minimum completeness for a sample to be corrected but the completeness correction
#[arg(long, default_value_t = 0.64)]
Expand Down
70 changes: 54 additions & 16 deletions src/distances/distance_matrix.rs
Original file line number Diff line number Diff line change
Expand Up @@ -274,24 +274,32 @@ pub enum DistVec {
CoreAcc(Vec<SparseCoreAcc>),
}

/// A sparse distance matrix with a maximum of `knn` distances for each sample
/// A sparse distance matrix with a maximum of `knn` distances for each sample.
///
/// In self-query mode (one database), `query_names` is `None` and `ref_names` is
/// used for both row iteration and column index lookup.
///
/// In cross-query mode (two databases), `query_names` holds the query genome names
/// (one per row) and `ref_names` holds the reference genome names (indexed by
/// the stored neighbour index inside each distance item).
pub struct SparseDistanceMatrix<'a> {
/// Total number of distances
pub n_distances: usize,
/// Maximum number of distances kept per sample: k smallest distances
pub knn: usize,
jaccard: DistType,
distances: DistVec,
/// Reference genome names — used as column labels (neighbour index lookup)
ref_names: Vec<&'a str>,
/// Query genome names — used as row labels in cross-query mode (`None` in self-query mode)
query_names: Option<Vec<&'a str>>,
}

impl<'a> SparseDistanceMatrix<'a> {
/// Create a new sparse distance matrix for a [`MultiSketch`] keeping the
/// minimum `knn` distances with distance options set by `jaccard`
/// Self-query constructor: one database, rows and columns are the same set.
pub fn new(ref_sketches: &'a MultiSketch, knn: usize, jaccard: DistType) -> Self {
let n_distances = ref_sketches.number_samples_loaded() * knn;

// Pre-allocate distances
let distances = match jaccard {
DistType::CoreAcc => DistVec::CoreAcc(vec![SparseCoreAcc(0, 0.0, 0.0); n_distances]),
DistType::Jaccard(_, _, _) => {
Expand All @@ -305,6 +313,35 @@ impl<'a> SparseDistanceMatrix<'a> {
jaccard,
distances,
ref_names: Self::sketch_names(ref_sketches),
query_names: None,
}
}

/// Cross-query constructor: two databases.
/// Rows are query genomes; column indices index into the reference genome list.
pub fn new_cross_query(
ref_sketches: &'a MultiSketch,
query_sketches: &'a MultiSketch,
knn: usize,
jaccard: DistType,
) -> Self {
let n_query = query_sketches.number_samples_loaded();
let n_distances = n_query * knn;

let distances = match jaccard {
DistType::CoreAcc => DistVec::CoreAcc(vec![SparseCoreAcc(0, 0.0, 0.0); n_distances]),
DistType::Jaccard(_, _, _) => {
DistVec::Jaccard(vec![SparseJaccard(0, 0.0); n_distances])
}
};

Self {
n_distances,
knn,
jaccard,
distances,
ref_names: Self::sketch_names(ref_sketches),
query_names: Some(Self::sketch_names(query_sketches)),
}
}

Expand All @@ -322,37 +359,38 @@ impl<'a> Distances<'a> for SparseDistanceMatrix<'a> {

impl fmt::Display for SparseDistanceMatrix<'_> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let mut ref_name_iter = self.ref_names.iter();
let mut ref_name = ref_name_iter.next().unwrap();
// In cross-query mode use query genome names as row labels;
// in self-query mode fall back to ref_names.
let query_names = self.query_names.as_deref().unwrap_or(&self.ref_names);
let mut row_name_iter = query_names.iter();
let mut row_name = row_name_iter.next().unwrap();
let mut k = 0;
match &self.distances {
DistVec::Jaccard(dists) => {
for dist_item in dists {
k += 1;
if k > self.knn {
ref_name = ref_name_iter.next().unwrap();
row_name = row_name_iter.next().unwrap();
k = 1;
}
let query_name = self.ref_names[dist_item.0];
// If fewer items than knn, padded with ref=query and dist=1.0 values
// which are skipped here
// TODO: more rust-like way of doing this would be to have
// SparseJaccard as an enum with an empty value
if dist_item.1 < 1.0_f32 || query_name != *ref_name {
writeln!(f, "{ref_name}\t{query_name}\t{}", dist_item.1)?;
// dist_item.0 is always an index into ref_names
let col_name = self.ref_names[dist_item.0];
// Padding entries (dist == 1.0, col == row) are skipped
if dist_item.1 < 1.0_f32 || col_name != *row_name {
writeln!(f, "{row_name}\t{col_name}\t{}", dist_item.1)?;
}
}
}
DistVec::CoreAcc(dists) => {
for dist_item in dists {
k += 1;
if k > self.knn {
ref_name = ref_name_iter.next().unwrap();
row_name = row_name_iter.next().unwrap();
k = 1;
}
writeln!(
f,
"{ref_name}\t{}\t{}\t{}",
"{row_name}\t{}\t{}\t{}",
self.ref_names[dist_item.0], dist_item.1, dist_item.2,
)?;
}
Expand Down
7 changes: 4 additions & 3 deletions src/distances/jaccard.rs
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,8 @@ pub fn core_acc_dist(
query_sketches: &MultiSketch,
ref_sketch_idx: usize,
query_sketch_idx: usize,
completeness_vec: Option<&Vec<f64>>,
ref_completeness_vec: Option<&Vec<f64>>,
query_completeness_vec: Option<&Vec<f64>>,
completeness_cutoff: f64,
) -> (f32, f32) {
if ref_sketches.kmer_lengths().len() < 2 {
Expand All @@ -74,8 +75,8 @@ pub fn core_acc_dist(
let tolerance = (2.0_f64 / ((ref_sketches.sketch_size * u64::BITS as u64) as f64)).ln();
//let tolerance = -100.0_f32;
for (k_idx, k) in ref_sketches.kmer_lengths().iter().enumerate() {
let c1 = completeness_vec.map(|cv| cv[ref_sketch_idx]);
let c2 = completeness_vec.map(|cv| cv[query_sketch_idx]);
let c1 = ref_completeness_vec.map(|cv| cv[ref_sketch_idx]);
let c2 = query_completeness_vec.map(|cv| cv[query_sketch_idx]);
let y = jaccard_index(
ref_sketches.get_sketch_slice(ref_sketch_idx, k_idx),
query_sketches.get_sketch_slice(query_sketch_idx, k_idx),
Expand Down
123 changes: 111 additions & 12 deletions src/distances/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,7 @@ pub fn self_dists_all<'a>(
i,
j,
completeness_vec,
completeness_vec,
completeness_cutoff,
);
dist_slice[dist_idx * 2] = dist.0;
Expand Down Expand Up @@ -208,6 +209,7 @@ pub fn self_dists_knn<'a>(
i,
j,
completeness_vec,
completeness_vec,
completeness_cutoff,
);
let dist_item = SparseCoreAcc(j, dists.0, dists.1);
Expand All @@ -221,15 +223,16 @@ pub fn self_dists_knn<'a>(
sp_distances
}

/// Self query mode (dense, all distances)
pub fn self_query_dists_all<'a>(
/// Cross-query mode (dense, all distances)
pub fn cross_dists_all<'a>(
ref_sketches: &'a MultiSketch,
query_sketches: &'a MultiSketch,
n: usize,
nq: usize,
n_query: usize,
dist_type: DistType,
quiet: bool,
completeness_vec: Option<&Vec<f64>>,
ref_completeness_vec: Option<&Vec<f64>>,
query_completeness_vec: Option<&Vec<f64>>,
completeness_cutoff: f64,
) -> DistanceMatrix<'a> {
let mut distances = DistanceMatrix::new(ref_sketches, Some(query_sketches), dist_type);
Expand All @@ -245,14 +248,11 @@ pub fn self_query_dists_all<'a>(
.for_each(|(chunk_idx, dist_slice)| {
// Get first i, j index for the chunk
let start_dist_idx = chunk_idx * CHUNK_SIZE;
let (mut i, mut j) = calc_query_indices(start_dist_idx, nq);
let (mut i, mut j) = calc_query_indices(start_dist_idx, n_query);
for dist_idx in 0..CHUNK_SIZE {
if let Some((k_idx, k_f64)) = k_vals {
// If completeness_vec is Some, extract the value at index i (or j) from the inner vector.
// If completeness_vec is None, the result will also be None.
// This uses Option::map to safely access the completeness value for each sample.
let c1 = completeness_vec.map(|cv| cv[i]);
let c2 = completeness_vec.map(|cv| cv[j]);
let c1 = ref_completeness_vec.map(|cv| cv[i]);
let c2 = query_completeness_vec.map(|cv| cv[j]);
let j_index = jaccard_index(
ref_sketches.get_sketch_slice(i, k_idx),
query_sketches.get_sketch_slice(j, k_idx),
Expand All @@ -273,7 +273,8 @@ pub fn self_query_dists_all<'a>(
query_sketches,
i,
j,
completeness_vec,
ref_completeness_vec,
query_completeness_vec,
completeness_cutoff,
);
dist_slice[dist_idx * 2] = dist.0;
Expand All @@ -282,7 +283,7 @@ pub fn self_query_dists_all<'a>(

// Move to next index
j += 1;
if j >= nq {
if j >= n_query {
i += 1;
j = 0;
// End of all dists reached (final chunk)
Expand All @@ -295,6 +296,104 @@ pub fn self_query_dists_all<'a>(
distances
}

/// Cross-query mode with kNN filtering.
///
/// For each query genome, computes distances to all `n` reference genomes and
/// retains only the `knn` nearest neighbours using a priority queue (max-heap
/// capped at `knn`). Output has `n_query × knn` entries — one row per query genome.
///
/// This is the cross-database analogue of [`self_dists_knn`].
pub fn cross_dists_knn<'a>(
ref_sketches: &'a MultiSketch,
query_sketches: &'a MultiSketch,
n: usize,
n_query: usize,
knn: usize,
dist_type: DistType,
quiet: bool,
ref_completeness_vec: Option<&Vec<f64>>,
query_completeness_vec: Option<&Vec<f64>>,
completeness_cutoff: f64,
) -> SparseDistanceMatrix<'a> {
if n == 0 {
panic!("Reference database has no loaded samples");
}
if n_query == 0 {
panic!("Query database has no loaded samples");
}
// Can't have more neighbours than there are reference genomes
let knn = knn.min(n);
let mut sp_distances =
SparseDistanceMatrix::new_cross_query(ref_sketches, query_sketches, knn, dist_type);
let k_vals = sp_distances.k_vals();
let ani = sp_distances.ani();
let progress_bar = get_progress_bar(n_query, BAR_PERCENT, quiet);
match sp_distances.dists_mut() {
DistVec::Jaccard(distances) => {
let (k_idx, k_f64) = k_vals.unwrap();
distances
.par_chunks_mut(knn)
.progress_with(progress_bar)
.enumerate()
.for_each(|(qi, row_dist_slice)| {
let mut heap = BinaryHeap::with_capacity(knn + 1);
let qi_sketch = query_sketches.get_sketch_slice(qi, k_idx);
for ri in 0..n {
let c1 = query_completeness_vec.map(|cv| cv[qi]);
let c2 = ref_completeness_vec.map(|cv| cv[ri]);
let dist = jaccard_index(
qi_sketch,
ref_sketches.get_sketch_slice(ri, k_idx),
ref_sketches.sketchsize64,
c1,
c2,
completeness_cutoff,
);
let dist_f32 = if ani {
(1.0_f64 - ani_pois(dist, k_f64)) as f32
} else {
(1.0_f64 - dist) as f32
};
push_heap(&mut heap, SparseJaccard(ri, dist_f32), knn);
}
if ani {
heap.into_sorted_vec().iter().zip(row_dist_slice).for_each(
|(inverse_ani, output_ani)| {
*output_ani =
SparseJaccard(inverse_ani.0, 1.0_f32 - inverse_ani.1);
},
);
} else {
row_dist_slice.clone_from_slice(&heap.into_sorted_vec());
}
});
}
DistVec::CoreAcc(distances) => {
distances
.par_chunks_mut(knn)
.progress_with(progress_bar)
.enumerate()
.for_each(|(qi, row_dist_slice)| {
let mut heap = BinaryHeap::with_capacity(knn + 1);
for ri in 0..n {
let dists = core_acc_dist(
ref_sketches,
query_sketches,
ri,
qi,
ref_completeness_vec,
query_completeness_vec,
completeness_cutoff,
);
push_heap(&mut heap, SparseCoreAcc(ri, dists.0, dists.1), knn);
}
row_dist_slice.clone_from_slice(&heap.into_sorted_vec());
});
}
}
sp_distances
}

/// Same as [`self_dists_knn`], but also using an inverted_index to precluster
/// to reduce the number of comparisons
pub fn self_dists_knn_precluster<'a>(
Expand Down
Loading
Loading