use crate::donnelly_stem_layout::donnelly_get_idx_v2_branchless;
use crate::stem_strategies::donnelly_4::DonnellyFullArith;
use crate::traits::Axis;
use crate::StemStrategy;
use aligned_vec::AVec;
use cmov::Cmov;
pub type Donnelly3X86F64 = Donnelly<3, 64, 8>;
pub type Donnelly4X86F32 = Donnelly<4, 64, 4>;
pub type Donnelly4M2F64 = Donnelly<4, 128, 8>;
pub type Donnelly5M2F32 = Donnelly<5, 128, 4>;
#[derive(Copy, Clone)]
pub struct Donnelly<const L: u32, const CL: u32, const VB: u32> {
minor_level: u32,
min_idx: u32,
maj_idx: u32,
base_maj: u32,
}
impl<const L: u32, const CL: u32, const VB: u32> Donnelly<L, CL, VB> {
#[inline(never)]
const fn items_per_line() -> u32 {
CL / VB
}
#[inline(never)]
const fn line_mask() -> u32 {
Self::items_per_line() - 1
}
}
impl<const L: u32, const CL: u32, const VB: u32> StemStrategy for Donnelly<L, CL, VB> {
#[inline(never)]
fn new_query() -> Self {
debug_assert!(L >= 2 && L <= 8);
debug_assert!(CL > VB);
Self {
minor_level: 0,
min_idx: 0,
maj_idx: 0,
base_maj: 0,
}
}
#[inline(never)]
fn get_child_idx(&mut self, is_right_child: bool, curr_idx: usize) -> usize {
let result =
donnelly_get_idx_v2_branchless(curr_idx as u32, is_right_child, self.minor_level);
self.minor_level += 1;
self.minor_level.cmovnz(&0, u8::from(self.minor_level == 3));
result as usize
}
#[inline(never)]
fn get_both_child_idx(&mut self, curr_idx: usize) -> (usize, usize) {
let left = donnelly_get_idx_v2_branchless(curr_idx as u32, false, self.minor_level);
let right = donnelly_get_idx_v2_branchless(curr_idx as u32, true, self.minor_level);
self.minor_level += 1;
self.minor_level.cmovnz(&0, u8::from(self.minor_level == 3));
(left as usize, right as usize)
}
fn get_closer_and_further_child_idx(
&mut self,
curr_idx: usize,
is_right_child: bool,
) -> (usize, usize) {
let (left, right) = self.get_both_child_idx(curr_idx);
if is_right_child {
(right, left)
} else {
(left, right)
}
}
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);
}
}
}
#[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 = Donnelly::<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);
}
}