burn-router 0.22.0

Multi-backend router decorator for the Burn framework
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;

/// Type alias for `<R as RouterChannel>::Client`.
pub type Client<R> = <R as RouterChannel>::Client;
pub(crate) static CLIENTS: RouterClientLocator = RouterClientLocator::new();

type Key = (core::any::TypeId, DeviceId);

/// Define how to interact with the interpreter.
pub trait RouterClient: Clone + Send + Sync + Sized {
    /// Device type.
    type Device: DeviceOps;

    /// Register a new tensor operation to be executed by the interpreter (server).
    fn register_op(&self, op: OperationIr);
    /// Register a new tensor operation to be executed by the interpreter (server).
    ///
    /// Returns the new (uninitialized) output tensor(s) generated by the registered operation.
    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
    }
    /// Read the values contained by a tensor.
    fn read_tensor_async(&self, tensor: TensorIr) -> DynFut<Result<TensorData, ExecutionError>>;
    /// Sync the interpreter, ensure that all computations are finished.
    fn sync(&self) -> Result<(), ExecutionError>;
    /// Eagerly submit the operations registered so far without waiting for them to complete.
    ///
    /// Unlike [`sync`](Self::sync), this does not block on results — it only ensures the
    /// buffered operations are handed off for execution (and, for the remote backend, sent to
    /// the server) instead of sitting in a local buffer.
    fn flush(&self);
    /// Create a new (uninitialized) empty tensor and returns its corresponding [tensor id](TensorId).
    fn create_empty_handle(&self) -> TensorId;
    /// Create a new [RouterTensor] from the tensor data.
    fn register_tensor_data(&self, data: TensorData) -> RouterTensor<Self>;
    /// Get the current device used by all operations handled by this client.
    fn device(&self) -> Self::Device;
    /// Seed the interpreter.
    fn seed(&self, seed: u64);
    /// Returns the supported data type usage set
    fn dtype_usage(&self, dtype: DType) -> burn_backend::DTypeUsageSet;
    /// Open a profiling window on the interpreter, where the calling stream
    /// stands — see [`Backend::profile_start`](burn_backend::Backend::profile_start).
    ///
    /// `None`, the default, from an interpreter that opens no windows.
    fn profile_start(&self) -> Result<Option<ProfileToken>, ExecutionError> {
        Ok(None)
    }
    /// Close the window `token` where the calling stream stands, flushing the
    /// interpreter's backend first when `options` ask for it.
    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",
        ))
    }

    /// Drop the window `token` without measuring it, for a caller that will
    /// never reach [`profile_end`](Self::profile_end) — see
    /// [`Backend::profile_abandon`](burn_backend::Backend::profile_abandon).
    ///
    /// The default closes the window and discards the measurement, which any
    /// interpreter that opens one can already do. An interpreter that can say
    /// so more cheaply — a remote one, where the close is a round trip and
    /// the caller is unwinding — does that instead.
    fn profile_abandon(&self, token: ProfileToken) {
        let _ = self.profile_end(token, ProfileOptions::default());
    }

    /// Register a reusable group of operations (in relative form) under `graph_id` *and* run its
    /// first invocation with `bindings`, so it can later be replayed by id with
    /// [`execute_graph`](Self::execute_graph).
    ///
    /// Registration always coincides with the first execution, so they're combined to save a
    /// round-trip on a cache miss. Used by the fusion layer to avoid re-sending a recurring
    /// op-graph.
    fn register_and_execute_graph(
        &self,
        graph_id: GraphId,
        relative_graph: Vec<OperationIr>,
        bindings: GraphBindings,
    );

    /// Replay a previously [registered](Self::register_and_execute_graph) graph with the given
    /// concrete bindings.
    fn execute_graph(&self, graph_id: GraphId, bindings: GraphBindings);

    /// Register `new_id` as an alias of `src_id` — a second handle over the same backing buffer.
    ///
    /// Used by the fusion layer's cross-stream sharing (see
    /// [`FusionRuntime::alias_handle`](burn_fusion::FusionRuntime::alias_handle)): when a tensor is
    /// shared to another stream, that stream's view needs its own id so that consuming it (a
    /// `ReadWrite` last-use) frees only this alias, leaving the original handle valid. The server
    /// clones the source handle (an `Arc`-style refcount on the device buffer) under `new_id`.
    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>>>>,
}

/// Get the client currently associated with `device`.
///
/// An existing scoped or unscoped registration is returned as-is. On a cache miss, the channel's
/// [`RouterChannel::init_client`] implementation creates an unscoped client that remains cached in
/// the global locator. Use [`register_scoped_client`] when the caller, rather than the locator,
/// must control the client's lifetime.
pub fn get_client<R: RouterChannel>(device: &R::Device) -> Client<R> {
    CLIENTS.client::<R>(device)
}

/// Guard owning a router client's registration in the global locator.
///
/// The locator stores a clone of the registered client; this guard stores the corresponding lookup
/// key. Dropping the guard removes that clone. Tensor handles remain safe because they own client
/// clones independently of the locator. Most router clients are process-long lived and use
/// [`get_client`]; lifecycle-owned channels use [`register_scoped_client`] and retain this guard.
#[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);
    }
}

/// Register `client` with a locator entry owned by the returned guard.
///
/// This inserts the caller-created client directly and does not call
/// [`RouterChannel::init_client`]. While the guard is alive, [`get_client`] for the same channel and
/// device returns a clone of this client. Dropping the guard removes the locator entry, allowing a
/// later scope to associate the same device with a different client.
///
/// Unlike [`get_client`], this requires the device not to have a registered client already. The
/// guard gives one lifecycle owner responsibility for cleanup and prevents exposing unrestricted
/// client removal to downstream crates. Returns `None` when the device already has a client.
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,
        })
}

/// Initialize a new client for the given device.
///
/// If a (global) seed was previously set, the client seed is set.
fn new_client<R: RouterChannel>(device: &R::Device) -> Client<R> {
    R::init_client(device)
}

impl RouterClientLocator {
    /// Create a new client locator.
    pub const fn new() -> Self {
        Self {
            clients: Mutex::new(None),
        }
    }

    /// Get the router client for the given device.
    ///
    /// If a client isn't already initialized, it is created.
    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!(),
        }
    }

    /// Register a client with a unique lifecycle owner.
    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)
    }

    /// Remove the client identified by a scoped registration guard.
    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));
        }
    }
}