use std::collections::HashSet;
use std::marker::PhantomData;
use std::sync::atomic::{AtomicU64, Ordering};
use burn_backend::DType;
use burn_backend::ops::FloatTensorOps;
use burn_fusion::stream::{Context, Operation, OrderedExecution};
use burn_fusion::{
ExecutionError, FuserProperties, FuserStatus, FusionBackend, FusionRuntime, NumOperations,
OperationFuser, OperationRan, Optimization,
};
use burn_ir::{
BackendIr, CustomOpIr, GraphBindings, GraphId, Handle, HandleContainer, OperationIr, ScalarIr,
TensorHandle, TensorId, TensorIr, TensorStatus,
};
use serde::{Deserialize, Serialize};
use burn_std::config::config;
use crate::{BackendRouter, Graph, RouterChannel, RouterClient, RouterTensor, get_client};
static GRAPH_ID_COUNTER: AtomicU64 = AtomicU64::new(0);
fn next_graph_id() -> GraphId {
GraphId(GRAPH_ID_COUNTER.fetch_add(1, Ordering::Relaxed))
}
impl<R: RouterChannel> BackendIr for BackendRouter<R> {
type Handle = RouterTensor<R::Client>;
fn float_tensor(handle: TensorHandle<Self::Handle>) -> burn_backend::tensor::FloatTensor<Self> {
handle.handle
}
fn int_tensor(handle: TensorHandle<Self::Handle>) -> burn_backend::tensor::IntTensor<Self> {
handle.handle
}
fn bool_tensor(handle: TensorHandle<Self::Handle>) -> burn_backend::tensor::BoolTensor<Self> {
handle.handle
}
fn quantized_tensor(
handle: TensorHandle<Self::Handle>,
) -> burn_backend::tensor::QuantizedTensor<Self> {
handle.handle
}
fn float_tensor_handle(tensor: burn_backend::tensor::FloatTensor<Self>) -> Self::Handle {
tensor
}
fn int_tensor_handle(tensor: burn_backend::tensor::IntTensor<Self>) -> Self::Handle {
tensor
}
fn bool_tensor_handle(tensor: burn_backend::tensor::BoolTensor<Self>) -> Self::Handle {
tensor
}
fn quantized_tensor_handle(
tensor: burn_backend::tensor::QuantizedTensor<Self>,
) -> Self::Handle {
tensor
}
}
impl<R: RouterChannel> FusionBackend for BackendRouter<R> {
type FusionRuntime = RouterFusionRuntime<R>;
type FullPrecisionBackend = Self;
fn cast_float(tensor: burn_backend::tensor::FloatTensor<Self>, dtype: DType) -> Self::Handle {
Self::float_cast(tensor, dtype.into())
}
}
pub struct RouterFusionRuntime<R: RouterChannel> {
_p: PhantomData<R>,
}
impl<R: RouterChannel> core::fmt::Debug for RouterFusionRuntime<R> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str("RouterFusionRuntime")
}
}
impl<R: RouterChannel> FusionRuntime for RouterFusionRuntime<R> {
type OptimizationState = RouterGraphExecutionState;
type Optimization = RouterGraphExecution<R>;
type FusionHandle = RouterTensor<R::Client>;
type FusionDevice = R::Device;
fn fusers(device: R::Device) -> Vec<Box<dyn OperationFuser<Self::Optimization>>> {
vec![Box::new(RouterFuser::<R>::new(device))]
}
fn alias_handle(handle: &RouterTensor<R::Client>) -> RouterTensor<R::Client> {
let id = handle.client.create_empty_handle();
handle.client.register_alias(id, handle.id);
RouterTensor::new(
id,
handle.shape.clone(),
handle.dtype,
handle.client.clone(),
)
}
fn free_handle(
handles: &mut HandleContainer<RouterTensor<R::Client>>,
tensor: &TensorIr,
ran: OperationRan,
) {
if tensor.status != TensorStatus::ReadWrite {
return;
}
if ran == OperationRan::No {
handles.free(tensor);
return;
}
if let Some(Handle::Existing(handle)) = handles.remove_handle(tensor.id) {
handle
.count
.fetch_add(1, core::sync::atomic::Ordering::Relaxed);
}
}
}
pub struct RouterFuser<R: RouterChannel> {
device: R::Device,
ops: Vec<OperationIr>,
closed_on_init: bool,
savings: ReplaySavings,
score: u64,
score_max: u64,
num_since_max_unchanged: usize,
max_graph_size: Option<usize>,
growth_patience: usize,
}
impl<R: RouterChannel> RouterFuser<R> {
fn new(device: R::Device) -> Self {
let cfg = config();
let fusion = cfg.fusion();
Self {
device,
ops: Vec::new(),
closed_on_init: false,
savings: ReplaySavings::default(),
score: 0,
score_max: 0,
num_since_max_unchanged: 0,
max_graph_size: fusion.max_graph_size,
growth_patience: fusion.growth_patience,
}
}
fn score(&self) -> u64 {
const FACTOR_SAVED: u64 = 100; const FREE_OPS: usize = 64; const PENALTY_PER_OP: u64 = 0;
let benefit = self.savings.percent() * FACTOR_SAVED;
let penalty = (self.ops.len().saturating_sub(FREE_OPS) as u64) * PENALTY_PER_OP;
benefit.saturating_sub(penalty) + 1
}
}
impl<R: RouterChannel> Clone for RouterFuser<R> {
fn clone(&self) -> Self {
Self {
device: self.device.clone(),
ops: self.ops.clone(),
closed_on_init: self.closed_on_init,
savings: self.savings.clone(),
score: self.score,
score_max: self.score_max,
num_since_max_unchanged: self.num_since_max_unchanged,
max_graph_size: self.max_graph_size,
growth_patience: self.growth_patience,
}
}
}
impl<R: RouterChannel> OperationFuser<RouterGraphExecution<R>> for RouterFuser<R> {
fn fuse(&mut self, operation: &OperationIr) {
if self.closed_on_init {
return;
}
if let OperationIr::Init(_) = operation {
self.closed_on_init = true;
return;
}
self.savings.add(operation);
self.ops.push(operation.clone());
self.score = self.score();
if self.score > self.score_max {
self.score_max = self.score;
self.num_since_max_unchanged = 0;
} else {
self.num_since_max_unchanged += 1;
}
}
fn finish(&mut self) -> RouterGraphExecution<R> {
self.savings = ReplaySavings::default();
let ops = core::mem::take(&mut self.ops);
RouterGraphExecution::new(ops, self.device.clone())
}
fn reset(&mut self) {
self.ops.clear();
self.closed_on_init = false;
self.savings = ReplaySavings::default();
}
fn status(&self) -> FuserStatus {
let over_max = self.max_graph_size.is_some_and(|max| self.len() > max);
if self.closed_on_init || self.num_since_max_unchanged >= self.growth_patience || over_max {
FuserStatus::Closed
} else {
FuserStatus::Open
}
}
fn properties(&self) -> FuserProperties {
FuserProperties {
score: self.score,
ready: self.ops.len() > 1,
}
}
fn len(&self) -> usize {
self.ops.len()
}
fn clone_dyn(&self) -> Box<dyn OperationFuser<RouterGraphExecution<R>>> {
Box::new(self.clone())
}
}
pub struct CustomOperation<R: RouterChannel> {
ir: CustomOpIr,
device: R::Device,
}
impl<R: RouterChannel> core::fmt::Debug for CustomOperation<R> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("CustomOperation")
.field("id", &self.ir.id)
.finish()
}
}
impl<R: RouterChannel> CustomOperation<R> {
pub fn new(ir: CustomOpIr, device: R::Device) -> Self {
Self { ir, device }
}
}
impl<R: RouterChannel> Operation<RouterFusionRuntime<R>> for CustomOperation<R> {
fn execute(
&self,
handles: &mut HandleContainer<RouterTensor<R::Client>>,
) -> Result<(), ExecutionError> {
let client = get_client::<R>(&self.device);
let inputs: Vec<TensorIr> = self
.ir
.inputs
.iter()
.map(|input| handles.get_handle(&input.id, &input.status).into_ir())
.collect();
let outputs: Vec<RouterTensor<R::Client>> = self
.ir
.outputs
.iter()
.map(|out| {
RouterTensor::new(
client.create_empty_handle(),
out.shape.clone(),
out.dtype,
client.clone(),
)
})
.collect();
client.register_op(OperationIr::Custom(CustomOpIr {
id: self.ir.id.clone(),
inputs,
outputs: outputs.iter().map(|out| out.to_ir_out()).collect(),
scalars: self.ir.scalars.clone(),
}));
for (out, tensor) in self.ir.outputs.iter().zip(outputs) {
handles.register_handle(out.id, tensor);
}
Ok(())
}
}
pub struct RouterGraphExecution<R: RouterChannel> {
graph: Graph,
device: R::Device,
graph_id: Option<GraphId>,
_p: PhantomData<R>,
}
#[derive(Serialize, Deserialize)]
pub struct RouterGraphExecutionState {
operations: Vec<OperationIr>,
}
impl<R: RouterChannel> RouterGraphExecution<R> {
fn new(graph: Vec<OperationIr>, device: R::Device) -> Self {
Self {
graph: Graph::new(graph),
device,
graph_id: None,
_p: PhantomData,
}
}
}
#[derive(Clone, Default)]
struct ReplaySavings {
baseline: u64,
dims: HashSet<usize>,
referenced: HashSet<TensorId>,
produced: HashSet<TensorId>,
consumed: HashSet<TensorId>,
inputs: usize,
outputs: usize,
}
impl ReplaySavings {
const TENSOR_BYTES: u64 = 10; const DIM_BYTES: u64 = 8;
const OP_BYTES: u64 = 8; const BINDING_BYTES: u64 = 16;
fn add(&mut self, op: &OperationIr) {
let nodes = op.nodes();
self.baseline += Self::OP_BYTES;
for tensor in nodes.iter().copied().chain(op.outputs()) {
self.baseline += Self::TENSOR_BYTES;
for dim in tensor.shape.iter() {
self.baseline += Self::DIM_BYTES;
self.dims.insert(*dim);
}
}
if let OperationIr::Drop(tensor) = op {
self.consume(tensor.id);
}
if !matches!(op, OperationIr::Init(_)) {
for tensor in op.outputs() {
self.produce(tensor.id);
}
}
for tensor in nodes {
self.reference(tensor.id);
if tensor.status == TensorStatus::ReadWrite {
self.consume(tensor.id);
}
}
}
fn reference(&mut self, id: TensorId) {
if self.referenced.insert(id) && !self.produced.contains(&id) {
self.inputs += 1;
}
}
fn produce(&mut self, id: TensorId) {
if self.produced.insert(id) {
if self.referenced.contains(&id) {
self.inputs -= 1;
}
if !self.consumed.contains(&id) {
self.outputs += 1;
}
}
}
fn consume(&mut self, id: TensorId) {
if self.consumed.insert(id) && self.produced.contains(&id) {
self.outputs -= 1;
}
}
fn percent(&self) -> u64 {
if self.baseline == 0 {
return 0;
}
let bindings = (self.inputs + self.outputs) as u64 * Self::BINDING_BYTES
+ self.dims.len() as u64 * Self::DIM_BYTES;
(self.baseline.saturating_sub(bindings) * 100 / self.baseline).min(100)
}
}
impl<R: RouterChannel> core::fmt::Debug for RouterGraphExecution<R> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("RouterGraphExecution")
.field("len", &self.graph.len())
.finish()
}
}
impl<R: RouterChannel> NumOperations for RouterGraphExecution<R> {
fn len(&self) -> usize {
self.graph.len()
}
fn name(&self) -> &'static str {
"RouterGraphExecution"
}
}
impl<R: RouterChannel> Optimization<RouterFusionRuntime<R>> for RouterGraphExecution<R> {
fn execute(
&mut self,
context: &mut Context<RouterTensor<R::Client>>,
_execution: &OrderedExecution<RouterFusionRuntime<R>>,
) {
let client = get_client::<R>(&self.device);
let mut tensors: Vec<(TensorId, TensorId)> =
Vec::with_capacity(self.graph.inputs().len() + self.graph.outputs().len());
for &input_id in self.graph.inputs() {
let global_id = match context.tensors.get(&input_id) {
Some(global) => global.id,
None => continue,
};
if let Some(handle) = context.handles.get_handle_ref(&global_id) {
tensors.push((input_id, handle.id()));
}
}
for &output_id in self.graph.outputs() {
let output = context
.tensors
.get(&output_id)
.map(|global| (global.id, global.shape.clone(), global.dtype));
if let Some((fusion_id, shape, dtype)) = output {
let concrete_id = client.create_empty_handle();
tensors.push((output_id, concrete_id));
let handle = RouterTensor::new(concrete_id, shape, dtype, client.clone());
context.handles.register_handle(fusion_id, handle);
}
}
let mut shapes = vec![0usize; context.shapes_relative2global.len()];
for (relative, concrete) in context.shapes_relative2global.iter() {
if *relative < shapes.len() {
shapes[*relative] = *concrete;
}
}
let mut scalars = vec![ScalarIr::UInt(0); context.scalars.len()];
for (scalar_id, value) in context.scalars.iter() {
let idx = scalar_id.value as usize;
if idx < scalars.len() {
scalars[idx] = *value;
}
}
let ranges = context.ranges.clone();
let bindings = GraphBindings {
tensors,
shapes,
scalars,
ranges,
};
match self.graph_id {
Some(id) => client.execute_graph(id, bindings),
None => {
let id = next_graph_id();
self.graph_id = Some(id);
client.register_and_execute_graph(id, self.graph.operations().to_vec(), bindings);
}
};
}
fn to_state(&self) -> RouterGraphExecutionState {
RouterGraphExecutionState {
operations: self.graph.operations().to_vec(),
}
}
fn from_state(device: &R::Device, state: RouterGraphExecutionState) -> Self {
Self::new(state.operations, device.clone())
}
}
#[cfg(test)]
mod tests {
use super::*;
use burn_backend::Shape;
use burn_ir::{CustomOpIr, GraphIr, InitOperationIr};
fn whole_graph_percent(ops: &[OperationIr]) -> u64 {
let mut baseline = 0u64;
let mut dims = HashSet::new();
for op in ops {
baseline += ReplaySavings::OP_BYTES;
for tensor in op.nodes().into_iter().chain(op.outputs()) {
baseline += ReplaySavings::TENSOR_BYTES;
for dim in tensor.shape.iter() {
baseline += ReplaySavings::DIM_BYTES;
dims.insert(*dim);
}
}
}
if baseline == 0 {
return 0;
}
let boundary = GraphIr::classify(ops);
let bindings = (boundary.inputs.len() + boundary.outputs.len()) as u64
* ReplaySavings::BINDING_BYTES
+ dims.len() as u64 * ReplaySavings::DIM_BYTES;
(baseline.saturating_sub(bindings) * 100 / baseline).min(100)
}
struct Rng(u64);
impl Rng {
fn below(&mut self, bound: u64) -> u64 {
self.0 ^= self.0 << 13;
self.0 ^= self.0 >> 7;
self.0 ^= self.0 << 17;
self.0 % bound
}
fn tensor(&mut self, pool: u64) -> TensorIr {
let id = TensorId::new(self.below(pool));
let dims: Vec<usize> = (0..=self.below(3))
.map(|_| 1 + self.below(6) as usize)
.collect();
let mut tensor = TensorIr::uninit(id, Shape::from(dims), DType::F32);
if self.below(4) == 0 {
tensor.status = TensorStatus::ReadWrite;
}
tensor
}
fn op(&mut self, pool: u64) -> OperationIr {
match self.below(6) {
0 => OperationIr::Drop(self.tensor(pool)),
1 => OperationIr::Init(InitOperationIr {
out: self.tensor(pool),
}),
_ => {
let inputs: Vec<TensorIr> =
(0..self.below(3)).map(|_| self.tensor(pool)).collect();
let outputs: Vec<TensorIr> =
(0..self.below(3)).map(|_| self.tensor(pool)).collect();
OperationIr::Custom(CustomOpIr::new("op", &inputs, &outputs))
}
}
}
}
#[test]
fn savings_match_the_whole_graph_after_every_op() {
let mut rng = Rng(0x9e37_79b9_7f4a_7c15);
for _ in 0..500 {
let pool = 2 + rng.below(8);
let mut ops = Vec::new();
let mut savings = ReplaySavings::default();
for _ in 0..40 {
let op = rng.op(pool);
savings.add(&op);
ops.push(op);
let boundary = GraphIr::classify(&ops);
assert_eq!(
(savings.inputs, savings.outputs),
(boundary.inputs.len(), boundary.outputs.len())
);
assert_eq!(savings.percent(), whole_graph_percent(&ops));
}
}
}
}