himada-dispatch 0.1.1

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

// ---------------------------------------------------------------------------
// conv1d_f64
// ---------------------------------------------------------------------------

pub fn conv1d_f64_scalar(input: &[f64], kernel: &[f64], output: &mut [f64]) {
    let klen = kernel.len();
    let olen = input.len().saturating_sub(klen).saturating_add(1);
    let olen = olen.min(output.len());
    for i in 0..olen {
        let mut sum = 0.0;
        for j in 0..klen {
            sum += input[i + j] * kernel[j];
        }
        output[i] = sum;
    }
}

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

#[cfg(target_arch = "x86_64")]
pub fn conv1d_f64_sse(input: &[f64], kernel: &[f64], output: &mut [f64]) {
    #[cfg(target_arch = "x86_64")]
    use std::arch::x86_64::*;
    let klen = kernel.len();
    let olen = input.len().saturating_sub(klen).saturating_add(1);
    let olen = olen.min(output.len());

    // SAFETY:
    // - `is_x86_feature_detected!("sse2")` confirms CPU support
    // - `i < olen`, `j + 2 <= klen` bounds ensure `i + j + 2 <= input.len()`, `j + 2 <= kernel.len()`,
    //   and `i < output.len()` — all pointer arithmetic is in-bounds
    // - `_mm_loadu_pd` tolerates 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 kernel elements (jj < klen)
    unsafe {
        if is_x86_feature_detected!("sse2") {
            for i in 0..olen {
                let mut j = 0;
                let mut vacc = _mm_setzero_pd();
                while j + 2 <= klen {
                    let vi = _mm_loadu_pd(input.as_ptr().add(i + j));
                    let vk = _mm_loadu_pd(kernel.as_ptr().add(j));
                    vacc = _mm_add_pd(vacc, _mm_mul_pd(vi, vk));
                    j += 2;
                }
                let tmp: [f64; 2] = std::mem::transmute::<_, [f64; 2]>(vacc);
                let mut sum = tmp[0] + tmp[1];
                for jj in j..klen {
                    sum += input[i + jj] * kernel[jj];
                }
                output[i] = sum;
            }
        }
    }
}

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

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

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

#[cfg(target_arch = "x86_64")]
pub fn conv1d_f64_avx2(input: &[f64], kernel: &[f64], output: &mut [f64]) {
    #[cfg(target_arch = "x86_64")]
    use std::arch::x86_64::*;
    let klen = kernel.len();
    let olen = input.len().saturating_sub(klen).saturating_add(1);
    let olen = olen.min(output.len());

    // SAFETY:
    // - `is_x86_feature_detected!("avx2")` confirms CPU support
    // - `i < olen`, `j + 4 <= klen` bounds ensure in-bounds pointer arithmetic
    // - `_mm256_loadu_pd` tolerates 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
    unsafe {
        if is_x86_feature_detected!("avx2") {
            for i in 0..olen {
                let mut j = 0;
                let mut vacc = _mm256_setzero_pd();
                while j + 4 <= klen {
                    let vi = _mm256_loadu_pd(input.as_ptr().add(i + j));
                    let vk = _mm256_loadu_pd(kernel.as_ptr().add(j));
                    vacc = _mm256_add_pd(vacc, _mm256_mul_pd(vi, vk));
                    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..klen {
                    sum += input[i + jj] * kernel[jj];
                }
                output[i] = sum;
            }
        }
    }
}

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

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

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

#[cfg(target_arch = "aarch64")]
pub fn conv1d_f64_neon(input: &[f64], kernel: &[f64], output: &mut [f64]) {
    #[cfg(target_arch = "aarch64")]
    use std::arch::aarch64::*;
    let klen = kernel.len();
    let olen = input.len().saturating_sub(klen).saturating_add(1);
    let olen = olen.min(output.len());

    // SAFETY:
    // - NEON is available on all aarch64 targets
    // - `i < olen`, `j + 2 <= klen` bounds ensure in-bounds pointer arithmetic
    // - `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
    unsafe {
        for i in 0..olen {
            let mut j = 0;
            let mut vacc = vdupq_n_f64(0.0);
            while j + 2 <= klen {
                let vi = vld1q_f64(input.as_ptr().add(i + j));
                let vk = vld1q_f64(kernel.as_ptr().add(j));
                vacc = vaddq_f64(vacc, vmulq_f64(vi, vk));
                j += 2;
            }
            let tmp: [f64; 2] = std::mem::transmute::<_, [f64; 2]>(vacc);
            let mut sum = tmp[0] + tmp[1];
            for jj in j..klen {
                sum += input[i + jj] * kernel[jj];
            }
            output[i] = sum;
        }
    }
}

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

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

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

// ---------------------------------------------------------------------------
// conv1d_f32
// ---------------------------------------------------------------------------

pub fn conv1d_f32_scalar(input: &[f32], kernel: &[f32], output: &mut [f32]) {
    let klen = kernel.len();
    let olen = input.len().saturating_sub(klen).saturating_add(1);
    let olen = olen.min(output.len());
    for i in 0..olen {
        let mut sum = 0.0;
        for j in 0..klen {
            sum += input[i + j] * kernel[j];
        }
        output[i] = sum;
    }
}

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

#[cfg(target_arch = "x86_64")]
pub fn conv1d_f32_sse(input: &[f32], kernel: &[f32], output: &mut [f32]) {
    #[cfg(target_arch = "x86_64")]
    use std::arch::x86_64::*;
    let klen = kernel.len();
    let olen = input.len().saturating_sub(klen).saturating_add(1);
    let olen = olen.min(output.len());

    // SAFETY:
    // - `is_x86_feature_detected!("sse2")` confirms CPU support
    // - `i < olen`, `j + 4 <= klen` bounds ensure in-bounds pointer arithmetic
    // - `_mm_loadu_ps` tolerates 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..olen {
                let mut j = 0;
                let mut vacc = _mm_setzero_ps();
                while j + 4 <= klen {
                    let vi = _mm_loadu_ps(input.as_ptr().add(i + j));
                    let vk = _mm_loadu_ps(kernel.as_ptr().add(j));
                    vacc = _mm_add_ps(vacc, _mm_mul_ps(vi, vk));
                    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..klen {
                    sum += input[i + jj] * kernel[jj];
                }
                output[i] = sum;
            }
        }
    }
}

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

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

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

#[cfg(target_arch = "x86_64")]
pub fn conv1d_f32_avx2(input: &[f32], kernel: &[f32], output: &mut [f32]) {
    #[cfg(target_arch = "x86_64")]
    use std::arch::x86_64::*;
    let klen = kernel.len();
    let olen = input.len().saturating_sub(klen).saturating_add(1);
    let olen = olen.min(output.len());

    // SAFETY:
    // - `is_x86_feature_detected!("avx2")` confirms CPU support
    // - `i < olen`, `j + 8 <= klen` bounds ensure in-bounds pointer arithmetic
    // - `_mm256_loadu_ps` tolerates 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..olen {
                let mut j = 0;
                let mut vacc = _mm256_setzero_ps();
                while j + 8 <= klen {
                    let vi = _mm256_loadu_ps(input.as_ptr().add(i + j));
                    let vk = _mm256_loadu_ps(kernel.as_ptr().add(j));
                    vacc = _mm256_add_ps(vacc, _mm256_mul_ps(vi, vk));
                    j += 8;
                }
                let tmp: [f32; 8] = std::mem::transmute::<_, [f32; 8]>(vacc);
                let mut sum = tmp.iter().sum::<f32>();
                for jj in j..klen {
                    sum += input[i + jj] * kernel[jj];
                }
                output[i] = sum;
            }
        }
    }
}

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

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

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

#[cfg(target_arch = "aarch64")]
pub fn conv1d_f32_neon(input: &[f32], kernel: &[f32], output: &mut [f32]) {
    #[cfg(target_arch = "aarch64")]
    use std::arch::aarch64::*;
    let klen = kernel.len();
    let olen = input.len().saturating_sub(klen).saturating_add(1);
    let olen = olen.min(output.len());

    // SAFETY:
    // - NEON is available on all aarch64 targets
    // - `i < olen`, `j + 4 <= klen` bounds ensure in-bounds pointer arithmetic
    // - `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..olen {
            let mut j = 0;
            let mut vacc = vdupq_n_f32(0.0);
            while j + 4 <= klen {
                let vi = vld1q_f32(input.as_ptr().add(i + j));
                let vk = vld1q_f32(kernel.as_ptr().add(j));
                vacc = vaddq_f32(vacc, vmulq_f32(vi, vk));
                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..klen {
                sum += input[i + jj] * kernel[jj];
            }
            output[i] = sum;
        }
    }
}

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

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

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