himada-integrations 0.1.0

Himada framework integrations — Candle, Burn, ONNX Runtime, ndarray
//! Candle integration — use Himada SIMD dispatch as a Candle backend.
//!
//! Candle is a minimalist ML framework by Hugging Face.
//! This module provides dot/matmul operations backed by Himada's
//! runtime-optimized SIMD dispatch.

use himada_dispatch::*;

/// Himada-backed compute engine for Candle tensors.
pub struct HimadaCandleBackend {
    dot_dispatch: Dispatch<DotKernel>,
    dot_f32_dispatch: Dispatch<DotF32Kernel>,
    matmul_dispatch: Dispatch<MatMulKernel>,
    matmul_f32_dispatch: Dispatch<MatMulF32Kernel>,
}

impl HimadaCandleBackend {
    pub fn new() -> Self {
        let dot_dispatch = Dispatch::new("candle_dot", vec![]);
        let dot_f32_dispatch = Dispatch::new("candle_dot_f32", vec![]);
        let matmul_dispatch = Dispatch::new("candle_matmul", vec![]);
        let matmul_f32_dispatch = Dispatch::new("candle_matmul_f32", vec![]);
        Self { dot_dispatch, dot_f32_dispatch, matmul_dispatch, matmul_f32_dispatch }
    }

    pub fn register_kernels(&mut self, dots: Vec<KernelInfo<DotKernel>>, matmuls: Vec<KernelInfo<MatMulKernel>>) {
        self.dot_dispatch = Dispatch::new("candle_dot", dots);
        self.matmul_dispatch = Dispatch::new("candle_matmul", matmuls);
    }

    pub fn register_kernels_f32(&mut self, dots: Vec<KernelInfo<DotF32Kernel>>, matmuls: Vec<KernelInfo<MatMulF32Kernel>>) {
        self.dot_f32_dispatch = Dispatch::new("candle_dot_f32", dots);
        self.matmul_f32_dispatch = Dispatch::new("candle_matmul_f32", matmuls);
    }

    pub fn dot(&mut self, a: &[f64], b: &[f64]) -> f64 {
        self.dot_dispatch.compute(a, b)
    }

    pub fn dot_f32(&mut self, a: &[f32], b: &[f32]) -> f32 {
        self.dot_f32_dispatch.compute(a, b)
    }

    pub fn matmul(&mut self, a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
        self.matmul_dispatch.compute(a, b, c, n)
    }

    pub fn matmul_f32(&mut self, a: &[f32], b: &[f32], c: &mut [f32], n: usize) {
        self.matmul_f32_dispatch.compute(a, b, c, n)
    }

    /// Select best kernels for all operations.
    pub fn optimize(&mut self) -> Result<usize, HwdnaError> {
        let mut count = 0;
        if self.dot_dispatch.select().is_ok() { count += 1; }
        if self.dot_f32_dispatch.select().is_ok() { count += 1; }
        if self.matmul_dispatch.select().is_ok() { count += 1; }
        if self.matmul_f32_dispatch.select().is_ok() { count += 1; }
        Ok(count)
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use himada_dispatch::kernels;
    use himada_core::HardwareDNA;

    #[test]
    fn test_candle_dot_consistency() {
        let mut backend = HimadaCandleBackend::new();
        backend.register_kernels(
            vec![KernelInfo {
                name: "scalar",
                func: kernels::dot_scalar as DotKernel,
                is_supported: |_: &HardwareDNA| true,
                thermal_priority: 0,
            }],
            vec![KernelInfo {
                name: "scalar",
                func: kernels::matmul_scalar as MatMulKernel,
                is_supported: |_: &HardwareDNA| true,
                thermal_priority: 0,
            }],
        );
        backend.optimize().unwrap();

        let a: Vec<f64> = (0..100).map(|i| i as f64).collect();
        let b: Vec<f64> = (0..100).map(|i| (i * 2) as f64).collect();

        let himada_result = backend.dot(&a, &b);
        let expected: f64 = a.iter().zip(&b).map(|(x, y)| x * y).sum();
        assert!((himada_result - expected).abs() < 1e-10,
            "candle dot mismatch: himada={} expected={}", himada_result, expected);

        // Test matmul
        let n = 8;
        let ma: Vec<f64> = (0..n*n).map(|i| i as f64).collect();
        let mb: Vec<f64> = (0..n*n).map(|i| (i * 3) as f64).collect();
        let mut mc = vec![0.0; n*n];
        backend.matmul(&ma, &mb, &mut mc, n);

        // Reference
        let mut expected_c = vec![0.0; n*n];
        for i in 0..n {
            for j in 0..n {
                let mut s = 0.0;
                for k in 0..n {
                    s += ma[i * n + k] * mb[k * n + j];
                }
                expected_c[i * n + j] = s;
            }
        }
        for i in 0..n*n {
            assert!((mc[i] - expected_c[i]).abs() < 1e-10,
                "candle matmul mismatch at {i}: himada={} expected={}", mc[i], expected_c[i]);
        }
    }

    #[test]
    fn test_candle_dot_f32() {
        let mut backend = HimadaCandleBackend::new();
        backend.register_kernels_f32(
            vec![KernelInfo {
                name: "scalar",
                func: kernels::dot_f32_scalar as DotF32Kernel,
                is_supported: |_: &HardwareDNA| true,
                thermal_priority: 0,
            }],
            vec![KernelInfo {
                name: "scalar",
                func: kernels::matmul_f32_scalar as MatMulF32Kernel,
                is_supported: |_: &HardwareDNA| true,
                thermal_priority: 0,
            }],
        );
        backend.optimize().unwrap();

        let a: Vec<f32> = (0..50).map(|i| i as f32).collect();
        let b: Vec<f32> = (0..50).map(|i| (i * 2) as f32).collect();
        let result = backend.dot_f32(&a, &b);
        let expected: f32 = a.iter().zip(&b).map(|(x, y)| x * y).sum();
        assert!((result - expected).abs() < 1e-6);
    }
}