use crate::{
alphabet::{Symbol, SymbolAlphabet},
search::SearchPtr,
simd_instructions::{masked_popcount, simd_and, simd_andnot, simd_or, Vec256},
};
use mem_dbg::MemSize;
use serde::{Deserialize, Serialize};
#[derive(Clone, Serialize, Deserialize, Debug, PartialEq, PartialOrd, Eq, Ord, Hash, Default, MemSize)]
#[repr(align(32))]
pub (crate) struct NucleotideBwtBlock {
bit_vectors: [Vec256; Self::NUM_BIT_VECTORS],
milestones: [u64; Self::NUM_MILESTONES],
}
#[derive(Clone, Serialize, Deserialize, Debug, PartialEq, PartialOrd, Eq, Ord, Hash, Default, MemSize)]
#[repr(align(32))]
pub (crate) struct AminoBwtBlock {
bit_vectors: [Vec256; Self::NUM_BIT_VECTORS],
milestones: [u64; Self::NUM_MILESTONES],
}
impl NucleotideBwtBlock {
pub (crate) const NUM_MILESTONES: usize = 8;
pub (crate) const NUM_BIT_VECTORS: usize = 3;
pub (crate) fn new() -> Self {
NucleotideBwtBlock {
bit_vectors: [Vec256::new(); Self::NUM_BIT_VECTORS],
milestones: [0; Self::NUM_MILESTONES],
}
}
pub (crate) fn from_data(
bit_vectors: [Vec256; Self::NUM_BIT_VECTORS],
milestones: [u64; Self::NUM_MILESTONES],
) -> Self {
NucleotideBwtBlock {
bit_vectors,
milestones,
}
}
pub (crate) fn symbol_at(&self, position_block: u64)->Symbol{
let mut bit_vector_encoding: u64 = 0;
for bit in 0..self.bit_vectors.len() {
let bit_value = self.bit_vectors[bit].extract_bit(&position_block);
bit_vector_encoding |= bit_value << bit;
}
Symbol::new_bit_vector(SymbolAlphabet::Nucleotide, bit_vector_encoding as u8)
}
pub (crate) fn set_symbol_at(&mut self, symbol:&Symbol, position_in_block: u64){
let mut encoded_symbol = symbol.bit_vector();
let mut bit_vector_idx = 0;
while encoded_symbol != 0 {
if encoded_symbol & 0x1 == 1 {
self.bit_vectors[bit_vector_idx].set_bit(&position_in_block);
}
encoded_symbol >>= 1;
bit_vector_idx += 1;
}
}
#[inline]
pub (crate) fn set_milestones(&mut self, values: &Vec<u64>) {
debug_assert!(values.len() >= SymbolAlphabet::Nucleotide.cardinality() as usize);
unsafe{
for milestone_idx in 0..SymbolAlphabet::Nucleotide.cardinality() as usize {
*self.milestones.get_unchecked_mut(milestone_idx) = *values.get_unchecked(milestone_idx);
}
}
}
#[inline]
pub (crate) fn milestone(&self, symbol: &Symbol) -> u64 {
unsafe{
return *self.milestones.get_unchecked(symbol.index() as usize);
}
}
pub (crate) fn milestones(&self) -> &[u64] {
&self.milestones
}
pub (crate) fn bit_vectors(&self) -> &[Vec256] {
&self.bit_vectors
}
#[inline]
pub (crate) fn global_occurrence(&self, local_query_position: u64, symbol: &Symbol) -> u64 {
unsafe{
let milestone_count = self.milestone(&symbol);
let vec0 = self.bit_vectors.get_unchecked(0).to_simd();
let vec1 = self.bit_vectors.get_unchecked(1).to_simd();
let vec2 = self.bit_vectors.get_unchecked(2).to_simd();
let occurrence_vector = match &symbol.index() {
1 => simd_and(vec1, vec2), 2 => simd_and(vec0, vec2), 3 => simd_and(vec0, vec1), 4 => simd_andnot(vec2, simd_andnot(vec0, vec1)), 5 => simd_andnot(vec2, simd_andnot(vec1, vec0)), _ => {
panic!("illegal letter index given in global occurrence function symbol idx given: {}", symbol.index());
} };
let popcount = masked_popcount(occurrence_vector, local_query_position);
return milestone_count + popcount as u64;
}
}
}
impl AminoBwtBlock {
pub (crate) const NUM_MILESTONES: usize = 24;
pub (crate) const NUM_BIT_VECTORS: usize = 5;
pub (crate) fn new() -> Self {
AminoBwtBlock {
bit_vectors: [Vec256::new(); Self::NUM_BIT_VECTORS],
milestones: [0; Self::NUM_MILESTONES],
}
}
pub (crate) fn from_data(
bit_vectors: [Vec256; Self::NUM_BIT_VECTORS],
milestones: [u64; Self::NUM_MILESTONES],
) -> Self {
AminoBwtBlock {
milestones,
bit_vectors,
}
}
pub (crate) fn symbol_at(&self, position_block: u64)->Symbol{
let mut bit_vector_encoding: u64 = 0;
unsafe{
for bit in 0..self.bit_vectors.len() {
let bit_value = self.bit_vectors.get_unchecked(bit).extract_bit(&position_block);
bit_vector_encoding |= bit_value << bit;
}
}
Symbol::new_bit_vector(SymbolAlphabet::Amino, bit_vector_encoding as u8)
}
pub (crate) fn set_symbol_at(&mut self, symbol:&Symbol, position_in_block: u64){
let mut encoded_symbol = symbol.bit_vector();
let mut bit_vector_idx = 0;
unsafe{
while encoded_symbol != 0 {
if encoded_symbol & 0x1 == 1 {
self.bit_vectors.get_unchecked_mut(bit_vector_idx).set_bit(&position_in_block);
}
encoded_symbol >>= 1;
bit_vector_idx += 1;
}
}
}
#[inline]
pub (crate) fn set_milestones(&mut self, values: &Vec<u64>) {
debug_assert!(values.len() >= SymbolAlphabet::Amino.cardinality() as usize);
for milestone_idx in 0..SymbolAlphabet::Amino.cardinality() as usize {
unsafe{
*self.milestones.get_unchecked_mut(milestone_idx) = *values.get_unchecked(milestone_idx);
}
}
}
#[inline]
pub (crate) fn milestone(&self, symbol: &Symbol) -> u64 {
unsafe{
return *self.milestones.get_unchecked(symbol.index() as usize);
}
}
pub (crate) fn milestones(&self) -> &[u64; Self::NUM_MILESTONES] {
&self.milestones
}
pub (crate) fn bit_vectors(&self) -> &[Vec256; Self::NUM_BIT_VECTORS] {
&self.bit_vectors
}
#[inline]
pub (crate) fn global_occurrence(&self, local_query_position: SearchPtr, symbol: &Symbol) -> SearchPtr {
let milestone_count = self.milestone(symbol);
unsafe{
let vec0 = self.bit_vectors.get_unchecked(0).to_simd();
let vec1 = self.bit_vectors.get_unchecked(1).to_simd();
let vec2 = self.bit_vectors.get_unchecked(2).to_simd();
let vec3 = self.bit_vectors.get_unchecked(3).to_simd();
let vec4 = self.bit_vectors.get_unchecked(4).to_simd();
let occurrence_vector = match symbol.index() {
1 => simd_and(vec2,simd_andnot(vec4, vec3)), 2 => simd_andnot(vec3, simd_and(simd_and(vec0, vec1), vec2)), 3 => simd_andnot(vec4, simd_and(vec0, vec1)), 4 => simd_andnot(vec4, simd_and(vec1, vec2)), 5 => simd_andnot(vec0, simd_and(simd_and(vec1, vec2), vec3)), 6 => simd_andnot(vec2, simd_andnot(vec0, vec4)), 7 => simd_andnot(vec2, simd_and(vec0, simd_and(vec1, vec3))), 8 => simd_andnot(vec2, simd_andnot(vec1, vec4)), 9 => simd_andnot(vec1, simd_andnot(vec3, vec4)), 10 => simd_andnot(vec1, simd_andnot(vec0, vec4)), 11 => simd_andnot(vec1, simd_and(vec3, simd_and(vec2, vec0))), 12 => simd_andnot(simd_or(vec0, vec1), simd_andnot(vec2, vec3)), 13 => simd_and(vec3, simd_andnot(vec4, vec0)), 14 => simd_andnot(simd_or(vec0, vec1), simd_andnot(vec3, vec2)), 15 => simd_andnot(vec2, simd_andnot(vec3, vec4)), 16 => simd_and(vec1, simd_andnot(vec4, vec3)), 17 => simd_and(vec0, simd_andnot(vec4, vec2)), 18 => simd_andnot(vec3, simd_andnot(vec0, vec4)), 19 => simd_andnot(simd_or(vec1, vec2), simd_andnot(vec3, vec0)), 20 => simd_and(simd_and(vec0, vec1), simd_and(vec2, vec3)), 21 => simd_andnot(simd_or(vec0, vec2), simd_andnot(vec3, vec1)), _ => {
panic!("illegal letter index given in global occurrence function");
}
};
let popcount = masked_popcount(occurrence_vector, local_query_position);
return milestone_count + popcount as u64;
}
}
}
#[derive(Clone, Serialize, Deserialize, Debug, PartialEq, PartialOrd, Eq, Ord, Hash, MemSize)]
#[serde(untagged)]
pub (crate) enum Bwt {
Nucleotide(Vec<NucleotideBwtBlock>),
Amino(Vec<AminoBwtBlock>),
}
impl Bwt {
pub (crate) const NUM_SYMBOLS_PER_BLOCK: u64 = 256;
pub (crate) fn set_symbol_at(&mut self, bwt_position: &SearchPtr, symbol: &Symbol) {
let bwt_block_idx = bwt_position / Self::NUM_SYMBOLS_PER_BLOCK;
let position_in_block = bwt_position % Self::NUM_SYMBOLS_PER_BLOCK;
unsafe{
match self{
Bwt::Nucleotide(vec) => vec.get_unchecked_mut(bwt_block_idx as usize).set_symbol_at(symbol, position_in_block),
Bwt::Amino(vec) => vec.get_unchecked_mut(bwt_block_idx as usize).set_symbol_at(symbol, position_in_block),
}
}
}
pub(crate) fn num_blocks(bwt_len: u64) -> usize {
bwt_len.div_ceil(Self::NUM_SYMBOLS_PER_BLOCK as u64) as usize
}
pub (crate) fn symbol_at(&self, bwt_position: &SearchPtr) -> Symbol {
let position_block_idx = bwt_position / Self::NUM_SYMBOLS_PER_BLOCK;
let position_in_block = bwt_position % Self::NUM_SYMBOLS_PER_BLOCK;
unsafe{
match &self {
Bwt::Nucleotide(vec) => {
let bwt_block = vec.get_unchecked(position_block_idx as usize);
bwt_block.symbol_at(position_in_block)
}
Bwt::Amino(vec) => {
let bwt_block = vec.get_unchecked(position_block_idx as usize);
bwt_block.symbol_at(position_in_block)
}
}
}
}
pub (crate) fn set_milestones(&mut self, block_idx: usize, counts: &Vec<u64>) {
unsafe{
match self {
Bwt::Nucleotide(vec) => vec.get_unchecked_mut(block_idx).set_milestones(counts),
Bwt::Amino(vec) => vec.get_unchecked_mut(block_idx).set_milestones(counts),
}
}
}
pub (crate) fn global_occurrence(
&self,
pointer_global_position: SearchPtr,
symbol: &Symbol,
) -> SearchPtr {
let block_idx: u64 = pointer_global_position / Self::NUM_SYMBOLS_PER_BLOCK;
let local_query_position: u64 = pointer_global_position % Self::NUM_SYMBOLS_PER_BLOCK;
unsafe{
match self {
Bwt::Nucleotide(vec) => {
vec.get_unchecked(block_idx as usize).global_occurrence(local_query_position, symbol)
}
Bwt::Amino(vec) => {
vec.get_unchecked(block_idx as usize).global_occurrence(local_query_position, symbol)
}
}
}
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use rand::{Rng, SeedableRng};
use crate::{alphabet::{Symbol, SymbolAlphabet}, simd_instructions::Vec256};
use super::{AminoBwtBlock, NucleotideBwtBlock};
#[test]
fn mock_nucleotide_empty_bwt_block_test() {
let mock_bwt_block = NucleotideBwtBlock::from_data(
[Vec256::new();NucleotideBwtBlock::NUM_BIT_VECTORS],
[
1000u64, 2000u64, 3000u64, 4000u64, 5000u64, 6000u64, 7000u64, 8000u64,
],
);
for symbol_idx in 1..6 {
for position in 0..256 {
let occurrence_count = mock_bwt_block.global_occurrence(
position,
&Symbol::new_index(crate::alphabet::SymbolAlphabet::Nucleotide, symbol_idx),
);
assert_eq!(occurrence_count, mock_bwt_block.milestones[symbol_idx as usize],
"nucleotide occurrence did not exactly match milestone in empty bwt block, count for sym {}, pos {}.", symbol_idx, position);
}
}
}
#[test]
fn mock_nucleotide_preset_bwt_block_test(){
let mut mock_bwt_block = NucleotideBwtBlock::from_data(
[Vec256::new();NucleotideBwtBlock::NUM_BIT_VECTORS],
[
1000u64, 2000u64, 3000u64, 4000u64, 5000u64, 6000u64, 7000u64, 8000u64,
],
);
let mut seeded_rng = rand::rngs::StdRng::seed_from_u64(2);
let mut counts:HashMap<(u64,u8), u64> = HashMap::new();
let mut current_counts:Vec<u64> = mock_bwt_block.milestones.to_vec();
for position in 0..super::Bwt::NUM_SYMBOLS_PER_BLOCK{
let symbol_idx = seeded_rng.gen_range(0..SymbolAlphabet::Nucleotide.cardinality());
let symbol = &crate::alphabet::Symbol::new_index(SymbolAlphabet::Nucleotide, symbol_idx as u8);
mock_bwt_block.set_symbol_at(symbol, position);
current_counts[symbol_idx as usize] += 1;
for idx in 0..SymbolAlphabet::Nucleotide.cardinality(){
counts.insert((position, idx as u8), current_counts[idx as usize]);
}
}
for symbol_idx in 1..SymbolAlphabet::Nucleotide.cardinality() {
for position in 0..256 {
let occurrence_count = mock_bwt_block.global_occurrence(
position,
&Symbol::new_index(crate::alphabet::SymbolAlphabet::Nucleotide, symbol_idx),
);
let expected_value = counts.get(&(position, symbol_idx)).expect("failed to get value from hash table");
assert_eq!(occurrence_count, *expected_value,
"nucleotide occurrence did not exactly match milestone in randomized bwt block, count for sym {}, pos {}.", symbol_idx, position);
}
}
}
#[test]
fn mock_amino_empty_bwt_block_test() {
let mock_bwt_block = AminoBwtBlock::from_data(
[Vec256::new();AminoBwtBlock::NUM_BIT_VECTORS],
[
1000u64, 2000u64, 3000u64, 4000u64, 5000u64, 6000u64, 7000u64, 8000u64,
9000u64, 10000u64, 11000u64, 12000u64, 13000u64, 14000u64, 15000u64, 16000u64,
17000u64, 18000u64, 19000u64, 20000u64, 21000u64, 22000u64, 23000u64, 24000u64,
],
);
for symbol_idx in 1..6 {
for position in 0..256 {
let occurrence_count = mock_bwt_block.global_occurrence(
position,
&Symbol::new_index(crate::alphabet::SymbolAlphabet::Amino, symbol_idx),
);
assert_eq!(occurrence_count, mock_bwt_block.milestones[symbol_idx as usize],
"amino occurrence did not exactly match milestone in empty bwt block, count for sym {}, pos {}.", symbol_idx, position);
}
}
}
#[test]
fn mock_amino_preset_bwt_block_test(){
let mut mock_bwt_block = AminoBwtBlock::from_data(
[Vec256::new();AminoBwtBlock::NUM_BIT_VECTORS],
[
1000u64, 2000u64, 3000u64, 4000u64, 5000u64, 6000u64, 7000u64, 8000u64,
9000u64, 10000u64, 11000u64, 12000u64, 13000u64, 14000u64, 15000u64, 16000u64,
17000u64, 18000u64, 19000u64, 20000u64, 21000u64, 22000u64, 23000u64, 24000u64,
],
);
let mut counts:HashMap<(u64,u8), u64> = HashMap::new();
let mut current_counts:Vec<u64> = mock_bwt_block.milestones.to_vec();
let mut seeded_rng = rand::rngs::StdRng::seed_from_u64(6);
for position in 0..super::Bwt::NUM_SYMBOLS_PER_BLOCK{
let symbol_idx = seeded_rng.gen_range(0..SymbolAlphabet::Amino.cardinality());
let symbol = &crate::alphabet::Symbol::new_index(SymbolAlphabet::Amino, symbol_idx as u8);
mock_bwt_block.set_symbol_at(symbol, position);
current_counts[symbol_idx as usize] += 1;
for idx in 0..SymbolAlphabet::Amino.cardinality(){
counts.insert((position, idx), current_counts[idx as usize]);
}
}
for symbol_idx in 1..SymbolAlphabet::Amino.cardinality() {
for position in 0..256 {
let occurrence_count = mock_bwt_block.global_occurrence(
position,
&Symbol::new_index(crate::alphabet::SymbolAlphabet::Amino, symbol_idx),
);
let expected_value = counts.get(&(position, symbol_idx)).expect("failed to get value from hash table");
assert_eq!(occurrence_count, *expected_value,
"amino occurrence did not exactly match milestone in randomized bwt block, count for sym {}, pos {}.", symbol_idx, position);
}
}
}
}