Skip to main content

Backend

Trait Backend 

pub trait Backend:
    Sized
    + BackendTypes
    + FloatTensorOps<Self>
    + BoolTensorOps<Self>
    + IntTensorOps<Self>
    + ModuleOps<Self>
    + ActivationOps<Self>
    + QTensorOps<Self>
    + TransactionOps<Self>
    + DistributedOps<Self>
    + Clone
    + Default
    + Send
    + Sync
    + Debug
    + 'static {
Show 21 methods // Required methods fn name(device: &Self::Device) -> String; fn seed(device: &Self::Device, seed: u64); fn flush(_device: &Self::Device) -> Result<(), ExecutionError>; fn dtype_usage(device: &Self::Device, dtype: DType) -> EnumSet<DTypeUsage>; fn device_count(type_id: u16) -> usize; // Provided methods fn ad_enabled(_device: &Self::Device) -> bool { ... } fn memory_persistent_allocations<Output, Input, Func>( device: &Self::Device, input: Input, func: Func, ) -> Output where Output: Send, Input: Send, Func: Fn(Input) -> Output + Send { ... } fn memory_cleanup(device: &Self::Device) { ... } fn memory_pool_report( device: &Self::Device, ) -> Option<Vec<SlicedPoolReport>> { ... } fn memory_pool_usage(device: &Self::Device) -> Option<MemoryPoolUsage> { ... } fn sync(_device: &Self::Device) -> Result<(), ExecutionError> { ... } fn profile<O>( device: &Self::Device, options: ProfileOptions, func: impl FnOnce() -> O + Send, ) -> Result<(O, ProfileDuration), ExecutionError> where O: Send + 'static { ... } fn profile_start( _device: &Self::Device, ) -> Result<Option<ProfileToken>, ExecutionError> { ... } fn profile_end( _device: &Self::Device, _token: ProfileToken, _options: ProfileOptions, ) -> Result<ProfileDuration, ExecutionError> { ... } fn profile_abandon(device: &Self::Device, token: ProfileToken) { ... } fn graph_prepare(_device: &Self::Device) -> Result<(), ExecutionError> { ... } fn graph_start_capture(_device: &Self::Device) -> Result<(), ExecutionError> { ... } fn graph_stop_capture( _device: &Self::Device, ) -> Result<Self::GraphPrimitive, ExecutionError> { ... } unsafe fn graph_replay( _device: &Self::Device, _graph: &Self::GraphPrimitive, ) -> Result<(), ExecutionError> { ... } fn staging<'a, Iter>(_data: Iter, _device: &Self::Device) where Iter: Iterator<Item = &'a mut TensorData> { ... } fn supports_dtype(device: &Self::Device, dtype: DType) -> bool { ... }
}
Expand description

This trait defines all types and functions needed for a backend to be used with burn.

§Design

This trait aims to be as unopinionated as possible and allows implementations to define their own types and patterns. Therefore, there are few pre-defined abstractions baked into this trait.

Backends must define their own tensor types for each data type: float, int, and bool. Since we minimize assumptions, we chose to separate these types, as they are used in different contexts. However, some backends may have a generic tensor type that is used for all data types.

§Eager Mode

Because burn supports dynamic graphs, the backend trait is designed around kernel implementations that can be called without any mutable context or graph. This may not be ideal for backends that want to configure their computational graphs and execute them multiple times.

To implement this kind of backend, channels could be used to communicate with a backend server thread to build the computation graphs and re-execute the ones that are repeated, with some form of cache. Once that pattern has matured, a graph mode backend trait could be extracted from it, allowing other backends of the same kind to be quickly integrated with burn. This pattern could also be used to create an operation fusion trait, which allows backends to define what kind of graph structures can be fused into one operation.

§Multi-Threaded

Backend tensor types are all Clone + Send, which allows them to be safely sent between threads. It is recommended to wrap tensors with Arc, which avoids copying the tensor’s buffer. Note that it is still possible to mutate and reuse tensors’ buffer without locking; see the next section on the Mutable API.

§Mutable API

There is no mutable or inplace operation API to implement, but that does not mean that backends cannot support them. Using try_unwrap and get_mut allows backends to have access to an owned or mutable reference to their tensor buffer data structure if the tensor is not shared. In that case, backends can dispatch to their owned inplace operations for better performance.

§Documentation

Most of the documentation for each function can be found on the user API Tensor struct in the burn-tensor crate. For modules, public functions are often created, which can be used by burn-core modules.

Required Methods§

fn name(device: &Self::Device) -> String

Name of the backend.

fn seed(device: &Self::Device, seed: u64)

Seeds the backend on the specified device.

There is no guarantee that only the specified device will be seeded, but it is guaranteed that at least the specified device will be seeded.

In all cases, this should ensure deterministic execution for a single-threaded program.

fn flush(_device: &Self::Device) -> Result<(), ExecutionError>

Flush any pending operation of the backend.

§Errors

Returns an ExecutionError when the pending operations cannot be dispatched, e.g. on a device that is poisoned.

fn dtype_usage(device: &Self::Device, dtype: DType) -> EnumSet<DTypeUsage>

Returns the DTypeUsageSet for the given DType on the specified device.

fn device_count(type_id: u16) -> usize

Returns the number of devices available on this backend. device is a reference device used to determine the underlying backend that should be queried. A CUDA device will return all devices available to CUDA, a Vulkan device will return all devices available to Vulkan, etc.

Provided Methods§

fn ad_enabled(_device: &Self::Device) -> bool

If autodiff is enabled.

fn memory_persistent_allocations<Output, Input, Func>( device: &Self::Device, input: Input, func: Func, ) -> Output
where Output: Send, Input: Send, Func: Fn(Input) -> Output + Send,

Sets the current allocation mode to persistent.

fn memory_cleanup(device: &Self::Device)

Manually triggers a memory cleanup on the given device.

fn memory_pool_report(device: &Self::Device) -> Option<Vec<SlicedPoolReport>>

The dynamic pools’ measured state, in the order allocations are routed through them. None on a backend that does not report one, or whose stream has failed.

One entry per pool that carves pages, the pools a growth left behind last. Pools of other kinds — an allocation’s own page, the metadata churn — are left out.

fn memory_pool_usage(device: &Self::Device) -> Option<MemoryPoolUsage>

The device allocator’s current state. None on a backend that does not report one, or whose stream has failed.

fn sync(_device: &Self::Device) -> Result<(), ExecutionError>

Sync the backend, ensure that all computation are finished.

fn profile<O>( device: &Self::Device, options: ProfileOptions, func: impl FnOnce() -> O + Send, ) -> Result<(O, ProfileDuration), ExecutionError>
where O: Send + 'static,

Measure how long the device spends on the work func puts on the calling stream, in device time.

The window opens where the stream is when the call is made and closes where the stream is when func returns: work the stream still owed from before falls in, and work a backend queues past the end (a batching backend’s last operations, unless options flush) falls out. Nothing is waited on — the ProfileDuration resolves later, when the device has stamped both ends — so windows nest without the inner ones being charged to the outer. Work on other streams is not kept out, and not counted. A window that nothing ran in reads as no time.

The default is profile_system_time: wall-clock time between two syncs, for a backend with no device clock to read. That one does wait, and an inner window’s syncs are charged to the outer.

§Errors

The device refused to open or close the window, or work inside it failed and took the measurement with it. func has run by then — a window that could not be opened does not cancel the work it was asked to measure — and its output is lost with the error, as it would be on the read that the failure surfaces on without a window.

A failure that the device only reports later is not here. The measurement is resolved after this returns, so anything the device learns in between — and everything a remote server reports, which travels back with the measurement rather than ahead of it — arrives as a window that resolves to no measurement, with the reason in the log. A caller that must distinguish “nothing ran” from “the server failed” cannot do it from the Result alone.

fn profile_start( _device: &Self::Device, ) -> Result<Option<ProfileToken>, ExecutionError>

Open a profiling window at the calling stream’s current position, to be closed with profile_end from the same stream.

For a caller that cannot bracket the work in a closure: a backend that forwards operations to be executed on another thread opens and closes the window from that thread, in order with the operations.

None from a backend that opens no windows and measures only with profile — the default — so the caller can bracket with profile_system_time instead.

fn profile_end( _device: &Self::Device, _token: ProfileToken, _options: ProfileOptions, ) -> Result<ProfileDuration, ExecutionError>

Close the window token at the calling stream’s current position.

When options flush, the work the backend still holds queued for the stream executes first, so it falls inside the window. A backend that forwards the close passes options along, so a queue further down the chain — a remote server’s fusion, say — is flushed too.

Errors on a backend whose profile_start hands out no token.

fn profile_abandon(device: &Self::Device, token: ProfileToken)

Drop the window token opened without measuring it, for a caller that will never reach profile_end.

An open window is not free, and the cost is not paid once: a backend holds a start event, keeps timestamp writes on, or retains command buffers for as long as one is open, and on wgpu every later pass keeps rewriting the live window’s end slot. So a window whose caller unwound between the two calls is abandoned rather than left, which is what profile_with_tokens does on the panic path.

Cannot fail and answers nothing: it is called while a panic is already unwinding, where there is nobody left to tell. The default closes the window and discards the measurement, which every backend can already do; one that can drop a window without recording an end does that instead.

fn graph_prepare(_device: &Self::Device) -> Result<(), ExecutionError>

Prepare device for an upcoming graph capture: route allocations into a stable pool so every buffer allocated before graph_stop_capture can be pinned. Call before the warmup run. No-op by default.

See burn_graph — the closure-based capture helper drives this whole sequence.

fn graph_start_capture(_device: &Self::Device) -> Result<(), ExecutionError>

Begin recording launches on device into a graph (see graph_stop_capture). Errors on backends without hardware graph support, so callers fall back to re-running.

fn graph_stop_capture( _device: &Self::Device, ) -> Result<Self::GraphPrimitive, ExecutionError>

Stop recording and return the captured graph, ready to graph_replay.

unsafe fn graph_replay( _device: &Self::Device, _graph: &Self::GraphPrimitive, ) -> Result<(), ExecutionError>

Replay a captured graph — one dispatch re-running the recorded launches against their original buffers.

§Safety

The replay dispatches raw device work against the exact buffers recorded at capture time, with nothing tracking whether those buffers are still valid. The caller must guarantee, for every tensor the captured closure read or wrote:

  • its buffer is still alive — no tensor referenced by the graph has been freed (and its memory possibly reallocated) since capture;
  • it is not concurrently read or written by work on another stream or thread while the replay executes;
  • input refreshes and output reads are issued on the stream the graph was captured on, so they order correctly against the replay.

fn staging<'a, Iter>(_data: Iter, _device: &Self::Device)
where Iter: Iterator<Item = &'a mut TensorData>,

Marks the given data as being used as a staging buffer for transfer between CPU and accelerators like GPUs.

The given data might be transferred to pinned memory or another format to improve data transfer speed.

fn supports_dtype(device: &Self::Device, dtype: DType) -> bool

Whether the type is fully supported by the specified device for general operations.

A type is considered supported if it can be used for the full suite of tensor operations, including storage, conversion, and basic arithmetic.

Returning false does not necessarily mean the device cannot handle the type at all. For instance, a device might support a type only for specialized hardware acceleration (e.g., matrix multiplication) but lack general arithmetic support. Such types should return false here as they are not globally supported.

Dyn Compatibility§

This trait is not dyn compatible.

In older versions of Rust, dyn compatibility was called "object safety".

Implementors§

§

impl Backend for Dispatch

§

impl Backend for Flex

§

impl<B, C> Backend for Autodiff<B, C>

§

impl<B> Backend for Fusion<B>
where B: FusionBackend,