use crate::donnelly_stem_layout::donnelly_get_idx_v2_branchless;
use crate::traits::Axis;
use crate::StemStrategy;
use aligned_vec::AVec;
use cmov::Cmov;
#[derive(Copy, Clone)]
pub struct DonnellyFullArithCombined<
const L: u32,
const CACHELINE_BYTES: u32,
const VALUE_BYTES: u32,
> {
combined_idx: u32,
minor_index: u32,
}
impl<const L: u32, const CL: u32, const VB: u32> StemStrategy
for DonnellyFullArithCombined<L, CL, VB>
{
#[inline(always)]
fn new_query() -> Self {
debug_assert!(L >= 2 && L <= 8);
Self {
combined_idx: 0,
minor_index: 0,
}
}
#[inline(always)]
fn get_child_idx(&mut self, is_right_child: bool, _curr_idx: usize) -> usize {
let (new_combined, new_minor, idx) =
Self::get_child_idx_pure(self.combined_idx, self.minor_index, is_right_child);
self.combined_idx = new_combined;
self.minor_index = new_minor;
idx
}
fn get_both_child_idx(&mut self, _child_idx: usize) -> (usize, usize) {
unimplemented!()
}
fn get_closer_and_further_child_idx(
&mut self,
curr_idx: usize,
is_right_child: bool,
) -> (usize, usize) {
unimplemented!()
}
fn get_initial_idx() -> usize {
0
}
fn get_stem_node_count_from_leaf_node_count(leaf_node_count: usize) -> usize {
if leaf_node_count < 2 {
0
} else {
leaf_node_count.next_power_of_two() - 1
}
}
fn stem_node_padding_factor() -> usize {
5
}
fn trim_unneeded_stems<A: Axis>(stems: &mut AVec<A>, max_stem_level: usize) {
if !stems.is_empty() {
let mut level: usize = 0;
let mut minor_level: u64 = 0;
let mut stem_idx = 0;
loop {
let val = &stems[stem_idx];
let is_right_child = val.is_finite();
stem_idx = donnelly_get_idx_v2_branchless(
stem_idx as u32,
is_right_child,
minor_level as u32,
) as usize;
level += 1;
minor_level += 1;
minor_level.cmovnz(&0, u8::from(minor_level == 3));
if level == max_stem_level {
break;
}
}
stems.truncate(stem_idx + 1);
}
}
}
impl<const L: u32, const CL: u32, const VB: u32> DonnellyFullArithCombined<L, CL, VB> {
const LOG2_L: u32 = const_ceil_log2::<L>();
#[inline(always)]
pub fn get_child_idx_pure(
combined_idx: u32, minor_index: u32, is_right_child: bool,
) -> (
u32, /*new_combined*/
u32, /*new_minor*/
usize, /*child idx*/
) {
let bits_for_minor: u32 = const_ceil_log2::<L>();
debug_assert!(L >= 2 && L <= 8);
let right_flag: u32 = is_right_child as u32;
let major_idx: u32 = combined_idx >> bits_for_minor;
let base_major: u32 = major_idx << L; let base_majorL: u32 = major_idx << (2 * L);
let incr = ((minor_index << 1) + 1) + right_flag;
let new_minor = minor_index + incr;
let same_block = base_major + new_minor;
let min_row = minor_index.wrapping_sub(L - 1);
let next_term = ((min_row << 1) + 1 + right_flag) << L;
let next_block = base_majorL + next_term;
let new_combined = combined_idx + 1;
let wrap_mask = (1u32 << bits_for_minor) - 1;
let wrapped = (new_combined & wrap_mask) == 0;
let result = if wrapped { next_block } else { same_block };
(new_combined, new_minor, result as usize)
}
}
const fn const_ceil_log2<const L: u32>() -> u32 {
let mut v = L - 1;
let mut bits = 0;
while v > 0 {
v >>= 1;
bits += 1;
}
bits
}
#[inline(never)]
pub fn calc_child_idx(
combined_idx: u32,
minor_index: u32,
is_right_child: bool,
) -> (u32, u32, usize) {
DonnellyFullArithCombined::<3, 64, 4>::get_child_idx_pure(
combined_idx,
minor_index,
is_right_child,
)
}
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
#[rstest]
#[case(vec![], 0)]
#[case(vec![false], 1)] #[case(vec![true], 2)] #[case(vec![false, false], 3)] #[case(vec![false, true], 4)] #[case(vec![true, false], 5)] #[case(vec![true, true], 6)] #[case(vec![false, false, false], 8)] #[case(vec![false, false, true], 16)] #[case(vec![false, true, false], 24)] #[case(vec![false, true, true], 32)] #[case(vec![true, false, false], 40)] #[case(vec![true, false, true], 48)] #[case(vec![true, true, false], 56)] #[case(vec![true, true, true], 64)] #[case(vec![false, false, false, false], 9)] #[case(vec![false, false, false, true], 10)] #[case(vec![false, false, false, false, false], 11)] #[case(vec![false, false, false, false, true], 12)] #[case(vec![false, false, false, true, false], 13)] #[case(vec![false, false, false, true, true], 14)] #[case(vec![false, false, false, false, false, false], 72)] #[case(vec![false, false, false, false, false, true], 80)] #[case(vec![false, false, false, false, true, false], 88)] #[case(vec![false, false, false, false, true, true], 96)] #[case(vec![false, false, false, true, false, false], 104)] #[case(vec![false, false, false, true, false, true], 112)] #[case(vec![false, false, false, true, true, false], 120)] #[case(vec![false, false, false, true, true, true], 128)] #[case(vec![false, false, true, false], 17)] #[case(vec![false, false, true, true], 18)] #[case(vec![false, false, true, false, false], 19)] #[case(vec![false, false, true, false, true], 20)] #[case(vec![false, false, true, true, false], 21)] #[case(vec![false, false, true, true, true], 22)] #[case(vec![false, false, true, false, false, false], 136)] #[case(vec![false, false, true, false, false, true], 144)] #[case(vec![false, false, true, false, true, false], 152)] #[case(vec![false, false, true, false, true, true], 160)] #[case(vec![false, false, true, true, false, false], 168)] #[case(vec![false, false, true, true, false, true], 176)] #[case(vec![false, false, true, true, true, false], 184)] #[case(vec![false, false, true, true, true, true], 192)] fn donnelly_v2_get_child_idx_produces_correct_values(
#[case] input: Vec<bool>,
#[case] expected: usize,
) {
let mut stem_strat = DonnellyFullArithCombined::<3, 64, 8>::new_query();
let mut result = 0;
input.iter().for_each(|selection| {
result = stem_strat.get_child_idx(*selection, result);
});
assert_eq!(result, expected);
}
}