use super::{QVector, QVectorIterator};
use crate::utils::{prefetch_read_NTA, select_in_word_u128};
use mem_dbg::{MemDbg, MemSize};
use num_traits::int::PrimInt;
use num_traits::{AsPrimitive, Unsigned};
use serde::{Deserialize, Serialize};
use crate::{AccessQuad, RankQuad, SelectQuad, WTSupport};
mod rs_support_plain;
use crate::qvector::rs_qvector::rs_support_plain::RSSupportPlain;
pub type RSQVector256 = RSQVector<RSSupportPlain<256>>;
pub type RSQVector512 = RSQVector<RSSupportPlain<512>>;
#[derive(Default, Clone, PartialEq, Debug, Serialize, MemSize, MemDbg, Deserialize)]
pub struct RSQVector<S> {
qv: QVector,
rs_support: S,
n_occs_smaller: [usize; 5], }
impl<S> RSQVector<S> {
pub fn iter(&self) -> QVectorIterator<&QVector> {
self.qv.iter()
}
}
impl<S: RSSupport> From<QVector> for RSQVector<S> {
fn from(qv: QVector) -> Self {
let rank_support = S::new(&qv);
let mut n_occs_smaller = [0; 5];
for c in qv.iter() {
n_occs_smaller[c as usize] += 1;
}
let mut prev = n_occs_smaller[0];
n_occs_smaller[0] = 0;
for i in 1..5 {
let tmp = n_occs_smaller[i];
n_occs_smaller[i] = n_occs_smaller[i - 1] + prev;
prev = tmp;
}
Self {
qv,
rs_support: rank_support,
n_occs_smaller,
}
}
}
impl<S: RSSupport> RSQVector<S> {
pub fn new<T>(v: &[T]) -> Self
where
T: Unsigned + Copy,
QVector: FromIterator<T>,
{
let qv: QVector = v.iter().copied().collect();
Self::from(qv)
}
#[inline]
fn select_intra_block(&self, symbol: u8, i: usize, pos: usize) -> usize {
let line_id = pos >> 8;
let mut result = 0;
let mut i = i - 1;
for j in 0..if S::BLOCK_SIZE == 256 { 1 } else { 2 } {
let (word_0, word_1) =
unsafe { self.qv.data.get_unchecked(line_id + j).normalize(symbol) };
let cnt_0 = word_0.count_ones() as usize;
if cnt_0 > i {
let p = select_in_word_u128(word_0, i as u64) as usize;
return result + p;
} else {
i -= cnt_0;
result += 128;
}
let cnt_1 = word_1.count_ones() as usize;
if cnt_1 > i {
return result + select_in_word_u128(word_1, i as u64) as usize;
} else {
i -= cnt_1;
result += 128;
}
}
0
}
#[inline]
fn rank_intra_block(&self, symbol: u8, i: usize) -> usize {
debug_assert!(
symbol <= 3,
"RSQVector indexes only four symbols in [0, 3]."
);
debug_assert!(
S::BLOCK_SIZE == 256 || S::BLOCK_SIZE == 512,
"RSQVector supports only blocks of size 256 or 512."
);
if S::BLOCK_SIZE == 256 {
let data_line_id = i >> 8;
let offset = i & 255;
let rank = if let Some(d) = self.qv.data.get(data_line_id) {
unsafe { d.rank_unchecked(symbol, offset) }
} else {
0
};
return rank;
}
if S::BLOCK_SIZE == 512 {
let block_id = i >> 9;
let offset_in_block = i & 511;
let offset_in_first_block = if offset_in_block <= 256 {
offset_in_block
} else {
256
};
let mut rank = if let Some(d) = self.qv.data.get(block_id * 2) {
unsafe { d.rank_unchecked(symbol, offset_in_first_block) }
} else {
0
};
if offset_in_block > 256 {
rank += if let Some(d) = self.qv.data.get(block_id * 2 + 1) {
unsafe { d.rank_unchecked(symbol, offset_in_block - 256) }
} else {
0
};
}
return rank;
}
0
}
pub fn len(&self) -> usize {
self.qv.len()
}
pub fn is_empty(&self) -> bool {
self.qv.len() == 0
}
}
impl<S> AccessQuad for RSQVector<S> {
#[inline]
unsafe fn get_unchecked(&self, i: usize) -> u8 {
self.qv.get_unchecked(i)
}
#[inline]
fn get(&self, i: usize) -> Option<u8> {
self.qv.get(i)
}
}
impl<S: RSSupport> RankQuad for RSQVector<S> {
#[inline(always)]
fn rank(&self, symbol: u8, i: usize) -> Option<usize> {
if i > self.qv.len() {
return None;
}
Some(unsafe { self.rank_unchecked(symbol, i) })
}
#[inline(always)]
unsafe fn rank_unchecked(&self, symbol: u8, i: usize) -> usize {
debug_assert!(symbol <= 3);
self.rs_support.rank_block(symbol, i) + self.rank_intra_block(symbol, i)
}
}
impl<S: RSSupport> SelectQuad for RSQVector<S> {
#[inline]
fn select(&self, symbol: u8, i: usize) -> Option<usize> {
if symbol > 3 || unsafe { self.occs_unchecked(symbol) } <= i {
return None;
}
let (mut pos, rank) = self.rs_support.select_block(symbol, i + 1);
pos += self.select_intra_block(symbol, i - rank + 1, pos);
Some(pos)
}
#[inline]
unsafe fn select_unchecked(&self, symbol: u8, i: usize) -> usize {
debug_assert!(symbol <= 3);
debug_assert!(i > 0);
debug_assert!(self.occs(symbol) <= Some(i));
self.select(symbol, i).unwrap()
}
}
impl<S: RSSupport> WTSupport for RSQVector<S> {
#[inline(always)]
fn occs(&self, symbol: u8) -> Option<usize> {
if symbol > 3 {
return None;
}
Some(unsafe { self.occs_unchecked(symbol) })
}
#[inline(always)]
unsafe fn occs_unchecked(&self, symbol: u8) -> usize {
debug_assert!(symbol <= 3, "Symbols are in [0, 3].");
self.n_occs_smaller[(symbol + 1) as usize] - self.n_occs_smaller[symbol as usize]
}
#[inline(always)]
fn occs_smaller(&self, symbol: u8) -> Option<usize> {
if symbol > 3 {
return None;
}
Some(unsafe { self.occs_smaller_unchecked(symbol) })
}
#[inline(always)]
unsafe fn occs_smaller_unchecked(&self, symbol: u8) -> usize {
debug_assert!(symbol <= 3, "Symbols are in [0, 3].");
self.n_occs_smaller[symbol as usize]
}
#[inline(always)]
unsafe fn rank_block_unchecked(&self, symbol: u8, i: usize) -> usize {
self.rs_support.rank_block(symbol, i)
}
#[inline(always)]
fn prefetch_info(&self, pos: usize) {
self.rs_support.prefetch(pos)
}
#[inline(always)]
fn prefetch_data(&self, pos: usize) {
let line_id = pos >> 8;
prefetch_read_NTA(&self.qv.data, line_id);
if S::BLOCK_SIZE == 512 {
prefetch_read_NTA(&self.qv.data, if line_id > 0 { line_id - 1 } else { 0 });
}
}
}
impl<S> AsRef<RSQVector<S>> for RSQVector<S> {
fn as_ref(&self) -> &RSQVector<S> {
self
}
}
impl<S> IntoIterator for RSQVector<S> {
type IntoIter = QVectorIterator<QVector>;
type Item = u8;
fn into_iter(self) -> Self::IntoIter {
self.qv.into_iter()
}
}
impl<'a, S> IntoIterator for &'a RSQVector<S> {
type IntoIter = QVectorIterator<&'a QVector>;
type Item = u8;
fn into_iter(self) -> Self::IntoIter {
self.qv.iter()
}
}
pub trait RSSupport {
const BLOCK_SIZE: usize;
fn new(qv: &QVector) -> Self;
fn rank_block(&self, symbol: u8, i: usize) -> usize;
fn select_block(&self, symbol: u8, i: usize) -> (usize, usize);
fn prefetch(&self, pos: usize);
}
impl<T, S: RSSupport> FromIterator<T> for RSQVector<S>
where
T: PrimInt + AsPrimitive<u8>,
{
fn from_iter<I>(iter: I) -> Self
where
I: IntoIterator<Item = T>,
{
Self::from(QVector::from_iter(iter))
}
}
#[cfg(test)]
#[generic_tests::define]
mod tests {
use super::*;
use std::iter;
#[test]
fn test_just_one_data_line<D>()
where
D: From<QVector> + AccessQuad + RankQuad + SelectQuad + WTSupport,
{
let qv: QVector = [0].into_iter().cycle().take(256).collect(); let rsqv = D::from(qv);
assert_eq!(rsqv.rank(0, 256), Some(256));
assert_eq!(rsqv.rank(1, 256), Some(0));
assert_eq!(rsqv.rank(2, 256), Some(0));
assert_eq!(rsqv.rank(3, 256), Some(0));
}
#[test]
fn test_small<D>()
where
D: From<QVector> + AccessQuad + RankQuad + SelectQuad + WTSupport,
{
let qv: QVector = [0, 1, 2, 3].into_iter().cycle().take(10000).collect();
let rsqv = D::from(qv.clone());
for c in 0..4 {
assert_eq!(rsqv.occs(c), Some(10000 / 4));
}
assert_eq!(rsqv.occs(4), None);
for c in 0..4 {
assert_eq!(rsqv.occs_smaller(c), Some((10000 / 4) * (c as usize)));
}
for (i, c) in qv.iter().enumerate() {
assert_eq!(rsqv.get(i), Some(c));
}
for i in 0..qv.len() {
let r = rsqv.rank(0, i);
let extra = if i % 4 > 0 { 1 } else { 0 };
assert_eq!(r, Some(i / 4 + extra));
}
for symbol in 0..4 {
let r = rsqv.rank(symbol, qv.len());
let cnt = qv.iter().filter(|x| *x == symbol).count();
assert_eq!(r, Some(cnt));
let r = rsqv.rank(symbol, qv.len() + 1);
assert_eq!(r, None);
}
for (i, c) in qv.iter().enumerate() {
let rank = rsqv.rank(c, i).unwrap();
let s = rsqv.select(c, rank).unwrap();
assert_eq!(s, i);
}
}
#[test]
fn test_boundaries<D>()
where
D: From<QVector> + AccessQuad + RankQuad + SelectQuad,
{
for n in [
100,
255,
256,
257,
511,
512,
513,
1024,
1025,
2047,
2048,
2049,
256 * 8 - 1,
256 * 8,
256 * 8 + 1,
512 * 8,
512 * 8 + 1,
] {
for symbol in 0..3u8 {
let qv: QVector = iter::repeat(symbol).take(n).collect();
let rsqv = D::from(qv.clone());
for i in 0..qv.len() + 1 {
if i < qv.len() {
assert_eq!(rsqv.get(i), Some(symbol));
}
assert_eq!(rsqv.rank(symbol, i), Some(i));
assert_eq!(rsqv.rank(3, i), Some(0));
}
}
}
}
#[test]
fn test_from_wt<D>()
where
D: From<QVector> + AccessQuad + RankQuad + SelectQuad,
{
let n = 1025 * 2;
let qv: QVector = (0..n).map(|x| if x < n / 2 { 0 } else { 1 }).collect();
let rsqv = D::from(qv);
for i in 0..n {
if i < n / 2 {
assert_eq!(rsqv.get(i), Some(0));
assert_eq!(rsqv.rank(0, i), Some(i));
assert_eq!(rsqv.rank(1, i), Some(0));
} else {
assert_eq!(rsqv.get(i), Some(1));
assert_eq!(rsqv.rank(0, i), Some(n / 2));
assert_eq!(rsqv.rank(1, i), Some(i - n / 2));
}
}
}
#[instantiate_tests(<RSQVector256>)]
mod testp256 {}
#[instantiate_tests(<RSQVector512>)]
mod testp512 {}
}