use crate::{
alphabet::{Symbol, SymbolAlphabet},
search::SearchPtr,
simd_instructions::{Vec256, SimdVec256},
};
use serde::{Deserialize, Serialize};
#[derive(Clone, Serialize, Deserialize, Debug, PartialEq, PartialOrd, Eq, Ord, Hash, Default)]
#[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)]
#[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);
for milestone_idx in 0..SymbolAlphabet::Nucleotide.cardinality() as usize {
self.milestones[milestone_idx] = values[milestone_idx];
}
}
#[inline]
pub (crate) fn milestone(&self, symbol: &Symbol) -> u64 {
return self.milestones[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 {
let milestone_count = self.milestone(&symbol);
let vec0 = SimdVec256::from(self.bit_vectors[0]);
let vec1 = SimdVec256::from(self.bit_vectors[1]);
let vec2 = SimdVec256::from(self.bit_vectors[2]);
let occurrence_vector = match &symbol.index() {
1 => vec2.and(&vec1), 2 => vec2.and(&vec0), 3 => vec1.and(&vec0), 4 => vec2.andnot(&vec0.andnot(&vec1)), 5 => vec2.andnot(&vec1.andnot(&vec0)), _ => {
panic!("illegal letter index given in global occurrence function symbol idx given: {}", symbol.index());
} };
let popcount = occurrence_vector.masked_popcount(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;
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::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;
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::Amino.cardinality() as usize);
for milestone_idx in 0..SymbolAlphabet::Amino.cardinality() as usize {
self.milestones[milestone_idx] = values[milestone_idx];
}
}
#[inline]
pub (crate) fn milestone(&self, symbol: &Symbol) -> u64 {
return self.milestones[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);
let vec0 = SimdVec256::from(self.bit_vectors[0]);
let vec1 = SimdVec256::from(self.bit_vectors[1]);
let vec2 = SimdVec256::from(self.bit_vectors[2]);
let vec3 = SimdVec256::from(self.bit_vectors[3]);
let vec4 = SimdVec256::from(self.bit_vectors[4]);
let occurrence_vector = match symbol.index() {
1 => vec3.and(&vec4.andnot(&vec2)), 2 => vec3.andnot(&vec2).and(&vec1.and(&vec0)), 3 => vec1.and(&vec4.andnot(&vec0)), 4 => vec4.andnot(&vec2.and(&vec1)), 5 => vec0.andnot(&vec3).and(&vec2.and(&vec1)), 6 => vec2.andnot(&vec0.andnot(&vec4)), 7 => vec2.andnot(&vec3).and(&vec1.and(&vec0)), 8 => vec2.andnot(&vec1.andnot(&vec4)), 9 => vec3.andnot(&vec1.andnot(&vec4)), 10 => vec1.andnot(&vec0.andnot(&vec4)), 11 => vec1.andnot(&vec3).and(&vec2.and(&vec0)), 12 => vec0.or(&vec1).andnot(&vec2.andnot(&vec3)), 13 => vec3.and(&vec4.andnot(&vec0)), 14 => vec3.or(&vec1).andnot(&vec0.andnot(&vec2)), 15 => vec3.andnot(&vec2.andnot(&vec4)), 16 => vec3.and(&vec4.andnot(&vec1)), 17 => vec2.and(&vec4.andnot(&vec0)), 18 => vec3.andnot(&vec0.andnot(&vec4)), 19 => vec3.or(&vec2).andnot(&vec1.andnot(&vec0)), 20 => vec3.and(&vec2).and(&vec1.and(&vec0)), 21 => vec0.or(&vec2).andnot(&vec3.andnot(&vec1)), _ => {
panic!("illegal letter index given in global occurrence function");
}
};
let popcount = occurrence_vector.masked_popcount(local_query_position);
return milestone_count + popcount as u64;
}
}
#[derive(Clone, Serialize, Deserialize, Debug, PartialEq, PartialOrd, Eq, Ord, Hash)]
#[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;
match self{
Bwt::Nucleotide(vec) => vec[bwt_block_idx as usize].set_symbol_at(symbol, position_in_block),
Bwt::Amino(vec) => vec[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;
match &self {
Bwt::Nucleotide(vec) => {
let bwt_block = &vec[position_block_idx as usize];
bwt_block.symbol_at(position_in_block)
}
Bwt::Amino(vec) => {
let bwt_block = &vec[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>) {
match self {
Bwt::Nucleotide(vec) => vec[block_idx].set_milestones(counts),
Bwt::Amino(vec) => vec[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;
match self {
Bwt::Nucleotide(vec) => {
vec[block_idx as usize].global_occurrence(local_query_position, symbol)
}
Bwt::Amino(vec) => {
vec[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);
}
}
}
}