use crate::CubeDevice;
use burn_cubecl_fusion::optim::elemwise::{self, ElementWiseFuser, ElemwiseOptimization};
use burn_cubecl_fusion::optim::matmul::{self, MatmulFuser, MatmulOptimization};
use burn_cubecl_fusion::optim::nhwc_relayout::{self, NHWCRelayoutFuser, NHWCRelayoutOptimization};
use burn_cubecl_fusion::optim::reduce::{self, ReduceFuser, ReduceOptimization, ReduceSettings};
use burn_cubecl_fusion::optim::reduce_broadcasted::{
self, ReduceBroadcastedFuser, ReduceBroadcastedOptimization,
};
use burn_cubecl_fusion::optim::{CubeOptimization, CubeOptimizationState, FusedOperation};
use burn_fusion::OperationFuser;
use core::any::Any;
use std::sync::{Mutex, OnceLock};
pub struct CubeFuser {
fuser: Box<dyn OperationFuser<CubeOptimization>>,
}
impl CubeFuser {
pub fn new(fuser: impl OperationFuser<CubeOptimization> + 'static) -> Self {
Self {
fuser: Box::new(fuser),
}
}
}
pub trait OptimizationProvider: Send + Sync + 'static {
type Operation: FusedOperation;
fn name(&self) -> &str {
Self::Operation::NAME
}
fn fuser(&self, device: &CubeDevice) -> CubeFuser;
fn restore(&self, device: &CubeDevice, state: &CubeOptimizationState) -> CubeOptimization {
CubeOptimization::new(Self::Operation::from_state(device, state.decode()))
}
}
trait DynProvider: Send + Sync {
fn fuser(&self, device: &CubeDevice) -> CubeFuser;
fn restore(&self, device: &CubeDevice, state: &CubeOptimizationState) -> CubeOptimization;
}
impl<P: OptimizationProvider> DynProvider for P {
fn fuser(&self, device: &CubeDevice) -> CubeFuser {
OptimizationProvider::fuser(self, device)
}
fn restore(&self, device: &CubeDevice, state: &CubeOptimizationState) -> CubeOptimization {
OptimizationProvider::restore(self, device, state)
}
}
struct ElemwiseProvider;
struct MatmulProvider;
struct ReduceProvider;
struct ReduceBroadcastedProvider;
struct NHWCRelayoutProvider;
impl OptimizationProvider for ElemwiseProvider {
type Operation = ElemwiseOptimization;
fn fuser(&self, device: &CubeDevice) -> CubeFuser {
CubeFuser::new(ElementWiseFuser::new(device.clone()))
}
}
impl OptimizationProvider for MatmulProvider {
type Operation = MatmulOptimization;
fn fuser(&self, device: &CubeDevice) -> CubeFuser {
CubeFuser::new(MatmulFuser::new(device.clone()))
}
}
impl OptimizationProvider for ReduceProvider {
type Operation = ReduceOptimization;
fn fuser(&self, device: &CubeDevice) -> CubeFuser {
CubeFuser::new(ReduceFuser::new(device.clone(), ReduceSettings::Always))
}
}
impl OptimizationProvider for ReduceBroadcastedProvider {
type Operation = ReduceBroadcastedOptimization;
fn fuser(&self, device: &CubeDevice) -> CubeFuser {
CubeFuser::new(ReduceBroadcastedFuser::new(device.clone()))
}
}
impl OptimizationProvider for NHWCRelayoutProvider {
type Operation = NHWCRelayoutOptimization;
fn fuser(&self, device: &CubeDevice) -> CubeFuser {
CubeFuser::new(NHWCRelayoutFuser::new(device.clone()))
}
}
pub const BUILTIN_NAMES: [&str; 5] = [
elemwise::NAME,
matmul::NAME,
reduce::NAME,
reduce_broadcasted::NAME,
nhwc_relayout::NAME,
];
fn builtins() -> Vec<(String, Slot)> {
vec![
slot(ElemwiseProvider),
slot(MatmulProvider),
slot(ReduceProvider),
slot(ReduceBroadcastedProvider),
slot(NHWCRelayoutProvider),
]
}
fn slot(provider: impl OptimizationProvider) -> (String, Slot) {
let name = provider.name().to_string();
(name, Box::new(ProviderSlot(Box::new(provider))))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RegistryError {
ServiceRunning,
DuplicateOptimization {
name: String,
},
}
impl core::fmt::Display for RegistryError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::ServiceRunning => write!(
f,
"the fusion backend service is already running; \
register or remove fusion optimizations at the start of your program, \
before the first tensor operation on the fusion backend"
),
Self::DuplicateOptimization { name } => write!(
f,
"a fusion optimization named `{name}` is already registered"
),
}
}
}
impl std::error::Error for RegistryError {}
pub fn register(provider: impl OptimizationProvider) -> Result<(), RegistryError> {
let (name, slot) = slot(provider);
registry().lock().unwrap().register(name, slot)
}
pub fn remove(name: &str) -> Result<(), RegistryError> {
registry().lock().unwrap().remove(name)
}
pub(crate) fn restore(device: &CubeDevice, state: CubeOptimizationState) -> CubeOptimization {
let registry = registry().lock().unwrap();
registry
.provider(&state.name)
.map(|slot| downcast(slot).restore(device, &state))
.unwrap_or_else(|| {
panic!(
"no fusion optimization named `{}` is registered; register its \
provider before restoring serialized execution plans",
state.name,
)
})
}
pub(crate) fn fusers(device: &CubeDevice) -> Vec<Box<dyn OperationFuser<CubeOptimization>>> {
let mut registry = registry().lock().unwrap();
registry
.start()
.providers
.iter()
.map(|(_, slot)| downcast(slot).fuser(device).fuser)
.collect()
}
fn registry() -> &'static Mutex<Registry> {
static REGISTRY: OnceLock<Mutex<Registry>> = OnceLock::new();
REGISTRY.get_or_init(|| Mutex::new(Registry::seeded(builtins)))
}
fn downcast(slot: &Slot) -> &dyn DynProvider {
slot.downcast_ref::<ProviderSlot>()
.expect("every slot in the registry holds a provider")
.0
.as_ref()
}
struct ProviderSlot(Box<dyn DynProvider>);
type Slot = Box<dyn Any + Send + Sync>;
struct Registry {
providers: Vec<(String, Slot)>,
started: bool,
}
impl Registry {
fn seeded(defaults: impl FnOnce() -> Vec<(String, Slot)>) -> Self {
Self {
providers: defaults(),
started: false,
}
}
fn register(&mut self, name: String, slot: Slot) -> Result<(), RegistryError> {
self.ensure_open()?;
if self.providers.iter().any(|(other, _)| *other == name) {
return Err(RegistryError::DuplicateOptimization { name });
}
self.providers.push((name, slot));
Ok(())
}
fn remove(&mut self, name: &str) -> Result<(), RegistryError> {
self.ensure_open()?;
self.providers.retain(|(other, _)| other != name);
Ok(())
}
fn start(&mut self) -> &Self {
self.started = true;
self
}
fn provider(&self, name: &str) -> Option<&Slot> {
self.providers
.iter()
.find(|(other, _)| other == name)
.map(|(_, slot)| slot)
}
fn ensure_open(&self) -> Result<(), RegistryError> {
if self.started {
return Err(RegistryError::ServiceRunning);
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn slot() -> Slot {
Box::new(())
}
fn empty() -> Registry {
Registry::seeded(Vec::new)
}
fn seeded() -> Vec<(String, Slot)> {
vec![("builtin".into(), slot())]
}
#[test]
fn register_then_remove_round_trips() {
let mut registry = empty();
registry
.register("custom".into(), slot())
.expect("first registration succeeds");
registry.remove("custom").expect("removal succeeds");
registry
.register("custom".into(), slot())
.expect("re-registration after removal succeeds");
}
#[test]
fn duplicate_names_are_rejected() {
let mut registry = empty();
registry.register("custom".into(), slot()).unwrap();
assert_eq!(
registry.register("custom".into(), slot()),
Err(RegistryError::DuplicateOptimization {
name: "custom".into()
})
);
}
#[test]
fn a_started_registry_is_sealed() {
let mut registry = empty();
registry.start();
assert_eq!(
registry.register("custom".into(), slot()),
Err(RegistryError::ServiceRunning)
);
assert_eq!(
registry.remove("builtin"),
Err(RegistryError::ServiceRunning)
);
}
#[test]
fn provider_lookup_by_name() {
let mut registry = empty();
registry.register("custom".into(), slot()).unwrap();
assert!(registry.provider("custom").is_some());
assert!(registry.provider("unknown").is_none());
}
#[test]
fn defaults_seed_the_registry() {
let registry = Registry::seeded(seeded);
assert!(registry.provider("builtin").is_some());
}
#[test]
fn defaults_are_removable_and_reserve_their_name() {
let mut registry = Registry::seeded(seeded);
assert_eq!(
registry.register("builtin".into(), slot()),
Err(RegistryError::DuplicateOptimization {
name: "builtin".into()
})
);
registry.remove("builtin").unwrap();
assert!(registry.provider("builtin").is_none());
registry
.register("builtin".into(), slot())
.expect("the name is free after removal");
}
}