burn-router 0.22.0

Multi-backend router decorator for the Burn framework
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
//! Generic client-side graph caching for any [router backend](crate::BackendRouter).
//!
//! Wrapping a router backend as [`Fusion`](burn_fusion::Fusion) turns recurring groups of tensor
//! operations into reusable, client-cached graphs. A single greedy [`RouterFuser`] accumulates
//! every operation up to the next upload and is drained at a sync point, so one graph covers each
//! connected block between syncs and uploads. On execution the group is registered once on the
//! backend (via the [`RouterClient`]) and thereafter invoked by id with only the changing
//! bindings. For the remote backend this means a recurring computation (e.g. a model block)
//! crosses the network once instead of every step.
//!
//! `burn-fusion` names its hook `Optimization`; here that hook *is* a cached graph execution
//! ([`RouterGraphExecution`]), to distinguish it from a compute backend's kernel fusion.

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))
}

// The router backend already implements `Backend`; these two impls add the `BackendIr` +
// `FusionBackend` glue so it can be wrapped as `Fusion<BackendRouter<R>>`. All four tensor
// primitives are `RouterTensor`, so the handle conversions are the identity (mirroring the
// `impl BackendIr for Fusion<B>` in burn-fusion).
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())
    }
}

/// The [fusion runtime](FusionRuntime) for a [router backend](BackendRouter).
///
/// Its [handle](FusionRuntime::FusionHandle) is a [`RouterTensor`] — a lightweight reference to a
/// backend-resident tensor plus the client to reach it — and its "optimization" is a backend-cached
/// op-graph (see [`RouterGraphExecution`]).
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> {
        // The router handle is a thin id into a server-side tensor, so a bare `clone()` would keep
        // the same server id — every cross-stream alias would collapse onto one server handle, and
        // the first stream to consume it (`ReadWrite`) would free it for the others. Instead mint a
        // fresh id and have the server register it as an alias of the same buffer (a refcounted
        // clone), so each stream's view frees independently. Mirrors local backends, where the
        // default `clone()` already yields an independent `HandleContainer` entry over a shared
        // `Arc` buffer.
        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,
    ) {
        // Only `ReadWrite` (last-use) nodes are freed here, matching `HandleContainer::free`.
        if tensor.status != TensorStatus::ReadWrite {
            return;
        }
        // An operation that never ran was never replayed, so the server still holds this tensor
        // and the client `Drop` is the only thing that will free it. Letting it through is what
        // keeps a failed or skipped operation from stranding a buffer on the server for good.
        if ran == OperationRan::No {
            handles.free(tensor);
            return;
        }
        // A replayed op (a `Drop`, or a `ReadWrite` input of a compute op) already popped this
        // handle server-side. So remove the client entry WITHOUT letting `RouterTensor::drop`
        // register a *second*, redundant `Drop` for the same id. Bumping the refcount makes that
        // drop a no-op.
        if let Some(Handle::Existing(handle)) = handles.remove_handle(tensor.id) {
            handle
                .count
                .fetch_add(1, core::sync::atomic::Ordering::Relaxed);
        }
    }
}

/// A greedy operation fuser that records every operation up to the next upload.
///
/// While [`status`](OperationFuser::status) reports [`FuserStatus::Open`], the fusion engine keeps
/// deferring in lazy mode and drains the queue at a sync point (a read, `sync`, or `flush`). At
/// that point [`finish`](OperationFuser::finish) yields a single [`RouterGraphExecution`] covering
/// the whole accumulated (connected) block.
///
/// It closes on an [`Init`](OperationIr::Init) without taking it: the data already went out when
/// the tensor was created, and a queued `Init` holds back a drop from another thread until this
/// stream executes, which on a thread that only uploads never happens.
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,
    /// Close the graph once it reaches this many ops, if set (`FusionConfig::max_graph_size`).
    max_graph_size: Option<usize>,
    /// Close the graph once the score hasn't reached a new max for this many consecutive ops
    /// (`FusionConfig::growth_patience`).
    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,
        }
    }

    /// Value-based fusion score for the currently accumulated ops.
    ///
    /// Benefit: estimated % of serialized bytes caching saves per replay (the relative graph vs the
    /// per-replay bindings, see [`ReplaySavings`]), weighted by `FACTOR_SAVED`. Cost: a
    /// penalty that grows with op count past `FREE_OPS`, so the score *peaks* at a "right-sized"
    /// graph and then decays once the per-op overhead outweighs the savings.
    ///
    /// Floored at 1 for any non-empty graph: a score of 0 makes `find_best_optimization_index`
    /// treat the graph as "don't fuse" (streaming every op unfused → the cache never replays). The
    /// `+1` doesn't move the argmax, so it doesn't change which size wins.
    fn score(&self) -> u64 {
        const FACTOR_SAVED: u64 = 100; // weight per % point saved → benefit in 0..=10_000
        const FREE_OPS: usize = 64; // ops below this are not penalized
        const PENALTY_PER_OP: u64 = 0; // penalty per op beyond FREE_OPS

        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) {
        // `Block::optimize` drains as many ops as the fuser took, so they must be a prefix.
        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())
    }
}

/// Unfused execution of a single [custom op](OperationIr::Custom) on a router backend.
///
/// Every built-in op reaches the backend through a `BackendRouter` method, so when the fuser leaves
/// it unfused (a one-op segment — e.g. a source op whose output is read right away) `burn-fusion`
/// runs that method op by op. A custom op has no such method, so it needs this handler: it is the
/// [`Operation`] registered alongside the op's IR, and runs only on the unfused path. When the op is
/// instead fused into a [`RouterGraphExecution`] it ships as part of the graph and this never runs.
///
/// It mirrors what a built-in op does on that path: resolve the fused input handles to their
/// backend tensor ids, ship the op through the [`RouterClient`], and bind the outputs so downstream
/// ops find them. (Without it the op would silently do nothing and the backend would never create
/// the output handle — "Should have handle for tensor ..." on the next access.)
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> {
    /// Create the unfused handler for `ir`, executed on `device`.
    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);

        // Map each fused input handle to its backend tensor id. `into_ir` carries the
        // refcount/free semantics, so a last-use (`ReadWrite`) input still frees correctly.
        let inputs: Vec<TensorIr> = self
            .ir
            .inputs
            .iter()
            .map(|input| handles.get_handle(&input.id, &input.status).into_ir())
            .collect();

        // Mint fresh backend ids for the outputs.
        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();

        // Ship the op to the backend with the translated ids; scalars travel unchanged.
        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(),
        }));

        // Bind the outputs under their fused ids so any downstream op resolves them.
        for (out, tensor) in self.ir.outputs.iter().zip(outputs) {
            handles.register_handle(out.id, tensor);
        }

        Ok(())
    }
}

/// A reusable group of operations, registered once on the backend and thereafter invoked by id.
///
/// The recorded [`graph`](Graph) is in *relative* form (positional tensor ids, relative
/// shape-dim ids, scalar placeholders), which is invariant across invocations — that is what lets
/// the backend cache it and the client reuse it. On each [`execute`](Optimization::execute) only
/// the concrete bindings are computed from the [`Context`] and sent; the (large) graph itself
/// travels only on the first invocation.
pub struct RouterGraphExecution<R: RouterChannel> {
    graph: Graph,
    device: R::Device,
    /// Backend-side id, assigned and registered on the first execution and reused afterwards.
    graph_id: Option<GraphId>,
    _p: PhantomData<R>,
}

/// Serializable state for a [`RouterGraphExecution`].
///
/// The backend id is intentionally not serialized — it is per-connection state, so a deserialized
/// graph re-registers itself on first use.
#[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,
        }
    }
}

/// What a replay of the cached graph saves over sending it again, kept per operation because the
/// fuser asks after each one. Scalars and ranges travel either way, so they cancel in the ratio.
#[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; // id (≈8) + dtype + status
    const DIM_BYTES: u64 = 8;
    const OP_BYTES: u64 = 8; // variant tag + small bookkeeping
    const BINDING_BYTES: u64 = 16; // a (relative id, concrete id) binding entry

    /// Adds `op` under the boundary rules of [`GraphIr::classify`](burn_ir::GraphIr::classify).
    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;
        }
    }

    /// The share of the baseline a replay saves, in `0..=100`.
    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);

        // Walk only the precomputed boundary — never every tensor — so a replay is O(boundary).
        // Inputs reuse their resident concrete id; surviving outputs get a fresh id and a handle.
        // Intermediates are never touched here: the replay owns their ids and derives their shapes
        // from the shape table below.
        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() {
            // Extract owned metadata first so the `context.tensors` borrow ends before we touch
            // `context.handles`.
            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);
            }
        }

        // Dense shape-dim table indexed by relative dim id (dim ids are dense `0..N`).
        let mut shapes = vec![0usize; context.shapes_relative2global.len()];
        for (relative, concrete) in context.shapes_relative2global.iter() {
            if *relative < shapes.len() {
                shapes[*relative] = *concrete;
            }
        }

        // Concrete scalar values, indexed by their placeholder id.
        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;
            }
        }

        // Concrete slice ranges, indexed by their placeholder id (carried in a relative range's
        // `start`). Cheap to clone — a few `Slice`s per slice op — and, like scalars, they can
        // change between invocations of the same cached graph, so they travel every time.
        let ranges = context.ranges.clone();

        let bindings = GraphBindings {
            tensors,
            shapes,
            scalars,
            ranges,
        };
        match self.graph_id {
            // Already registered: replay by id, sending only the bindings.
            Some(id) => client.execute_graph(id, bindings),
            // First invocation: register the relative graph and execute it in a single round-trip.
            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};

    /// The estimate recomputed over the whole graph, which the savings must match after every op.
    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 {
            // Few ids, so ops read, produce and drop the same tensors in every order.
            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));
            }
        }
    }
}