use super::{narrow, widen};
pub(crate) fn widen_f16(src: &[u16], dst: &mut [f32]) {
#[cfg(target_arch = "x86_64")]
{
if has_f16c() {
unsafe {
widen_f16_x86(src, dst);
}
return;
}
}
widen_f16_scalar(src, dst);
}
pub(crate) fn narrow_f16(src: &[f32], dst: &mut [u16]) {
#[cfg(target_arch = "x86_64")]
{
if has_f16c() {
unsafe {
narrow_f16_x86(src, dst);
}
return;
}
}
narrow_f16_scalar(src, dst);
}
#[inline]
fn widen_f16_scalar(src: &[u16], dst: &mut [f32]) {
let n = src.len().min(dst.len());
for i in 0..n {
dst[i] = f32::from_bits(widen::<5, 10>(u32::from(src[i])));
}
}
#[inline]
fn narrow_f16_scalar(src: &[f32], dst: &mut [u16]) {
let n = src.len().min(dst.len());
for i in 0..n {
dst[i] = narrow::<5, 10>(src[i].to_bits()) as u16;
}
}
pub(crate) fn widen_bf16(src: &[u16], dst: &mut [f32]) {
let n = src.len().min(dst.len());
for i in 0..n {
dst[i] = f32::from_bits(u32::from(src[i]) << 16);
}
}
pub(crate) fn narrow_bf16(src: &[f32], dst: &mut [u16]) {
let n = src.len().min(dst.len());
for i in 0..n {
dst[i] = round_f32_to_bf16(src[i].to_bits());
}
}
#[inline]
fn round_f32_to_bf16(bits: u32) -> u16 {
let rounded = bits.wrapping_add(0x7FFF).wrapping_add((bits >> 16) & 1) >> 16;
let high = bits >> 16;
let nan = high | u32::from(high & 0x7F == 0);
let is_nan = (bits & 0x7F80_0000 == 0x7F80_0000) & (bits & 0x007F_FFFF != 0);
(if is_nan { nan } else { rounded }) as u16
}
#[cfg(target_arch = "x86_64")]
#[inline(always)]
fn has_f16c() -> bool {
#[cfg(feature = "std")]
{
std::is_x86_feature_detected!("f16c")
}
#[cfg(not(feature = "std"))]
{
cfg!(target_feature = "f16c")
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "f16c")]
unsafe fn widen_f16_x86(src: &[u16], dst: &mut [f32]) {
use core::arch::x86_64::{_mm256_cvtph_ps, _mm256_storeu_ps, _mm_loadu_si128};
let n = src.len().min(dst.len());
let mut i = 0;
while i + 8 <= n {
let packed = _mm_loadu_si128(src.as_ptr().add(i).cast());
let widened = _mm256_cvtph_ps(packed);
_mm256_storeu_ps(dst.as_mut_ptr().add(i), widened);
i += 8;
}
while i < n {
dst[i] = f32::from_bits(widen::<5, 10>(u32::from(src[i])));
i += 1;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "f16c")]
unsafe fn narrow_f16_x86(src: &[f32], dst: &mut [u16]) {
use core::arch::x86_64::{
_mm256_cvtps_ph, _mm256_loadu_ps, _mm_storeu_si128, _MM_FROUND_TO_NEAREST_INT,
};
let n = src.len().min(dst.len());
let mut i = 0;
while i + 8 <= n {
let values = _mm256_loadu_ps(src.as_ptr().add(i));
let narrowed = _mm256_cvtps_ph::<{ _MM_FROUND_TO_NEAREST_INT }>(values);
_mm_storeu_si128(dst.as_mut_ptr().add(i).cast(), narrowed);
i += 8;
}
while i < n {
dst[i] = narrow::<5, 10>(src[i].to_bits()) as u16;
i += 1;
}
}