use generic_array::{
GenericArray,
typenum::{U256, Unsigned},
};
use super::*;
use super::compress::{CompressRow, CompressTable};
pub static EXPAND8: GenericArray<CompressRow, U256> = build_expand_table8();
const fn build_expand_table8() -> GenericArray<CompressRow, U256> {
let lanes = 8;
let patterns = 256;
let mut table: GenericArray<CompressRow, U256> = unsafe { core::mem::MaybeUninit::zeroed().assume_init() };
let rows = table.as_mut_slice();
let mut m = 0;
while m < patterns {
let row = rows[m].0.as_mut_slice();
let mut pos = 0;
let mut i = 0;
while i < lanes {
if (m >> i) & 1 == 1 {
row[i] = pos as u8;
pos += 1;
}
i += 1;
}
let count = pos as u8;
let mut i = 0;
while i < lanes {
if (m >> i) & 1 == 0 {
row[i] = pos as u8;
pos += 1;
}
i += 1;
}
rows[m].1 = count;
m += 1;
}
table
}
#[inline(always)]
pub fn expand_permute<R>(value: Storage<R>, mask: Storage<R::Mask>) -> Storage<R>
where
R: WidenIndexRegister<Lanes: CompressTable>,
{
unsafe { expand_permute8_raw::<R>(value, mask) }
}
#[inline(always)]
pub unsafe fn expand_permute8_raw<R: WidenIndexRegister>(value: Storage<R>, mask: Storage<R::Mask>) -> Storage<R> {
let bm = unsafe { <R::Mask as MaskRegister>::native_bitmask(mask).unwrap_unchecked() } as usize;
R::permutev_row(value, &unsafe { EXPAND8.get_unchecked(bm) }.0)
}
#[inline(always)]
pub fn expand_permute_wide<R>(value: Storage<R>, mask: Storage<R::Mask>) -> Storage<R>
where
R: Register,
{
const {
assert!(
<R::Lanes as Unsigned>::USIZE % 8 == 0,
"expand_permute_wide requires a lane count that is a multiple of 8"
);
assert!(
<R::Lanes as Unsigned>::USIZE <= 64,
"expand_permute_wide requires <= 64 lanes (native_bitmask bound)"
);
}
let n = <R::Lanes as Unsigned>::USIZE;
let groups = n / 8;
let bm = unsafe { <R::Mask as MaskRegister>::native_bitmask(mask).unwrap_unchecked() };
let mut counts = [0u8; 8];
let mut total = 0usize;
for group in 0..groups {
let bmg = ((bm >> (group * 8)) & 0xFF) as usize;
let cnt = unsafe { EXPAND8.get_unchecked(bmg) }.1;
counts[group] = cnt;
total += cnt as usize;
}
let mut g: GenericArray<u32, R::Lanes> = GenericArray::default();
let mut base = 0usize;
for group in 0..groups {
let out_base = group * 8;
let bmg = ((bm >> out_base) & 0xFF) as usize;
let cnt = counts[group] as usize;
let ubase = total + (out_base - base);
let row = &unsafe { EXPAND8.get_unchecked(bmg) }.0;
for j in 0..8 {
let r = row[j] as usize;
let src = if r < cnt { base + r } else { ubase + (r - cnt) };
unsafe { *g.get_unchecked_mut(out_base + j) = src as u32 };
}
base += cnt;
}
R::permutev(value, g)
}
pub fn expand_default<R: Register>(value: Storage<R>, mask: Storage<R::Mask>) -> Storage<R> {
let n = <R::Lanes as Unsigned>::USIZE;
let src = R::as_slice(&value);
let mut result = value;
let dst = R::as_mut_slice(&mut result);
let mut pos = 0;
for i in 0..n {
if <R::Mask as MaskRegister>::test(mask, i) {
dst[i] = src[pos];
pos += 1;
}
}
for i in 0..n {
if !<R::Mask as MaskRegister>::test(mask, i) {
dst[i] = src[pos];
pos += 1;
}
}
result
}
pub fn expand_z_default<R: Register>(value: Storage<R>, mask: Storage<R::Mask>) -> Storage<R> {
let n = <R::Lanes as Unsigned>::USIZE;
let src = R::as_slice(&value);
let mut result = R::EMPTY;
let dst = R::as_mut_slice(&mut result);
let mut pos = 0;
for i in 0..n {
if <R::Mask as MaskRegister>::test(mask, i) {
dst[i] = src[pos];
pos += 1;
}
}
result
}
#[cfg(test)]
mod tests {
use super::*;
use crate::register::array::ArrayRegister;
fn expand_oracle<const N: usize>(data: &[i32], bits: u64) -> [i32; 64] {
let mut expected = [0i32; 64];
let mut pos = 0;
for lane in 0..N {
if (bits >> lane) & 1 == 1 {
expected[lane] = data[pos];
pos += 1;
}
}
for lane in 0..N {
if (bits >> lane) & 1 == 0 {
expected[lane] = data[pos];
pos += 1;
}
}
expected
}
fn make_mask<const N: usize>(bits: u64) -> Storage<<ArrayRegister<i32, N> as CoreRegister>::Mask>
where
generic_array::typenum::Const<N>: generic_array::IntoArrayLength,
ArrayRegister<i32, N>: Register<Element = i32>,
{
let mut sel = [0i32; 64];
for lane in 0..N {
sel[lane] = ((bits >> lane) & 1) as i32;
}
<ArrayRegister<i32, N>>::into_mask(<ArrayRegister<i32, N>>::new(
GenericArray::from_slice(&sel[..N]).clone(),
))
}
fn make_value<const N: usize>() -> Storage<ArrayRegister<i32, N>>
where
generic_array::typenum::Const<N>: generic_array::IntoArrayLength,
ArrayRegister<i32, N>: Register<Element = i32>,
{
let mut data = [0i32; 64];
for i in 0..N {
data[i] = ((i + 1) * 10) as i32;
}
<ArrayRegister<i32, N>>::new(GenericArray::from_slice(&data[..N]).clone())
}
fn check<const N: usize>()
where
generic_array::typenum::Const<N>: generic_array::IntoArrayLength,
ArrayRegister<i32, N>: WidenIndexRegister<Element = i32, Lanes: CompressTable>,
{
type R<const N: usize> = ArrayRegister<i32, N>;
let mut data = [0i32; 64];
for i in 0..N {
data[i] = ((i + 1) * 10) as i32;
}
let value = make_value::<N>();
for bits in 0u64..(1 << N) {
let mask = make_mask::<N>(bits);
let expected = expand_oracle::<N>(&data, bits);
let got = expand_permute::<R<N>>(value, mask);
let got = <R<N>>::as_slice(&got);
for lane in 0..N {
assert_eq!(got[lane], expected[lane], "N={N} bits={bits:b} lane={lane}");
}
}
}
#[test]
fn expand_permute_exhaustive() {
check::<2>();
check::<4>();
check::<8>();
}
fn check_wide<const N: usize>(patterns: impl Iterator<Item = u64>)
where
generic_array::typenum::Const<N>: generic_array::IntoArrayLength,
ArrayRegister<i32, N>: Register<Element = i32>,
{
type R<const N: usize> = ArrayRegister<i32, N>;
let mut data = [0i32; 64];
for i in 0..N {
data[i] = ((i + 1) * 10) as i32;
}
let value = make_value::<N>();
for bits in patterns {
let mask = make_mask::<N>(bits);
let expected = expand_oracle::<N>(&data, bits);
let got = expand_permute_wide::<R<N>>(value, mask);
let got = <R<N>>::as_slice(&got);
for lane in 0..N {
assert_eq!(got[lane], expected[lane], "N={N} bits={bits:b} lane={lane}");
}
}
}
#[test]
fn expand_permute_wide_correct() {
check_wide::<8>(0..(1u64 << 8));
check_wide::<16>(0..(1u64 << 16));
let mut s = 0x9E3779B97F4A7C15u64;
let mut rng = move || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
let structured = [
0u64,
u64::MAX,
0x5555_5555_5555_5555,
0xAAAA_AAAA_AAAA_AAAA,
0x0000_0000_FFFF_FFFF,
0xFFFF_FFFF_0000_0000,
0x00FF_00FF_00FF_00FF,
0x0101_0101_0101_0101,
0x8080_8080_8080_8080,
];
check_wide::<32>(structured.into_iter().chain((0..2000).map(move |_| rng())));
let mut s = 0xDEADBEEFCAFEBABEu64;
let mut rng = move || {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
s
};
check_wide::<64>(structured.into_iter().chain((0..2000).map(move |_| rng())));
}
#[test]
fn round_trip_exhaustive() {
type R = ArrayRegister<i32, 8>;
let value = make_value::<8>();
for bits in 0u64..256 {
let mask = make_mask::<8>(bits);
let there = expand_permute::<R>(compress_permute::<R>(value, mask), mask);
let back = compress_permute::<R>(expand_permute::<R>(value, mask), mask);
assert_eq!(
<R>::as_slice(&there),
<R>::as_slice(&value),
"expand(compress(v)) bits={bits:b}"
);
assert_eq!(
<R>::as_slice(&back),
<R>::as_slice(&value),
"compress(expand(v)) bits={bits:b}"
);
}
}
#[test]
fn round_trip_wide_16() {
type R = ArrayRegister<i32, 16>;
let value = make_value::<16>();
for bits in 0u64..(1 << 16) {
let mask = make_mask::<16>(bits);
let there = expand_permute_wide::<R>(compress_permute_wide::<R>(value, mask), mask);
assert_eq!(
<R>::as_slice(&there),
<R>::as_slice(&value),
"expand(compress(v)) bits={bits:b}"
);
}
}
}