Attribute Macro backend_extension
#[backend_extension]Expand description
Generates the Dispatch implementation for a backend extension trait.
The backend comes from one routing tensor, preferring a float. The autodiff contexts of all tensor-bearing inputs are merged; disabled inputs act as constants, while enabled inputs must share a gradient-checkpointing strategy.
Struct and enum inputs derive ExtensionType and use #[extension_type] on the corresponding
method argument. Autodiff support for custom operations requires a handwritten implementation of the extension trait
for Autodiff<B, C>; this macro does not generate backward passes.
§Backend selection
Execution backend selectors include Cube, Flex, NdArray, LibTorch, and Remote.
Cube covers the CubeCL runtimes, including WGPU, CUDA, ROCm, and CPU; Wgpu and Cuda
are not selectors. Gate a selector with, for example, Cube: cfg(feature = "wgpu");
this condition refers to the consuming crate’s features. Enable Burn’s extension feature
and the backend features you target.
Implement the extension trait on each selected backend. Add Autodiff to generate routing
to your Autodiff<B, C> implementation. A default trait body composing differentiable
operations can supply that implementation; the macro does not derive a custom backward pass.
Calling an extension on an unlisted runtime backend panics.
Expose a high-level Tensor<D> wrapper by calling the generated Dispatch implementation
with Tensor::into_dispatch and wrapping its output with Tensor::from_dispatch.
§Fusion
Add Fusion (optionally Fusion: cfg(...)) to also generate a lazy implementation
for Fusion<B>. Enable Burn’s fusion feature and choose a behavior for each method:
#[fusion(dtype = lhs, shape = lhs)]: describe a single tensor output using field expressions.#[fusion(meta = callable)]: compute output metadata now and defer the inner backend call.#[fusion(default)]: inherit the trait’s existing default body.
Choose exactly one form. Output metadata is computed before registration; execution is deferred.
§Field expressions
Tensor names refer to DType values in dtype and &Shape values in shape.
Extension arguments refer to borrowed extension metadata; ordinary arguments are borrowed.
Both fields are required.
A bare operand copies its shape. Function calls such as shape = output_shape(lhs, rhs)
and inline blocks return an owned Shape. Use a block for an inline calculation;
standalone closures are not invoked automatically.
#[backend_extension(Cube, Fusion)]
pub trait MyExtension: Backend {
#[fusion(dtype = input, shape = {
let mut shape = input.clone();
shape.swap(0, 1);
shape
})]
fn transpose_2d(input: FloatTensor<Self>) -> FloatTensor<Self>;
}§Complete metadata
meta accepts a function path or closure receiving borrowed burn::backend::fusion::custom::TensorSpec
values, extension metadata, and ordinary arguments in declaration order. Its result mirrors the
outputs: specs for tensors, tuples for tuples, and metadata types generated by ExtensionType with
#[extension_type(fusion)] for structs and enums. Enum metadata selects the output variant.
For example, #[fusion(meta = |x| (x.clone(), x.clone()))] describes two tensors matching x,
while #[fusion(meta = |cache| cache.clone())] preserves a structured input’s layout.
Tuple elements must themselves be tensors, derived extension values, or tuples of those types.
Plain scalar returns and tuples such as (FloatTensor<Self>, u32) are unsupported; put ordinary
output fields in a struct or enum deriving ExtensionType with Fusion enabled.
For struct and enum outputs, the callback supplies the actual return values for non-tensor
fields. Fusion returns these values without waiting for execution and discards the backend’s
values for those fields without comparison. They must equal what a direct backend call returns.
If metadata supplies count = 7 but the backend returns count = 8, enabling Fusion changes the
result. That is an incorrect extension implementation, and the generated code does not detect it.
Tensor shapes, dtypes, and enum variants must also agree with direct backend execution.
For example, a derived output struct with a tensor and a count: usize field can return the
tensor and its element count:
#[fusion(meta = |input| CountedMetadata {
tensor: input.clone(),
count: input.shape.num_elements(),
})]
fn counted(input: FloatTensor<Self>) -> Counted<Self>;The backend must return the same count.
A count of nonzero elements depends on tensor contents and cannot be computed from TensorSpec.
Return content-dependent values as tensors, or write a Fusion implementation that waits for
computation to finish before returning them.
§Optimizer integration
The operation ID defaults to the method name; use id = "custom_matmul" to override it.
Integer parameters (8–64 bits, usize, isize), f32, f64, and bool are exposed to custom
optimizers automatically, in declaration order. These are host values, not tensor contents.
Other ordinary arguments and extension fields are captured for execution only.
Mark aliases or types convertible to burn::backend::Scalar with #[fusion(scalar)].
For custom encodings, annotate the parameter with #[fusion(scalar = strategy.to_code())].
The expression returns a value convertible to Scalar and runs before inputs are consumed.
The original argument is still passed to the backend; marked and inferred scalars share
declaration order.
§Requirements
Methods using field expressions or meta must have:
- A synchronous, non-generic signature.
- At least one input tensor, directly or inside an owned extension value. Primitive inputs may be borrowed.
- Owned ordinary arguments implementing
Clone + Send + Sync + 'static. - Output metadata computable without tensor readback.
Debug builds check dtype categories and tensor devices, and validate output dtypes and devices. During execution, output enum variants are checked in all builds before any output handles are published. Output shapes and ordinary field values are never compared with backend results, even in debug builds.
Custom kernels are opaque unless a custom optimizer recognizes their IR.
For other signatures, use an existing default body or omit Fusion from #[backend_extension]
and implement the trait for Fusion<B> manually.