use generic_array::{ArrayLength, GenericArray, typenum, typenum::Unsigned};
use super::*;
use crate::sort::SortOrder;
use crate::swizzle::SwizzleIndices;
#[inline(always)]
fn cmp_merge<R, O, I, const KEEP_MAX: u64>(v: Storage<R>) -> Storage<R>
where
R: NumericRegister,
O: SortOrder,
I: SwizzleIndices<R::Lanes>,
{
let v_shuf = R::permutev_const::<I>(v);
let v_first = O::first::<R>(v, v_shuf);
let v_last = O::last::<R>(v, v_shuf);
let keep_max = <R::Mask as MaskRegister>::from_native_bitmask(KEEP_MAX);
R::blendv(keep_max, v_first, v_last)
}
macro_rules! layer_indices {
($name:ident, [$($i:expr),* $(,)?]) => {
struct $name<N: ArrayLength>(core::marker::PhantomData<N>);
impl<N: ArrayLength> SwizzleIndices<N> for $name<N> {
const INDICES: GenericArray<u32, N> = const {
let idxs = [$($i as u32),*];
assert!(N::USIZE == idxs.len(), "layer width must equal the lane count");
unsafe { generic_array::const_transmute::<_, GenericArray<u32, N>>(idxs) }
};
}
};
}
#[inline(always)]
pub fn bitonic_clean_2<R, O>(v: Storage<R>) -> Storage<R>
where
O: SortOrder,
R: NumericRegister<Lanes = typenum::U2>,
{
layer_indices!(S1, [1, 0]);
cmp_merge::<R, O, S1<R::Lanes>, 0b10>(v)
}
#[inline(always)]
pub fn bitonic_clean_4<R, O>(v: Storage<R>) -> Storage<R>
where
O: SortOrder,
R: NumericRegister<Lanes = typenum::U4>,
{
layer_indices!(S2, [2, 3, 0, 1]);
layer_indices!(S1, [1, 0, 3, 2]);
let v = cmp_merge::<R, O, S2<R::Lanes>, 0b1100>(v);
cmp_merge::<R, O, S1<R::Lanes>, 0b1010>(v)
}
#[inline(always)]
pub fn bitonic_clean_8<R, O>(v: Storage<R>) -> Storage<R>
where
O: SortOrder,
R: NumericRegister<Lanes = typenum::U8>,
{
layer_indices!(S4, [4, 5, 6, 7, 0, 1, 2, 3]);
layer_indices!(S2, [2, 3, 0, 1, 6, 7, 4, 5]);
layer_indices!(S1, [1, 0, 3, 2, 5, 4, 7, 6]);
let v = cmp_merge::<R, O, S4<R::Lanes>, 0b1111_0000>(v);
let v = cmp_merge::<R, O, S2<R::Lanes>, 0b1100_1100>(v);
cmp_merge::<R, O, S1<R::Lanes>, 0b1010_1010>(v)
}
#[inline(always)]
pub fn sort_2<R, O>(v: Storage<R>) -> Storage<R>
where
O: SortOrder,
R: NumericRegister<Lanes = typenum::U2>,
{
layer_indices!(L0, [1, 0]);
cmp_merge::<R, O, L0<R::Lanes>, 0b10>(v)
}
#[inline(always)]
pub fn sort_4<R, O>(v: Storage<R>) -> Storage<R>
where
O: SortOrder,
R: NumericRegister<Lanes = typenum::U4>,
{
layer_indices!(L0, [2, 3, 0, 1]); layer_indices!(L1, [1, 0, 3, 2]); layer_indices!(L2, [0, 2, 1, 3]);
let v = cmp_merge::<R, O, L0<R::Lanes>, 0b1100>(v);
let v = cmp_merge::<R, O, L1<R::Lanes>, 0b1010>(v);
cmp_merge::<R, O, L2<R::Lanes>, 0b0100>(v)
}
#[inline(always)]
pub fn sort_8<R, O>(v: Storage<R>) -> Storage<R>
where
O: SortOrder,
R: NumericRegister<Lanes = typenum::U8>,
{
layer_indices!(H1, [1, 0, 3, 2, 5, 4, 7, 6]); layer_indices!(H2, [2, 3, 0, 1, 6, 7, 4, 5]); layer_indices!(H3, [0, 2, 1, 3, 4, 6, 5, 7]); layer_indices!(RevHi, [0, 1, 2, 3, 7, 6, 5, 4]);
let v = cmp_merge::<R, O, H1<R::Lanes>, 0b1010_1010>(v);
let v = cmp_merge::<R, O, H2<R::Lanes>, 0b1100_1100>(v);
let v = cmp_merge::<R, O, H3<R::Lanes>, 0b0100_0100>(v);
let v = R::permutev_const::<RevHi<R::Lanes>>(v);
bitonic_clean_8::<R, O>(v)
}
macro_rules! ce_chunks {
($c:ident: $(($i:tt, $j:tt))+) => {$(
{
let lo = O::first::<R>($c[$i], $c[$j]);
$c[$j] = O::last::<R>($c[$i], $c[$j]);
$c[$i] = lo;
}
)+};
}
macro_rules! clean_chunks {
($c:ident: $($i:tt)+) => {$(
$c[$i] = R::bitonic_clean_by::<O>($c[$i]);
)+};
}
#[inline(always)]
pub fn bitonic_clean_array_2<R: NumericRegister, O: SortOrder>(mut c: [Storage<R>; 2]) -> [Storage<R>; 2] {
ce_chunks!(c: (0, 1));
clean_chunks!(c: 0 1);
c
}
#[inline(always)]
pub fn bitonic_clean_array_4<R: NumericRegister, O: SortOrder>(mut c: [Storage<R>; 4]) -> [Storage<R>; 4] {
ce_chunks!(c: (0, 2)(1, 3));
ce_chunks!(c: (0, 1)(2, 3));
clean_chunks!(c: 0 1 2 3);
c
}
#[inline(always)]
pub fn bitonic_clean_array_8<R: NumericRegister, O: SortOrder>(mut c: [Storage<R>; 8]) -> [Storage<R>; 8] {
ce_chunks!(c: (0, 4)(1, 5)(2, 6)(3, 7));
ce_chunks!(c: (0, 2)(1, 3)(4, 6)(5, 7));
ce_chunks!(c: (0, 1)(2, 3)(4, 5)(6, 7));
clean_chunks!(c: 0 1 2 3 4 5 6 7);
c
}
#[inline(always)]
pub fn bitonic_clean_array_16<R: NumericRegister, O: SortOrder>(mut c: [Storage<R>; 16]) -> [Storage<R>; 16] {
ce_chunks!(c: (0, 8)(1, 9)(2, 10)(3, 11)(4, 12)(5, 13)(6, 14)(7, 15));
ce_chunks!(c: (0, 4)(1, 5)(2, 6)(3, 7)(8, 12)(9, 13)(10, 14)(11, 15));
ce_chunks!(c: (0, 2)(1, 3)(4, 6)(5, 7)(8, 10)(9, 11)(12, 14)(13, 15));
ce_chunks!(c: (0, 1)(2, 3)(4, 5)(6, 7)(8, 9)(10, 11)(12, 13)(14, 15));
clean_chunks!(c: 0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15);
c
}
#[inline(always)]
fn merge_runs_1<R: NumericRegister, O: SortOrder>(a: Storage<R>, b: Storage<R>) -> [Storage<R>; 2] {
bitonic_clean_array_2::<R, O>([a, R::reverse(b)])
}
#[inline(always)]
fn merge_runs_2<R: NumericRegister, O: SortOrder>(a: [Storage<R>; 2], b: [Storage<R>; 2]) -> [Storage<R>; 4] {
bitonic_clean_array_4::<R, O>([a[0], a[1], R::reverse(b[1]), R::reverse(b[0])])
}
#[inline(always)]
fn merge_runs_4<R: NumericRegister, O: SortOrder>(a: [Storage<R>; 4], b: [Storage<R>; 4]) -> [Storage<R>; 8] {
bitonic_clean_array_8::<R, O>([
a[0],
a[1],
a[2],
a[3],
R::reverse(b[3]),
R::reverse(b[2]),
R::reverse(b[1]),
R::reverse(b[0]),
])
}
#[inline(always)]
fn merge_runs_8<R: NumericRegister, O: SortOrder>(a: [Storage<R>; 8], b: [Storage<R>; 8]) -> [Storage<R>; 16] {
bitonic_clean_array_16::<R, O>([
a[0],
a[1],
a[2],
a[3],
a[4],
a[5],
a[6],
a[7],
R::reverse(b[7]),
R::reverse(b[6]),
R::reverse(b[5]),
R::reverse(b[4]),
R::reverse(b[3]),
R::reverse(b[2]),
R::reverse(b[1]),
R::reverse(b[0]),
])
}
const MERGE_MAX_LANES: usize = 16;
#[inline(always)]
fn pair_stage<R, O, S>(v: Storage<R>) -> Storage<R>
where
R: NumericRegister,
O: SortOrder,
S: crate::sort::PairStage<R::Lanes>,
{
let partner = R::permutev_const::<S::Indices>(v);
let first = O::first::<R>(v, partner);
let last = O::last::<R>(v, partner);
let keep_last = <R::Mask as MaskRegister>::from_native_bitmask(S::KEEP_LAST);
R::blendv(keep_last, first, last)
}
macro_rules! columnar_chunks {
($name:ident, $n:literal; $([$(($a:literal, $b:literal)),* $(,)?])*) => {
#[inline(always)]
fn $name<R: NumericRegister, O: SortOrder>(mut c: [Storage<R>; $n]) -> [Storage<R>; $n] {
$($(
{
const _: () = assert!($a < $b, "comparator must be written (lo, hi)");
const _: () = assert!($b < $n, "comparator index out of range");
let lo = O::first::<R>(c[$a], c[$b]);
c[$b] = O::last::<R>(c[$a], c[$b]);
c[$a] = lo;
}
)*)*
c
}
};
}
columnar_chunks! {
columns_2, 2;
[(0,1)]
}
columnar_chunks! {
columns_4, 4;
[(0,2),(1,3)]
[(0,1),(2,3)]
[(1,2)]
}
columnar_chunks! {
columns_8, 8;
[(0,2),(1,3),(4,6),(5,7)]
[(0,4),(1,5),(2,6),(3,7)]
[(0,1),(2,3),(4,5),(6,7)]
[(2,4),(3,5)]
[(1,4),(3,6)]
[(1,2),(3,4),(5,6)]
}
columnar_chunks! {
columns_16, 16;
[(0,13),(1,12),(2,15),(3,14),(4,8),(5,6),(7,11),(9,10)]
[(0,5),(1,7),(2,9),(3,4),(6,13),(8,14),(10,15),(11,12)]
[(0,1),(2,3),(4,5),(6,8),(7,9),(10,11),(12,13),(14,15)]
[(0,2),(1,3),(4,10),(5,11),(6,7),(8,9),(12,14),(13,15)]
[(1,2),(3,12),(4,6),(5,7),(8,10),(9,11),(13,14)]
[(1,4),(2,6),(5,8),(7,10),(9,13),(11,14)]
[(2,4),(3,6),(9,12),(11,13)]
[(3,5),(6,8),(7,9),(10,12)]
[(3,4),(5,6),(7,8),(9,10),(11,12)]
[(6,7),(8,9)]
}
macro_rules! tail_fn {
($name:ident, $c:literal, [$($d:literal),*]) => {
#[inline(always)]
fn $name<R: NumericRegister, O: SortOrder>(v: Storage<R>) -> Storage<R> {
let v = pair_stage::<R, O, crate::sort::RevPairs<$c>>(v);
$( let v = pair_stage::<R, O, crate::sort::Distance<$d>>(v); )*
v
}
};
}
tail_fn!(tail_c2, 2, []);
tail_fn!(tail_c4, 4, [1]);
tail_fn!(tail_c8, 8, [2, 1]);
tail_fn!(tail_c16, 16, [4, 2, 1]);
#[inline(always)]
pub fn sort_lanes<R: NumericRegister, O: SortOrder>(v: Storage<R>) -> Storage<R> {
if const { <R::Lanes as Unsigned>::USIZE > MERGE_MAX_LANES } {
return sort_any::<R, O>(v);
}
let v = if const { <R::Lanes as Unsigned>::USIZE >= 2 } {
tail_c2::<R, O>(v)
} else {
v
};
let v = if const { <R::Lanes as Unsigned>::USIZE >= 4 } {
tail_c4::<R, O>(v)
} else {
v
};
let v = if const { <R::Lanes as Unsigned>::USIZE >= 8 } {
tail_c8::<R, O>(v)
} else {
v
};
if const { <R::Lanes as Unsigned>::USIZE >= 16 } {
tail_c16::<R, O>(v)
} else {
v
}
}
#[inline(always)]
pub fn bitonic_clean_lanes<R: NumericRegister, O: SortOrder>(v: Storage<R>) -> Storage<R> {
if const { <R::Lanes as Unsigned>::USIZE > MERGE_MAX_LANES } {
return sort_any::<R, O>(v);
}
let v = if const { <R::Lanes as Unsigned>::USIZE >= 16 } {
pair_stage::<R, O, crate::sort::Distance<8>>(v)
} else {
v
};
let v = if const { <R::Lanes as Unsigned>::USIZE >= 8 } {
pair_stage::<R, O, crate::sort::Distance<4>>(v)
} else {
v
};
let v = if const { <R::Lanes as Unsigned>::USIZE >= 4 } {
pair_stage::<R, O, crate::sort::Distance<2>>(v)
} else {
v
};
if const { <R::Lanes as Unsigned>::USIZE >= 2 } {
pair_stage::<R, O, crate::sort::Distance<1>>(v)
} else {
v
}
}
macro_rules! merge_fn {
(
$name:ident, $n:literal, $c:literal, $tail:ident,
[ $( [ $( ($lo:tt, $hi:tt) )+ ] )+ ],
[ $($chunk:tt)+ ]
) => {
#[inline(always)]
fn $name<R: NumericRegister, O: SortOrder>(c: &mut [Storage<R>; $n]) {
$(
$( c[$hi] = R::permutev_const::<crate::sort::RevIdx<$c, R::Lanes>>(c[$hi]); )+
$(
{
let lo = O::first::<R>(c[$lo], c[$hi]);
c[$hi] = O::last::<R>(c[$lo], c[$hi]);
c[$lo] = lo;
}
)+
)+
$( c[$chunk] = $tail::<R, O>(c[$chunk]); )+
}
};
}
macro_rules! merge_group {
(
$n:literal, [ $($chunk:tt)+ ], $f2:ident, $f4:ident, $f8:ident, $f16:ident,
$($stage:tt)+
) => {
merge_fn!($f2, $n, 2, tail_c2, [$($stage)+], [$($chunk)+]);
merge_fn!($f4, $n, 4, tail_c4, [$($stage)+], [$($chunk)+]);
merge_fn!($f8, $n, 8, tail_c8, [$($stage)+], [$($chunk)+]);
merge_fn!($f16, $n, 16, tail_c16, [$($stage)+], [$($chunk)+]);
};
}
merge_group! {
2, [0 1], merge_2_c2, merge_2_c4, merge_2_c8, merge_2_c16,
[(0,1)]
}
merge_group! {
4, [0 1 2 3], merge_4_c2, merge_4_c4, merge_4_c8, merge_4_c16,
[(0,3)(1,2)]
[(0,1)(2,3)]
}
merge_group! {
8, [0 1 2 3 4 5 6 7], merge_8_c2, merge_8_c4, merge_8_c8, merge_8_c16,
[(0,7)(1,6)(2,5)(3,4)]
[(0,3)(1,2)(4,7)(5,6)]
[(0,1)(2,3)(4,5)(6,7)]
}
merge_group! {
16, [0 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15],
merge_16_c2, merge_16_c4, merge_16_c8, merge_16_c16,
[(0,15)(1,14)(2,13)(3,12)(4,11)(5,10)(6,9)(7,8)]
[(0,7)(1,6)(2,5)(3,4)(8,15)(9,14)(10,13)(11,12)]
[(0,3)(1,2)(4,7)(5,6)(8,11)(9,10)(12,15)(13,14)]
[(0,1)(2,3)(4,5)(6,7)(8,9)(10,11)(12,13)(14,15)]
}
macro_rules! merged_array_fn {
($name:ident, $n:literal, $cols:ident, $f2:ident, $f4:ident, $f8:ident, $f16:ident) => {
#[inline(always)]
fn $name<R: NumericRegister, O: SortOrder>(c: [Storage<R>; $n]) -> [Storage<R>; $n] {
let mut c = $cols::<R, O>(c);
if const { <R::Lanes as Unsigned>::USIZE >= 2 } {
$f2::<R, O>(&mut c);
if const { <R::Lanes as Unsigned>::USIZE >= 4 } {
$f4::<R, O>(&mut c);
if const { <R::Lanes as Unsigned>::USIZE >= 8 } {
$f8::<R, O>(&mut c);
if const { <R::Lanes as Unsigned>::USIZE >= 16 } {
$f16::<R, O>(&mut c);
}
}
}
}
c
}
};
}
merged_array_fn!(
merged_array_2,
2,
columns_2,
merge_2_c2,
merge_2_c4,
merge_2_c8,
merge_2_c16
);
merged_array_fn!(
merged_array_4,
4,
columns_4,
merge_4_c2,
merge_4_c4,
merge_4_c8,
merge_4_c16
);
merged_array_fn!(
merged_array_8,
8,
columns_8,
merge_8_c2,
merge_8_c4,
merge_8_c8,
merge_8_c16
);
merged_array_fn!(
merged_array_16,
16,
columns_16,
merge_16_c2,
merge_16_c4,
merge_16_c8,
merge_16_c16
);
#[inline(always)]
pub fn bitonic_array_2<R: NumericRegister, O: SortOrder>(c: [Storage<R>; 2]) -> [Storage<R>; 2] {
merge_runs_1::<R, O>(R::sort_by::<O>(c[0]), R::sort_by::<O>(c[1]))
}
#[inline(always)]
pub fn bitonic_array_4<R: NumericRegister, O: SortOrder>(c: [Storage<R>; 4]) -> [Storage<R>; 4] {
merge_runs_2::<R, O>(
bitonic_array_2::<R, O>([c[0], c[1]]),
bitonic_array_2::<R, O>([c[2], c[3]]),
)
}
#[inline(always)]
pub fn bitonic_array_8<R: NumericRegister, O: SortOrder>(c: [Storage<R>; 8]) -> [Storage<R>; 8] {
merge_runs_4::<R, O>(
bitonic_array_4::<R, O>([c[0], c[1], c[2], c[3]]),
bitonic_array_4::<R, O>([c[4], c[5], c[6], c[7]]),
)
}
#[inline(always)]
pub fn bitonic_array_16<R: NumericRegister, O: SortOrder>(c: [Storage<R>; 16]) -> [Storage<R>; 16] {
merge_runs_8::<R, O>(
bitonic_array_8::<R, O>([c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]]),
bitonic_array_8::<R, O>([c[8], c[9], c[10], c[11], c[12], c[13], c[14], c[15]]),
)
}
macro_rules! array_sort_fn {
($(#[$attr:meta])* $name:ident, $n:literal, $fast:ident, $slow:ident) => {
$(#[$attr])*
#[inline(always)]
pub fn $name<R: NumericRegister, O: SortOrder>(c: [Storage<R>; $n]) -> [Storage<R>; $n] {
if const { <R::Lanes as Unsigned>::USIZE <= MERGE_MAX_LANES } {
$fast::<R, O>(c)
} else {
$slow::<R, O>(c)
}
}
};
}
array_sort_fn!(
sort_array_2, 2, merged_array_2, bitonic_array_2
);
array_sort_fn!(
sort_array_4, 4, merged_array_4, bitonic_array_4
);
array_sort_fn!(
sort_array_8, 8, merged_array_8, bitonic_array_8
);
array_sort_fn!(
sort_array_16, 16, merged_array_16, bitonic_array_16
);
#[inline(always)]
pub fn sort_any<R: NumericRegister, O: SortOrder>(mut value: Storage<R>) -> Storage<R> {
let s = R::as_mut_slice(&mut value);
#[inline(always)]
fn cas<T: PartialOrd>(s: &mut [T], i: usize, j: usize) {
if s[i] > s[j] {
s.swap(i, j);
}
}
#[rustfmt::skip]
let () = match s.len() {
2 => cas(s, 0, 1),
4 => {
cas(s, 0, 1); cas(s, 2, 3); cas(s, 0, 2); cas(s, 1, 3); cas(s, 1, 2); },
_ => {
for i in 1..s.len() {
let mut j = i;
while j > 0 && s[j - 1] > s[j] {
s.swap(j - 1, j);
j -= 1;
}
}
}
};
if const { !O::IS_ASCENDING } {
return R::reverse(value);
}
value
}