use crate::utils::{msb, stable_partition_of_4};
use crate::{
AccessUnsigned, OccsRangeUnsigned, RankUnsigned, SelectUnsigned, WTIterator, WTSupport,
};
use crate::{QVector, QVectorBuilder};
use mem_dbg::{MemDbg, MemSize};
use serde::{Deserialize, Serialize};
use std::marker::PhantomData;
use num_traits::{AsPrimitive, PrimInt, Unsigned};
use std::ops::{Bound, Range, RangeBounds, Shl, Shr};
pub mod huffqwt;
mod prefetch_support;
use crate::quadwt::prefetch_support::PrefetchSupport;
pub trait RSforWT: From<QVector> + WTSupport + MemSize + MemDbg + Default {}
impl<T> RSforWT for T where T: From<QVector> + WTSupport + MemSize + MemDbg + Default {}
pub trait WTIndexable:
Unsigned + PrimInt + Ord + Shr<usize> + Shl<usize> + AsPrimitive<usize>
{
}
impl<T> WTIndexable for T
where
T: Unsigned + PrimInt + Ord + Shr<usize> + Shl<usize> + AsPrimitive<usize> + AsPrimitive<u8>,
u8: AsPrimitive<T>,
usize: AsPrimitive<T>,
{
}
#[derive(Default, Clone, PartialEq, Debug, Serialize, MemSize, MemDbg, Deserialize)]
pub struct QWaveletTree<T, RS, const WITH_PREFETCH_SUPPORT: bool = false> {
n: usize, n_levels: usize, sigma: T, qvs: Vec<RS>, prefetch_support: Option<Vec<PrefetchSupport>>,
}
impl<T, RS, const WITH_PREFETCH_SUPPORT: bool> QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>
where
T: WTIndexable,
usize: AsPrimitive<T>,
RS: RSforWT,
{
#[must_use]
pub fn new(sequence: &mut [T]) -> Self {
if sequence.is_empty() {
return Self {
n: 0,
n_levels: 0,
sigma: T::zero(),
qvs: vec![RS::default()],
prefetch_support: None,
};
}
let sigma = *sequence.iter().max().unwrap();
let log_sigma = msb(sigma) + 1; let n_levels = log_sigma.div_ceil(2) as usize;
let mut prefetch_support = Vec::with_capacity(n_levels);
let mut qvs = Vec::<RS>::with_capacity(n_levels);
let mut shift = 2 * (n_levels - 1);
for _level in 0..n_levels {
let mut cur_qv = QVectorBuilder::with_capacity(sequence.len());
for &symbol in sequence.iter() {
let two_bits: u8 = ((symbol >> shift).as_() & 3) as u8; cur_qv.push(two_bits);
}
let qv = cur_qv.build();
if WITH_PREFETCH_SUPPORT {
let pfs = PrefetchSupport::new(&qv, 11); prefetch_support.push(pfs);
}
qvs.push(RS::from(qv));
stable_partition_of_4(sequence, shift);
if shift >= 2 {
shift -= 2;
}
}
qvs.shrink_to_fit();
Self {
n: sequence.len(),
n_levels,
sigma,
qvs,
prefetch_support: if WITH_PREFETCH_SUPPORT {
Some(prefetch_support)
} else {
None
},
}
}
#[must_use]
pub fn len(&self) -> usize {
self.n
}
#[must_use]
pub fn sigma(&self) -> Option<T> {
if self.is_empty() {
None
} else {
Some(self.sigma)
}
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.n == 0
}
#[must_use]
pub fn n_levels(&self) -> usize {
self.n_levels
}
pub fn iter(
&self,
) -> WTIterator<
T,
QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>,
&QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>,
> {
WTIterator {
i: 0,
end: self.len(),
qwt: self,
_phantom: PhantomData,
}
}
#[inline]
unsafe fn rank_prefetch_superblocks_unchecked(&self, symbol: T, i: usize) -> usize {
if !WITH_PREFETCH_SUPPORT {
return 0;
}
if let Some(ref prefetch_support) = self.prefetch_support {
let mut shift: i64 = (2 * (self.n_levels - 1)) as i64;
let mut range = 0..i;
self.qvs[0].prefetch_data(range.end);
self.qvs[0].prefetch_info(range.start);
self.qvs[0].prefetch_info(range.end);
#[allow(clippy::needless_range_loop)]
for level in 0..self.n_levels - 1 {
let two_bits: u8 = ((symbol >> shift as usize).as_() & 3) as u8;
let offset = self.qvs[level].occs_smaller_unchecked(two_bits);
let rank_start =
prefetch_support[level].approx_rank_unchecked(two_bits, range.start);
let rank_end = prefetch_support[level].approx_rank_unchecked(two_bits, range.end);
range = (rank_start + offset)..(rank_end + offset);
self.qvs[level + 1].prefetch_info(range.start);
self.qvs[level + 1].prefetch_info(range.start + 2048);
self.qvs[level + 1].prefetch_info(range.end);
self.qvs[level + 1].prefetch_info(range.end + 2048);
if level > 0 {
self.qvs[level + 1].prefetch_info(range.start + 2 * 2048);
self.qvs[level + 1].prefetch_info(range.end + 2 * 2048);
self.qvs[level + 1].prefetch_info(range.end + 3 * 2048);
}
shift -= 2;
}
return range.end - range.start;
}
0
}
#[inline(always)]
#[must_use]
pub fn rank_prefetch(&self, symbol: T, i: usize) -> Option<usize> {
if i > self.n || symbol > self.sigma {
return None;
}
Some(unsafe { self.rank_prefetch_unchecked(symbol, i) })
}
#[must_use]
#[inline(always)]
pub unsafe fn rank_prefetch_unchecked(&self, symbol: T, i: usize) -> usize {
if WITH_PREFETCH_SUPPORT {
let _ = self.rank_prefetch_superblocks_unchecked(symbol, i);
}
let mut range = 0..i;
let mut shift: i64 = (2 * (self.n_levels - 1)) as i64;
const BLOCK_SIZE: usize = 256;
self.qvs[0].prefetch_data(range.start);
self.qvs[0].prefetch_data(range.end);
for level in 0..self.n_levels - 1 {
let two_bits: u8 = ((symbol >> shift as usize).as_() & 3) as u8;
let offset = self.qvs[level].occs_smaller_unchecked(two_bits);
let rank_start = self.qvs[level].rank_block_unchecked(two_bits, range.start);
let rank_end = self.qvs[level].rank_block_unchecked(two_bits, range.end);
range = (rank_start + offset)..(rank_end + offset);
self.qvs[level + 1].prefetch_data(range.start);
self.qvs[level + 1].prefetch_data(range.start + BLOCK_SIZE);
self.qvs[level + 1].prefetch_data(range.end);
self.qvs[level + 1].prefetch_data(range.end + BLOCK_SIZE);
for i in 0..level {
self.qvs[level + 1].prefetch_data(range.end + 2 * BLOCK_SIZE + i * BLOCK_SIZE);
}
shift -= 2;
}
self.rank_unchecked(symbol, i)
}
}
impl<T, RS, const WITH_PREFETCH_SUPPORT: bool> OccsRangeUnsigned
for QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>
where
T: WTIndexable,
usize: AsPrimitive<T>,
RS: RSforWT,
{
type Iter<'a>
= OccsRangeIter<'a, T, RS, WITH_PREFETCH_SUPPORT>
where
Self: 'a;
fn occs_range<R: RangeBounds<usize>>(&self, range: R) -> Option<Self::Iter<'_>> {
let start = match range.start_bound() {
Bound::Included(start) => *start,
Bound::Excluded(start) => *start + 1,
Bound::Unbounded => 0,
};
let end = match range.end_bound() {
Bound::Included(end) => *end + 1,
Bound::Excluded(end) => *end,
Bound::Unbounded => self.n,
};
if end > self.n || start > end {
return None;
}
Some(unsafe { self.occs_range_unchecked(start..end) })
}
unsafe fn occs_range_unchecked(&self, range: Range<usize>) -> Self::Iter<'_> {
if range.start == range.end {
let stack = vec![];
return OccsRangeIter { tree: self, stack };
}
let mut stack = Vec::with_capacity(self.n_levels * 3 + 1);
stack.push(OccsRangeFrame {
range,
level: 0,
bit_path: 0,
});
OccsRangeIter { tree: self, stack }
}
}
pub struct OccsRangeIter<'a, T, RS, const WITH_PREFETCH_SUPPORT: bool> {
tree: &'a QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>,
stack: Vec<OccsRangeFrame>,
}
struct OccsRangeFrame {
range: Range<usize>,
level: usize,
bit_path: usize,
}
impl<'a, T, RS, const WITH_PREFETCH_SUPPORT: bool> Iterator
for OccsRangeIter<'a, T, RS, WITH_PREFETCH_SUPPORT>
where
T: WTIndexable,
usize: AsPrimitive<T>,
RS: RSforWT,
{
type Item = (T, usize);
#[inline]
fn next(&mut self) -> Option<Self::Item> {
while let Some(cur) = self.stack.pop() {
if cur.level == self.tree.n_levels {
return Some((cur.bit_path.as_(), cur.range.end - cur.range.start));
}
let qv = unsafe { self.tree.qvs.get_unchecked(cur.level) };
for bit in (0..4u8).rev() {
let lo = unsafe { qv.rank_unchecked(bit, cur.range.start) };
let hi = unsafe { qv.rank_unchecked(bit, cur.range.end) };
if hi == lo {
continue;
}
let offset = unsafe { qv.occs_smaller_unchecked(bit) };
let frame = OccsRangeFrame {
range: offset + lo..offset + hi,
level: cur.level + 1,
bit_path: (cur.bit_path << 2) | (bit as usize),
};
self.stack.push(frame);
}
}
None
}
}
impl<T, RS, const WITH_PREFETCH_SUPPORT: bool> RankUnsigned
for QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>
where
T: WTIndexable,
usize: AsPrimitive<T>,
RS: RSforWT,
{
#[inline(always)]
fn rank(&self, symbol: Self::Item, i: usize) -> Option<usize> {
if i > self.n || symbol > self.sigma {
return None;
}
Some(unsafe { self.rank_unchecked(symbol, i) })
}
#[inline(always)]
unsafe fn rank_unchecked(&self, symbol: Self::Item, i: usize) -> usize {
let mut shift: i64 = (2 * (self.n_levels - 1)) as i64;
let mut cur_i = i;
let mut cur_p = 0;
for level in 0..self.n_levels - 1 {
let two_bits: u8 = ((symbol >> shift as usize).as_() & 3) as u8;
let offset = unsafe { self.qvs[level].occs_smaller_unchecked(two_bits) };
cur_p = self.qvs[level].rank_unchecked(two_bits, cur_p) + offset;
cur_i = self.qvs[level].rank_unchecked(two_bits, cur_i) + offset;
shift -= 2;
}
let two_bits: u8 = ((symbol >> shift as usize).as_() & 3) as u8;
cur_i = self.qvs[self.n_levels - 1].rank_unchecked(two_bits, cur_i);
cur_p = self.qvs[self.n_levels - 1].rank_unchecked(two_bits, cur_p);
cur_i - cur_p
}
}
impl<T, RS, const WITH_PREFETCH_SUPPORT: bool> AccessUnsigned
for QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>
where
T: WTIndexable,
usize: AsPrimitive<T>,
RS: RSforWT,
{
type Item = T;
#[inline(always)]
fn get(&self, i: usize) -> Option<Self::Item> {
if i >= self.n {
return None;
}
Some(unsafe { self.get_unchecked(i) })
}
#[inline(always)]
unsafe fn get_unchecked(&self, i: usize) -> Self::Item {
let mut result = T::zero();
let mut cur_i = i;
for level in 0..self.n_levels - 1 {
self.qvs[level].prefetch_info(cur_i); let symbol = self.qvs[level].get_unchecked(cur_i);
result = (result << 2) | (symbol as usize).as_();
let offset = unsafe { self.qvs[level].occs_smaller_unchecked(symbol) };
cur_i = self.qvs[level].rank_unchecked(symbol, cur_i) + offset;
}
let symbol = self.qvs[self.n_levels - 1].get_unchecked(cur_i);
(result << 2) | (symbol as usize).as_()
}
}
impl<T, RS, const WITH_PREFETCH_SUPPORT: bool> SelectUnsigned
for QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>
where
T: WTIndexable,
usize: AsPrimitive<T>,
RS: RSforWT,
{
#[inline(always)]
fn select(&self, symbol: Self::Item, i: usize) -> Option<usize> {
if symbol > self.sigma {
return None;
}
let mut path_off = Vec::with_capacity(self.n_levels);
let mut rank_path_off = Vec::with_capacity(self.n_levels);
let mut b = 0;
let mut shift: i64 = 2 * (self.n_levels - 1) as i64;
for level in 0..self.n_levels {
path_off.push(b);
let two_bits = (symbol >> shift as usize).as_() & 3;
let rank_b = self.qvs[level].rank(two_bits as u8, b)?;
b = rank_b + unsafe { self.qvs[level].occs_smaller_unchecked(two_bits as u8) };
shift -= 2;
rank_path_off.push(rank_b);
}
shift = 0;
let mut result = i;
for level in (0..self.n_levels).rev() {
b = path_off[level];
let rank_b = rank_path_off[level];
let two_bits = (symbol >> shift as usize).as_() & 3;
result = self.qvs[level].select(two_bits as u8, rank_b + result)? - b;
shift += 2;
}
Some(result)
}
#[inline(always)]
unsafe fn select_unchecked(&self, symbol: Self::Item, i: usize) -> usize {
self.select(symbol, i).unwrap()
}
}
impl<T, RS, const WITH_PREFETCH_SUPPORT: bool> AsRef<QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>>
for QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>
{
fn as_ref(&self) -> &QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT> {
self
}
}
impl<T, RS, const WITH_PREFETCH_SUPPORT: bool> IntoIterator
for QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>
where
T: WTIndexable,
usize: AsPrimitive<T>,
RS: RSforWT,
{
type IntoIter = WTIterator<
T,
QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>,
QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>,
>;
type Item = T;
fn into_iter(self) -> Self::IntoIter {
WTIterator {
i: 0,
end: self.len(),
qwt: self,
_phantom: PhantomData,
}
}
}
impl<'a, T, RS, const WITH_PREFETCH_SUPPORT: bool> IntoIterator
for &'a QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>
where
T: WTIndexable,
usize: AsPrimitive<T>,
RS: RSforWT,
{
type IntoIter = WTIterator<
T,
QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>,
&'a QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>,
>;
type Item = T;
fn into_iter(self) -> Self::IntoIter {
self.iter()
}
}
impl<T, RS, const WITH_PREFETCH_SUPPORT: bool> FromIterator<T>
for QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>
where
T: WTIndexable,
usize: AsPrimitive<T>,
RS: RSforWT,
{
fn from_iter<I>(iter: I) -> Self
where
I: IntoIterator<Item = T>,
{
QWaveletTree::new(&mut iter.into_iter().collect::<Vec<T>>())
}
}
impl<T, RS, const WITH_PREFETCH_SUPPORT: bool> From<Vec<T>>
for QWaveletTree<T, RS, WITH_PREFETCH_SUPPORT>
where
T: WTIndexable,
usize: AsPrimitive<T>,
RS: RSforWT,
{
fn from(mut v: Vec<T>) -> Self {
QWaveletTree::new(&mut v[..])
}
}
#[cfg(test)]
mod tests;