#[cfg(feature = "gpu")]
mod tests {
use haystackfm::alphabet::DnaSequence;
use haystackfm::{BidirFmIndex, FmIndexConfig, MemHit, SeqId};
use pollster::FutureExt as _;
use std::collections::HashSet;
fn cpu_config() -> FmIndexConfig {
FmIndexConfig {
sa_sample_rate: 1,
use_gpu: false,
..Default::default()
}
}
fn build(seqs: &[&str]) -> BidirFmIndex {
let dna: Vec<DnaSequence> = seqs
.iter()
.map(|s| DnaSequence::from_str(s).unwrap())
.collect();
BidirFmIndex::build_cpu(&dna, &cpu_config()).unwrap()
}
fn seq(s: &str) -> DnaSequence {
DnaSequence::from_str(s).unwrap()
}
fn cpu_mems_with_positions(
idx: &BidirFmIndex,
query: &[u8],
min_len: usize,
) -> Vec<(usize, usize, HashSet<(SeqId, u32)>)> {
idx.find_mems(query, min_len, true)
.into_iter()
.map(|m| {
let positions: HashSet<(SeqId, u32)> = m.positions.iter().copied().collect();
(m.query_start, m.query_end, positions)
})
.collect()
}
fn cpu_smems_with_positions(
idx: &BidirFmIndex,
query: &[u8],
min_len: usize,
) -> Vec<(usize, usize, HashSet<(SeqId, u32)>)> {
idx.find_smems(query, min_len, true)
.into_iter()
.map(|m| {
let positions: HashSet<(SeqId, u32)> = m.positions.iter().copied().collect();
(m.query_start, m.query_end, positions)
})
.collect()
}
fn gpu_mems_with_positions(hits: &[MemHit]) -> Vec<(usize, usize, HashSet<(SeqId, u32)>)> {
hits.iter()
.map(|h| {
let positions: HashSet<(SeqId, u32)> = h.positions.iter().copied().collect();
(h.query_start as usize, h.query_end as usize, positions)
})
.collect()
}
fn sort_mem_tuples(v: &mut Vec<(usize, usize, HashSet<(SeqId, u32)>)>) {
v.sort_by_key(|(qs, qe, _)| (*qs, *qe));
}
#[test]
fn smem_positions_single_ref() {
let idx = build(&["ACGTACGT"]);
let boundaries = idx.seq_boundaries().to_vec();
let q = seq("ACGT");
let mut cpu = cpu_smems_with_positions(&idx, q.as_slice(), 1);
let gpu_raw = idx
.find_smems_gpu(&[q], 1, &boundaries, 1024)
.block_on()
.unwrap();
let mut gpu = gpu_mems_with_positions(&gpu_raw[0]);
sort_mem_tuples(&mut cpu);
sort_mem_tuples(&mut gpu);
assert_eq!(cpu.len(), gpu.len(), "MEM count differs");
for (i, (c, g)) in cpu.iter().zip(gpu.iter()).enumerate() {
assert_eq!(c.0, g.0, "MEM {i}: query_start");
assert_eq!(c.1, g.1, "MEM {i}: query_end");
assert_eq!(
c.2, g.2,
"MEM {i}: positions mismatch\ncpu={:?}\ngpu={:?}",
c.2, g.2
);
}
}
#[test]
fn smem_positions_multi_ref() {
let idx = build(&["ACGTACGT", "ACGT", "TTTTACGT"]);
let boundaries = idx.seq_boundaries().to_vec();
let q = seq("ACGT");
let mut cpu = cpu_smems_with_positions(&idx, q.as_slice(), 1);
let gpu_raw = idx
.find_smems_gpu(&[q], 1, &boundaries, 1024)
.block_on()
.unwrap();
let mut gpu = gpu_mems_with_positions(&gpu_raw[0]);
sort_mem_tuples(&mut cpu);
sort_mem_tuples(&mut gpu);
assert_eq!(cpu.len(), gpu.len());
for (i, (c, g)) in cpu.iter().zip(gpu.iter()).enumerate() {
assert_eq!(c.2, g.2, "MEM {i}: positions\ncpu={:?}\ngpu={:?}", c.2, g.2);
}
}
#[test]
fn mem_positions_single_ref() {
let idx = build(&["ACGTACGT"]);
let boundaries = idx.seq_boundaries().to_vec();
let q = seq("ACGT");
let mut cpu = cpu_mems_with_positions(&idx, q.as_slice(), 1);
let gpu_raw = idx
.find_mems_gpu(&[q], 1, &boundaries, 1024)
.block_on()
.unwrap();
let mut gpu = gpu_mems_with_positions(&gpu_raw[0]);
sort_mem_tuples(&mut cpu);
sort_mem_tuples(&mut gpu);
assert_eq!(cpu.len(), gpu.len());
for (i, (c, g)) in cpu.iter().zip(gpu.iter()).enumerate() {
assert_eq!(c.2, g.2, "MEM {i}: positions\ncpu={:?}\ngpu={:?}", c.2, g.2);
}
}
#[test]
fn mem_positions_multi_ref_batch() {
let idx = build(&["ACGTACGT", "CGTTAGCC", "AAACGT"]);
let boundaries = idx.seq_boundaries().to_vec();
let queries = vec![seq("ACGT"), seq("CGT"), seq("AAA")];
let cpu: Vec<_> = queries
.iter()
.map(|q| {
let mut v = cpu_mems_with_positions(&idx, q.as_slice(), 1);
sort_mem_tuples(&mut v);
v
})
.collect();
let gpu_raw = idx
.find_mems_gpu(&queries, 1, &boundaries, 1024)
.block_on()
.unwrap();
for (qi, (cpu_q, gpu_q)) in cpu.iter().zip(gpu_raw.iter()).enumerate() {
let mut gpu_v = gpu_mems_with_positions(gpu_q);
sort_mem_tuples(&mut gpu_v);
assert_eq!(cpu_q.len(), gpu_v.len(), "query {qi}: MEM count");
for (i, (c, g)) in cpu_q.iter().zip(gpu_v.iter()).enumerate() {
assert_eq!(
c.2, g.2,
"query {qi} MEM {i}: positions\ncpu={:?}\ngpu={:?}",
c.2, g.2
);
}
}
}
#[test]
fn mem_positions_empty_query() {
let idx = build(&["ACGTACGT"]);
let boundaries = idx.seq_boundaries().to_vec();
let q = seq("A"); let gpu_raw = idx
.find_mems_gpu(&[q], 4, &boundaries, 1024)
.block_on()
.unwrap();
assert!(gpu_raw[0].is_empty());
}
#[test]
fn mem_positions_no_ref_boundaries_skips_resolve() {
let idx = build(&["ACGTACGT"]);
let q = seq("ACGT");
let gpu_raw = idx.find_mems_gpu(&[q], 1, &[], 1024).block_on().unwrap();
assert!(!gpu_raw[0].is_empty());
for hit in &gpu_raw[0] {
assert!(hit.positions.is_empty());
}
}
#[test]
fn mem_truncation_flag_set_when_capped() {
let idx = build(&["ACGTACGTACGTACGT"]); let boundaries = idx.seq_boundaries().to_vec();
let q = seq("ACGT");
let gpu_raw = idx
.find_mems_gpu(&[q], 1, &boundaries, 1)
.block_on()
.unwrap();
let hit = &gpu_raw[0][0];
assert_eq!(hit.positions.len(), 1);
assert!(hit.truncated);
}
}