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 DonnellyFullArith<const L: u32, const CL: u32, const VB: u32> {
minor: u32, min_idx: u32, maj_idx: u32, base_maj: u32, base_majL: u32, }
impl<const L: u32, const CL: u32, const VB: u32> StemStrategy for DonnellyFullArith<L, CL, VB> {
#[inline(always)]
fn new_query() -> Self {
Self {
minor: 0,
min_idx: 0,
maj_idx: 0,
base_maj: 0,
base_majL: 0,
}
}
#[inline(always)]
fn get_child_idx(&mut self, is_right: bool, _curr_idx: usize) -> usize {
self.step(is_right)
}
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> DonnellyFullArith<L, CL, VB> {
#[inline(always)]
const fn items_per_line() -> u32 {
CL / VB
}
#[inline(always)]
const fn log2_items_per_line() -> u32 {
Self::items_per_line().ilog2()
}
#[inline(always)]
const fn last_row_start() -> u32 {
(1u32 << (L - 1)) - 1
}
#[inline(always)]
fn step_pure(
mut minor_lvl: u32, mut min_idx: u32, mut maj_idx: u32, mut base_maj: u32, mut base_majL: u32, is_right: bool,
) -> (usize, u32, u32, u32, u32, u32) {
debug_assert!(L >= 2 && L <= 8);
let right = is_right as u32;
let t = minor_lvl.wrapping_add(1);
let wrapped = (t == L) as u32; minor_lvl = t.wrapping_sub(wrapped * L);
let child_same = (min_idx << 1).wrapping_add(1).wrapping_add(right);
let same = base_maj.wrapping_add(child_same);
let r = min_idx.wrapping_sub(Self::last_row_start());
let child_next = (r << 1).wrapping_add(1).wrapping_add(right);
let base_step = wrapped << Self::log2_items_per_line();
let base_maj_nxt = base_maj.wrapping_add(base_step);
let next = base_maj_nxt.wrapping_add(child_next);
let m = 0u32.wrapping_sub(wrapped); let res = ((same & !m) | (next & m)) as usize;
min_idx = (child_same & !m) | (child_next & m);
maj_idx = maj_idx.wrapping_add(wrapped);
base_maj = base_maj_nxt;
let base_step_L = wrapped << (Self::log2_items_per_line() + L);
base_majL = base_majL.wrapping_add(base_step_L);
(res, minor_lvl, min_idx, maj_idx, base_maj, base_majL)
}
#[inline(always)]
fn step(&mut self, is_right: bool) -> usize {
let (child_idx, minor, min_idx, maj_idx, base_maj, base_majL) = Self::step_pure(
self.minor,
self.min_idx,
self.maj_idx,
self.base_maj,
self.base_majL,
is_right,
);
self.minor = minor;
self.min_idx = min_idx;
self.maj_idx = maj_idx;
self.base_maj = base_maj;
self.base_majL = base_majL;
child_idx
}
}
#[inline(never)]
pub fn calc_child_idx(
minor: u32,
min_idx: u32,
maj_idx: u32,
base_maj: u32,
base_majL: u32,
is_right: bool,
) -> (usize, u32, u32, u32, u32, u32) {
DonnellyFullArith::<3, 64, 4>::step_pure(minor, min_idx, maj_idx, base_maj, base_majL, is_right)
}
#[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 = DonnellyFullArith::<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);
}
}