use std::ops::Range;
#[derive(Debug, Clone)]
pub struct BlockPartition {
pub block_ranges: Vec<Range<usize>>,
}
impl BlockPartition {
pub fn num_blocks(&self) -> usize {
self.block_ranges.len()
}
pub fn num_features(&self) -> usize {
self.block_ranges.last().map_or(0, |r| r.end)
}
pub fn feature_to_block(&self) -> Vec<u32> {
let p = self.num_features();
let mut mapping = vec![0u32; p];
for (b, range) in self.block_ranges.iter().enumerate() {
for j in range.clone() {
mapping[j] = b as u32;
}
}
mapping
}
}
impl BlockPartition {
pub fn regular(num_features: usize, block_size: usize) -> Self {
assert!(block_size > 0, "block_size must be > 0");
assert!(num_features > 0, "num_features must be > 0");
let num_blocks = num_features.div_ceil(block_size);
let block_ranges: Vec<Range<usize>> = (0..num_blocks)
.map(|b| {
let start = b * block_size;
let end = ((b + 1) * block_size).min(num_features);
start..end
})
.collect();
Self { block_ranges }
}
pub fn from_boundaries(boundaries: &[usize], num_features: usize) -> Self {
assert!(!boundaries.is_empty(), "need at least one boundary");
assert_eq!(boundaries[0], 0, "first boundary must be 0");
let num_blocks = boundaries.len();
let block_ranges: Vec<Range<usize>> = (0..num_blocks)
.map(|b| {
let start = boundaries[b];
let end = if b + 1 < num_blocks {
boundaries[b + 1]
} else {
num_features
};
start..end
})
.collect();
Self { block_ranges }
}
pub fn build_hierarchy(num_features: usize, block_size: usize) -> Vec<Self> {
assert!(block_size > 1, "block_size must be > 1");
let mut levels = Vec::new();
let mut current = num_features;
while current > block_size {
let partition = Self::regular(current, block_size);
let next = partition.num_blocks();
levels.push(partition);
current = next;
}
levels
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_regular_even() {
let part = BlockPartition::regular(100, 10);
assert_eq!(part.num_blocks(), 10);
assert_eq!(part.num_features(), 100);
assert_eq!(part.block_ranges[0], 0..10);
assert_eq!(part.block_ranges[9], 90..100);
}
#[test]
fn test_regular_uneven() {
let part = BlockPartition::regular(103, 10);
assert_eq!(part.num_blocks(), 11);
assert_eq!(part.block_ranges[10], 100..103);
}
#[test]
fn test_from_boundaries() {
let part = BlockPartition::from_boundaries(&[0, 20, 50, 80], 100);
assert_eq!(part.num_blocks(), 4);
assert_eq!(part.block_ranges[0], 0..20);
assert_eq!(part.block_ranges[1], 20..50);
assert_eq!(part.block_ranges[2], 50..80);
assert_eq!(part.block_ranges[3], 80..100);
}
#[test]
fn test_build_hierarchy() {
let levels = BlockPartition::build_hierarchy(50, 100);
assert!(levels.is_empty());
let levels = BlockPartition::build_hierarchy(1000, 100);
assert_eq!(levels.len(), 1);
assert_eq!(levels[0].num_features(), 1000);
assert_eq!(levels[0].num_blocks(), 10);
let levels = BlockPartition::build_hierarchy(1_000_000, 100);
assert_eq!(levels.len(), 2);
assert_eq!(levels[0].num_features(), 1_000_000);
assert_eq!(levels[0].num_blocks(), 10_000);
assert_eq!(levels[1].num_features(), 10_000);
assert_eq!(levels[1].num_blocks(), 100);
}
#[test]
fn test_contiguous_coverage() {
let part = BlockPartition::regular(57, 10);
let mut covered = [false; 57];
for range in &part.block_ranges {
for j in range.clone() {
assert!(!covered[j], "feature {} covered twice", j);
covered[j] = true;
}
}
assert!(covered.iter().all(|&c| c), "not all features covered");
}
}