pub mod cpu;
#[cfg(feature = "gpu")]
pub mod gpu;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SuffixArray {
pub data: Vec<u32>,
}
impl SuffixArray {
pub fn len(&self) -> usize {
self.data.len()
}
pub fn is_empty(&self) -> bool {
self.data.is_empty()
}
}
const SA_MARKER_SB_WORDS: usize = 8;
const SA_RECORD_STRIDE: usize = 14;
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct SampledSuffixArray {
word_data: Vec<u8>,
sa_vals: Vec<u32>,
pub sample_rate: u32,
text_len: u32,
}
impl SampledSuffixArray {
pub(crate) fn from_full(sa: &SuffixArray, sample_rate: u32, force_sampled: &[u32]) -> Self {
let n = sa.data.len();
let num_words = n.div_ceil(64);
let forced: std::collections::HashSet<u32> = force_sampled.iter().copied().collect();
let mut bitvector = vec![0u64; num_words];
let mut sa_vals = Vec::new();
for (i, &sa_val) in sa.data.iter().enumerate() {
if sa_val.is_multiple_of(sample_rate) || forced.contains(&sa_val) {
bitvector[i / 64] |= 1u64 << (i % 64);
sa_vals.push(sa_val);
}
}
let mut word_data = vec![0u8; num_words * SA_RECORD_STRIDE];
let mut cumulative = 0u32;
let mut sb_base = 0u32;
for (w, &word) in bitvector.iter().enumerate() {
if w % SA_MARKER_SB_WORDS == 0 {
sb_base = cumulative;
}
let delta = (cumulative - sb_base) as u16;
let rec = &mut word_data[w * SA_RECORD_STRIDE..(w + 1) * SA_RECORD_STRIDE];
rec[0..8].copy_from_slice(&word.to_ne_bytes());
rec[8..12].copy_from_slice(&sb_base.to_ne_bytes());
rec[12..14].copy_from_slice(&delta.to_ne_bytes());
cumulative += word.count_ones();
}
Self {
word_data,
sa_vals,
sample_rate,
text_len: n as u32,
}
}
#[inline]
fn word_at(&self, word_idx: usize) -> u64 {
let off = word_idx * SA_RECORD_STRIDE;
debug_assert!(off + 8 <= self.word_data.len());
unsafe {
self.word_data
.as_ptr()
.add(off)
.cast::<u64>()
.read_unaligned()
}
}
#[inline]
fn sb_count_at(&self, word_idx: usize) -> u32 {
let off = word_idx * SA_RECORD_STRIDE + 8;
debug_assert!(off + 4 <= self.word_data.len());
unsafe {
self.word_data
.as_ptr()
.add(off)
.cast::<u32>()
.read_unaligned()
}
}
#[inline]
fn delta_at(&self, word_idx: usize) -> u32 {
let off = word_idx * SA_RECORD_STRIDE + 12;
debug_assert!(off + 2 <= self.word_data.len());
unsafe {
self.word_data
.as_ptr()
.add(off)
.cast::<u16>()
.read_unaligned() as u32
}
}
#[inline]
pub(crate) fn prefetch(&self, i: u32) {
let off = (i as usize / 64) * SA_RECORD_STRIDE;
if off < self.word_data.len() {
crate::prefetch::prefetch_read(unsafe { self.word_data.as_ptr().add(off) });
}
}
#[inline]
pub fn is_sampled(&self, i: u32) -> bool {
let i = i as usize;
(self.word_at(i / 64) >> (i % 64)) & 1 == 1
}
#[inline]
pub fn get(&self, i: u32) -> Option<u32> {
let i = i as usize;
let word_idx = i / 64;
let bit_offset = i % 64;
let word = self.word_at(word_idx);
if (word >> bit_offset) & 1 == 0 {
return None;
}
let mask = if bit_offset == 0 {
0
} else {
(1u64 << bit_offset) - 1
};
let rank =
self.sb_count_at(word_idx) + self.delta_at(word_idx) + (word & mask).count_ones();
Some(self.sa_vals[rank as usize])
}
pub fn to_flat_vec(&self, n: usize) -> Vec<u32> {
let mut flat = vec![u32::MAX; n];
let mut rank = 0usize;
let num_words = self.word_data.len() / SA_RECORD_STRIDE;
for word_idx in 0..num_words {
let mut w = self.word_at(word_idx);
while w != 0 {
let bit = w.trailing_zeros() as usize;
let pos = word_idx * 64 + bit;
if pos < n {
flat[pos] = self.sa_vals[rank];
rank += 1;
}
w &= w - 1;
}
}
flat
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::alphabet::encode_char;
use crate::bwt::cpu::build_bwt;
use crate::suffix_array::cpu::build_suffix_array;
fn encode(s: &str) -> Vec<u8> {
use crate::alphabet::SENTINEL;
let mut v: Vec<u8> = s.chars().map(|c| encode_char(c).unwrap()).collect();
if v.last() != Some(&SENTINEL) {
v.push(SENTINEL);
}
v
}
fn make_sa(text: &str) -> (SuffixArray, Vec<u8>) {
let encoded = encode(text);
let sa = build_suffix_array(&encoded);
(sa, encoded)
}
#[test]
fn test_sampled_sa_get_matches_full() {
let (sa, _) = make_sa("ACGTACGTACGT");
let sample_rate = 4;
let ssa = SampledSuffixArray::from_full(&sa, sample_rate, &[]);
for (i, &sa_val) in sa.data.iter().enumerate() {
if sa_val % sample_rate == 0 {
assert_eq!(ssa.get(i as u32), Some(sa_val), "row {i}");
assert!(ssa.is_sampled(i as u32), "row {i} should be sampled");
} else {
assert_eq!(ssa.get(i as u32), None, "row {i}");
assert!(!ssa.is_sampled(i as u32), "row {i} should not be sampled");
}
}
}
#[test]
fn test_sampled_sa_to_flat_vec() {
let (sa, _) = make_sa("ACGTACGTACGT");
let sample_rate = 4;
let ssa = SampledSuffixArray::from_full(&sa, sample_rate, &[]);
let n = sa.data.len();
let flat = ssa.to_flat_vec(n);
for (i, &sa_val) in sa.data.iter().enumerate() {
if sa_val % sample_rate == 0 {
assert_eq!(flat[i], sa_val, "flat[{i}]");
} else {
assert_eq!(flat[i], u32::MAX, "flat[{i}] should be sentinel");
}
}
}
#[test]
fn test_sampled_sa_rate_1_covers_all() {
let (sa, _) = make_sa("ACGTACGT");
let ssa = SampledSuffixArray::from_full(&sa, 1, &[]);
for (i, &sa_val) in sa.data.iter().enumerate() {
assert_eq!(ssa.get(i as u32), Some(sa_val));
}
}
#[test]
fn test_sampled_sa_long_text_spans_multiple_words() {
let s = "ACGT".repeat(50);
let (sa, _) = make_sa(&s);
let sample_rate = 8;
let ssa = SampledSuffixArray::from_full(&sa, sample_rate, &[]);
let n = sa.data.len();
let flat = ssa.to_flat_vec(n);
for (i, &sa_val) in sa.data.iter().enumerate() {
if sa_val % sample_rate == 0 {
assert_eq!(flat[i], sa_val);
assert_eq!(ssa.get(i as u32), Some(sa_val));
} else {
assert_eq!(flat[i], u32::MAX);
assert_eq!(ssa.get(i as u32), None);
}
}
}
#[test]
fn test_locate_via_lf_walk_uses_ssa() {
let text = encode("ACGTACGTACGT");
let sa = build_suffix_array(&text);
let bwt = build_bwt(&text, &sa);
let _ = bwt; let ssa = SampledSuffixArray::from_full(&sa, 4, &[]);
let sampled_count = sa.data.iter().filter(|&&v| v % 4 == 0).count();
let found: usize = (0..sa.data.len() as u32).filter_map(|i| ssa.get(i)).count();
assert_eq!(found, sampled_count);
}
}