Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

Tensor

The public tensor type is Tensor<const D: usize, K = Float>: D is its rank and K is Float, Int, or Bool. Backend and element precision are runtime properties, selected through Device and DType. Models and tensor functions do not have a backend generic parameter.

From backend generics to runtime selection

Previously, Tensor<B, D, K> carried the backend in its Rust type. Model fields and functions propagated B: Backend, and selecting another backend instantiated that generic code for another backend type. The current API keeps rank and kind in the type while moving backend selection to runtime values. The same Tensor<2> type can represent a tensor on Flex, CUDA, or WGPU.

Cargo features determine which backends are compiled into an application. A Device selects the compute resource and supplies creation settings; existing tensors carry the information needed to route subsequent operations. Multiple enabled backends can be used side by side, with the same model code targeting each of them. There is no single process-wide backend that every tensor must use. Combining tensors still requires compatible devices; transfers are explicit.

The Backend trait remains the implementation contract below dispatch. Concrete backends and decorators such as Autodiff<B, C> and Fusion<B> still compose using backend types. The change moves those types out of ordinary application signatures and into the implementation layers.

Tensor Operations

Operations follow this path:

Tensor<D, K> → BridgeTensor → DispatchTensor → backend primitive

The layers have distinct responsibilities:

  • burn-tensor/src/tensor/api defines the rank- and kind-checked API, shape checks, and user documentation. base.rs covers common operations, numeric.rs covers numeric kinds, and float.rs, int.rs, and bool.rs contain kind-specific operations. Activations and neural-network operations also have function APIs in tensor/activation and tensor/module.rs.
  • burn-tensor/src/bridge stores tensor primitives opaquely and implements kind-specific forwarding. Thin generic methods call non-generic helpers taking BridgeTensor, so downstream monomorphization does not repeatedly resolve the backend implementation types. Follow the *_impl pattern described in the burn-tensor crate documentation when adding operations.
  • burn-dispatch selects a backend from runtime tensors and devices. It also routes autodiff and checkpointing contexts. Built-in forwarding uses #[backend_dispatch]; extensions use #[backend_extension].
  • burn-backend defines BackendTypes, Backend, and operation traits implemented by concrete backends and decorators. These low-level APIs still use backend generics and tensor primitive aliases.

A new kernel operation may need changes at every layer. An operation expressed entirely in terms of existing tensor operations may need only a public composition. See Adding a New Operation.

Why the bridge is opaque

Runtime selection and type erasure solve different problems. DispatchTensor represents runtime selection with an enum containing the concrete primitives of the enabled backends, plus autodiff context. Putting that enum directly into the public Tensor would still expose its nested backend types to the compiler when compiling downstream code.

Instead, Tensor stores a BridgeTensor. Its private BridgeTensorVariant distinguishes float, integer, boolean, and quantized float values, each holding a dispatch tensor. That variant lives inside aligned opaque storage generated by burn_std::obfuscate!. The public Device similarly hides its DispatchDevice representation. These wrappers erase the concrete field types at the public boundary; unwrapping them does not itself transfer tensor data between devices.

Rust specializes generic tensor methods for the ranks and kinds used by an application. To keep that work small, these methods pass opaque BridgeTensor handles to non-generic bridge methods or *_impl helper functions. Those functions contain the calls into dispatch and are compiled once in burn-tensor. This lets application code specialize the thin public methods without repeatedly compiling the dispatch logic or resolving the concrete backend types behind it.

When adding operations, preserve both boundaries: keep backend primitives behind the opaque representation and outline dispatch-facing work into non-generic functions. See the burn-tensor contributor notes for the helper pattern.

Following an addition through the stack

For an ordinary floating-point lhs + rhs, the path is:

  1. The operator implementation calls Tensor::add in tensor/api/numeric.rs. The public method checks shape compatibility and forwards the two bridge handles to the numeric operation for their kind.
  2. The Float implementation in bridge/ops/float.rs unwraps the bridge values and calls Dispatch::float_add. This layer also handles the distinct paths for quantized float operands.
  3. The implementation in burn-dispatch/src/ops/tensor.rs uses #[backend_dispatch] to generate routing to the selected backend's B::float_add. Routing uses the runtime tensor variants and autodiff context; it does not pick a new device for each operation. Creation operations instead dispatch from their supplied device.
  4. The selected backend or decorator handles the primitive operation. Autodiff can record backward steps and fusion can defer execution while collecting operations. Reaching this layer does not necessarily launch a kernel immediately.
  5. The returned primitive is wrapped into a dispatch tensor, then a bridge tensor, then the public Tensor<D, Float> result.

CUDA, ROCm, WGPU, and the CubeCL CPU runtime share the Cube dispatch variant. Within that backend, the tensor's device and runtime client select the runtime. Other backends, such as Flex, have their own dispatch variants. See Backend for the primitive contract and decorators.

Static and runtime checks

Rank and tensor kind remain compile-time API properties. Shapes, dtypes, device compatibility, and autodiff association are runtime properties. In particular, a public Tensor<D> type no longer proves B: AutodiffBackend; when the autodiff feature is enabled, backward() validates graph participation at runtime. Backend extension code can still use backend trait bounds below the public boundary.

Ownership and autodiff

Tensor handles can be cloned and sent across threads. Operations usually take owned tensors, which lets a backend reuse storage when no other handle references it. Cloning a handle does not imply a copy of its allocation, nor does it copy an autodiff tape.

Devices provide defaults for newly created tensors; each tensor retains its own autodiff context. to_device preserves the source context and records a differentiable transfer for tracked floats. Use autodiff or without_autodiff to change association explicitly. Graph participation, gradient retention, and whether recorded backward steps are still available are distinct concepts; see the Burn Book autodiff chapter.