#[cfg(not(feature = "std"))]
use alloc::vec::Vec;
#[cfg(feature = "std")]
use std::vec::Vec;
pub(crate) fn decode_into(zd: &[u16], initial: i16, q: u8, out: &mut Vec<i16>) {
let q = q & 15;
#[cfg(all(feature = "simd-avx2", target_arch = "x86_64"))]
{
unsafe { decode_avx2(zd, initial, q, out) };
}
#[cfg(all(
feature = "simd-ssse3",
not(feature = "simd-avx2"),
target_arch = "x86_64"
))]
{
unsafe { decode_sse2(zd, initial, q, out) };
}
#[cfg(all(feature = "simd-neon", target_arch = "aarch64"))]
{
unsafe { decode_neon(zd, initial, q, out) };
}
#[cfg(all(
feature = "simd-auto",
not(any(feature = "simd-avx2", feature = "simd-ssse3", feature = "simd-neon")),
feature = "std",
target_arch = "x86_64"
))]
{
if is_x86_feature_detected!("avx2") {
unsafe { decode_avx2(zd, initial, q, out) };
} else {
unsafe { decode_sse2(zd, initial, q, out) };
}
}
#[cfg(all(
feature = "simd-auto",
not(any(feature = "simd-avx2", feature = "simd-ssse3", feature = "simd-neon")),
target_arch = "aarch64"
))]
{
unsafe { decode_neon(zd, initial, q, out) };
}
#[cfg(not(any(
all(feature = "simd-avx2", target_arch = "x86_64"),
all(
feature = "simd-ssse3",
not(feature = "simd-avx2"),
target_arch = "x86_64"
),
all(feature = "simd-neon", target_arch = "aarch64"),
all(
feature = "simd-auto",
not(any(feature = "simd-avx2", feature = "simd-ssse3", feature = "simd-neon")),
feature = "std",
target_arch = "x86_64"
),
all(
feature = "simd-auto",
not(any(feature = "simd-avx2", feature = "simd-ssse3", feature = "simd-neon")),
target_arch = "aarch64"
)
)))]
decode_scalar(zd, initial, q, out);
}
#[inline]
fn unzigzag_one(code: u16) -> i16 {
((code >> 1) as i16) ^ -((code & 1) as i16)
}
#[allow(dead_code)]
fn decode_scalar(zd: &[u16], initial: i16, q: u8, out: &mut Vec<i16>) {
out.reserve(zd.len());
let mut acc = initial;
for &code in zd {
acc = acc.wrapping_add(unzigzag_one(code));
out.push(acc << q);
}
}
#[cfg(target_arch = "x86_64")]
#[allow(dead_code)]
unsafe fn decode_sse2(zd: &[u16], initial: i16, q: u8, out: &mut Vec<i16>) {
use core::arch::x86_64::*;
let n = zd.len();
out.reserve(n);
let base = out.len();
let simd_n = (n / 8) * 8;
let mut acc = initial;
let mut i = 0usize;
let q_vec = unsafe { _mm_cvtsi32_si128(q as i32) };
while i < simd_n {
let result = unsafe {
let v = _mm_loadu_si128(zd.as_ptr().add(i) as *const __m128i);
let one = _mm_set1_epi16(1);
let zero = _mm_setzero_si128();
let low_bit = _mm_and_si128(v, one);
let sign = _mm_sub_epi16(zero, low_bit);
let shifted = _mm_srli_epi16(v, 1);
let delta = _mm_xor_si128(shifted, sign);
let d = _mm_add_epi16(delta, _mm_slli_si128(delta, 2));
let d = _mm_add_epi16(d, _mm_slli_si128(d, 4));
let d = _mm_add_epi16(d, _mm_slli_si128(d, 8));
_mm_add_epi16(d, _mm_set1_epi16(acc))
};
unsafe {
let out_ptr = out.as_mut_ptr().add(base + i) as *mut __m128i;
acc = _mm_extract_epi16(result, 7) as i16;
_mm_storeu_si128(out_ptr, _mm_sll_epi16(result, q_vec));
}
i += 8;
}
unsafe {
out.set_len(base + simd_n);
}
for &code in &zd[simd_n..] {
acc = acc.wrapping_add(unzigzag_one(code));
out.push(acc << q);
}
}
#[cfg(target_arch = "x86_64")]
#[allow(dead_code)]
#[target_feature(enable = "avx2")]
unsafe fn decode_avx2(zd: &[u16], initial: i16, q: u8, out: &mut Vec<i16>) {
use core::arch::x86_64::*;
let n = zd.len();
out.reserve(n);
let base = out.len();
let simd_n = (n / 16) * 16;
let mut acc = initial;
let mut i = 0usize;
let q_vec = _mm_cvtsi32_si128(q as i32);
while i < simd_n {
let result = unsafe {
let v = _mm256_loadu_si256(zd.as_ptr().add(i) as *const __m256i);
let one = _mm256_set1_epi16(1);
let zero = _mm256_setzero_si256();
let low_bit = _mm256_and_si256(v, one);
let sign = _mm256_sub_epi16(zero, low_bit);
let shifted = _mm256_srli_epi16(v, 1);
let delta = _mm256_xor_si256(shifted, sign);
let d = _mm256_add_epi16(delta, _mm256_slli_si256(delta, 2));
let d = _mm256_add_epi16(d, _mm256_slli_si256(d, 4));
let d = _mm256_add_epi16(d, _mm256_slli_si256(d, 8));
let lo = _mm256_extracti128_si256(d, 0);
let low_total = _mm_extract_epi16(lo, 7) as i16;
let bridge_hi = _mm_set1_epi16(low_total);
let bridge = _mm256_inserti128_si256(_mm256_setzero_si256(), bridge_hi, 1);
let d = _mm256_add_epi16(d, bridge);
_mm256_add_epi16(d, _mm256_set1_epi16(acc))
};
unsafe {
let out_ptr = out.as_mut_ptr().add(base + i) as *mut __m256i;
let hi = _mm256_extracti128_si256(result, 1);
acc = _mm_extract_epi16(hi, 7) as i16;
_mm256_storeu_si256(out_ptr, _mm256_sll_epi16(result, q_vec));
}
i += 16;
}
unsafe {
out.set_len(base + simd_n);
}
for &code in &zd[simd_n..] {
acc = acc.wrapping_add(unzigzag_one(code));
out.push(acc << q);
}
}
#[cfg(target_arch = "aarch64")]
#[allow(dead_code)]
unsafe fn decode_neon(zd: &[u16], initial: i16, q: u8, out: &mut Vec<i16>) {
use core::arch::aarch64::*;
let n = zd.len();
out.reserve(n);
let base = out.len();
let simd_n = (n / 8) * 8;
let mut acc = initial;
let mut i = 0usize;
let q_vec = unsafe { vdupq_n_s16(q as i16) };
while i < simd_n {
let result = unsafe {
let v = vld1q_u16(zd.as_ptr().add(i));
let one = vdupq_n_u16(1);
let zero16 = vdupq_n_u16(0);
let low_bit = vandq_u16(v, one);
let sign = vsubq_u16(zero16, low_bit);
let shifted = vshrq_n_u16(v, 1);
let delta = vreinterpretq_s16_u16(veorq_u16(shifted, sign));
let zero = vdupq_n_s16(0);
let d = vaddq_s16(delta, vextq_s16(zero, delta, 7));
let d = vaddq_s16(d, vextq_s16(zero, d, 6));
let d = vaddq_s16(d, vextq_s16(zero, d, 4));
vaddq_s16(d, vdupq_n_s16(acc))
};
unsafe {
acc = vgetq_lane_s16(result, 7);
vst1q_s16(out.as_mut_ptr().add(base + i), vshlq_s16(result, q_vec));
}
i += 8;
}
unsafe {
out.set_len(base + simd_n);
}
for &code in &zd[simd_n..] {
acc = acc.wrapping_add(unzigzag_one(code));
out.push(acc << q);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(not(feature = "std"))]
use alloc::vec;
#[cfg(not(feature = "std"))]
use alloc::vec::Vec;
fn reference(zd: &[u16], initial: i16, q: u8) -> Vec<i16> {
let mut acc = initial;
zd.iter()
.map(|&code| {
acc = acc.wrapping_add(unzigzag_one(code));
acc << (q & 15)
})
.collect()
}
fn dispatch_matches_reference(zd: &[u16], q: u8) {
let mut out = Vec::new();
decode_into(zd, 0, q, &mut out);
assert_eq!(out, reference(zd, 0, q));
}
#[test]
fn empty() {
dispatch_matches_reference(&[], 0);
}
#[test]
fn small_tail_only() {
for n in 0..16 {
let zd: Vec<u16> = (0..n as u16).map(|i| i * 3 + 1).collect();
dispatch_matches_reference(&zd, 0);
}
}
#[test]
fn one_full_block_plus_tail() {
for n in [13u16, 29] {
let zd: Vec<u16> = (0..n).map(|i| i * 37 % 257).collect();
dispatch_matches_reference(&zd, 0);
}
}
#[test]
fn large_input() {
let zd: Vec<u16> = (0..1000u32).map(|i| ((i * 6151) % 65536) as u16).collect();
dispatch_matches_reference(&zd, 0);
}
#[test]
fn extremes() {
let zd = vec![0u16, 1, 65535, 65534, 32768, 32767];
dispatch_matches_reference(&zd, 0);
}
#[test]
fn nonzero_initial_carry() {
let zd: Vec<u16> = (0..20u16).map(|i| i * 5 + 2).collect();
let mut out = Vec::new();
decode_into(&zd, 1234, 0, &mut out);
assert_eq!(out, reference(&zd, 1234, 0));
}
#[test]
fn nonzero_shift() {
for q in [0u8, 1, 3, 5, 15] {
let zd: Vec<u16> = (0..37u16).map(|i| i * 91 % 401).collect();
dispatch_matches_reference(&zd, q);
}
}
#[test]
fn shift_amount_masked_to_avoid_panic() {
let zd: Vec<u16> = (0..20u16).map(|i| i * 7 + 1).collect();
let mut out = Vec::new();
decode_into(&zd, 0, 255, &mut out);
assert_eq!(out, reference(&zd, 0, 255));
}
#[cfg(all(target_arch = "x86_64", feature = "simd-avx2"))]
#[test]
fn avx2_matches_reference_directly() {
for n in 0..=31usize {
for q in [0u8, 1, 5, 15] {
let zd: Vec<u16> = (0..n as u32).map(|i| ((i * 6151) % 65536) as u16).collect();
let mut out = Vec::new();
unsafe { decode_avx2(&zd, 0, q, &mut out) };
assert_eq!(out, reference(&zd, 0, q), "n={n} q={q}");
}
}
}
#[cfg(all(target_arch = "x86_64", feature = "simd-avx2"))]
#[test]
fn avx2_nonzero_initial_carry() {
let zd: Vec<u16> = (0..40u16).map(|i| i * 5 + 2).collect();
let mut out = Vec::new();
unsafe { decode_avx2(&zd, 1234, 0, &mut out) };
assert_eq!(out, reference(&zd, 1234, 0));
}
#[cfg(all(
target_arch = "x86_64",
any(feature = "simd-auto", feature = "simd-ssse3")
))]
#[test]
fn sse2_matches_reference_directly() {
for q in [0u8, 1, 5] {
let zd: Vec<u16> = (0..37u16).map(|i| i * 91 % 401).collect();
let mut out = Vec::new();
unsafe { decode_sse2(&zd, 0, q, &mut out) };
assert_eq!(out, reference(&zd, 0, q));
}
}
#[cfg(all(
target_arch = "aarch64",
any(feature = "simd-auto", feature = "simd-neon")
))]
#[test]
fn neon_matches_reference_directly() {
for q in [0u8, 1, 5] {
let zd: Vec<u16> = (0..37u16).map(|i| i * 91 % 401).collect();
let mut out = Vec::new();
unsafe { decode_neon(&zd, 0, q, &mut out) };
assert_eq!(out, reference(&zd, 0, q));
}
}
}