use himada_dispatch::*;
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)
}
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);
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);
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);
}
}