pub mod cpu;
#[cfg(feature = "gpu")]
pub mod gpu;
use crate::alphabet::ALPHABET_SIZE;
pub const BLOCK_SIZE: u32 = 64;
pub const SUPERBLOCK_SIZE: u32 = 512;
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct OccTable {
num_lanes: u8,
num_planes: u8,
symbol_to_lane: [u8; ALPHABET_SIZE],
lane_to_symbol: [u8; ALPHABET_SIZE],
superblock_checkpoints: Vec<u32>,
block_data: Vec<u8>,
block_stride: usize,
#[serde(default)]
encoding: OccEncoding,
pub text_len: u32,
}
const NO_LANE: u8 = u8::MAX;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub enum OccEncoding {
#[default]
Bitplane,
OneHot,
}
fn num_planes_for(num_lanes: u8) -> u8 {
if num_lanes <= 1 {
0
} else {
(u8::BITS - (num_lanes - 1).leading_zeros()) as u8
}
}
impl OccTable {
pub fn from_parts(
num_lanes: u8,
symbol_to_lane: [u8; ALPHABET_SIZE],
superblock_checkpoints: Vec<u32>,
block_deltas: Vec<u16>,
lane_data: Vec<u64>,
text_len: u32,
encoding: OccEncoding,
) -> Self {
let mut lane_to_symbol = [0u8; ALPHABET_SIZE];
for (c, &lane) in symbol_to_lane.iter().enumerate() {
if lane != NO_LANE {
lane_to_symbol[lane as usize] = c as u8;
}
}
let num_planes = num_planes_for(num_lanes);
let num_lanes_usize = num_lanes as usize;
let num_planes_usize = num_planes as usize;
let lane_data_width = match encoding {
OccEncoding::Bitplane => num_planes_usize,
OccEncoding::OneHot => num_lanes_usize,
};
let block_stride = num_lanes_usize * 4 + num_lanes_usize * 2 + lane_data_width * 8;
let num_blocks = block_deltas.len().checked_div(num_lanes_usize).unwrap_or(0);
let blocks_per_sb = (SUPERBLOCK_SIZE / BLOCK_SIZE) as usize;
let mut block_data = vec![0u8; num_blocks * block_stride];
for b in 0..num_blocks {
let sb = b / blocks_per_sb.max(1);
let rec = &mut block_data[b * block_stride..(b + 1) * block_stride];
for lane in 0..num_lanes_usize {
let sb_count = superblock_checkpoints[sb * num_lanes_usize + lane];
rec[lane * 4..lane * 4 + 4].copy_from_slice(&sb_count.to_ne_bytes());
}
let deltas_off = num_lanes_usize * 4;
for lane in 0..num_lanes_usize {
let delta = block_deltas[b * num_lanes_usize + lane];
rec[deltas_off + lane * 2..deltas_off + lane * 2 + 2]
.copy_from_slice(&delta.to_ne_bytes());
}
let lane_data_off = deltas_off + num_lanes_usize * 2;
for w in 0..lane_data_width {
let word = lane_data[b * lane_data_width + w];
rec[lane_data_off + w * 8..lane_data_off + w * 8 + 8]
.copy_from_slice(&word.to_ne_bytes());
}
}
Self {
num_lanes,
num_planes,
symbol_to_lane,
lane_to_symbol,
superblock_checkpoints,
block_data,
block_stride,
encoding,
text_len,
}
}
#[inline]
fn block_base(&self, block: usize) -> usize {
block * self.block_stride
}
#[inline]
fn sb_count_at(&self, base: usize, lane: usize) -> u32 {
let off = base + lane * 4;
debug_assert!(off + 4 <= self.block_data.len());
unsafe {
self.block_data
.as_ptr()
.add(off)
.cast::<u32>()
.read_unaligned()
}
}
#[inline]
fn delta_at(&self, base: usize, lane: usize) -> u32 {
let num_lanes = self.num_lanes as usize;
let off = base + num_lanes * 4 + lane * 2;
debug_assert!(off + 2 <= self.block_data.len());
unsafe {
self.block_data
.as_ptr()
.add(off)
.cast::<u16>()
.read_unaligned() as u32
}
}
#[inline]
fn word_at(&self, base: usize, w: usize) -> u64 {
let num_lanes = self.num_lanes as usize;
let off = base + num_lanes * 4 + num_lanes * 2 + w * 8;
debug_assert!(off + 8 <= self.block_data.len());
unsafe {
self.block_data
.as_ptr()
.add(off)
.cast::<u64>()
.read_unaligned()
}
}
pub fn num_lanes(&self) -> u8 {
self.num_lanes
}
#[inline]
pub(crate) fn prefetch_block(&self, pos: u32) {
let block = (pos / BLOCK_SIZE) as usize;
let base = self.block_base(block);
if base < self.block_data.len() {
crate::prefetch::prefetch_read(unsafe { self.block_data.as_ptr().add(base) });
}
}
#[inline]
fn lane_mask(&self, base: usize, lane: usize) -> u64 {
if self.encoding == OccEncoding::OneHot {
if self.num_lanes as usize == 0 {
return 0;
}
return self.word_at(base, lane);
}
let num_planes = self.num_planes as usize;
if num_planes == 0 {
return u64::MAX;
}
let mut mask = u64::MAX;
for p in 0..num_planes {
let plane_val = self.word_at(base, p);
mask &= if (lane >> p) & 1 == 1 {
plane_val
} else {
!plane_val
};
}
mask
}
#[inline]
fn lane_at(&self, base: usize, offset: u32) -> u8 {
if self.encoding == OccEncoding::OneHot {
let num_lanes = self.num_lanes as usize;
for lane in 0..num_lanes {
if (self.word_at(base, lane) >> offset) & 1 == 1 {
return lane as u8;
}
}
return 0;
}
let num_planes = self.num_planes as usize;
if num_planes == 0 {
return 0;
}
let mut lane = 0u8;
for p in 0..num_planes {
let bit = (self.word_at(base, p) >> offset) & 1;
lane |= (bit as u8) << p;
}
lane
}
pub fn rank(&self, c: u8, i: u32) -> u32 {
if i == 0 {
return 0;
}
let lane = self.symbol_to_lane[c as usize];
if lane == NO_LANE {
return 0;
}
self.rank_at_lane(lane as usize, i)
}
#[inline]
fn rank_at_lane(&self, lane: usize, i: u32) -> u32 {
let pos = i - 1;
let block = (pos / BLOCK_SIZE) as usize;
let offset = pos % BLOCK_SIZE;
let base = self.block_base(block);
let sb_count = self.sb_count_at(base, lane);
let delta = self.delta_at(base, lane);
let lane_bits = self.lane_mask(base, lane);
let mask = if offset == 63 {
u64::MAX
} else {
(1u64 << (offset + 1)) - 1
};
sb_count + delta + (lane_bits & mask).count_ones()
}
#[inline]
pub fn rank_pair(&self, c: u8, lo: u32, hi: u32) -> (u32, u32) {
let lane = self.symbol_to_lane[c as usize];
if lane == NO_LANE {
return (0, 0);
}
let lane = lane as usize;
if lo != 0 {
self.prefetch_block(lo - 1);
}
if hi != 0 {
self.prefetch_block(hi - 1);
}
let rlo = if lo == 0 {
0
} else {
self.rank_at_lane(lane, lo)
};
let rhi = if hi == 0 {
0
} else {
self.rank_at_lane(lane, hi)
};
(rlo, rhi)
}
pub fn rank_many(&self, queries: &[(u8, u32)], out: &mut [u32]) {
debug_assert_eq!(queries.len(), out.len());
for &(_, i) in queries {
if i != 0 {
self.prefetch_block(i - 1);
}
}
for (slot, &(c, i)) in out.iter_mut().zip(queries) {
*slot = self.rank(c, i);
}
}
pub fn symbol_at(&self, pos: u32) -> u8 {
let block = (pos / BLOCK_SIZE) as usize;
let offset = pos % BLOCK_SIZE;
let base = self.block_base(block);
self.lane_to_symbol[self.lane_at(base, offset) as usize]
}
#[inline]
fn lane_and_mask_bitplane(&self, base: usize, offset: u32) -> (u8, u64) {
let num_planes = self.num_planes as usize;
if num_planes == 0 {
return (0, u64::MAX);
}
let mut lane = 0u8;
let mut mask = u64::MAX;
for p in 0..num_planes {
let plane_val = self.word_at(base, p);
let bit = (plane_val >> offset) & 1;
lane |= (bit as u8) << p;
mask &= if bit == 1 { plane_val } else { !plane_val };
}
(lane, mask)
}
#[inline]
pub fn lf_step(&self, pos: u32) -> (u8, u32) {
let block = (pos / BLOCK_SIZE) as usize;
let offset = pos % BLOCK_SIZE;
let base = self.block_base(block);
let (lane, lane_bits) = if self.encoding == OccEncoding::Bitplane {
self.lane_and_mask_bitplane(base, offset)
} else {
let lane = self.lane_at(base, offset);
(lane, self.lane_mask(base, lane as usize))
};
let lane = lane as usize;
let symbol = self.lane_to_symbol[lane];
let sb_count = self.sb_count_at(base, lane);
let delta = self.delta_at(base, lane);
let mask = (1u64 << offset) - 1; let rank = sb_count + delta + (lane_bits & mask).count_ones();
(symbol, rank)
}
pub fn reconstruct_bwt_u32(&self) -> Vec<u32> {
let num_blocks = self.block_data.len() / self.block_stride.max(1);
let mut out = Vec::with_capacity(num_blocks * BLOCK_SIZE as usize);
for b in 0..num_blocks {
let base = self.block_base(b);
for offset in 0..BLOCK_SIZE {
let sym = self.lane_to_symbol[self.lane_at(base, offset) as usize];
out.push(sym as u32);
}
}
out.truncate(self.text_len as usize);
out
}
#[cfg(feature = "gpu")]
pub fn flat_block_checkpoints(&self) -> Vec<[u32; ALPHABET_SIZE]> {
let blocks_per_sb = (SUPERBLOCK_SIZE / BLOCK_SIZE) as usize;
let num_blocks = self.block_data.len() / self.block_stride.max(1);
let num_lanes = self.num_lanes as usize;
(0..num_blocks)
.map(|b| {
let sb = b / blocks_per_sb;
let base = self.block_base(b);
let mut combined = [0u32; ALPHABET_SIZE];
for (c, slot) in combined.iter_mut().enumerate() {
let lane = self.symbol_to_lane[c];
if lane == NO_LANE {
continue;
}
let lane = lane as usize;
*slot = self.superblock_checkpoints[sb * num_lanes + lane]
+ self.delta_at(base, lane);
}
combined
})
.collect()
}
#[cfg(feature = "gpu")]
pub fn bitvectors_full16(&self) -> Vec<[u64; ALPHABET_SIZE]> {
let num_blocks = self.block_data.len() / self.block_stride.max(1);
(0..num_blocks)
.map(|b| {
let base = self.block_base(b);
let mut bv = [0u64; ALPHABET_SIZE];
for (c, slot) in bv.iter_mut().enumerate() {
let lane = self.symbol_to_lane[c];
if lane == NO_LANE {
continue;
}
*slot = self.lane_mask(base, lane as usize);
}
bv
})
.collect()
}
}