use crate::types::{Bf16, Bf4, Bf8, F32, F4, F8};
#[cfg(target_arch = "x86_64")]
use super::arch::{has_avx2, has_avx512bw, has_avx512f};
#[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))]
use super::unsafe_intrinsics;
#[cfg(not(target_arch = "aarch64"))]
use crate::convert::{widen_finite, widen_finite_high_word, widen_high_word};
#[inline]
pub fn unpack_bf8_to_bf16(packed: &[Bf8], unpacked: &mut [Bf16]) {
#[cfg(target_arch = "x86_64")]
{
if has_avx512bw() {
unsafe {
unsafe_intrinsics::avx512::unpack_bf8_to_bf16(packed, unpacked);
}
return;
}
if has_avx2() {
unsafe {
unsafe_intrinsics::avx2::unpack_bf8_to_bf16(packed, unpacked);
}
return;
}
}
#[cfg(target_arch = "aarch64")]
{
unsafe {
unsafe_intrinsics::neon::unpack_bf8_to_bf16(packed, unpacked);
}
return;
}
#[cfg(not(target_arch = "aarch64"))]
{
let len = packed.len();
let n = len.min(unpacked.len());
for i in 0..n {
unpacked[i] = Bf16(widen_high_word::<5, 2>(u32::from(packed[i].0)));
}
}
}
#[inline]
pub fn unpack_bf4_to_bf16(packed: &[Bf4], unpacked: &mut [Bf16]) {
#[cfg(target_arch = "x86_64")]
{
if has_avx512bw() {
unsafe {
unsafe_intrinsics::avx512::unpack_bf4_to_bf16(packed, unpacked);
}
return;
}
if has_avx2() {
unsafe {
unsafe_intrinsics::avx2::unpack_bf4_to_bf16(packed, unpacked);
}
return;
}
}
#[cfg(target_arch = "aarch64")]
{
unsafe {
unsafe_intrinsics::neon::unpack_bf4_to_bf16(packed, unpacked);
}
return;
}
#[cfg(not(target_arch = "aarch64"))]
{
let len = packed.len();
let n = len.min(unpacked.len());
for i in 0..n {
let b = packed[i].0;
unpacked[i] = Bf16(widen_finite_high_word::<2, 1>(b as u32));
}
}
}
#[inline]
pub fn unpack_bf4_to_bf16_packed(packed: &[u8], unpacked: &mut [Bf16]) {
#[cfg(target_arch = "x86_64")]
{
if has_avx512bw() {
unsafe {
unsafe_intrinsics::avx512::unpack_bf4_to_bf16_packed(packed, unpacked);
}
return;
}
if has_avx2() {
unsafe {
unsafe_intrinsics::avx2::unpack_bf4_to_bf16_packed(packed, unpacked);
}
return;
}
}
#[cfg(target_arch = "aarch64")]
{
unsafe {
unsafe_intrinsics::neon::unpack_bf4_to_bf16_packed(packed, unpacked);
}
return;
}
#[cfg(not(target_arch = "aarch64"))]
{
let len = packed.len();
let n = len.min(unpacked.len() / 2);
for i in 0..n {
let byte = packed[i];
let b1 = byte & 0x0f;
let b2 = (byte >> 4) & 0x0f;
unpacked[2 * i] = Bf16(widen_finite_high_word::<2, 1>(b1 as u32));
unpacked[2 * i + 1] = Bf16(widen_finite_high_word::<2, 1>(b2 as u32));
}
}
}
#[inline]
pub fn unpack_f4_to_f32(packed: &[F4], unpacked: &mut [F32]) {
#[cfg(target_arch = "x86_64")]
{
if has_avx512f() {
unsafe {
unsafe_intrinsics::avx512::unpack_f4_to_f32(packed, unpacked);
}
return;
}
if has_avx2() {
unsafe {
unsafe_intrinsics::avx2::unpack_f4_to_f32(packed, unpacked);
}
return;
}
}
#[cfg(target_arch = "aarch64")]
{
unsafe {
unsafe_intrinsics::neon::unpack_f4_to_f32(packed, unpacked);
}
return;
}
#[cfg(not(target_arch = "aarch64"))]
{
let len = packed.len().min(unpacked.len());
for i in 0..len {
unpacked[i] = F32(packed[i].to_f32());
}
}
}
#[inline]
pub fn unpack_f4_to_f32_packed(packed: &[u8], unpacked: &mut [F32]) {
#[cfg(target_arch = "x86_64")]
{
if has_avx512f() {
unsafe {
unsafe_intrinsics::avx512::unpack_f4_to_f32_packed(packed, unpacked);
}
return;
}
if has_avx2() {
unsafe {
unsafe_intrinsics::avx2::unpack_f4_to_f32_packed(packed, unpacked);
}
return;
}
}
#[cfg(target_arch = "aarch64")]
{
unsafe {
unsafe_intrinsics::neon::unpack_f4_to_f32_packed(packed, unpacked);
}
return;
}
#[cfg(not(target_arch = "aarch64"))]
{
let len = packed.len();
let n = len.min(unpacked.len() / 2);
for i in 0..n {
let byte = packed[i];
let b1 = byte & 0x0f;
let b2 = (byte >> 4) & 0x0f;
unpacked[2 * i] = F32(F4(b1).to_f32());
unpacked[2 * i + 1] = F32(F4(b2).to_f32());
}
}
}
#[inline]
pub fn unpack_f8_to_f32(packed: &[F8], unpacked: &mut [F32]) {
#[cfg(target_arch = "x86_64")]
{
if has_avx512f() {
unsafe {
unsafe_intrinsics::avx512::unpack_f8_to_f32(packed, unpacked);
}
return;
}
if has_avx2() {
unsafe {
unsafe_intrinsics::avx2::unpack_f8_to_f32(packed, unpacked);
}
return;
}
}
#[cfg(target_arch = "aarch64")]
{
unsafe {
unsafe_intrinsics::neon::unpack_f8_to_f32(packed, unpacked);
}
return;
}
#[cfg(not(target_arch = "aarch64"))]
{
static TABLE_BITS: [u32; 256] = {
let mut t = [0u32; 256];
let mut idx = 0;
while idx < 256 {
t[idx] = widen_finite::<4, 3>(idx as u32);
idx += 1;
}
t
};
let len = packed.len().min(unpacked.len());
for i in 0..len {
unpacked[i] = F32(f32::from_bits(TABLE_BITS[packed[i].0 as usize]));
}
}
}