use himada_dispatch::*;
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);
}
pub fn register(&self) -> Result<(), HwdnaError> {
Ok(())
}
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 }
pub fn dot(&mut self, a: &[f64], b: &[f64]) -> f64 {
self.dot_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)
}
}
#[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);
}
}