use himada_core::HardwareDNA;
pub fn gemv_f64_scalar(alpha: f64, a: &[f64], x: &[f64], beta: f64, y: &mut [f64], rows: usize, cols: usize) {
for i in 0..rows {
let mut sum = 0.0;
for j in 0..cols {
sum += a[i * cols + j] * x[j];
}
y[i] = alpha * sum + beta * y[i];
}
}
pub fn gemv_f64_supported(_: &HardwareDNA) -> bool { true }
#[cfg(target_arch = "x86_64")]
pub fn gemv_f64_sse(alpha: f64, a: &[f64], x: &[f64], beta: f64, y: &mut [f64], rows: usize, cols: usize) {
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
unsafe {
if is_x86_feature_detected!("sse2") {
for i in 0..rows {
let mut j = 0;
let mut vacc = _mm_setzero_pd();
while j + 2 <= cols {
let va = _mm_loadu_pd(a.as_ptr().add(i * cols + j));
let vx = _mm_loadu_pd(x.as_ptr().add(j));
vacc = _mm_add_pd(vacc, _mm_mul_pd(va, vx));
j += 2;
}
let tmp: [f64; 2] = std::mem::transmute::<_, [f64; 2]>(vacc);
let mut sum = tmp[0] + tmp[1];
for jj in j..cols {
sum += a[i * cols + jj] * x[jj];
}
y[i] = alpha * sum + beta * y[i];
}
}
}
}
#[cfg(target_arch = "x86_64")]
pub fn gemv_f64_sse_supported(dna: &HardwareDNA) -> bool {
dna.cpu.features.iter().any(|f| f == "SSE2")
}
#[cfg(not(target_arch = "x86_64"))]
pub fn gemv_f64_sse(_: f64, _: &[f64], _: &[f64], _: f64, _: &mut [f64], _: usize, _: usize) {}
#[cfg(not(target_arch = "x86_64"))]
pub fn gemv_f64_sse_supported(_: &HardwareDNA) -> bool { false }
#[cfg(target_arch = "x86_64")]
pub fn gemv_f64_avx2(alpha: f64, a: &[f64], x: &[f64], beta: f64, y: &mut [f64], rows: usize, cols: usize) {
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
unsafe {
if is_x86_feature_detected!("avx2") {
for i in 0..rows {
let mut j = 0;
let mut vacc = _mm256_setzero_pd();
while j + 4 <= cols {
let va = _mm256_loadu_pd(a.as_ptr().add(i * cols + j));
let vx = _mm256_loadu_pd(x.as_ptr().add(j));
vacc = _mm256_add_pd(vacc, _mm256_mul_pd(va, vx));
j += 4;
}
let tmp: [f64; 4] = std::mem::transmute::<_, [f64; 4]>(vacc);
let mut sum = tmp[0] + tmp[1] + tmp[2] + tmp[3];
for jj in j..cols {
sum += a[i * cols + jj] * x[jj];
}
y[i] = alpha * sum + beta * y[i];
}
}
}
}
#[cfg(target_arch = "x86_64")]
pub fn gemv_f64_avx2_supported(dna: &HardwareDNA) -> bool {
dna.cpu.features.iter().any(|f| f == "AVX2")
}
#[cfg(not(target_arch = "x86_64"))]
pub fn gemv_f64_avx2(_: f64, _: &[f64], _: &[f64], _: f64, _: &mut [f64], _: usize, _: usize) {}
#[cfg(not(target_arch = "x86_64"))]
pub fn gemv_f64_avx2_supported(_: &HardwareDNA) -> bool { false }
#[cfg(target_arch = "aarch64")]
pub fn gemv_f64_neon(alpha: f64, a: &[f64], x: &[f64], beta: f64, y: &mut [f64], rows: usize, cols: usize) {
#[cfg(target_arch = "aarch64")]
use std::arch::aarch64::*;
unsafe {
for i in 0..rows {
let mut j = 0;
let mut vacc = vdupq_n_f64(0.0);
while j + 2 <= cols {
let va = vld1q_f64(a.as_ptr().add(i * cols + j));
let vx = vld1q_f64(x.as_ptr().add(j));
vacc = vaddq_f64(vacc, vmulq_f64(va, vx));
j += 2;
}
let tmp: [f64; 2] = std::mem::transmute::<_, [f64; 2]>(vacc);
let mut sum = tmp[0] + tmp[1];
for jj in j..cols {
sum += a[i * cols + jj] * x[jj];
}
y[i] = alpha * sum + beta * y[i];
}
}
}
#[cfg(target_arch = "aarch64")]
pub fn gemv_f64_neon_supported(_: &HardwareDNA) -> bool { true }
#[cfg(not(target_arch = "aarch64"))]
pub fn gemv_f64_neon(_: f64, _: &[f64], _: &[f64], _: f64, _: &mut [f64], _: usize, _: usize) {}
#[cfg(not(target_arch = "aarch64"))]
pub fn gemv_f64_neon_supported(_: &HardwareDNA) -> bool { false }
pub fn gemv_f32_scalar(alpha: f32, a: &[f32], x: &[f32], beta: f32, y: &mut [f32], rows: usize, cols: usize) {
for i in 0..rows {
let mut sum = 0.0;
for j in 0..cols {
sum += a[i * cols + j] * x[j];
}
y[i] = alpha * sum + beta * y[i];
}
}
pub fn gemv_f32_supported(_: &HardwareDNA) -> bool { true }
#[cfg(target_arch = "x86_64")]
pub fn gemv_f32_sse(alpha: f32, a: &[f32], x: &[f32], beta: f32, y: &mut [f32], rows: usize, cols: usize) {
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
unsafe {
if is_x86_feature_detected!("sse2") {
for i in 0..rows {
let mut j = 0;
let mut vacc = _mm_setzero_ps();
while j + 4 <= cols {
let va = _mm_loadu_ps(a.as_ptr().add(i * cols + j));
let vx = _mm_loadu_ps(x.as_ptr().add(j));
vacc = _mm_add_ps(vacc, _mm_mul_ps(va, vx));
j += 4;
}
let tmp: [f32; 4] = std::mem::transmute::<_, [f32; 4]>(vacc);
let mut sum = tmp[0] + tmp[1] + tmp[2] + tmp[3];
for jj in j..cols {
sum += a[i * cols + jj] * x[jj];
}
y[i] = alpha * sum + beta * y[i];
}
}
}
}
#[cfg(target_arch = "x86_64")]
pub fn gemv_f32_sse_supported(dna: &HardwareDNA) -> bool {
dna.cpu.features.iter().any(|f| f == "SSE2")
}
#[cfg(not(target_arch = "x86_64"))]
pub fn gemv_f32_sse(_: f32, _: &[f32], _: &[f32], _: f32, _: &mut [f32], _: usize, _: usize) {}
#[cfg(not(target_arch = "x86_64"))]
pub fn gemv_f32_sse_supported(_: &HardwareDNA) -> bool { false }
#[cfg(target_arch = "x86_64")]
pub fn gemv_f32_avx2(alpha: f32, a: &[f32], x: &[f32], beta: f32, y: &mut [f32], rows: usize, cols: usize) {
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
unsafe {
if is_x86_feature_detected!("avx2") {
for i in 0..rows {
let mut j = 0;
let mut vacc = _mm256_setzero_ps();
while j + 8 <= cols {
let va = _mm256_loadu_ps(a.as_ptr().add(i * cols + j));
let vx = _mm256_loadu_ps(x.as_ptr().add(j));
vacc = _mm256_add_ps(vacc, _mm256_mul_ps(va, vx));
j += 8;
}
let tmp: [f32; 8] = std::mem::transmute::<_, [f32; 8]>(vacc);
let mut sum = tmp.iter().sum::<f32>();
for jj in j..cols {
sum += a[i * cols + jj] * x[jj];
}
y[i] = alpha * sum + beta * y[i];
}
}
}
}
#[cfg(target_arch = "x86_64")]
pub fn gemv_f32_avx2_supported(dna: &HardwareDNA) -> bool {
dna.cpu.features.iter().any(|f| f == "AVX2")
}
#[cfg(not(target_arch = "x86_64"))]
pub fn gemv_f32_avx2(_: f32, _: &[f32], _: &[f32], _: f32, _: &mut [f32], _: usize, _: usize) {}
#[cfg(not(target_arch = "x86_64"))]
pub fn gemv_f32_avx2_supported(_: &HardwareDNA) -> bool { false }
#[cfg(target_arch = "aarch64")]
pub fn gemv_f32_neon(alpha: f32, a: &[f32], x: &[f32], beta: f32, y: &mut [f32], rows: usize, cols: usize) {
#[cfg(target_arch = "aarch64")]
use std::arch::aarch64::*;
unsafe {
for i in 0..rows {
let mut j = 0;
let mut vacc = vdupq_n_f32(0.0);
while j + 4 <= cols {
let va = vld1q_f32(a.as_ptr().add(i * cols + j));
let vx = vld1q_f32(x.as_ptr().add(j));
vacc = vaddq_f32(vacc, vmulq_f32(va, vx));
j += 4;
}
let tmp: [f32; 4] = std::mem::transmute::<_, [f32; 4]>(vacc);
let mut sum = tmp[0] + tmp[1] + tmp[2] + tmp[3];
for jj in j..cols {
sum += a[i * cols + jj] * x[jj];
}
y[i] = alpha * sum + beta * y[i];
}
}
}
#[cfg(target_arch = "aarch64")]
pub fn gemv_f32_neon_supported(_: &HardwareDNA) -> bool { true }
#[cfg(not(target_arch = "aarch64"))]
pub fn gemv_f32_neon(_: f32, _: &[f32], _: &[f32], _: f32, _: &mut [f32], _: usize, _: usize) {}
#[cfg(not(target_arch = "aarch64"))]
pub fn gemv_f32_neon_supported(_: &HardwareDNA) -> bool { false }