himada-dispatch 0.1.0

Adaptive SIMD dispatch for Himada — auto-selects fastest kernel at runtime
use himada_core::HardwareDNA;

// ---------------------------------------------------------------------------
// gemv_f64: y = alpha * A * x + beta * y
// ---------------------------------------------------------------------------

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::*;

    // SAFETY:
    // - `is_x86_feature_detected!("sse2")` confirms CPU support
    // - All pointer arithmetic is bounded by `rows`, `cols`; `i * cols + j` < a.len(),
    //   `j < x.len()`, `i < y.len()` (caller contract ensures sufficient slice lengths)
    // - `_mm_loadu_pd` / `_mm_storeu_pd` tolerate unaligned pointers
    // - `_mm_setzero_pd`, `_mm_add_pd`, `_mm_mul_pd` are pure compute intrinsics
    // - `std::mem::transmute` from __m128d to [f64; 2] is safe; both are 16 bytes
    // - Scalar fallback for remainder (j < cols)
    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::*;

    // SAFETY:
    // - `is_x86_feature_detected!("avx2")` confirms CPU support
    // - All pointer arithmetic is bounded by `rows`, `cols`
    // - `_mm256_loadu_pd` / `_mm256_storeu_pd` tolerate unaligned pointers
    // - `_mm256_setzero_pd`, `_mm256_add_pd`, `_mm256_mul_pd` are pure compute intrinsics
    // - `std::mem::transmute` from __m256d to [f64; 4] is safe; both are 32 bytes
    // - Scalar fallback for remainder (j < cols)
    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::*;

    // SAFETY:
    // - NEON is available on all aarch64 targets
    // - All pointer arithmetic is bounded by `rows`, `cols`
    // - `vld1q_f64` tolerates unaligned pointers on aarch64
    // - `vdupq_n_f64`, `vaddq_f64`, `vmulq_f64` are pure compute intrinsics
    // - `std::mem::transmute` from float64x2_t to [f64; 2] is safe; both are 16 bytes
    // - Scalar fallback for remainder (j < cols)
    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 }

// ---------------------------------------------------------------------------
// gemv_f32: y = alpha * A * x + beta * y
// ---------------------------------------------------------------------------

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::*;

    // SAFETY:
    // - `is_x86_feature_detected!("sse2")` confirms CPU support
    // - All pointer arithmetic is bounded by `rows`, `cols`
    // - `_mm_loadu_ps` / `_mm_storeu_ps` tolerate unaligned pointers
    // - `_mm_setzero_ps`, `_mm_add_ps`, `_mm_mul_ps` are pure compute intrinsics
    // - `std::mem::transmute` from __m128 to [f32; 4] is safe; both are 16 bytes
    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::*;

    // SAFETY:
    // - `is_x86_feature_detected!("avx2")` confirms CPU support
    // - All pointer arithmetic is bounded by `rows`, `cols`
    // - `_mm256_loadu_ps` / `_mm256_storeu_ps` tolerate unaligned pointers
    // - `_mm256_setzero_ps`, `_mm256_add_ps`, `_mm256_mul_ps` are pure compute intrinsics
    // - `std::mem::transmute` from __m256 to [f32; 8] is safe; both are 32 bytes
    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::*;

    // SAFETY:
    // - NEON is available on all aarch64 targets
    // - All pointer arithmetic is bounded by `rows`, `cols`
    // - `vld1q_f32` tolerates unaligned pointers on aarch64
    // - `vdupq_n_f32`, `vaddq_f32`, `vmulq_f32` are pure compute intrinsics
    // - `std::mem::transmute` from float32x4_t to [f32; 4] is safe; both are 16 bytes
    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 }