#[cfg(target_arch = "x86_64")]
use super::horizontal::horizontal_sum_256;
#[cfg(target_arch = "x86_64")]
use super::is_avx2_fma_available;
#[inline(always)]
pub fn simd_dot_f32(a: &[f32], b: &[f32], len: usize) -> f32 {
#[cfg(target_arch = "aarch64")]
{
unsafe { neon_dot_f32(a, b, len) }
}
#[cfg(target_arch = "x86_64")]
{
if is_avx2_fma_available() {
unsafe { avx2_dot_f32(a, b, len) }
} else {
scalar_dot_f32(a, b, len)
}
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
{
unsafe { wasm32_simd128_dot_f32(a, b, len) }
}
#[cfg(not(any(
target_arch = "aarch64",
target_arch = "x86_64",
all(target_arch = "wasm32", target_feature = "simd128")
)))]
{
scalar_dot_f32(a, b, len)
}
}
#[inline(always)]
pub fn simd_fma_row(weight_row: &[f32], input: &[f32], len: usize) -> f32 {
simd_dot_f32(weight_row, input, len)
}
#[inline(always)]
#[allow(dead_code)]
pub(super) fn scalar_dot_f32(a: &[f32], b: &[f32], len: usize) -> f32 {
let mut acc = [0.0f32; 4];
let chunks = len / 4;
let mut i = 0;
for _ in 0..chunks {
unsafe {
acc[0] = (*a.get_unchecked(i)).mul_add(*b.get_unchecked(i), acc[0]);
acc[1] = (*a.get_unchecked(i + 1)).mul_add(*b.get_unchecked(i + 1), acc[1]);
acc[2] = (*a.get_unchecked(i + 2)).mul_add(*b.get_unchecked(i + 2), acc[2]);
acc[3] = (*a.get_unchecked(i + 3)).mul_add(*b.get_unchecked(i + 3), acc[3]);
}
i += 4;
}
let mut sum = acc.iter().sum::<f32>();
while i < len {
unsafe {
sum = (*a.get_unchecked(i)).mul_add(*b.get_unchecked(i), sum);
}
i += 1;
}
sum
}
#[cfg(target_arch = "aarch64")]
#[inline]
unsafe fn neon_dot_f32(a: &[f32], b: &[f32], len: usize) -> f32 {
use core::arch::aarch64::{vaddq_f32, vaddvq_f32, vdupq_n_f32, vfmaq_f32, vld1q_f32};
unsafe {
let mut acc0 = vdupq_n_f32(0.0);
let mut acc1 = vdupq_n_f32(0.0);
let mut acc2 = vdupq_n_f32(0.0);
let mut acc3 = vdupq_n_f32(0.0);
let mut i = 0;
let chunks4 = len / 16;
for _ in 0..chunks4 {
acc0 = vfmaq_f32(
acc0,
vld1q_f32(a.as_ptr().add(i)),
vld1q_f32(b.as_ptr().add(i)),
);
acc1 = vfmaq_f32(
acc1,
vld1q_f32(a.as_ptr().add(i + 4)),
vld1q_f32(b.as_ptr().add(i + 4)),
);
acc2 = vfmaq_f32(
acc2,
vld1q_f32(a.as_ptr().add(i + 8)),
vld1q_f32(b.as_ptr().add(i + 8)),
);
acc3 = vfmaq_f32(
acc3,
vld1q_f32(a.as_ptr().add(i + 12)),
vld1q_f32(b.as_ptr().add(i + 12)),
);
i += 16;
}
let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3)));
let mut acc_rem = vdupq_n_f32(0.0);
let remaining = (len - i) / 4;
for _ in 0..remaining {
acc_rem = vfmaq_f32(
acc_rem,
vld1q_f32(a.as_ptr().add(i)),
vld1q_f32(b.as_ptr().add(i)),
);
i += 4;
}
sum += vaddvq_f32(acc_rem);
while i < len {
sum += *a.get_unchecked(i) * *b.get_unchecked(i);
i += 1;
}
sum
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
#[inline]
unsafe fn avx2_dot_f32(a: &[f32], b: &[f32], len: usize) -> f32 {
use core::arch::x86_64::{_mm256_add_ps, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_setzero_ps};
unsafe {
let mut acc0 = _mm256_setzero_ps();
let mut acc1 = _mm256_setzero_ps();
let mut acc2 = _mm256_setzero_ps();
let mut acc3 = _mm256_setzero_ps();
let mut i = 0;
let chunks4 = len / 32;
for _ in 0..chunks4 {
acc0 = _mm256_fmadd_ps(
_mm256_loadu_ps(a.as_ptr().add(i)),
_mm256_loadu_ps(b.as_ptr().add(i)),
acc0,
);
acc1 = _mm256_fmadd_ps(
_mm256_loadu_ps(a.as_ptr().add(i + 8)),
_mm256_loadu_ps(b.as_ptr().add(i + 8)),
acc1,
);
acc2 = _mm256_fmadd_ps(
_mm256_loadu_ps(a.as_ptr().add(i + 16)),
_mm256_loadu_ps(b.as_ptr().add(i + 16)),
acc2,
);
acc3 = _mm256_fmadd_ps(
_mm256_loadu_ps(a.as_ptr().add(i + 24)),
_mm256_loadu_ps(b.as_ptr().add(i + 24)),
acc3,
);
i += 32;
}
let mut sum = horizontal_sum_256(_mm256_add_ps(
_mm256_add_ps(acc0, acc1),
_mm256_add_ps(acc2, acc3),
));
let mut acc = _mm256_setzero_ps();
let remaining = (len - i) / 8;
for _ in 0..remaining {
let va = _mm256_loadu_ps(a.as_ptr().add(i));
let vb = _mm256_loadu_ps(b.as_ptr().add(i));
acc = _mm256_fmadd_ps(va, vb, acc);
i += 8;
}
sum += horizontal_sum_256(acc);
while i < len {
sum += *a.get_unchecked(i) * *b.get_unchecked(i);
i += 1;
}
sum
}
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
#[inline]
unsafe fn wasm32_simd128_dot_f32(a: &[f32], b: &[f32], len: usize) -> f32 {
use core::arch::wasm32::{f32x4_add, f32x4_extract_lane, f32x4_mul, f32x4_splat, v128_load};
unsafe {
let mut acc0 = f32x4_splat(0.0);
let mut acc1 = f32x4_splat(0.0);
let mut acc2 = f32x4_splat(0.0);
let mut acc3 = f32x4_splat(0.0);
let mut i = 0usize;
let chunks4 = len / 16;
for _ in 0..chunks4 {
acc0 = f32x4_add(
f32x4_mul(
v128_load(a.as_ptr().add(i).cast()),
v128_load(b.as_ptr().add(i).cast()),
),
acc0,
);
acc1 = f32x4_add(
f32x4_mul(
v128_load(a.as_ptr().add(i + 4).cast()),
v128_load(b.as_ptr().add(i + 4).cast()),
),
acc1,
);
acc2 = f32x4_add(
f32x4_mul(
v128_load(a.as_ptr().add(i + 8).cast()),
v128_load(b.as_ptr().add(i + 8).cast()),
),
acc2,
);
acc3 = f32x4_add(
f32x4_mul(
v128_load(a.as_ptr().add(i + 12).cast()),
v128_load(b.as_ptr().add(i + 12).cast()),
),
acc3,
);
i += 16;
}
let s01 = f32x4_add(acc0, acc1);
let s23 = f32x4_add(acc2, acc3);
let s = f32x4_add(s01, s23);
let mut sum = f32x4_extract_lane::<0>(s)
+ f32x4_extract_lane::<1>(s)
+ f32x4_extract_lane::<2>(s)
+ f32x4_extract_lane::<3>(s);
let mut acc = f32x4_splat(0.0);
let remaining = (len - i) / 4;
for _ in 0..remaining {
acc = f32x4_add(
f32x4_mul(
v128_load(a.as_ptr().add(i).cast()),
v128_load(b.as_ptr().add(i).cast()),
),
acc,
);
i += 4;
}
sum += f32x4_extract_lane::<0>(acc)
+ f32x4_extract_lane::<1>(acc)
+ f32x4_extract_lane::<2>(acc)
+ f32x4_extract_lane::<3>(acc);
while i < len {
sum += *a.get_unchecked(i) * *b.get_unchecked(i);
i += 1;
}
sum
}
}
#[inline(always)]
pub fn simd_outer_product_acc(acc: &mut [f32], a: &[f32], b: &[f32], m: usize, n: usize) {
#[cfg(target_arch = "aarch64")]
{
unsafe { neon_outer_product_acc(acc, a, b, m, n) }
}
#[cfg(target_arch = "x86_64")]
{
if is_avx2_fma_available() {
unsafe { avx2_outer_product_acc(acc, a, b, m, n) }
} else {
scalar_outer_product_acc(acc, a, b, m, n)
}
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
{
unsafe { wasm32_simd128_outer_product_acc(acc, a, b, m, n) }
}
#[cfg(not(any(
target_arch = "aarch64",
target_arch = "x86_64",
all(target_arch = "wasm32", target_feature = "simd128")
)))]
{
scalar_outer_product_acc(acc, a, b, m, n)
}
}
#[inline(always)]
#[allow(dead_code)]
pub(super) fn scalar_outer_product_acc(acc: &mut [f32], a: &[f32], b: &[f32], m: usize, n: usize) {
for i in 0..m {
let ai = unsafe { *a.get_unchecked(i) };
let row = &mut acc[i * n..i * n + n];
for j in 0..n {
unsafe {
let bj = *b.get_unchecked(j);
*row.get_unchecked_mut(j) = ai.mul_add(bj, *row.get_unchecked(j));
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[inline]
unsafe fn neon_outer_product_acc(acc: &mut [f32], a: &[f32], b: &[f32], m: usize, n: usize) {
use core::arch::aarch64::{vfmaq_f32, vld1q_dup_f32, vld1q_f32, vst1q_f32};
unsafe {
let n_chunks = n / 4;
for i in 0..m {
let ai = *a.get_unchecked(i);
let va = vld1q_dup_f32(&ai);
let row = &mut acc[i * n..i * n + n];
let mut j = 0;
for _ in 0..n_chunks {
let vacc = vld1q_f32(row.as_ptr().add(j));
let vb = vld1q_f32(b.as_ptr().add(j));
let vresult = vfmaq_f32(vacc, va, vb);
vst1q_f32(row.as_mut_ptr().add(j), vresult);
j += 4;
}
while j < n {
*row.get_unchecked_mut(j) += ai * *b.get_unchecked(j);
j += 1;
}
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
#[inline]
unsafe fn avx2_outer_product_acc(acc: &mut [f32], a: &[f32], b: &[f32], m: usize, n: usize) {
use core::arch::x86_64::{
_mm256_broadcast_ss, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_storeu_ps,
};
unsafe {
let n_chunks8 = n / 8;
for i in 0..m {
let ai = *a.get_unchecked(i);
let va = _mm256_broadcast_ss(&ai);
let row = &mut acc[i * n..i * n + n];
let mut j = 0;
for _ in 0..n_chunks8 {
let vacc = _mm256_loadu_ps(row.as_ptr().add(j));
let vb = _mm256_loadu_ps(b.as_ptr().add(j));
let vresult = _mm256_fmadd_ps(va, vb, vacc);
_mm256_storeu_ps(row.as_mut_ptr().add(j), vresult);
j += 8;
}
while j < n {
*row.get_unchecked_mut(j) += ai * *b.get_unchecked(j);
j += 1;
}
}
}
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
#[inline]
unsafe fn wasm32_simd128_outer_product_acc(
acc: &mut [f32],
a: &[f32],
b: &[f32],
m: usize,
n: usize,
) {
use core::arch::wasm32::{f32x4_add, f32x4_mul, f32x4_splat, v128_load, v128_store};
unsafe {
let simd_n = n / 4 * 4;
for i in 0..m {
let ai = *a.get_unchecked(i);
let vai = f32x4_splat(ai);
let row = &mut acc[i * n..i * n + n];
let mut j = 0;
while j < simd_n {
let vacc = v128_load(row.as_ptr().add(j) as *const _);
let vb = v128_load(b.as_ptr().add(j) as *const _);
let r = f32x4_add(vacc, f32x4_mul(vai, vb));
v128_store(row.as_mut_ptr().add(j) as *mut _, r);
j += 4;
}
for jj in simd_n..n {
*row.get_unchecked_mut(jj) += ai * *b.get_unchecked(jj);
}
}
}
}
#[inline(always)]
pub fn simd_outer_product_acc_scaled(
acc: &mut [f32],
scale: f32,
a: &[f32],
b: &[f32],
m: usize,
n: usize,
) {
#[cfg(target_arch = "aarch64")]
{
unsafe { neon_outer_product_acc_scaled(acc, scale, a, b, m, n) }
}
#[cfg(target_arch = "x86_64")]
{
if is_avx2_fma_available() {
unsafe { avx2_outer_product_acc_scaled(acc, scale, a, b, m, n) }
} else {
scalar_outer_product_acc_scaled(acc, scale, a, b, m, n)
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
{
scalar_outer_product_acc_scaled(acc, scale, a, b, m, n)
}
}
#[inline(always)]
#[allow(dead_code)]
pub(super) fn scalar_outer_product_acc_scaled(
acc: &mut [f32],
scale: f32,
a: &[f32],
b: &[f32],
m: usize,
n: usize,
) {
for i in 0..m {
let ai = scale.mul_add(unsafe { *a.get_unchecked(i) }, 0.0);
let row = &mut acc[i * n..i * n + n];
for j in 0..n {
unsafe {
let bj = *b.get_unchecked(j);
*row.get_unchecked_mut(j) = ai.mul_add(bj, *row.get_unchecked(j));
}
}
}
}
#[cfg(target_arch = "aarch64")]
#[inline]
unsafe fn neon_outer_product_acc_scaled(
acc: &mut [f32],
scale: f32,
a: &[f32],
b: &[f32],
m: usize,
n: usize,
) {
use core::arch::aarch64::{vfmaq_f32, vld1q_dup_f32, vld1q_f32, vst1q_f32};
unsafe {
let n_chunks = n / 4;
for i in 0..m {
let ai_scaled = scale * *a.get_unchecked(i);
let va = vld1q_dup_f32(&ai_scaled);
let row = &mut acc[i * n..i * n + n];
let mut j = 0;
for _ in 0..n_chunks {
let vacc = vld1q_f32(row.as_ptr().add(j));
let vb = vld1q_f32(b.as_ptr().add(j));
let vresult = vfmaq_f32(vacc, va, vb);
vst1q_f32(row.as_mut_ptr().add(j), vresult);
j += 4;
}
while j < n {
*row.get_unchecked_mut(j) += ai_scaled * *b.get_unchecked(j);
j += 1;
}
}
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
#[inline]
unsafe fn avx2_outer_product_acc_scaled(
acc: &mut [f32],
scale: f32,
a: &[f32],
b: &[f32],
m: usize,
n: usize,
) {
use core::arch::x86_64::{
_mm256_broadcast_ss, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_storeu_ps,
};
unsafe {
let n_chunks8 = n / 8;
for i in 0..m {
let ai_scaled = scale * *a.get_unchecked(i);
let va = _mm256_broadcast_ss(&ai_scaled);
let row = &mut acc[i * n..i * n + n];
let mut j = 0;
for _ in 0..n_chunks8 {
let vacc = _mm256_loadu_ps(row.as_ptr().add(j));
let vb = _mm256_loadu_ps(b.as_ptr().add(j));
let vresult = _mm256_fmadd_ps(va, vb, vacc);
_mm256_storeu_ps(row.as_mut_ptr().add(j), vresult);
j += 8;
}
while j < n {
*row.get_unchecked_mut(j) += ai_scaled * *b.get_unchecked(j);
j += 1;
}
}
}
}
#[inline(always)]
pub fn simd_matvec(acc: &mut [f32], mat: &[f32], vec: &[f32], rows: usize, cols: usize) {
for r in 0..rows {
let row_off = r * cols;
unsafe {
*acc.get_unchecked_mut(r) = simd_dot_f32(&mat[row_off..row_off + cols], vec, cols);
}
}
}
#[inline(always)]
pub fn simd_matmul_rows(
output: &mut [f32],
weight: &[f32],
input: &[f32],
rows: usize,
cols: usize,
) {
for r in 0..rows {
let row_off = r * cols;
unsafe {
*output.get_unchecked_mut(r) =
simd_dot_f32(&weight[row_off..row_off + cols], input, cols);
}
}
}
#[inline]
pub fn simd_matmul_rows_parallel(
output: &mut [f32],
weight: &[f32],
input: &[f32],
rows: usize,
cols: usize,
) {
const PARALLEL_ROWS_MIN: usize = 512;
if rows < PARALLEL_ROWS_MIN {
simd_matmul_rows(output, weight, input, rows, cols);
} else {
use rayon::prelude::*;
const PARALLEL_CHUNK_ROWS: usize = 256;
output
.par_chunks_mut(PARALLEL_CHUNK_ROWS)
.enumerate()
.for_each(|(chunk_idx, out_chunk)| {
let start_row = chunk_idx * PARALLEL_CHUNK_ROWS;
for (local_r, out) in out_chunk.iter_mut().enumerate() {
let r = start_row + local_r;
let row_off = r * cols;
*out = simd_dot_f32(&weight[row_off..row_off + cols], input, cols);
}
});
}
}
#[inline(always)]
pub fn simd_matmul_relu_rows(
output: &mut [f32],
weight: &[f32],
input: &[f32],
rows: usize,
cols: usize,
) {
for r in 0..rows {
let row_off = r * cols;
let sum = simd_dot_f32(&weight[row_off..row_off + cols], input, cols);
unsafe {
*output.get_unchecked_mut(r) = sum.max(0.0);
}
}
}
#[inline]
pub fn simd_dot_f16_f32(w_f16: &[half::f16], x_f32: &[f32], len: usize) -> f32 {
#[cfg(target_arch = "aarch64")]
{
unsafe { neon_dot_f16_f32(w_f16, x_f32, len) }
}
#[cfg(not(target_arch = "aarch64"))]
{
scalar_dot_f16_f32(w_f16, x_f32, len)
}
}
#[inline(always)]
#[allow(dead_code)]
pub(super) fn scalar_dot_f16_f32(w: &[half::f16], x: &[f32], len: usize) -> f32 {
let mut acc = [0.0f32; 4];
let chunks = len / 4;
let mut i = 0;
for _ in 0..chunks {
unsafe {
acc[0] = (*w.get_unchecked(i))
.to_f32()
.mul_add(*x.get_unchecked(i), acc[0]);
acc[1] = (*w.get_unchecked(i + 1))
.to_f32()
.mul_add(*x.get_unchecked(i + 1), acc[1]);
acc[2] = (*w.get_unchecked(i + 2))
.to_f32()
.mul_add(*x.get_unchecked(i + 2), acc[2]);
acc[3] = (*w.get_unchecked(i + 3))
.to_f32()
.mul_add(*x.get_unchecked(i + 3), acc[3]);
}
i += 4;
}
let mut sum = acc.iter().sum::<f32>();
while i < len {
unsafe {
sum = (*w.get_unchecked(i))
.to_f32()
.mul_add(*x.get_unchecked(i), sum);
}
i += 1;
}
sum
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
#[inline]
unsafe fn fcvtl_4x(w_ptr: *const u16, out: *mut f32) {
let mut _v0: f32;
let mut _v1: f32;
unsafe {
core::arch::asm!(
"ldr d0, [{addr}]",
"fcvtl v1.4s, v0.4h",
"str q1, [{out}]",
addr = in(reg) w_ptr,
out = in(reg) out,
lateout("v0") _v0,
lateout("v1") _v1,
options(nostack, preserves_flags),
);
}
}
#[cfg(target_arch = "aarch64")]
#[inline]
unsafe fn neon_dot_f16_f32(w: &[half::f16], x: &[f32], len: usize) -> f32 {
use core::arch::aarch64::{vaddq_f32, vaddvq_f32, vdupq_n_f32, vfmaq_f32, vld1q_f32};
unsafe {
let mut acc0 = vdupq_n_f32(0.0);
let mut acc1 = vdupq_n_f32(0.0);
let mut acc2 = vdupq_n_f32(0.0);
let mut acc3 = vdupq_n_f32(0.0);
let mut i = 0;
let w_ptr = w.as_ptr() as *const u16;
let mut buf = [0.0f32; 4];
let chunks16 = len / 16;
for _ in 0..chunks16 {
fcvtl_4x(w_ptr.add(i), buf.as_mut_ptr());
let w0 = vld1q_f32(buf.as_ptr());
fcvtl_4x(w_ptr.add(i + 4), buf.as_mut_ptr());
let w1 = vld1q_f32(buf.as_ptr());
fcvtl_4x(w_ptr.add(i + 8), buf.as_mut_ptr());
let w2 = vld1q_f32(buf.as_ptr());
fcvtl_4x(w_ptr.add(i + 12), buf.as_mut_ptr());
let w3 = vld1q_f32(buf.as_ptr());
acc0 = vfmaq_f32(acc0, w0, vld1q_f32(x.as_ptr().add(i)));
acc1 = vfmaq_f32(acc1, w1, vld1q_f32(x.as_ptr().add(i + 4)));
acc2 = vfmaq_f32(acc2, w2, vld1q_f32(x.as_ptr().add(i + 8)));
acc3 = vfmaq_f32(acc3, w3, vld1q_f32(x.as_ptr().add(i + 12)));
i += 16;
}
let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3)));
let chunks4 = (len - i) / 4;
let mut acc_rem = vdupq_n_f32(0.0);
for _ in 0..chunks4 {
fcvtl_4x(w_ptr.add(i), buf.as_mut_ptr());
let w0 = vld1q_f32(buf.as_ptr());
acc_rem = vfmaq_f32(acc_rem, w0, vld1q_f32(x.as_ptr().add(i)));
i += 4;
}
sum += vaddvq_f32(acc_rem);
while i < len {
sum += (*w.get_unchecked(i)).to_f32() * *x.get_unchecked(i);
i += 1;
}
sum
}
}
#[inline(always)]
pub fn simd_matmul_f16_f32_rows(
output: &mut [f32],
weight_f16: &[half::f16],
input_f32: &[f32],
rows: usize,
cols: usize,
) {
for r in 0..rows {
let row_off = r * cols;
unsafe {
*output.get_unchecked_mut(r) =
simd_dot_f16_f32(&weight_f16[row_off..row_off + cols], input_f32, cols);
}
}
}
#[inline]
pub fn simd_matmul_f16_f32_rows_parallel(
output: &mut [f32],
weight_f16: &[half::f16],
input_f32: &[f32],
rows: usize,
cols: usize,
) {
const PARALLEL_ROWS_MIN: usize = 512;
if rows < PARALLEL_ROWS_MIN {
simd_matmul_f16_f32_rows(output, weight_f16, input_f32, rows, cols);
} else {
use rayon::prelude::*;
const PARALLEL_CHUNK_ROWS: usize = 256;
output
.par_chunks_mut(PARALLEL_CHUNK_ROWS)
.enumerate()
.for_each(|(chunk_idx, out_chunk)| {
let start_row = chunk_idx * PARALLEL_CHUNK_ROWS;
for (local_r, out) in out_chunk.iter_mut().enumerate() {
let r = start_row + local_r;
let row_off = r * cols;
*out = simd_dot_f16_f32(&weight_f16[row_off..row_off + cols], input_f32, cols);
}
});
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon,fp16,fhm")]
#[inline]
unsafe fn fmlal_lo_inline(
acc: core::arch::aarch64::float32x4_t,
a: core::arch::aarch64::uint16x8_t,
b: core::arch::aarch64::uint16x8_t,
) -> core::arch::aarch64::float32x4_t {
let mut acc = acc;
unsafe {
core::arch::asm!
(
"fmlal v0.4s, v1.4h, v2.4h",
inout("v0") acc,
in("v1") a,
in("v2") b,
lateout("v1") _,
lateout("v2") _,
options(pure, nomem, nostack, preserves_flags),
);
acc
}
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon,fp16,fhm")]
#[inline]
unsafe fn fmlal_hi_inline(
acc: core::arch::aarch64::float32x4_t,
a: core::arch::aarch64::uint16x8_t,
b: core::arch::aarch64::uint16x8_t,
) -> core::arch::aarch64::float32x4_t {
let mut acc = acc;
unsafe {
core::arch::asm!
(
"fmlal2 v0.4s, v1.4h, v2.4h",
inout("v0") acc,
in("v1") a,
in("v2") b,
lateout("v1") _,
lateout("v2") _,
options(pure, nomem, nostack, preserves_flags),
);
acc
}
}
#[inline(always)]
pub fn simd_dot_f16_f16(w: &[half::f16], x: &[half::f16], len: usize) -> f32 {
#[cfg(target_arch = "aarch64")]
{
if std::arch::is_aarch64_feature_detected!("fp16")
&& std::arch::is_aarch64_feature_detected!("fhm")
{
unsafe { neon_dot_f16_f16(w, x, len) }
} else {
scalar_dot_f16_f16(w, x, len)
}
}
#[cfg(not(target_arch = "aarch64"))]
{
scalar_dot_f16_f16(w, x, len)
}
}
#[inline(always)]
#[allow(dead_code)]
pub(super) fn scalar_dot_f16_f16(w: &[half::f16], x: &[half::f16], len: usize) -> f32 {
let mut acc = [0.0f32; 4];
let chunks = len / 4;
let mut i = 0;
for _ in 0..chunks {
unsafe {
acc[0] = (*w.get_unchecked(i))
.to_f32()
.mul_add((*x.get_unchecked(i)).to_f32(), acc[0]);
acc[1] = (*w.get_unchecked(i + 1))
.to_f32()
.mul_add((*x.get_unchecked(i + 1)).to_f32(), acc[1]);
acc[2] = (*w.get_unchecked(i + 2))
.to_f32()
.mul_add((*x.get_unchecked(i + 2)).to_f32(), acc[2]);
acc[3] = (*w.get_unchecked(i + 3))
.to_f32()
.mul_add((*x.get_unchecked(i + 3)).to_f32(), acc[3]);
}
i += 4;
}
let mut sum = acc.iter().sum::<f32>();
while i < len {
unsafe {
sum = (*w.get_unchecked(i))
.to_f32()
.mul_add((*x.get_unchecked(i)).to_f32(), sum);
}
i += 1;
}
sum
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon,fp16,fhm")]
#[inline]
unsafe fn neon_dot_f16_f16(w: &[half::f16], x: &[half::f16], len: usize) -> f32 {
use core::arch::aarch64::{
float32x4_t, uint16x8_t, vaddq_f32, vaddvq_f32, vdupq_n_f32, vld1q_u16,
};
unsafe {
let w_ptr = w.as_ptr() as *const u16;
let x_ptr = x.as_ptr() as *const u16;
let zero = vdupq_n_f32(0.0);
let mut acc0: float32x4_t = zero;
let mut acc1: float32x4_t = zero;
let mut acc2: float32x4_t = zero;
let mut acc3: float32x4_t = zero;
let mut i = 0;
let chunks16 = len / 16;
for _ in 0..chunks16 {
let w0: uint16x8_t = vld1q_u16(w_ptr.add(i));
let x0: uint16x8_t = vld1q_u16(x_ptr.add(i));
acc0 = fmlal_lo_inline(acc0, w0, x0); acc1 = fmlal_hi_inline(acc1, w0, x0);
let w1: uint16x8_t = vld1q_u16(w_ptr.add(i + 8));
let x1: uint16x8_t = vld1q_u16(x_ptr.add(i + 8));
acc2 = fmlal_lo_inline(acc2, w1, x1);
acc3 = fmlal_hi_inline(acc3, w1, x1);
i += 16;
}
let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3)));
let chunks8 = (len - i) / 8;
let mut acc_lo = zero;
let mut acc_hi = zero;
for _ in 0..chunks8 {
let w0: uint16x8_t = vld1q_u16(w_ptr.add(i));
let x0: uint16x8_t = vld1q_u16(x_ptr.add(i));
acc_lo = fmlal_lo_inline(acc_lo, w0, x0);
acc_hi = fmlal_hi_inline(acc_hi, w0, x0);
i += 8;
}
sum += vaddvq_f32(vaddq_f32(acc_lo, acc_hi));
while i < len {
sum += (*w.get_unchecked(i)).to_f32() * (*x.get_unchecked(i)).to_f32();
i += 1;
}
sum
}
}
#[inline(always)]
pub fn simd_matmul_f16_f16_rows(
output: &mut [f32],
weight_f16: &[half::f16],
input_f16: &[half::f16],
rows: usize,
cols: usize,
) {
for r in 0..rows {
let row_off = r * cols;
unsafe {
*output.get_unchecked_mut(r) =
simd_dot_f16_f16(&weight_f16[row_off..row_off + cols], input_f16, cols);
}
}
}
#[inline]
pub fn simd_matmul_f16_f16_rows_parallel(
output: &mut [f32],
weight_f16: &[half::f16],
input_f16: &[half::f16],
rows: usize,
cols: usize,
) {
const PARALLEL_ROWS_MIN: usize = 512;
if rows < PARALLEL_ROWS_MIN {
simd_matmul_f16_f16_rows(output, weight_f16, input_f16, rows, cols);
} else {
use rayon::prelude::*;
const PARALLEL_CHUNK_ROWS: usize = 256;
output
.par_chunks_mut(PARALLEL_CHUNK_ROWS)
.enumerate()
.for_each(|(chunk_idx, out_chunk)| {
let start_row = chunk_idx * PARALLEL_CHUNK_ROWS;
for (local_r, out) in out_chunk.iter_mut().enumerate() {
let r = start_row + local_r;
let row_off = r * cols;
*out = simd_dot_f16_f16(&weight_f16[row_off..row_off + cols], input_f16, cols);
}
});
}
}