#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) enum Float16Kind {
F16,
Bf16,
}
pub(super) fn float16_row_dot(kind: Float16Kind, row: &[u8], x: &[f32]) -> f32 {
let n = (row.len() / 2).min(x.len());
let (row, x) = (&row[..n * 2], &x[..n]);
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
match kind {
Float16Kind::F16 if is_x86_feature_detected!("f16c") => {
return unsafe { f16_row_dot_avx2(row, x) };
},
Float16Kind::Bf16 => {
return unsafe { bf16_row_dot_avx2(row, x) };
},
Float16Kind::F16 => {},
}
}
}
float16_row_dot_portable(kind, row, x)
}
#[inline]
fn bf16_bits_to_f32(bits: u16) -> f32 {
f32::from_bits(u32::from(bits) << 16)
}
#[inline]
fn decode_float16(kind: Float16Kind, bits: u16) -> f32 {
match kind {
Float16Kind::F16 => crate::quantize::f16_to_f32_lut(bits),
Float16Kind::Bf16 => bf16_bits_to_f32(bits),
}
}
pub(super) fn float16_row_dot_portable(kind: Float16Kind, row: &[u8], x: &[f32]) -> f32 {
const CHUNK: usize = 64;
let mut buf = [0.0f32; CHUNK];
let mut lanes = [0.0f32; 8];
let mut tail = 0.0f32;
for (bytes, xs) in row.chunks(CHUNK * 2).zip(x.chunks(CHUNK)) {
let m = xs.len();
for (w, b) in buf[..m].iter_mut().zip(bytes.as_chunks::<2>().0) {
*w = decode_float16(kind, u16::from_le_bytes(*b));
}
let ((ws, w_rem), (xs8, x_rem)) = (buf[..m].as_chunks::<8>(), xs.as_chunks::<8>());
for (w8, x8) in ws.iter().zip(xs8) {
for l in 0..8 {
lanes[l] += w8[l] * x8[l];
}
}
tail += w_rem.iter().zip(x_rem).map(|(w, x)| w * x).sum::<f32>();
}
lanes.iter().sum::<f32>() + tail
}
#[cfg(target_arch = "x86_64")]
macro_rules! avx2_float16_row_dot {
($row:ident, $x:ident, $load8:ident, $decode1:expr) => {{
use std::arch::x86_64::{
_mm256_add_ps, _mm256_castps256_ps128, _mm256_extractf128_ps, _mm256_fmadd_ps,
_mm256_loadu_ps, _mm256_setzero_ps, _mm_add_ps, _mm_cvtss_f32, _mm_hadd_ps,
};
let n = $x.len();
let (rp, xp) = ($row.as_ptr(), $x.as_ptr());
let (mut a0, mut a1, mut a2, mut a3) =
(_mm256_setzero_ps(), _mm256_setzero_ps(), _mm256_setzero_ps(), _mm256_setzero_ps());
let mut i = 0;
while i + 32 <= n {
a0 = _mm256_fmadd_ps($load8(rp.add(2 * i)), _mm256_loadu_ps(xp.add(i)), a0);
a1 = _mm256_fmadd_ps($load8(rp.add(2 * i + 16)), _mm256_loadu_ps(xp.add(i + 8)), a1);
a2 = _mm256_fmadd_ps($load8(rp.add(2 * i + 32)), _mm256_loadu_ps(xp.add(i + 16)), a2);
a3 = _mm256_fmadd_ps($load8(rp.add(2 * i + 48)), _mm256_loadu_ps(xp.add(i + 24)), a3);
i += 32;
}
while i + 8 <= n {
a0 = _mm256_fmadd_ps($load8(rp.add(2 * i)), _mm256_loadu_ps(xp.add(i)), a0);
i += 8;
}
let acc = _mm256_add_ps(_mm256_add_ps(a0, a1), _mm256_add_ps(a2, a3));
let s = _mm_add_ps(_mm256_castps256_ps128(acc), _mm256_extractf128_ps::<1>(acc));
let s = _mm_hadd_ps(s, s);
let mut sum = _mm_cvtss_f32(_mm_hadd_ps(s, s));
for j in i..n {
sum += $decode1(u16::from_le_bytes([$row[2 * j], $row[2 * j + 1]])) * $x[j];
}
sum
}};
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,f16c")]
#[inline]
unsafe fn load8_f16(p: *const u8) -> std::arch::x86_64::__m256 {
use std::arch::x86_64::{__m128i, _mm256_cvtph_ps, _mm_loadu_si128};
unsafe { _mm256_cvtph_ps(_mm_loadu_si128(p.cast::<__m128i>())) }
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
#[inline]
unsafe fn load8_bf16(p: *const u8) -> std::arch::x86_64::__m256 {
use std::arch::x86_64::{
__m128i, _mm256_castsi256_ps, _mm256_cvtepu16_epi32, _mm256_slli_epi32, _mm_loadu_si128,
};
unsafe {
let wide = _mm256_cvtepu16_epi32(_mm_loadu_si128(p.cast::<__m128i>()));
_mm256_castsi256_ps(_mm256_slli_epi32::<16>(wide))
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma,f16c")]
unsafe fn f16_row_dot_avx2(row: &[u8], x: &[f32]) -> f32 {
unsafe { avx2_float16_row_dot!(row, x, load8_f16, crate::quantize::f16_to_f32_lut) }
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn bf16_row_dot_avx2(row: &[u8], x: &[f32]) -> f32 {
unsafe { avx2_float16_row_dot!(row, x, load8_bf16, bf16_bits_to_f32) }
}