himada-dispatch 0.1.0

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

/// Scalar matrix multiply: C += A * B  (all flat N×N)
pub fn matmul_scalar(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
    for i in 0..n {
        for k in 0..n {
            let aik = a[i * n + k];
            let row_b = k * n;
            let row_c = i * n;
            for j in 0..n {
                c[row_c + j] += aik * b[row_b + j];
            }
        }
    }
}

pub fn matmul_scalar_supported(_: &HardwareDNA) -> bool {
    true
}

/// Tiled (blocked) matrix multiply — cache-friendly.
pub fn matmul_tiled(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
    const T: usize = 64;
    for i in (0..n).step_by(T) {
        let imax = (i + T).min(n);
        for k in (0..n).step_by(T) {
            let kmax = (k + T).min(n);
            for j in (0..n).step_by(T) {
                let jmax = (j + T).min(n);
                for ii in i..imax {
                    for kk in k..kmax {
                        let aik = a[ii * n + kk];
                        let row_b = kk * n;
                        let row_c = ii * n;
                        for jj in j..jmax {
                            c[row_c + jj] += aik * b[row_b + jj];
                        }
                    }
                }
            }
        }
    }
}

pub fn matmul_tiled_supported(_: &HardwareDNA) -> bool {
    true
}

// ---------------------------------------------------------------------------
// x86 SSE – 2 doubles at a time
// ---------------------------------------------------------------------------

#[cfg(target_arch = "x86_64")]
pub fn matmul_sse(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
    #[cfg(target_arch = "x86_64")]
    use std::arch::x86_64::*;

    // SAFETY:
    // - All pointer arithmetic uses in-bounds indices: `i < n`, `k < n`,
    //   `row_b + j < b.len()`, `row_c + j < c.len()` because `a.len() >= n*n`,
    //   `b.len() >= n*n`, `c.len() >= n*n` (caller contract)
    // - `is_x86_feature_detected!("sse2")` confirms CPU support
    // - `_mm_loadu_pd` / `_mm_storeu_pd` tolerate unaligned pointers
    unsafe {
        for i in 0..n {
            for k in 0..n {
                let aik = a[i * n + k];
                let row_b = k * n;
                let row_c = i * n;
                let mut j = 0;
                if is_x86_feature_detected!("sse2") {
                    let vaik = _mm_set1_pd(aik);
                    while j + 2 <= n {
                        let vb = _mm_loadu_pd(b.as_ptr().add(row_b + j));
                        let vc = _mm_loadu_pd(c.as_ptr().add(row_c + j));
                        let vm = _mm_mul_pd(vaik, vb);
                        _mm_storeu_pd(c.as_mut_ptr().add(row_c + j), _mm_add_pd(vc, vm));
                        j += 2;
                    }
                }
                for jj in j..n {
                    c[row_c + jj] += aik * b[row_b + jj];
                }
            }
        }
    }
}

#[cfg(target_arch = "x86_64")]
pub fn matmul_sse_supported(dna: &HardwareDNA) -> bool {
    dna.cpu.features.iter().any(|f| f == "SSE2")
}

#[cfg(not(target_arch = "x86_64"))]
pub fn matmul_sse(_: &[f64], _: &[f64], _: &mut [f64], _: usize) {}

#[cfg(not(target_arch = "x86_64"))]
pub fn matmul_sse_supported(_: &HardwareDNA) -> bool { false }

// ---------------------------------------------------------------------------
// x86 AVX2 – 4 doubles at a time
// ---------------------------------------------------------------------------

#[cfg(target_arch = "x86_64")]
pub fn matmul_avx2(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
    #[cfg(target_arch = "x86_64")]
    use std::arch::x86_64::*;

    // SAFETY:
    // - All pointer arithmetic uses in-bounds indices: `i < n`, `k < n`,
    //   `row_b + j < b.len()`, `row_c + j < c.len()` because `a.len() >= n*n`,
    //   `b.len() >= n*n`, `c.len() >= n*n` (caller contract)
    // - `is_x86_feature_detected!("avx2")` confirms CPU support
    // - `_mm256_loadu_pd` / `_mm256_storeu_pd` tolerate unaligned pointers
    unsafe {
        for i in 0..n {
            for k in 0..n {
                let aik = a[i * n + k];
                let row_b = k * n;
                let row_c = i * n;
                let mut j = 0;
                if is_x86_feature_detected!("avx2") {
                    let vaik = _mm256_set1_pd(aik);
                    while j + 4 <= n {
                        let vb = _mm256_loadu_pd(b.as_ptr().add(row_b + j));
                        let vc = _mm256_loadu_pd(c.as_ptr().add(row_c + j));
                        let vm = _mm256_mul_pd(vaik, vb);
                        _mm256_storeu_pd(c.as_mut_ptr().add(row_c + j), _mm256_add_pd(vc, vm));
                        j += 4;
                    }
                }
                for jj in j..n {
                    c[row_c + jj] += aik * b[row_b + jj];
                }
            }
        }
    }
}

#[cfg(target_arch = "x86_64")]
pub fn matmul_avx2_supported(dna: &HardwareDNA) -> bool {
    dna.cpu.features.iter().any(|f| f == "AVX2")
}

#[cfg(not(target_arch = "x86_64"))]
pub fn matmul_avx2(_: &[f64], _: &[f64], _: &mut [f64], _: usize) {}

#[cfg(not(target_arch = "x86_64"))]
pub fn matmul_avx2_supported(_: &HardwareDNA) -> bool { false }

// ---------------------------------------------------------------------------
// ARM NEON – 2 doubles at a time
// ---------------------------------------------------------------------------

#[cfg(target_arch = "aarch64")]
pub fn matmul_neon(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
    #[cfg(target_arch = "aarch64")]
    use std::arch::aarch64::*;

    // SAFETY:
    // - All pointer arithmetic uses in-bounds indices: `i < n`, `k < n`,
    //   `row_b + j < b.len()`, `row_c + j < c.len()` because `a.len() >= n*n`,
    //   `b.len() >= n*n`, `c.len() >= n*n` (caller contract)
    // - NEON is available on all aarch64 targets
    // - `vld1q_f64` / `vst1q_f64` tolerate unaligned pointers on aarch64
    unsafe {
        for i in 0..n {
            for k in 0..n {
                let aik = a[i * n + k];
                let row_b = k * n;
                let row_c = i * n;
                let mut j = 0;
                if n >= 2 {
                    let vaik = vdupq_n_f64(aik);
                    while j + 2 <= n {
                        let vb = vld1q_f64(b.as_ptr().add(row_b + j));
                        let vc = vld1q_f64(c.as_ptr().add(row_c + j));
                        let vm = vmulq_f64(vaik, vb);
                        vst1q_f64(c.as_mut_ptr().add(row_c + j), vaddq_f64(vc, vm));
                        j += 2;
                    }
                }
                for jj in j..n {
                    c[row_c + jj] += aik * b[row_b + jj];
                }
            }
        }
    }
}

#[cfg(target_arch = "aarch64")]
pub fn matmul_neon_supported(_: &HardwareDNA) -> bool {
    true
}

#[cfg(not(target_arch = "aarch64"))]
pub fn matmul_neon(_: &[f64], _: &[f64], _: &mut [f64], _: usize) {}

#[cfg(not(target_arch = "aarch64"))]
pub fn matmul_neon_supported(_: &HardwareDNA) -> bool { false }

// ---------------------------------------------------------------------------
// Cache-aware auto-tiled — tile size from himada-core L1
// ---------------------------------------------------------------------------

use himada_core::profile;
use std::sync::OnceLock;

fn optimal_tile_size_n() -> usize {
    static TILE: OnceLock<usize> = OnceLock::new();
    *TILE.get_or_init(|| {
        let dna = profile::load_or_collect();
        let l1 = dna.caches.l1_data_size.max(16384);
        let elems = (l1 / (3 * 8)) as usize;
        let mut t = 1;
        while (t << 1) <= elems {
            t <<= 1;
        }
        t.clamp(8, 256)
    })
}

/// Matrix multiply with auto-sized tiles based on L1 cache size.
pub fn matmul_cache_tiled(a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
    let t = optimal_tile_size_n();
    for i in (0..n).step_by(t) {
        let imax = (i + t).min(n);
        for k in (0..n).step_by(t) {
            let kmax = (k + t).min(n);
            for j in (0..n).step_by(t) {
                let jmax = (j + t).min(n);
                for ii in i..imax {
                    for kk in k..kmax {
                        let aik = a[ii * n + kk];
                        let row_b = kk * n;
                        let row_c = ii * n;
                        for jj in j..jmax {
                            c[row_c + jj] += aik * b[row_b + jj];
                        }
                    }
                }
            }
        }
    }
}

pub fn matmul_cache_tiled_supported(_: &HardwareDNA) -> bool { true }