#[derive(Debug, Clone)]
pub(crate) struct LocalIDMapper {
bits: Vec<u64>,
bit_count: u64,
rank: Vec<u64>,
}
impl LocalIDMapper {
#[must_use]
pub(crate) fn new(words: &[u64], bit_count: u64) -> Self {
let num_512_blocks = usize::try_from(bit_count.div_ceil(512)).expect("fits");
let mut rank = vec![0u64; num_512_blocks + 1];
let padded_len = num_512_blocks * 8;
let mut bits = words.to_vec();
bits.resize(padded_len, 0u64);
let mut s = 0u64;
for (j, slot) in rank.iter_mut().take(num_512_blocks).enumerate() {
*slot = s;
let base = j * 8;
let block_count: u64 = bits[base..base + 8]
.iter()
.map(|w| u64::from(w.count_ones()))
.sum();
s += block_count;
}
rank[num_512_blocks] = s;
Self {
bits,
bit_count,
rank,
}
}
#[must_use]
#[inline]
pub(crate) fn global_id_count(&self) -> u64 {
self.bit_count
}
#[must_use]
#[inline]
pub(crate) fn local_id_count(&self) -> u64 {
self.rank.last().copied().unwrap_or(0)
}
#[must_use]
#[inline]
pub(crate) fn is_global_id_mapped(&self, global_id: u64) -> bool {
assert!(global_id < self.bit_count, "global_id out of bounds");
let word_idx = usize::try_from(global_id / 64).expect("fits");
(self.bits[word_idx] >> (global_id % 64)) & 1 == 1
}
#[must_use]
pub(crate) fn to_local(&self, global_id: u64) -> u64 {
assert!(global_id < self.bit_count, "global_id out of bounds");
assert!(
self.is_global_id_mapped(global_id),
"global_id is not mapped"
);
self.rank_of(global_id)
}
#[must_use]
pub(crate) fn to_local_or(&self, global_id: u64, invalid: u64) -> u64 {
if global_id >= self.bit_count {
return invalid;
}
let word_idx = usize::try_from(global_id / 64).expect("fits");
let bit_pos = global_id % 64;
if (self.bits[word_idx] >> bit_pos) & 1 == 0 {
return invalid;
}
self.rank_of(global_id)
}
fn rank_of(&self, global_id: u64) -> u64 {
let uint64_index = usize::try_from(global_id / 64).expect("fits");
let uint64_offset = global_id % 64;
let uint512_index = usize::try_from(global_id / 512).expect("fits");
let mut local_id = self.rank[uint512_index];
let block_start = uint512_index * 8;
for word in &self.bits[block_start..uint64_index] {
local_id += u64::from(word.count_ones());
}
let mask = if uint64_offset == 0 {
0u64
} else {
(1u64 << uint64_offset) - 1
};
local_id += u64::from((self.bits[uint64_index] & mask).count_ones());
local_id
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::internal::bitvec::BitVector;
fn brute_force_to_local(bits: &[bool], g: usize) -> u64 {
u64::try_from(bits[..g].iter().filter(|&&b| b).count()).unwrap()
}
fn build_mapper_from_bools(bits: &[bool]) -> LocalIDMapper {
let n = u64::try_from(bits.len()).unwrap();
let mut bv = BitVector::new(n);
for (i, &b) in bits.iter().enumerate() {
if b {
bv.set(u64::try_from(i).unwrap());
}
}
LocalIDMapper::new(bv.words(), n)
}
#[test]
fn all_mapped_small() {
let bits = vec![true; 10];
let mapper = build_mapper_from_bools(&bits);
assert_eq!(mapper.global_id_count(), 10);
assert_eq!(mapper.local_id_count(), 10);
for i in 0u64..10 {
assert!(mapper.is_global_id_mapped(i));
assert_eq!(mapper.to_local(i), i);
}
}
#[test]
fn none_mapped_small() {
let bits = vec![false; 10];
let mapper = build_mapper_from_bools(&bits);
assert_eq!(mapper.global_id_count(), 10);
assert_eq!(mapper.local_id_count(), 0);
for i in 0u64..10 {
assert!(!mapper.is_global_id_mapped(i));
assert_eq!(mapper.to_local_or(i, u64::MAX), u64::MAX);
}
}
#[test]
fn alternating_vs_brute_force() {
let n = 150usize;
let bits: Vec<bool> = (0..n).map(|i| i % 2 == 0).collect();
let mapper = build_mapper_from_bools(&bits);
assert_eq!(mapper.global_id_count(), u64::try_from(n).unwrap());
assert_eq!(
mapper.local_id_count(),
u64::try_from(bits.iter().filter(|&&b| b).count()).unwrap()
);
for i in 0..n {
let gi = u64::try_from(i).unwrap();
assert_eq!(mapper.is_global_id_mapped(gi), bits[i]);
let expected = brute_force_to_local(&bits, i);
if bits[i] {
assert_eq!(mapper.to_local(gi), expected, "to_local({i})");
assert_eq!(
mapper.to_local_or(gi, u64::MAX),
expected,
"to_local_or({i})"
);
} else {
assert_eq!(
mapper.to_local_or(gi, u64::MAX),
u64::MAX,
"to_local_or unmapped({i})"
);
}
}
}
#[test]
fn word_boundary_bits() {
let n = 130usize;
let mut bits = vec![false; n];
bits[0] = true;
bits[63] = true;
bits[64] = true;
bits[65] = true;
bits[127] = true;
bits[129] = true;
let mapper = build_mapper_from_bools(&bits);
assert_eq!(
mapper.local_id_count(),
u64::try_from(bits.iter().filter(|&&b| b).count()).unwrap()
);
for i in 0..n {
let gi = u64::try_from(i).unwrap();
assert_eq!(mapper.is_global_id_mapped(gi), bits[i]);
let expected = brute_force_to_local(&bits, i);
if bits[i] {
assert_eq!(mapper.to_local(gi), expected, "to_local({i})");
}
}
}
#[test]
fn block_boundary_512() {
let n = 600usize;
let mut bits = vec![false; n];
bits[510] = true;
bits[511] = true;
bits[512] = true;
bits[513] = true;
bits[599] = true;
let mapper = build_mapper_from_bools(&bits);
assert_eq!(
mapper.local_id_count(),
u64::try_from(bits.iter().filter(|&&b| b).count()).unwrap()
);
for i in 0..n {
let gi = u64::try_from(i).unwrap();
assert_eq!(mapper.is_global_id_mapped(gi), bits[i]);
let expected = brute_force_to_local(&bits, i);
if bits[i] {
assert_eq!(mapper.to_local(gi), expected, "to_local({i})");
}
}
}
#[test]
fn to_local_or_out_of_range() {
let bits = vec![true; 5];
let mapper = build_mapper_from_bools(&bits);
assert_eq!(mapper.to_local_or(5, 999), 999);
assert_eq!(mapper.to_local_or(100, 42), 42);
}
#[test]
fn large_span_multiple_blocks() {
let n = 1025usize;
let bits: Vec<bool> = (0..n).map(|i| i % 3 == 0).collect();
let mapper = build_mapper_from_bools(&bits);
for i in 0..n {
if bits[i] {
let gi = u64::try_from(i).unwrap();
let expected = brute_force_to_local(&bits, i);
assert_eq!(mapper.to_local(gi), expected, "large span to_local({i})");
}
}
}
#[test]
fn zero_length_mapper() {
let mapper = LocalIDMapper::new(&[], 0);
assert_eq!(mapper.global_id_count(), 0);
assert_eq!(mapper.local_id_count(), 0);
}
#[test]
fn single_bit_set() {
let n = 1_000_000u64;
let mut bv = BitVector::new(n);
bv.set(555_555);
let mapper = LocalIDMapper::new(bv.words(), n);
assert_eq!(mapper.local_id_count(), 1);
assert!(mapper.is_global_id_mapped(555_555));
assert_eq!(mapper.to_local(555_555), 0);
assert_eq!(mapper.to_local_or(555_554, u64::MAX), u64::MAX);
assert_eq!(mapper.to_local_or(555_556, u64::MAX), u64::MAX);
}
}