use crate::{RouterChannel, RouterTensor};
use alloc::boxed::Box;
use alloc::vec::Vec;
use burn_backend::{
DType, ProfileDuration, ProfileOptions, ProfileToken, TensorData,
backend::{DeviceId, DeviceOps, ExecutionError},
};
use burn_ir::{GraphBindings, GraphId, OperationIr, TensorId, TensorIr};
use burn_std::future::DynFut;
use core::{marker::PhantomData, ops::DerefMut};
use hashbrown::HashMap;
use spin::Mutex;
pub type Client<R> = <R as RouterChannel>::Client;
pub(crate) static CLIENTS: RouterClientLocator = RouterClientLocator::new();
type Key = (core::any::TypeId, DeviceId);
pub trait RouterClient: Clone + Send + Sync + Sized {
type Device: DeviceOps;
fn register_op(&self, op: OperationIr);
fn register(&self, op: OperationIr) -> Vec<RouterTensor<Self>> {
let out = op
.outputs()
.map(|output| {
RouterTensor::new(output.id, output.shape.clone(), output.dtype, self.clone())
})
.collect();
self.register_op(op);
out
}
fn read_tensor_async(&self, tensor: TensorIr) -> DynFut<Result<TensorData, ExecutionError>>;
fn sync(&self) -> Result<(), ExecutionError>;
fn flush(&self);
fn create_empty_handle(&self) -> TensorId;
fn register_tensor_data(&self, data: TensorData) -> RouterTensor<Self>;
fn device(&self) -> Self::Device;
fn seed(&self, seed: u64);
fn dtype_usage(&self, dtype: DType) -> burn_backend::DTypeUsageSet;
fn profile_start(&self) -> Result<Option<ProfileToken>, ExecutionError> {
Ok(None)
}
fn profile_end(
&self,
token: ProfileToken,
options: ProfileOptions,
) -> Result<ProfileDuration, ExecutionError> {
let _ = (token, options);
Err(ExecutionError::with_context(
"profiling windows are not supported by this interpreter",
))
}
fn profile_abandon(&self, token: ProfileToken) {
let _ = self.profile_end(token, ProfileOptions::default());
}
fn register_and_execute_graph(
&self,
graph_id: GraphId,
relative_graph: Vec<OperationIr>,
bindings: GraphBindings,
);
fn execute_graph(&self, graph_id: GraphId, bindings: GraphBindings);
fn register_alias(&self, new_id: TensorId, src_id: TensorId);
}
pub(crate) struct RouterClientLocator {
clients: Mutex<Option<HashMap<Key, Box<dyn core::any::Any + Send>>>>,
}
pub fn get_client<R: RouterChannel>(device: &R::Device) -> Client<R> {
CLIENTS.client::<R>(device)
}
#[must_use = "dropping the registration immediately unregisters its router client"]
pub struct RouterClientRegistration<R: RouterChannel> {
key: Key,
_channel: PhantomData<R>,
}
impl<R: RouterChannel> Drop for RouterClientRegistration<R> {
fn drop(&mut self) {
CLIENTS.remove(self.key);
}
}
pub fn register_scoped_client<R: RouterChannel>(
device: &R::Device,
client: Client<R>,
) -> Option<RouterClientRegistration<R>> {
CLIENTS
.register_scoped::<R>(device, client)
.map(|key| RouterClientRegistration {
key,
_channel: PhantomData,
})
}
fn new_client<R: RouterChannel>(device: &R::Device) -> Client<R> {
R::init_client(device)
}
impl RouterClientLocator {
pub const fn new() -> Self {
Self {
clients: Mutex::new(None),
}
}
pub fn client<R: RouterChannel + 'static>(&self, device: &R::Device) -> Client<R> {
let device_id = device.id();
let client_id = (core::any::TypeId::of::<R>(), device_id);
let mut clients = self.clients.lock();
if clients.is_none() {
let client = new_client::<R>(device);
Self::register_inner::<R>(client_id, client, &mut clients);
}
match clients.deref_mut() {
Some(clients) => match clients.get(&client_id) {
Some(client) => {
let client: &Client<R> = client.downcast_ref().unwrap();
client.clone()
}
None => {
let client = new_client::<R>(device);
let any = Box::new(client.clone());
clients.insert(client_id, any);
client
}
},
_ => unreachable!(),
}
}
fn register_scoped<R: RouterChannel + 'static>(
&self,
device: &R::Device,
client: Client<R>,
) -> Option<Key> {
let key = (core::any::TypeId::of::<R>(), device.id());
let mut clients = self.clients.lock();
let clients = clients.get_or_insert_with(HashMap::new);
if clients.contains_key(&key) {
return None;
}
clients.insert(key, Box::new(client));
Some(key)
}
fn remove(&self, key: Key) {
let mut clients = self.clients.lock();
if let Some(clients) = clients.as_mut() {
clients.remove(&key);
}
}
fn register_inner<R: RouterChannel + 'static>(
key: Key,
client: Client<R>,
clients: &mut Option<HashMap<Key, Box<dyn core::any::Any + Send>>>,
) {
if clients.is_none() {
*clients = Some(HashMap::new());
}
if let Some(clients) = clients {
if clients.contains_key(&key) {
panic!("Client already created for device {key:?}");
}
clients.insert(key, Box::new(client));
}
}
}