use alloc::string::{String, ToString};
use alloc::sync::Arc;
use burn_backend::backend::BackendTypes;
use burn_ir::{BackendIr, CustomOpIr, HandleContainer};
use hashbrown::HashMap;
pub type CustomOpHandler<B> = Arc<
dyn Fn(
&mut HandleContainer<<B as BackendIr>::Handle>,
&CustomOpIr,
&<B as BackendTypes>::Device,
) + Send
+ Sync,
>;
#[derive(Clone)]
pub struct CustomOpRegistry<B: BackendIr> {
handlers: HashMap<String, CustomOpHandler<B>>,
}
impl<B: BackendIr> Default for CustomOpRegistry<B> {
fn default() -> Self {
Self {
handlers: HashMap::new(),
}
}
}
impl<B: BackendIr> CustomOpRegistry<B> {
pub fn new() -> Self {
Self::default()
}
pub fn register<F>(&mut self, id: &str, handler: F)
where
F: Fn(&mut HandleContainer<B::Handle>, &CustomOpIr, &B::Device) + Send + Sync + 'static,
{
self.handlers.insert(id.to_string(), Arc::new(handler));
}
pub(crate) fn get(&self, id: &str) -> Option<&CustomOpHandler<B>> {
self.handlers.get(id)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::TensorInterpreter;
use burn_backend::{DType, Scalar, Shape, TensorData, ops::FloatTensorOps};
use burn_flex::Flex;
use burn_ir::{OperationIr, ScalarIr, TensorId, TensorIr};
use core::slice;
use std::sync::Mutex;
#[test]
fn custom_op_is_dispatched_with_scalars() {
let seen_scalars = Arc::new(Mutex::new(None));
let seen_scalars_handler = seen_scalars.clone();
let mut registry = CustomOpRegistry::<Flex>::new();
registry.register("scale", move |handles, ir, _device| {
let input = handles.get_float_tensor::<Flex>(&ir.inputs[0]);
let factor: Scalar = ir.scalars[0].into();
let output = Flex::float_mul_scalar(input, factor);
handles.register_float_tensor::<Flex>(&ir.outputs[0].id, output);
*seen_scalars_handler.lock().unwrap() = Some(ir.scalars.clone());
});
let mut interp = TensorInterpreter::<Flex>::with_custom_ops(Default::default(), registry);
let input = interp.register_tensor_data_desc(TensorData::from([2.0f32, 4.0]));
let output = TensorIr::uninit(TensorId::new(1_000_000), Shape::from([2]), DType::F32);
let desc =
CustomOpIr::with_scalars("scale", &[input], &[output], vec![ScalarIr::Float(3.0)]);
interp.register_op(OperationIr::Custom(desc));
assert_eq!(
seen_scalars.lock().unwrap().clone(),
Some(vec![ScalarIr::Float(3.0)])
);
}
#[test]
fn custom_op_can_create_tensors_from_the_device() {
use burn_backend::ops::IntTensorOps;
let mut registry = CustomOpRegistry::<Flex>::new();
registry.register("load", |handles, ir, device| {
let values: Vec<i64> = ir.scalars.iter().map(|s| s.elem::<i64>()).collect();
let n = values.len();
let tensor = Flex::int_from_data(TensorData::new(values, [n]), device);
handles.register_int_tensor::<Flex>(&ir.outputs[0].id, tensor);
});
let mut interp = TensorInterpreter::<Flex>::with_custom_ops(Default::default(), registry);
let out = TensorIr::uninit(TensorId::new(2_000_000), Shape::from([3]), DType::I64);
let desc = CustomOpIr::with_scalars(
"load",
&[],
slice::from_ref(&out),
vec![ScalarIr::Int(10), ScalarIr::Int(20), ScalarIr::Int(30)],
);
interp.register_op(OperationIr::Custom(desc));
let data = interp.get_tensor(&TensorIr {
status: burn_ir::TensorStatus::ReadOnly,
..out
});
match data {
burn_ir::HandleKind::Int(_) => {}
_ => panic!("expected an int tensor"),
}
}
#[test]
#[should_panic(expected = "No custom-op handler registered")]
fn unregistered_custom_op_panics() {
let mut interp = TensorInterpreter::<Flex>::new(Default::default());
let desc = CustomOpIr::new("missing", &[], &[]);
interp.register_op(OperationIr::Custom(desc));
}
}