rlnc-simdx 1.3.1

SIMD-accelerated Random Linear Network Coding over GF(2^8) — no_std, maximum SIMD
Documentation
//! AVX-512BW + SSSE3 nibble-split kernel (Tier 4 — 512-bit `vpshufb`).
//!
//! Uses aligned 512-bit load/store when buffers are 64-byte aligned.

#[cfg(target_arch = "x86")]
use core::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use core::arch::x86_64::*;

use crate::kernel::scalar::make_nibble_tables;

#[inline(always)]
fn both_aligned64(x: &[u8], y: &[u8]) -> bool {
    (x.as_ptr() as usize | y.as_ptr() as usize) & 63 == 0
}

/// AXPY using AVX-512BW `vpshufb`, aligned or unaligned.
///
/// # Safety
/// Requires `avx512f`, `avx512bw`, and `ssse3` target features.
#[target_feature(enable = "avx512f,avx512bw,ssse3")]
pub(crate) unsafe fn axpy_avx512_ssse3(c: u8, x: &[u8], y: &mut [u8]) {
    debug_assert_eq!(x.len(), y.len());
    if c == 0 {
        return;
    }
    if c == 1 {
        let len = x.len();
        let mut i = 0usize;
        while i + 64 <= len {
            let xv = _mm512_loadu_si512(x.as_ptr().add(i) as *const _);
            let yv = _mm512_loadu_si512(y.as_ptr().add(i) as *const _);
            _mm512_storeu_si512(y.as_mut_ptr().add(i) as *mut _, _mm512_xor_si512(yv, xv));
            i += 64;
        }
        if i + 32 <= len {
            let xv = _mm256_loadu_si256(x.as_ptr().add(i) as *const __m256i);
            let yv = _mm256_loadu_si256(y.as_ptr().add(i) as *const __m256i);
            _mm256_storeu_si256(
                y.as_mut_ptr().add(i) as *mut __m256i,
                _mm256_xor_si256(yv, xv),
            );
            i += 32;
        }
        if i + 16 <= len {
            let xv = _mm_loadu_si128(x.as_ptr().add(i) as *const __m128i);
            let yv = _mm_loadu_si128(y.as_ptr().add(i) as *const __m128i);
            _mm_storeu_si128(y.as_mut_ptr().add(i) as *mut __m128i, _mm_xor_si128(yv, xv));
            i += 16;
        }
        crate::kernel::scalar::xor_assign(&x[i..], &mut y[i..]);
        return;
    }

    let (lo_arr, hi_arr) = make_nibble_tables(c);
    let lo128 = _mm_loadu_si128(lo_arr.as_ptr() as *const __m128i);
    let hi128 = _mm_loadu_si128(hi_arr.as_ptr() as *const __m128i);
    let lo_tbl = _mm512_broadcast_i32x4(lo128);
    let hi_tbl = _mm512_broadcast_i32x4(hi128);
    let mask_lo = _mm512_set1_epi8(0x0Fi8);

    let len = x.len();
    let mut i = 0usize;

    macro_rules! axpy512_ssse3 {
        ($load:ident, $store:ident) => {
            while i + 64 <= len {
                let xv = $load(x.as_ptr().add(i) as *const _);
                let yv = $load(y.as_ptr().add(i) as *const _);
                let xlo = _mm512_and_si512(xv, mask_lo);
                let xhi = _mm512_and_si512(_mm512_srli_epi16(xv, 4), mask_lo);
                let mul = _mm512_xor_si512(
                    _mm512_shuffle_epi8(lo_tbl, xlo),
                    _mm512_shuffle_epi8(hi_tbl, xhi),
                );
                $store(y.as_mut_ptr().add(i) as *mut _, _mm512_xor_si512(yv, mul));
                i += 64;
            }
        };
    }

    if both_aligned64(x, y) {
        axpy512_ssse3!(_mm512_load_si512, _mm512_store_si512);
    } else {
        axpy512_ssse3!(_mm512_loadu_si512, _mm512_storeu_si512);
    }

    // Reuse the prepared tables for the AVX2 and SSSE3 tails.
    if i + 32 <= len {
        let lo256 = _mm256_broadcastsi128_si256(lo128);
        let hi256 = _mm256_broadcastsi128_si256(hi128);
        let mask256 = _mm256_set1_epi8(0x0Fi8);
        let xv = _mm256_loadu_si256(x.as_ptr().add(i) as *const __m256i);
        let yv = _mm256_loadu_si256(y.as_ptr().add(i) as *const __m256i);
        let xlo = _mm256_and_si256(xv, mask256);
        let xhi = _mm256_and_si256(_mm256_srli_epi16(xv, 4), mask256);
        let mul = _mm256_xor_si256(
            _mm256_shuffle_epi8(lo256, xlo),
            _mm256_shuffle_epi8(hi256, xhi),
        );
        _mm256_storeu_si256(
            y.as_mut_ptr().add(i) as *mut __m256i,
            _mm256_xor_si256(yv, mul),
        );
        i += 32;
    }
    if i + 16 <= len {
        let mask128 = _mm_set1_epi8(0x0Fi8);
        let xv = _mm_loadu_si128(x.as_ptr().add(i) as *const __m128i);
        let yv = _mm_loadu_si128(y.as_ptr().add(i) as *const __m128i);
        let xlo = _mm_and_si128(xv, mask128);
        let xhi = _mm_and_si128(_mm_srli_epi16(xv, 4), mask128);
        let mul = _mm_xor_si128(_mm_shuffle_epi8(lo128, xlo), _mm_shuffle_epi8(hi128, xhi));
        _mm_storeu_si128(
            y.as_mut_ptr().add(i) as *mut __m128i,
            _mm_xor_si128(yv, mul),
        );
        i += 16;
    }

    crate::kernel::scalar::axpy(c, &x[i..], &mut y[i..]);
}

/// Scale using AVX-512BW `vpshufb`, aligned or unaligned.
///
/// # Safety
/// Requires `avx512f`, `avx512bw`, and `ssse3` target features.
#[target_feature(enable = "avx512f,avx512bw,ssse3")]
pub(crate) unsafe fn scale_avx512_ssse3(c: u8, x: &[u8], y: &mut [u8]) {
    debug_assert_eq!(x.len(), y.len());
    if c == 0 {
        for yi in y.iter_mut() {
            *yi = 0;
        }
        return;
    }
    if c == 1 {
        y.copy_from_slice(x);
        return;
    }

    let (lo_arr, hi_arr) = make_nibble_tables(c);
    let lo128 = _mm_loadu_si128(lo_arr.as_ptr() as *const __m128i);
    let hi128 = _mm_loadu_si128(hi_arr.as_ptr() as *const __m128i);
    let lo_tbl = _mm512_broadcast_i32x4(lo128);
    let hi_tbl = _mm512_broadcast_i32x4(hi128);
    let mask_lo = _mm512_set1_epi8(0x0Fi8);

    let len = x.len();
    let mut i = 0usize;

    if both_aligned64(x, y) {
        while i + 64 <= len {
            let xv = _mm512_load_si512(x.as_ptr().add(i) as *const _);
            let xlo = _mm512_and_si512(xv, mask_lo);
            let xhi = _mm512_and_si512(_mm512_srli_epi16(xv, 4), mask_lo);
            _mm512_store_si512(
                y.as_mut_ptr().add(i) as *mut _,
                _mm512_xor_si512(
                    _mm512_shuffle_epi8(lo_tbl, xlo),
                    _mm512_shuffle_epi8(hi_tbl, xhi),
                ),
            );
            i += 64;
        }
    } else {
        while i + 64 <= len {
            let xv = _mm512_loadu_si512(x.as_ptr().add(i) as *const _);
            let xlo = _mm512_and_si512(xv, mask_lo);
            let xhi = _mm512_and_si512(_mm512_srli_epi16(xv, 4), mask_lo);
            _mm512_storeu_si512(
                y.as_mut_ptr().add(i) as *mut _,
                _mm512_xor_si512(
                    _mm512_shuffle_epi8(lo_tbl, xlo),
                    _mm512_shuffle_epi8(hi_tbl, xhi),
                ),
            );
            i += 64;
        }
    }

    if i + 32 <= len {
        let lo256 = _mm256_broadcastsi128_si256(lo128);
        let hi256 = _mm256_broadcastsi128_si256(hi128);
        let mask256 = _mm256_set1_epi8(0x0Fi8);
        let xv = _mm256_loadu_si256(x.as_ptr().add(i) as *const __m256i);
        let xlo = _mm256_and_si256(xv, mask256);
        let xhi = _mm256_and_si256(_mm256_srli_epi16(xv, 4), mask256);
        _mm256_storeu_si256(
            y.as_mut_ptr().add(i) as *mut __m256i,
            _mm256_xor_si256(
                _mm256_shuffle_epi8(lo256, xlo),
                _mm256_shuffle_epi8(hi256, xhi),
            ),
        );
        i += 32;
    }
    if i + 16 <= len {
        let mask128 = _mm_set1_epi8(0x0Fi8);
        let xv = _mm_loadu_si128(x.as_ptr().add(i) as *const __m128i);
        let xlo = _mm_and_si128(xv, mask128);
        let xhi = _mm_and_si128(_mm_srli_epi16(xv, 4), mask128);
        _mm_storeu_si128(
            y.as_mut_ptr().add(i) as *mut __m128i,
            _mm_xor_si128(_mm_shuffle_epi8(lo128, xlo), _mm_shuffle_epi8(hi128, xhi)),
        );
        i += 16;
    }

    crate::kernel::scalar::scale(c, &x[i..], &mut y[i..]);
}

/// In-place scale using AVX-512BW nibble-split.
///
/// # Safety
/// Requires `avx512f`, `avx512bw`, `ssse3`.
#[target_feature(enable = "avx512f,avx512bw,ssse3")]
pub(crate) unsafe fn scale_inplace_avx512_ssse3(c: u8, y: &mut [u8]) {
    if c == 0 {
        for yi in y.iter_mut() {
            *yi = 0;
        }
        return;
    }
    if c == 1 {
        return;
    }
    let (lo_arr, hi_arr) = make_nibble_tables(c);
    let lo128 = _mm_loadu_si128(lo_arr.as_ptr() as *const __m128i);
    let hi128 = _mm_loadu_si128(hi_arr.as_ptr() as *const __m128i);
    let lo_tbl = _mm512_broadcast_i32x4(lo128);
    let hi_tbl = _mm512_broadcast_i32x4(hi128);
    let mask_lo = _mm512_set1_epi8(0x0Fi8);
    let len = y.len();
    let mut i = 0usize;
    while i + 64 <= len {
        let yv = _mm512_loadu_si512(y.as_ptr().add(i) as *const _);
        let ylo = _mm512_and_si512(yv, mask_lo);
        let yhi = _mm512_and_si512(_mm512_srli_epi16(yv, 4), mask_lo);
        _mm512_storeu_si512(
            y.as_mut_ptr().add(i) as *mut _,
            _mm512_xor_si512(
                _mm512_shuffle_epi8(lo_tbl, ylo),
                _mm512_shuffle_epi8(hi_tbl, yhi),
            ),
        );
        i += 64;
    }
    if i + 32 <= len {
        let lo256 = _mm256_broadcastsi128_si256(lo128);
        let hi256 = _mm256_broadcastsi128_si256(hi128);
        let mask256 = _mm256_set1_epi8(0x0Fi8);
        let yv = _mm256_loadu_si256(y.as_ptr().add(i) as *const __m256i);
        let ylo = _mm256_and_si256(yv, mask256);
        let yhi = _mm256_and_si256(_mm256_srli_epi16(yv, 4), mask256);
        _mm256_storeu_si256(
            y.as_mut_ptr().add(i) as *mut __m256i,
            _mm256_xor_si256(
                _mm256_shuffle_epi8(lo256, ylo),
                _mm256_shuffle_epi8(hi256, yhi),
            ),
        );
        i += 32;
    }
    if i + 16 <= len {
        let mask128 = _mm_set1_epi8(0x0Fi8);
        let yv = _mm_loadu_si128(y.as_ptr().add(i) as *const __m128i);
        let ylo = _mm_and_si128(yv, mask128);
        let yhi = _mm_and_si128(_mm_srli_epi16(yv, 4), mask128);
        _mm_storeu_si128(
            y.as_mut_ptr().add(i) as *mut __m128i,
            _mm_xor_si128(_mm_shuffle_epi8(lo128, ylo), _mm_shuffle_epi8(hi128, yhi)),
        );
        i += 16;
    }
    crate::kernel::scalar::scale_inplace(c, &mut y[i..]);
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::AlignedBuffer;

    const TAIL_LENGTHS: [usize; 9] = [15, 16, 17, 31, 32, 33, 63, 64, 65];

    fn check_operations(x: &[u8], initial: &[u8], aligned: bool) {
        for c in [1u8, 0xA3] {
            let mut expected = initial.to_vec();
            crate::kernel::scalar::axpy(c, x, &mut expected);
            if aligned {
                let mut y = AlignedBuffer::from_slice(initial);
                unsafe { axpy_avx512_ssse3(c, x, &mut y) };
                assert_eq!(y.as_slice(), expected, "axpy c={c:#x}");
            } else {
                let mut backing = vec![0u8; initial.len() + 1];
                backing[1..].copy_from_slice(initial);
                unsafe { axpy_avx512_ssse3(c, x, &mut backing[1..]) };
                assert_eq!(&backing[1..], expected, "axpy c={c:#x}");
            }
        }

        let c = 0xA3;
        let mut expected = vec![0u8; x.len()];
        crate::kernel::scalar::scale(c, x, &mut expected);
        if aligned {
            let mut y = AlignedBuffer::zeroed(x.len());
            unsafe { scale_avx512_ssse3(c, x, &mut y) };
            assert_eq!(y.as_slice(), expected, "scale");
        } else {
            let mut backing = vec![0xA5u8; x.len() + 1];
            unsafe { scale_avx512_ssse3(c, x, &mut backing[1..]) };
            assert_eq!(&backing[1..], expected, "scale");
        }

        let mut expected = initial.to_vec();
        crate::kernel::scalar::scale_inplace(c, &mut expected);
        if aligned {
            let mut y = AlignedBuffer::from_slice(initial);
            unsafe { scale_inplace_avx512_ssse3(c, &mut y) };
            assert_eq!(y.as_slice(), expected, "scale_inplace");
        } else {
            let mut backing = vec![0u8; initial.len() + 1];
            backing[1..].copy_from_slice(initial);
            unsafe { scale_inplace_avx512_ssse3(c, &mut backing[1..]) };
            assert_eq!(&backing[1..], expected, "scale_inplace");
        }
    }

    #[test]
    fn axpy_matches_scalar() {
        if !is_x86_feature_detected!("avx512f") || !is_x86_feature_detected!("avx512bw") {
            return;
        }
        let c = 0xA3u8;
        let x: Vec<u8> = (0u8..=255).collect();
        let mut y_simd = vec![0x77u8; 256];
        let mut y_scalar = vec![0x77u8; 256];
        unsafe {
            axpy_avx512_ssse3(c, &x, &mut y_simd);
        }
        crate::kernel::scalar::axpy(c, &x, &mut y_scalar);
        assert_eq!(y_simd, y_scalar);
    }

    #[test]
    fn vector_tails_match_scalar_at_boundaries() {
        if !is_x86_feature_detected!("avx512f")
            || !is_x86_feature_detected!("avx512bw")
            || !is_x86_feature_detected!("ssse3")
        {
            return;
        }
        for len in TAIL_LENGTHS {
            let source: Vec<u8> = (0..len).map(|i| (i as u8).wrapping_mul(29)).collect();
            let initial: Vec<u8> = (0..len)
                .map(|i| (i as u8).wrapping_mul(17).wrapping_add(3))
                .collect();
            let aligned_source = AlignedBuffer::from_slice(&source);
            check_operations(&aligned_source, &initial, true);

            let mut unaligned_source = vec![0u8; len + 1];
            unaligned_source[1..].copy_from_slice(&source);
            assert_ne!(unaligned_source[1..].as_ptr() as usize & 63, 0);
            check_operations(&unaligned_source[1..], &initial, false);
        }
    }
}