use generic_array::typenum::Unsigned;
use crate::register::{CoreRegister, NumericRegister, Register, Storage};
macro_rules! lanes {
($r:ty) => {
<<$r as CoreRegister>::Lanes as Unsigned>::USIZE
};
}
macro_rules! use_ladder {
($r:ty) => {
const { <$r as Register>::HAS_NATIVE_ALIGN && lanes!($r).is_power_of_two() && lanes!($r) <= 64 }
};
}
#[rustfmt::skip]
macro_rules! forward_ladder {
($r:ty, $v:expr, $fill:expr, $op:path) => {{
let mut v = $v;
let f = $fill;
let () = match const { lanes!($r) } {
0 | 1 => {}
2 => { v = $op(v, <$r as Register>::align::<1>(f, v)); }
4 => { v = $op(v, <$r as Register>::align::<3>(f, v));
v = $op(v, <$r as Register>::align::<2>(f, v)); }
8 => { v = $op(v, <$r as Register>::align::<7>(f, v));
v = $op(v, <$r as Register>::align::<6>(f, v));
v = $op(v, <$r as Register>::align::<4>(f, v)); }
16 => { v = $op(v, <$r as Register>::align::<15>(f, v));
v = $op(v, <$r as Register>::align::<14>(f, v));
v = $op(v, <$r as Register>::align::<12>(f, v));
v = $op(v, <$r as Register>::align::<8>(f, v)); }
32 => { v = $op(v, <$r as Register>::align::<31>(f, v));
v = $op(v, <$r as Register>::align::<30>(f, v));
v = $op(v, <$r as Register>::align::<28>(f, v));
v = $op(v, <$r as Register>::align::<24>(f, v));
v = $op(v, <$r as Register>::align::<16>(f, v)); }
64 => { v = $op(v, <$r as Register>::align::<63>(f, v));
v = $op(v, <$r as Register>::align::<62>(f, v));
v = $op(v, <$r as Register>::align::<60>(f, v));
v = $op(v, <$r as Register>::align::<56>(f, v));
v = $op(v, <$r as Register>::align::<48>(f, v));
v = $op(v, <$r as Register>::align::<32>(f, v)); }
_ => unreachable!(),
};
v
}};
}
#[rustfmt::skip]
macro_rules! reverse_ladder {
($r:ty, $v:expr, $fill:expr, $op:path) => {{
let mut v = $v;
let f = $fill;
let () = {
if const { lanes!($r) > 1 } { v = $op(v, <$r as Register>::align::<1>(v, f)); }
if const { lanes!($r) > 2 } { v = $op(v, <$r as Register>::align::<2>(v, f)); }
if const { lanes!($r) > 4 } { v = $op(v, <$r as Register>::align::<4>(v, f)); }
if const { lanes!($r) > 8 } { v = $op(v, <$r as Register>::align::<8>(v, f)); }
if const { lanes!($r) > 16 } { v = $op(v, <$r as Register>::align::<16>(v, f)); }
if const { lanes!($r) > 32 } { v = $op(v, <$r as Register>::align::<32>(v, f)); }
};
v
}};
}
macro_rules! scalar_scan {
($name:ident, forward, $combine:expr) => {
#[inline(always)]
pub fn $name<R: NumericRegister>(mut value: Storage<R>) -> Storage<R> {
let s = R::as_mut_slice(&mut value);
let mut i = 1;
while i < s.len() {
let acc = s[i - 1];
s[i] = $combine(s[i], acc);
i += 1;
}
value
}
};
($name:ident, reverse, $combine:expr) => {
#[inline(always)]
pub fn $name<R: NumericRegister>(mut value: Storage<R>) -> Storage<R> {
let s = R::as_mut_slice(&mut value);
let mut i = s.len().saturating_sub(1);
while i > 0 {
i -= 1;
let acc = s[i + 1];
s[i] = $combine(s[i], acc);
}
value
}
};
}
scalar_scan!(scalar_prefix_sum, forward, |cur, acc| cur + acc);
scalar_scan!(scalar_prefix_min, forward, |cur, acc| if cur < acc { cur } else { acc });
scalar_scan!(scalar_prefix_max, forward, |cur, acc| if cur > acc { cur } else { acc });
scalar_scan!(scalar_reverse_prefix_sum, reverse, |cur, acc| cur + acc);
scalar_scan!(scalar_reverse_prefix_min, reverse, |cur, acc| if cur < acc {
cur
} else {
acc
});
scalar_scan!(scalar_reverse_prefix_max, reverse, |cur, acc| if cur > acc {
cur
} else {
acc
});
#[inline(always)]
fn first_lane<R: Register>(value: Storage<R>) -> Storage<R> {
R::broadcast::<0>(value)
}
#[inline(always)]
fn last_lane<R: Register>(value: Storage<R>) -> Storage<R> {
R::broadcast::<0>(R::reverse(value))
}
#[inline(always)]
pub fn prefix_sum<R: NumericRegister>(value: Storage<R>) -> Storage<R> {
if const { !use_ladder!(R) } {
return scalar_prefix_sum::<R>(value);
}
forward_ladder!(R, value, R::ZERO, R::add)
}
#[inline(always)]
pub fn prefix_min<R: NumericRegister>(value: Storage<R>) -> Storage<R> {
if const { !use_ladder!(R) } {
return scalar_prefix_min::<R>(value);
}
forward_ladder!(R, value, first_lane::<R>(value), R::min)
}
#[inline(always)]
pub fn prefix_max<R: NumericRegister>(value: Storage<R>) -> Storage<R> {
if const { !use_ladder!(R) } {
return scalar_prefix_max::<R>(value);
}
forward_ladder!(R, value, first_lane::<R>(value), R::max)
}
#[inline(always)]
pub fn reverse_prefix_sum<R: NumericRegister>(value: Storage<R>) -> Storage<R> {
if const { !use_ladder!(R) } {
return scalar_reverse_prefix_sum::<R>(value);
}
reverse_ladder!(R, value, R::ZERO, R::add)
}
#[inline(always)]
pub fn reverse_prefix_min<R: NumericRegister>(value: Storage<R>) -> Storage<R> {
if const { !use_ladder!(R) } {
return scalar_reverse_prefix_min::<R>(value);
}
reverse_ladder!(R, value, last_lane::<R>(value), R::min)
}
#[inline(always)]
pub fn reverse_prefix_max<R: NumericRegister>(value: Storage<R>) -> Storage<R> {
if const { !use_ladder!(R) } {
return scalar_reverse_prefix_max::<R>(value);
}
reverse_ladder!(R, value, last_lane::<R>(value), R::max)
}