himada-integrations 0.1.0

Himada framework integrations — Candle, Burn, ONNX Runtime, ndarray
//! ONNX Runtime integration — Himada as a custom execution provider.
//!
//! ONNX Runtime allows custom execution providers via its C API.
//! This module wraps Himada's dispatch behind the ORT EP interface.

use himada_dispatch::*;

/// Himada execution provider for ONNX Runtime.
pub struct HimadaOrtProvider {
    dot_dispatch: Dispatch<DotKernel>,
    matmul_dispatch: Dispatch<MatMulKernel>,
    name: String,
}

impl HimadaOrtProvider {
    pub fn new(name: &str) -> Self {
        Self {
            dot_dispatch: Dispatch::new("ort_dot", vec![]),
            matmul_dispatch: Dispatch::new("ort_matmul", vec![]),
            name: name.into(),
        }
    }

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

    /// Register with ONNX Runtime (stub — requires onnxruntime-sys).
    pub fn register(&self) -> Result<(), HwdnaError> {
        Ok(())
    }

    /// 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.matmul_dispatch.select().is_ok() { count += 1; }
        Ok(count)
    }

    pub fn name(&self) -> &str { &self.name }

    /// Execute a dot product using Himada's selected kernel.
    pub fn dot(&mut self, a: &[f64], b: &[f64]) -> f64 {
        self.dot_dispatch.compute(a, b)
    }

    /// Execute a matmul using Himada's selected kernel.
    pub fn matmul(&mut self, a: &[f64], b: &[f64], c: &mut [f64], n: usize) {
        self.matmul_dispatch.compute(a, b, c, n)
    }
}

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

    #[test]
    fn test_ort_provider_dot() {
        let mut provider = HimadaOrtProvider::new("himada_test");
        provider.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,
            }],
        );
        provider.optimize().unwrap();

        let a: Vec<f64> = (0..100).map(|i| i as f64).collect();
        let b: Vec<f64> = (0..100).map(|i| (i * 3) as f64).collect();
        let result = provider.dot(&a, &b);
        let expected: f64 = a.iter().zip(&b).map(|(x, y)| x * y).sum();
        assert!((result - expected).abs() < 1e-10);
    }
}