Skip to main content

backend_extension

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.