#[inline]
fn trans_bit_8x8(mut x: u64) -> u64 {
let t = (x ^ (x >> 7)) & 0x00AA_00AA_00AA_00AA;
x = x ^ t ^ (t << 7);
let t = (x ^ (x >> 14)) & 0x0000_CCCC_0000_CCCC;
x = x ^ t ^ (t << 14);
let t = (x ^ (x >> 28)) & 0x0000_0000_F0F0_F0F0;
x = x ^ t ^ (t << 28);
x
}
#[inline]
fn read_u64_le(b: &[u8], off: usize) -> u64 {
u64::from_le_bytes(b[off..off + 8].try_into().unwrap())
}
#[cfg(feature = "simd")]
macro_rules! dispatch {
($($tt:tt)*) => {
fearless_simd::dispatch!(crate::simd::level(), $($tt)*)
};
}
pub(crate) fn trans_byte_elem(input: &[u8], out: &mut [u8], size: usize, elem_size: usize) {
if elem_size == 1 {
out[..size].copy_from_slice(&input[..size]);
return;
}
#[cfg(feature = "simd")]
let done = dispatch!(s => simd::trans_byte_elem(s, input, out, size, elem_size));
#[cfg(not(feature = "simd"))]
let done = 0;
trans_byte_elem_from(done, input, out, size, elem_size);
}
fn trans_byte_elem_from(from: usize, input: &[u8], out: &mut [u8], size: usize, elem_size: usize) {
let mut ii = from;
while ii + 7 < size {
for jj in 0..elem_size {
for kk in 0..8 {
out[jj * size + ii + kk] = input[ii * elem_size + kk * elem_size + jj];
}
}
ii += 8;
}
let mut ii = ii.max(size - size % 8);
while ii < size {
for jj in 0..elem_size {
out[jj * size + ii] = input[ii * elem_size + jj];
}
ii += 1;
}
}
pub(crate) fn untrans_byte_elem(input: &[u8], out: &mut [u8], size: usize, elem_size: usize) {
if elem_size == 1 {
out[..size].copy_from_slice(&input[..size]);
return;
}
#[cfg(feature = "simd")]
let done = dispatch!(s => simd::untrans_byte_elem(s, input, out, size, elem_size));
#[cfg(not(feature = "simd"))]
let done = 0;
untrans_byte_elem_from(done, input, out, size, elem_size);
}
fn untrans_byte_elem_from(
from: usize,
input: &[u8],
out: &mut [u8],
size: usize,
elem_size: usize,
) {
for jj in 0..elem_size {
let plane = &input[jj * size..(jj + 1) * size];
for ii in from..size {
out[ii * elem_size + jj] = plane[ii];
}
}
}
fn trans_bit_byte(input: &[u8], out: &mut [u8], nbyte: usize) {
#[cfg(feature = "simd")]
let done = dispatch!(s => simd::trans_bit_byte(s, input, out, nbyte));
#[cfg(not(feature = "simd"))]
let done = 0;
trans_bit_byte_from(done, input, out, nbyte);
}
fn trans_bit_byte_from(from: usize, input: &[u8], out: &mut [u8], nbyte: usize) {
let nbyte_bitrow = nbyte / 8;
for ii in from / 8..nbyte_bitrow {
let mut x = trans_bit_8x8(read_u64_le(input, ii * 8));
for kk in 0..8 {
out[kk * nbyte_bitrow + ii] = x as u8;
x >>= 8;
}
}
}
fn trans_bitrow_eight(input: &[u8], out: &mut [u8], size: usize, elem_size: usize) {
let nbyte_bitrow = size / 8;
for ii in 0..8 {
for jj in 0..elem_size {
let src = (ii * elem_size + jj) * nbyte_bitrow;
let dst = (jj * 8 + ii) * nbyte_bitrow;
out[dst..dst + nbyte_bitrow].copy_from_slice(&input[src..src + nbyte_bitrow]);
}
}
}
fn trans_byte_bitrow(input: &[u8], out: &mut [u8], size: usize, elem_size: usize) {
#[cfg(feature = "simd")]
let done = dispatch!(s => simd::trans_byte_bitrow(s, input, out, size, elem_size));
#[cfg(not(feature = "simd"))]
let done = 0;
trans_byte_bitrow_from(done, input, out, size, elem_size);
}
fn trans_byte_bitrow_from(
from: usize,
input: &[u8],
out: &mut [u8],
size: usize,
elem_size: usize,
) {
let nbyte_row = size / 8;
for jj in 0..elem_size {
for ii in from..nbyte_row {
for kk in 0..8 {
out[ii * 8 * elem_size + jj * 8 + kk] = input[(jj * 8 + kk) * nbyte_row + ii];
}
}
}
}
fn shuffle_bit_eightelem(input: &[u8], out: &mut [u8], nbyte: usize, elem_size: usize) {
#[cfg(feature = "simd")]
let done = dispatch!(s => simd::shuffle_bit_eightelem(s, input, out, nbyte, elem_size));
#[cfg(not(feature = "simd"))]
let done = 0;
shuffle_bit_eightelem_from(done, input, out, nbyte, elem_size);
}
fn shuffle_bit_eightelem_from(
from: usize,
input: &[u8],
out: &mut [u8],
nbyte: usize,
elem_size: usize,
) {
let group = 8 * elem_size;
let mut p = from;
while p + 7 < nbyte {
let mut x = trans_bit_8x8(read_u64_le(input, p));
let base = p / group * group + p % group / 8;
for kk in 0..8 {
out[base + kk * elem_size] = x as u8;
x >>= 8;
}
p += 8;
}
}
pub(crate) struct BitshuffleScratch {
a: Vec<u8>,
b: Vec<u8>,
}
impl BitshuffleScratch {
pub(crate) fn new(nbyte: usize) -> Self {
Self {
a: vec![0u8; nbyte],
b: vec![0u8; nbyte],
}
}
}
pub(crate) fn bitshuffle_block_into(
input: &[u8],
scratch: &mut BitshuffleScratch,
out: &mut [u8],
elem_size: usize,
) {
let size = input.len() / elem_size;
debug_assert_eq!(size % 8, 0);
let nbyte = size * elem_size;
let (a, b) = (&mut scratch.a[..nbyte], &mut scratch.b[..nbyte]);
trans_byte_elem(input, a, size, elem_size);
trans_bit_byte(a, b, nbyte);
trans_bitrow_eight(b, &mut out[..nbyte], size, elem_size);
}
pub(crate) fn bitunshuffle_block_into(
input: &[u8],
scratch: &mut BitshuffleScratch,
out: &mut [u8],
elem_size: usize,
) {
let size = input.len() / elem_size;
debug_assert_eq!(size % 8, 0);
let nbyte = size * elem_size;
let tmp = &mut scratch.b[..nbyte];
trans_byte_bitrow(input, tmp, size, elem_size);
shuffle_bit_eightelem(tmp, &mut out[..nbyte], nbyte, elem_size);
}
#[cfg(feature = "simd")]
mod simd {
use fearless_simd::{prelude::*, Simd};
use fearless_simd_macros::simd;
#[inline(always)]
fn bit_transposed<S: Simd>(simd: S, x: S::u8s) -> S::u8s {
let mut x = S::u64s::from_bytes(x);
let m = S::u64s::splat(simd, 0x00AA_00AA_00AA_00AA);
let t = (x ^ (x >> 7)) & m;
x = x ^ t ^ (t << 7);
let m = S::u64s::splat(simd, 0x0000_CCCC_0000_CCCC);
let t = (x ^ (x >> 14)) & m;
x = x ^ t ^ (t << 14);
let m = S::u64s::splat(simd, 0x0000_0000_F0F0_F0F0);
let t = (x ^ (x >> 28)) & m;
x = x ^ t ^ (t << 28);
x.to_bytes()
}
#[inline(always)]
fn row_order<S: Simd>(simd: S) -> S::u8s {
let q = S::u8s::LEN / 8;
S::u8s::from_fn(simd, |i| (8 * (i % q) + i / q) as u8)
}
#[simd]
pub(super) fn trans_byte_elem<S: Simd>(
simd: S,
input: &[u8],
out: &mut [u8],
size: usize,
elem_size: usize,
) -> usize {
if !matches!(elem_size, 2 | 4 | 8) {
return 0;
}
let n = S::u8s::LEN;
let mut v = [S::u8s::splat(simd, 0); 8];
let mut ii = 0;
while ii + n <= size {
let chunk = &input[ii * elem_size..(ii + n) * elem_size];
for (k, x) in v[..elem_size].iter_mut().enumerate() {
*x = S::u8s::from_slice(simd, &chunk[k * n..(k + 1) * n]);
}
let mut stride = elem_size;
while stride > 1 {
let mut next = v;
for k in 0..elem_size / 2 {
let (lo, hi) = v[2 * k].deinterleave(v[2 * k + 1]);
next[k] = lo;
next[elem_size / 2 + k] = hi;
}
v = next;
stride /= 2;
}
for (k, x) in v[..elem_size].iter().enumerate() {
x.store_slice(&mut out[k * size + ii..k * size + ii + n]);
}
ii += n;
}
ii
}
#[simd]
pub(super) fn untrans_byte_elem<S: Simd>(
simd: S,
input: &[u8],
out: &mut [u8],
size: usize,
elem_size: usize,
) -> usize {
if !matches!(elem_size, 2 | 4 | 8) {
return 0;
}
let n = S::u8s::LEN;
let mut v = [S::u8s::splat(simd, 0); 8];
let mut ii = 0;
while ii + n <= size {
for (k, x) in v[..elem_size].iter_mut().enumerate() {
*x = S::u8s::from_slice(simd, &input[k * size + ii..k * size + ii + n]);
}
let mut stride = 1;
while stride < elem_size {
let mut prev = v;
for k in 0..elem_size / 2 {
let (lo, hi) = v[k].interleave(v[elem_size / 2 + k]);
prev[2 * k] = lo;
prev[2 * k + 1] = hi;
}
v = prev;
stride *= 2;
}
let chunk = &mut out[ii * elem_size..(ii + n) * elem_size];
for (k, x) in v[..elem_size].iter().enumerate() {
x.store_slice(&mut chunk[k * n..(k + 1) * n]);
}
ii += n;
}
ii
}
#[simd]
pub(super) fn trans_bit_byte<S: Simd>(
simd: S,
input: &[u8],
out: &mut [u8],
nbyte: usize,
) -> usize {
let n = S::u8s::LEN;
let q = n / 8;
let nbyte_bitrow = nbyte / 8;
let order = row_order(simd);
let mut ii = 0;
while ii + n <= nbyte {
let rows = bit_transposed(simd, S::u8s::from_slice(simd, &input[ii..ii + n]))
.swizzle_dyn(order);
for (r, bytes) in rows.as_slice().chunks_exact(q).enumerate() {
let o = r * nbyte_bitrow + ii / 8;
out[o..o + q].copy_from_slice(bytes);
}
ii += n;
}
ii
}
#[inline(always)]
fn zip<S: Simd>(a: S::u8s, b: S::u8s, w: usize, perm: S::u8s) -> (S::u8s, S::u8s) {
match w {
1 => a.interleave(b),
2 => {
let (lo, hi) = S::u16s::from_bytes(a).interleave(S::u16s::from_bytes(b));
(lo.to_bytes(), hi.to_bytes())
}
4 => {
let (lo, hi) = S::u32s::from_bytes(a).interleave(S::u32s::from_bytes(b));
(lo.to_bytes(), hi.to_bytes())
}
8 => {
let (lo, hi) = S::u64s::from_bytes(a).interleave(S::u64s::from_bytes(b));
(lo.to_bytes(), hi.to_bytes())
}
_ => {
let (lo, hi) = S::u64s::from_bytes(a).interleave(S::u64s::from_bytes(b));
(
lo.to_bytes().swizzle_dyn(perm),
hi.to_bytes().swizzle_dyn(perm),
)
}
}
}
#[inline(always)]
fn zip_perm<S: Simd>(simd: S, w: usize) -> S::u8s {
let u = (w / 8).max(1);
S::u8s::from_fn(simd, |i| {
let lane = i / 8;
let (t, s) = (lane / (2 * u), lane % (2 * u));
let src = if s < u {
2 * (t * u + s)
} else {
2 * (t * u + s - u) + 1
};
(src * 8 + i % 8) as u8
})
}
#[simd]
pub(super) fn trans_byte_bitrow<S: Simd>(
simd: S,
input: &[u8],
out: &mut [u8],
size: usize,
elem_size: usize,
) -> usize {
if !matches!(elem_size, 1 | 2 | 4 | 8) {
return 0;
}
let n = S::u8s::LEN;
let nbyte_row = size / 8;
let nrows = 8 * elem_size;
let perms = [zip_perm(simd, 16), zip_perm(simd, 32)];
let mut a = [S::u8s::splat(simd, 0); 64];
let mut b = [S::u8s::splat(simd, 0); 64];
let mut ii = 0;
while ii + n <= nbyte_row {
for (k, x) in a[..nrows].iter_mut().enumerate() {
let s = k * nbyte_row + ii;
*x = S::u8s::from_slice(simd, &input[s..s + n]);
}
let (mut src, mut dst) = (&mut a, &mut b);
let (mut g, mut c, mut w) = (nrows, 1, 1);
while g > 1 && w < n {
let perm = perms[usize::from(w == 32)];
for m in 0..g / 2 {
for j in 0..c {
let (lo, hi) =
zip::<S>(src[2 * m * c + j], src[(2 * m + 1) * c + j], w, perm);
dst[2 * m * c + 2 * j] = lo;
dst[2 * m * c + 2 * j + 1] = hi;
}
}
std::mem::swap(&mut src, &mut dst);
g /= 2;
c *= 2;
w *= 2;
}
for j in 0..c {
for rg in 0..g {
let o = ii * nrows + (j * g + rg) * n;
src[rg * c + j].store_slice(&mut out[o..o + n]);
}
}
ii += n;
}
ii
}
#[simd]
pub(super) fn shuffle_bit_eightelem<S: Simd>(
simd: S,
input: &[u8],
out: &mut [u8],
nbyte: usize,
elem_size: usize,
) -> usize {
let n = S::u8s::LEN;
let q = n / 8;
let group = 8 * elem_size;
if group % n != 0 && n % group != 0 {
return 0;
}
let mut ii = 0;
if group >= n {
let order = row_order(simd);
while ii + n <= nbyte {
let rows = bit_transposed(simd, S::u8s::from_slice(simd, &input[ii..ii + n]))
.swizzle_dyn(order);
let base = ii / group * group + ii % group / 8;
for (r, bytes) in rows.as_slice().chunks_exact(q).enumerate() {
let o = base + r * elem_size;
out[o..o + q].copy_from_slice(bytes);
}
ii += n;
}
} else {
let order = S::u8s::from_fn(simd, |i| {
let (t, rem) = (i / group, i % group);
let (r, e) = (rem / elem_size, rem % elem_size);
(8 * (t * elem_size + e) + r) as u8
});
while ii + n <= nbyte {
bit_transposed(simd, S::u8s::from_slice(simd, &input[ii..ii + n]))
.swizzle_dyn(order)
.store_slice(&mut out[ii..ii + n]);
ii += n;
}
}
ii
}
}
#[cfg(test)]
mod tests {
use super::*;
fn pattern(nbyte: usize) -> Vec<u8> {
(0..nbyte)
.map(|i| ((i as u32).wrapping_mul(2_654_435_761) >> 24) as u8)
.collect()
}
fn bitshuffle_reference(input: &[u8], elem_size: usize) -> Vec<u8> {
let n_elems = input.len() / elem_size;
let mut out = vec![0u8; input.len()];
for bit in 0..elem_size * 8 {
for elem in 0..n_elems {
let src_bit = (input[elem * elem_size + bit / 8] >> (bit % 8)) & 1;
let dst = bit * n_elems + elem;
out[dst / 8] |= src_bit << (dst % 8);
}
}
out
}
fn byte_transpose_reference(input: &[u8], size: usize, elem_size: usize) -> Vec<u8> {
let mut out = vec![0u8; size * elem_size];
for i in 0..size {
for k in 0..elem_size {
out[k * size + i] = input[i * elem_size + k];
}
}
out
}
const ELEM_SIZES: [usize; 8] = [1, 2, 3, 4, 5, 8, 12, 16];
const SIZES: [usize; 10] = [8, 16, 24, 64, 128, 136, 520, 1024, 1032, 4096];
#[test]
fn three_stage_transpose_matches_the_per_bit_reference() {
for elem_size in ELEM_SIZES {
for size in SIZES {
let input = pattern(size * elem_size);
let mut got = vec![0u8; input.len()];
let mut scratch = BitshuffleScratch::new(input.len());
bitshuffle_block_into(&input, &mut scratch, &mut got, elem_size);
assert_eq!(
got,
bitshuffle_reference(&input, elem_size),
"elem_size={elem_size} size={size}"
);
let mut back = vec![0u8; input.len()];
bitunshuffle_block_into(&got, &mut scratch, &mut back, elem_size);
assert_eq!(back, input, "inverse elem_size={elem_size} size={size}");
}
}
}
#[test]
fn byte_transpose_matches_the_reference_for_any_element_count() {
for elem_size in ELEM_SIZES {
for size in [1usize, 5, 7, 8, 9, 31, 33, 64, 100, 1000, 1031] {
let input = pattern(size * elem_size);
let mut got = vec![0u8; input.len()];
trans_byte_elem(&input, &mut got, size, elem_size);
let want = byte_transpose_reference(&input, size, elem_size);
assert_eq!(got, want, "trans elem_size={elem_size} size={size}");
let mut back = vec![0u8; input.len()];
untrans_byte_elem(&got, &mut back, size, elem_size);
assert_eq!(back, input, "untrans elem_size={elem_size} size={size}");
}
}
}
#[cfg(feature = "simd")]
#[test]
fn kernels_match_the_scalar_loops_on_every_level() {
use fearless_simd::Level;
let top = crate::simd::level();
let mut levels = vec![top, Level::baseline()];
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
levels.extend(top.as_avx2().map(Level::Avx2));
levels.extend(top.as_sse4_2().map(Level::Sse4_2));
levels.extend(top.as_sse2().map(Level::Sse2));
}
for elem_size in [1usize, 2, 3, 4, 8, 16] {
for size in SIZES {
let nbyte = size * elem_size;
let input = pattern(nbyte);
let mut want = vec![0u8; nbyte];
let mut got = vec![0u8; nbyte];
for &level in &levels {
let what = format!("{level:?} elem_size={elem_size} size={size}");
trans_byte_elem_from(0, &input, &mut want, size, elem_size);
got.fill(0);
let done = fearless_simd::dispatch!(level, s => simd::trans_byte_elem(s, &input, &mut got, size, elem_size));
assert_eq!(done % 8, 0, "{what} trans_byte_elem done");
trans_byte_elem_from(done, &input, &mut got, size, elem_size);
assert_eq!(got, want, "{what} trans_byte_elem");
untrans_byte_elem_from(0, &input, &mut want, size, elem_size);
got.fill(0);
let done = fearless_simd::dispatch!(level, s => simd::untrans_byte_elem(s, &input, &mut got, size, elem_size));
untrans_byte_elem_from(done, &input, &mut got, size, elem_size);
assert_eq!(got, want, "{what} untrans_byte_elem");
trans_bit_byte_from(0, &input, &mut want, nbyte);
got.fill(0);
let done = fearless_simd::dispatch!(level, s => simd::trans_bit_byte(s, &input, &mut got, nbyte));
assert_eq!(done % 8, 0, "{what} trans_bit_byte done");
trans_bit_byte_from(done, &input, &mut got, nbyte);
assert_eq!(got, want, "{what} trans_bit_byte");
trans_byte_bitrow_from(0, &input, &mut want, size, elem_size);
got.fill(0);
let done = fearless_simd::dispatch!(level, s => simd::trans_byte_bitrow(s, &input, &mut got, size, elem_size));
trans_byte_bitrow_from(done, &input, &mut got, size, elem_size);
assert_eq!(got, want, "{what} trans_byte_bitrow");
shuffle_bit_eightelem_from(0, &input, &mut want, nbyte, elem_size);
got.fill(0);
let done = fearless_simd::dispatch!(level, s => simd::shuffle_bit_eightelem(s, &input, &mut got, nbyte, elem_size));
assert_eq!(done % 8, 0, "{what} shuffle_bit_eightelem done");
shuffle_bit_eightelem_from(done, &input, &mut got, nbyte, elem_size);
assert_eq!(got, want, "{what} shuffle_bit_eightelem");
}
}
}
}
}