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

Overview

Welcome to The Burn Contributor's Book 👋

This book will help you get acquainted with the internals of the Burn deep learning framework and provide some detailed guidance on how to contribute to the project. Before opening a PR, please read the Contributing Guidelines.

We have crafted some sections for you:

  • Getting Started: Much like the Burn Book which targets users, we'll start with the fundamentals, guiding you through tasks like setting up the development environment, running tests, and what you should check prior to each commit.

  • Project Architecture: This section will give you an in-depth look at the architecture of Burn.

  • Guides: We provide some guides on how to do specific tasks, such as adding a new operations to Burn.

  • Frequently Encountered Issues: If you are running into an issue that has you stumped, this is the section to check out prior to asking on the Discord. It's a collection of errors encountered by contributors, what caused them, and how they were resolved.

As this book is geared towards contributors and not towards users of Burn, we'll assume you have a good understanding of software development, but will make efforts to explain anything outside of that scope, or at least provide links to resources that explain it better than we can.

How to read this book

Throughout this book, we maintain the following structure.

Linking

When referring to structures or functions within codebase, we provide permalinks to the lines in specific commits, and indicate them by the relative path of their parent file from the project root. For example this is a reference to the Tensor struct in crates/burn-tensor/src/tensor/api/base.rs

When some reference information is useful but is beyond the scope of contributing to Burn, we provide that information in a footnote. To build on the previous example, the Tensor mentioned is what's referred to as a newtype struct1.

Direct hyperlinks are for tools and resources that are not part of the Burn project, but are useful for contributing to it. For example, when working on implementing an operation for autodiff, it can be useful to use symbolab to calculate the left and right partial derivatives.


  1. For more information on newtype please refer to the Advanced Types chapter of the Rust Book ↩

Getting Started

This section is for setting up the environment and how to do basic development tasks such as running tests and checking your code before committing. If you need help with the process or run into issues, feel free to ask on the Discord server in the Development channels.

Setting up the environment

Depending on what part of the project you plan on contributing to, there are a couple of tools to install and commands to be familiar with.

General

During development, these commands can automatically address common formatting and lint issues:

  1. cargo fmt --all runs rustfmt on all files in the project.
  2. cargo clippy --fix runs Clippy and applies fixes for supported lint findings. It requires a clean Git state unless you pass --allow-dirty.

Before submitting a PR, run:

cargo run-checks

This is the common local validation command. It runs, in order:

  • formatting checks;
  • typo checks;
  • a dependency audit;
  • Clippy across the workspace;
  • a quick host no-std compilation check; and
  • backend tests in release mode using the Flex backend.

If your changes target another backend, override the default:

cargo run-checks --backend <backend>

The command is intended as a fast common baseline rather than the entire CI test matrix. You should also run tests relevant to the crates changed. CI runs the broader workspace, documentation, platform, feature, and backend combinations.

Want more detailed macro error diagnostics? This is especially useful for debugging tensor-related tests:

RUSTC_BOOTSTRAP=1 RUSTFLAGS="-Zmacro-backtrace" cargo run-checks

Updating the burn semver version

If for some reason you need to bump for the next version (though that should probably be left to the maintainers), edit the semantic version number in burn/Cargo.toml, and then run cargo update to update the lock file.

Contributing to either the Burn Book or Contributor Book

Both the Burn Book and the Contributor Book are built with mdBook. From the repository root, use the xtask commands to build them. These commands install mdBook if needed:

cargo xtask books burn build
cargo xtask books contributor build

To serve a book locally and open it in your browser:

cargo xtask books burn open
cargo xtask books contributor open

Run the command for the book you want to preview. Add --port 3000 to choose a port; otherwise, xtask selects one automatically.

Alternatively, if you want to install mdbook directly, run the following command1:

cargo install mdbook

For documentation-only changes, you can run cargo xtask check typos to check for misspellings without running the full local validation. This installs typos when needed. To apply suggested corrections to a book, run typos -w /path/to/book.


  1. You might also want to install cargo-update to easily keep your tools up to date, though it is in no way required. ↩

Configuring your editor

These steps are not required, and most of this isn't specific to Burn, but it's definitely helpful if you haven't already done it.

VSCode

Install the following extensions:

Setting up the Debugger

To use the debugger, follow these steps:

  1. Open Command Palette with Ctrl+Shift+P or F1 and type LLDB: Generate Launch Configurations from Cargo.toml then select it, this will generate a file that should be saved as .vscode/launch.json.
  2. Select the configuration from the "run and debug" side panel, then select the target from the list. Since this repo has debug = 0 in the root Cargo.toml to speed up compilation, you need replace it with debug = true in the root Cargo.toml when using a debugger and breakpoints with launch.json settings.
  3. Now you can enable breakpoints on code through IDE then start debugging the library/binary you want, like in the following example:

debug-options

If you're creating a new library or binary, keep in mind to repeat step 1 to always keep a fresh list of targets.

Have another editor? Open a PR!

Testing

Tensor operations

Shared tensor tests live in crates/burn-backend-tests/tests/tensor. Register new test modules in the corresponding mod.rs or tests/common/tensor.rs. The test executables reuse those modules across precisions; backend features select which runtime to test. Backend-specific implementation tests also live alongside the backend code.

Run the relevant shared suites using the backend aliases defined in crates/burn-backend-tests/.cargo/config.toml. Start from the repository root and change into the crate directory so Cargo discovers its aliases:

cd crates/burn-backend-tests
cargo test-flex --test tensor
cargo test-flex --test autodiff

These aliases run in release mode and select the backend features explicitly. Omit --test to run all test targets for that configuration, or append a test-name filter to narrow the run:

cargo test-flex
cargo test-flex --test tensor matmul

Choose the alias for the backend you are changing, such as cargo test-cuda, cargo test-vulkan, or cargo test-metal. Aliases for backends that support fusion enable it by default; their -no-fusion variants test without fusion. For example, run both configurations when changing CUDA operations or fusion behavior:

cargo test-cuda --test tensor
cargo test-cuda-no-fusion --test tensor

From the repository root, use cargo run-checks for the repository validation workflow. It defaults to Flex; select another backend with cargo run-checks --backend <backend> when working on backend-specific code.

Autodiff

Shared backward tests live in crates/burn-backend-tests/tests/autodiff and are registered through tests/common/autodiff.rs. Graph engine unit tests also live in burn-autodiff. For operations with multiple differentiable inputs, verify every input gradient.

Choose small inputs whose derivatives can be calculated independently. Create source leaves on an autodiff device and call require_grad() before the forward pass. Retrieve their gradients from the result of backward(). Check broadcasting and untracked inputs where relevant; a numerically correct forward pass does not establish a correct backward implementation.

You can also use PyTorch as a reference implementation to obtain expected outputs and gradients for Burn tests. For example, this small broadcasting case checks gradients for both operands:

import torch

x = torch.tensor([[1., 2.], [3., 4.]], requires_grad=True)
y = torch.tensor([5., 6.], requires_grad=True)
output = x * y
output.sum().backward()

print(output.detach().tolist())  # [[5.0, 12.0], [15.0, 24.0]]
print(x.grad.tolist())           # [[5.0, 6.0], [5.0, 6.0]]
print(y.grad.tolist())           # [4.0, 6.0]

Use the same inputs, operation parameters, and reduction in the Burn test, and record these values as expected data so the test does not depend on PyTorch. Here, the gradient for y sums contributions over the broadcast dimension. For other operations, check that the reference uses matching semantics and dtypes, and compare with an appropriate tolerance as described below.

Precision

Shared suites define FloatElem and IntElem aliases for each test executable and configure the device defaults before creating tensors. They are not associated types of a TestBackend. Use the aliases in expected data and use .elem() when a literal needs conversion.

For approximate floating-point comparisons, follow nearby tests:

actual.into_data().assert_approx_eq::<FloatElem>(&expected, Tolerance::default());

For integers, use IntElem and skip cases whose inputs cannot be represented by the selected dtype. Exercise additional precision targets when the change depends on dtype or numerical stability.

Project Architecture

This section documents most major architectural decisions with the reasoning behind them.

Sections

Module

Modules organize parameters into structures that can be optimized, saved, and loaded. #[derive(Module)] generates parameter traversal and training/validation conversions. A module does not force the declaration of the forward pass, leaving it up to the implementer to decide how it should be defined.

Configuration describes a module's structure and hyperparameters; records store its parameters separately.

Parameters and traversal

Param<T> gives a value an identity and supports lazy initialization. Tensor parameters use that identity to associate optimizer state and gradients. Param<Flag> represents module-owned control state, such as whether dropout or batch normalization behaves as during training.

Module::visit inspects parameters; Module::map transforms them. Visitors and mappers have hooks for float, integer, and boolean tensors, control flags, and module paths. Reparameterizations such as LoRA have nested parameters that participate in these traversals. param.base() reads the stored base; param.val() materializes the effective value, including a reparameterization.

Training and validation

Module includes the valid(&self) and train(self) transition hooks; the derive generates both alongside traversal. Training and validation use the same module type, and a module can contain parameters with different runtime autodiff contexts.

  • valid(&self) creates a validation snapshot with autodiff and training flags disabled. It keeps configured trainability and flag settings, folds reparameterizations into parameter values, and removes checkpointing strategies with the autodiff association.
  • train(self) enables autodiff and applies configured trainability and flags. It does not undo explicit freezing, reconstruct folded adapters, or restore discarded checkpointing strategies.
  • no_grad() persistently disables parameter gradients while leaving control flags unchanged.
  • freeze() and unfreeze() configure both gradients and flags; group variants target subtrees.

Keep the training module and use its valid() snapshot for validation. Inspect individual tensors with is_autodiff(), is_tracked(), and is_require_grad() rather than inferring training state from the trait or module.devices(). Device equality ignores autodiff settings, and the latter method deduplicates compute resources.

Optimization

Optimizer updates one tensor at a time from its gradient and optional state. State<D> implements Clone and RecordState, allowing tensors and scalars to be serialized independently of backend types.

ModuleOptimizer wraps these optimizers and manages module traversal, parameter groups, gradient lookup, device migration, and per-parameter state. Parameter groups can use different optimizers.

An update proceeds as follows:

  1. Run the model and call loss.backward().
  2. Convert the tensor gradients with GradientsParams::from_grads(grads, &model).
  3. Call optimizer.step(learning_rate, model, grads) to obtain the updated module.

The optimizer performs updates outside the autodiff graph and restores parameter trainability and checkpointing strategy afterward. A transferred tracked parameter is an intermediate, so use model.fork(&device) when the destination module should have independently optimizable leaves. Both to_device and fork preserve source autodiff context; use train() to enable it explicitly.

See the serialization chapter for ModuleRecord and OptimizerRecord.

Serialization

An important aspect of a deep learning framework is the ability to save and load training state to and from disk. Burn serializes its records with the burnpack format, a compact binary container implemented by the burn-pack crate.

Constraints

  1. Users should be able to add any field to a module, even fields that are not serializable.

    This can include constants, database connections, other module references, or any other information. Only the parameters (tensors) should be serialized; the structure of the module itself is encapsulated by its configuration (hyperparameters).

  2. Records should be decoupled from the backend in use.

    A record holds plain tensor data (TensorData), so weights saved with one backend can be loaded on another. Parameter initialization is lazy, so loading a record does not require eagerly materializing the module first.

  3. The format should be fast to load and embeddable.

    Tensor data is stored contiguously and aligned, so it can be read back with zero-copy / memory-mapped loading, and a record can be saved straight to bytes for no-std environments.

The burnpack format

The burn-pack crate is intentionally minimal and tensor-library-agnostic: it knows how to read and write the container format but has no notion of Burn modules. A burnpack file has three parts:

┌────────────────────────────────────────────────────────────┐
│ Header (fixed size)                                          │
│   magic "BURN", format version, metadata byte length         │
├────────────────────────────────────────────────────────────┤
│ Metadata (CBOR)                                              │
│   tensors : map<name, descriptor>                            │
│     dtype, shape, data_offsets, optional param_id            │
│   scalars  : map<name, typed scalar>                         │
│   metadata : map<string, string>  user key/value pairs       │
├────────────────────────────────────────────────────────────┤
│ Tensor data section                                          │
│   each tensor's bytes start on a 256-byte boundary so the    │
│   data can be sliced zero-copy / memory-mapped from a file   │
└────────────────────────────────────────────────────────────┘

All multi-byte integers are little-endian. Tensor entries carry an optional param_id used to preserve a parameter's identity across save/load. Besides tensors, a pack can store named typed scalars (integers, floats, booleans), which the optimizer and learning rate scheduler records use to persist their non-tensor state.

A pack is read back with burn_pack::Reader, which yields burn_pack::Tensor entries plus a scalar map. Tensor bytes stay lazy: a file-backed reader only touches the disk when an entry's data is actually used.

A pack is written with burn_pack::Writer, which is symmetric about this, because laziness is a property of Tensor rather than of the writer's input. Tensor::new carries bytes that are already resident; Tensor::deferred carries a byte length plus a provider that yields the data on demand. The writer reads only the metadata while planning, to compute every descriptor and offset before any I/O, then calls each provider once in write order, dropping each tensor's bytes before requesting the next. burn-store collects a module's parameters as deferred tensors, so each device readback waits until the writer reaches that tensor. Paired with Writer::write_to_file, which streams to disk, saving a large module costs one tensor of host memory at a time rather than the whole set. The in-memory sinks (into_bytes, write_into) still build the container as a whole.

Because the offset table is committed from the declared length before the bytes exist, a length that turns out to be wrong would misplace every later tensor. Two checks prevent that: planning rejects a length that disagrees with the tensor's own shape and dtype, and the write pass rejects bytes whose length differs from what was reserved. Quantized tensors are exempt from the first, since their packed values and inline scales are not a product of shape and dtype; that exception is why Tensor::deferred takes an explicit length rather than deriving one.

Writer::write_to_file builds the container in a scratch file beside the destination and renames it into place only once it is complete, so a save that fails partway leaves whatever was already at that path untouched. That covers disk failure (a full disk, a quota, an I/O error) for every caller, and provider failure for deferred tensors, whose bytes are produced mid-write. Writer::write_to_file_in_place truncates and rewrites the destination instead: it skips the fsync and the transient second copy, but a failed save leaves a truncated file.

The three record types

Higher layers bridge their state to and from burnpack through three record types. Each one can be serialized to a file (save / load, appending the .bpk extension when the path has none) or to an in-memory byte buffer (into_bytes / from_bytes).

ModuleRecord (burn-core, burn::store)

Holds a module's parameters: a flat list of (path, ParamId, TensorData) entries keyed by module path. It is produced and applied through the Module trait itself rather than a separate codegen type:

  • module.into_record() walks the module with a ModuleVisitor (the Collector), recording each float/int/bool parameter under its dotted path.
  • module.load_record(record) (or the fallible try_load_record) walks the module with a ModuleMapper that looks each parameter up by path and loads the matching tensor.

Load-time behavior is configured with builder methods on the record, ignored when saving:

  • allow_partial(bool) — tolerate module parameters absent from the record.
  • validate(bool) — toggle shape-mismatch / missing-tensor validation.
  • with_dtype_policy(..) / cast_to_module_dtype() — choose whether a parameter adopts the record's dtype (DTypePolicy::FromRecord, the default) or casts the data to the module parameter's current dtype (DTypePolicy::CastToModule).

The save-side dtype is not configurable: the record stores whatever dtype the module currently holds. The dtype applied on load is controlled by the record's DTypePolicy (.cast_to_module_dtype() / .with_dtype_policy(..)).

This module in burn-core is intentionally tiny — no filtering, key remapping, or adapters. The richer snapshot/import tooling (filtering, key remapping, PyTorch/SafeTensors cross-framework stores) lives in the burn-store crate.

OptimizerRecord (burn-optim)

Holds an optimizer's per-parameter state. Unlike a module record (keyed by module path), it is keyed per parameter: each parameter's state is decomposed into tensors named "{param_id}.{field}" plus a few typed scalar entries kept in the burnpack scalar map (including a __rank scalar so the state can be reconstructed without inferring rank from tensor shapes).

  • optimizer.to_record() flattens each parameter's DynState into tensors and scalars.
  • optimizer.load_record(record) reconstructs the states (no device argument: tensors load on the default device and migrate to each parameter's device on the next step).

LrSchedulerRecord (burn-optim)

Holds a learning rate scheduler's state, which is just a handful of scalars (step counters, current learning rate) and no tensors. Produced/applied through the LrScheduler trait's to_record() / load_record(). Composed schedulers nest their children's records under an index prefix (with_record / record).

Checkpointing in burn-train

During training, burn-train defines a Checkpoint trait (save(path) / load(path)) implemented for all three record types — ModuleRecord, OptimizerRecord, LrSchedulerRecord — and for () (a stateless no-op). The Checkpointer<R: Checkpoint> trait drives periodic saves; the FileCheckpointer writes each record to {name}-{epoch}.bpk under the experiment directory. This is how the model, optimizer, and scheduler are persisted and restored across epochs.

Notes

  • There is no Recorder, PrecisionSettings, Module::Record associated type, or #[derive(Record)] any more. All of those were part of the previous serde-based record system, which has been removed.
  • Cross-framework import/export (PyTorch .pt, SafeTensors) still lives in burn-store (PytorchStore, SafetensorsStore, BurnpackStore).

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.

Backend

The low-level Backend trait defines the implementation contract for tensor operations. Its BackendTypes supertrait supplies device, float, integer, boolean, quantized tensor, and captured graph primitive types. Backend implementations and decorators remain generic where useful; application tensors and modules select them at runtime through Device.

See Tensor for the Tensor → Bridge → Dispatch → Backend path.

Element types

Precision is a runtime property. BackendTypes no longer has associated FloatElem or IntElem types. Tensor primitives expose their dtype through metadata, and operations accept dtype or scalar information as needed.

Users can configure default dtypes with Device::configure(DeviceConfig::default()...) before creating tensors, request an explicit dtype when creating a tensor, or call tensor.cast(...). Device settings are initialized once; creating the first tensor uses the backend's defaults if configuration has not happened earlier. A backend reports which dtypes it supports through supports_dtype. See the Burn Book's Device Settings section for configuration and per-tensor dtype examples.

Do not assume every backend supports every precision. Backend kernels should handle the dtypes they advertise, and shared tests use the test target's FloatElem and IntElem aliases.

Operations

Backend operations are associated functions taking and returning primitives. Backends may enqueue work asynchronously and reuse uniquely owned storage. The public tensor API handles user-facing shape checks, while the bridge and dispatch layers route to the appropriate primitive operations.

CubeCL runtimes share burn_cubecl::CubeBackend; the device selects CUDA, ROCm, WGPU, or CPU. burn_cubecl::Cube wraps that backend with Fusion when the fusion feature is enabled. Runtime facade names such as Cuda and Wgpu are aliases of Cube, not distinct extension selectors.

Runtime dispatch uses a fixed backend catalog defined in crates/burn-backend-extension/src/catalog.rs, shared by generated and handwritten routing. #[backend_extension] extends the operations available on supported backends; it does not provide a mechanism for registering additional backends.

Autodiff

burn_autodiff::Autodiff<B, C> decorates a backend with first-order reverse-mode differentiation, where C is a checkpoint strategy. The low-level AutodiffBackend trait defines backward and gradient access. At the public API, autodiff is a runtime context carried by tensors and supplied as a creation default by devices.

Enabling autodiff does not automatically track every input. Source leaves use require_grad(); operations depending on tracked inputs record backward steps. Backward consumes reachable steps, including steps shared by cloned handles. Repeated backward through consumed intermediates is rejected; parameter leaves can be reused in fresh forwards.

#[backend_extension(Autodiff, Cube)] generates runtime routing for an extension. The author still supplies its implementation for Autodiff<B, C> (a composition of differentiable operations or a custom backward pass). See the extension guide.

When writing a custom backward pass, obtain input NodeGuards with AutodiffTensor::node() or into_parts() and pass them to Backward::prepare. They must stay alive through child-step registration, but must not be stored in backward or checkpoint state. Access the primitive with primitive() or into_parts(); the old public fields are private.

Guides for Contributors

The following guides are meant to help contributors accomplish specific tasks, such as adding new operations to Burn.

Adding a New Operation to Burn

Choosing where the operation belongs

First consider the operation's intended users and scope:

  • General-purpose tensor operations shared across domains may belong in Burn's core tensor API.
  • Domain-specific operations belong in the corresponding extension crate, such as burn-vision, burn-linalg, or burn-signal. These crates keep specialized APIs and their backend implementations together, with capabilities enabled as needed.
  • Application-specific or experimental operations can live in your application or a separate crate. They do not need to be upstreamed to be used with Burn.

Then decide how to implement the operation. A composition of existing tensor operations can be exposed as a function or extension trait in any of these locations without changing the backend contract. If custom kernels are needed, a backend extension lets you define the operation and its backend implementations in an extension crate, including one maintained outside Burn.

The sections below describe adding an operation to the core tensor API and, when a new primitive is needed, its backend contract and routing. For domain or external extensions, follow the backend extension guide and the conventions of the crate that owns the operation.

Public tensor API and bridge

Add the method to the appropriate file in crates/burn-tensor/src/tensor/api: base.rs for common operations, numeric.rs for shared numeric operations, or float.rs, int.rs, and bool.rs for kind-specific operations. Neural-network operations and activations also have function APIs. Document shapes, broadcasting, dtype behavior, examples, and runtime preconditions. Add necessary validation in the tensor checks.

The public type is Tensor<D, K>, with an opaque BridgeTensor primitive. Keep generic method bodies thin: route through a non-generic *_impl helper where needed and the corresponding kind operations in crates/burn-tensor/src/bridge/ops. Do not expose backend types in ordinary public method signatures. See Tensor Architecture.

Backend contract and dispatch

Define the primitive operation in the relevant trait under crates/burn-backend/src/backend/ops. Shared names are prefixed by kind, such as float_powf and int_powf. A default implementation may compose existing primitive operations where appropriate. Dtypes are runtime values; there are no backend-associated float or integer element types.

Add forwarding in crates/burn-dispatch/src/ops. Built-in implementations use #[backend_dispatch], which selects the runtime backend and handles autodiff contexts. Operations needing custom routing can use #[backend_dispatch(skip)] and an explicit implementation; follow an existing operation with matching inputs and outputs.

For a complete existing path, trace Tensor::powf through the numeric bridge, Dispatch::float_powf, and FloatTensorOps::float_powf.

Concrete backends and decorators

Implement the operation on the supported concrete backends. For CubeCL kernels, the implementation is on CubeBackend; runtime-specific execution is selected by its device. Handle supported dtypes and layouts, including non-contiguous inputs where the operation permits them.

A new primitive also needs the applicable decorator and graph paths:

  • Autodiff: implement the derivative in crates/burn-autodiff/src/ops. Follow neighboring operations for Backward, saved state, checkpointing, broadcasting reductions, and tracked versus untracked inputs. The forward pass must save only the state needed by the backward computation.
  • IR: add the representation under crates/burn-ir/src/operation.rs when the operation is recorded or transmitted, including shape and scalar arguments.
  • Fusion: record the operation in crates/burn-fusion/src/ops. Kernel fusion support, when appropriate, also involves burn-cubecl-fusion; merely recording an operation does not make it fusible with its neighbors.
  • Router: record the operation in crates/burn-router/src/ops and add its execution to TensorInterpreter in crates/burn-router/src/interpreter.rs. The recorded IR and interpreter must agree on the operation's inputs, outputs, and metadata.
  • Remote and capture: verify the operation through these consumers of the router layer. burn-remote uses BackendRouter<RemoteChannel> and executes received operations through TensorInterpreter; burn-capture uses BackendRouter<CaptureChannel> to record them without execution. Ordinary operation support belongs in the shared router layer; changes in these crates are needed when the operation requires additional transport or capture handling.

Some operations have intentional backend limitations. Make them explicit in documentation and errors rather than assuming every runtime has the same capability.

Tests and documentation

Add forward and backward coverage to burn-backend-tests; see the testing guide. Verify shapes, values, broadcasting, and dtypes, plus empty inputs or non-contiguous layouts when relevant. Check every differentiable input and any saved state used by the backward pass.

Run the affected suites with the target backend, and run the repository validation workflow before submitting. Update the Burn Book if the operation adds a new user workflow or changes existing semantics. Keep examples aligned with the runtime device API.

Submitting Examples to Burn

This guide explains how to create and submit new examples to the Burn repository. Examples are a great way to demonstrate Burn's capabilities and help users understand how to use the framework effectively.

For a minimal working example, see the simple-regression example in the repository.

Repository Structure

The Burn repository is set up as a workspace, with examples located in the examples/ directory. Each example is a separate crate that can reuse workspace dependencies.

Creating a New Example

  1. Navigate to the examples directory:

    cd examples
    
  2. Create a new library crate:

    cargo new --lib <my-example>
    
  3. Update the example's Cargo.toml:

    [package]
    name = "<my-example>"
    version = "0.1.0"
    edition = "2021"
    readme = "README.md"
    # Remove this line if it exists
    # readme.workspace = true
    
    [dependencies]
    # Reuse workspace dependencies when available
    serde = { workspace = true }
    # Add example-specific dependencies
    burn = { path = "../../" }
    

Required Files and Structure

README.md

Each example must include a README.md file with:

  • A brief description of what the example demonstrates
  • A terminal command showing how to run the example
  • Any prerequisites or setup instructions

Example README structure:

# Example Name

Brief description of what this example demonstrates.

## Running the Example

```bash
cargo run -p <my-example> --example <my-example>
```

## Prerequisites

List any prerequisites here.

Source Code Structure

  • src/ directory: Contains the main implementation code
  • examples/ directory: Contains example code
    • <my-example>.rs: Example implementation

Resource Handling

  • Resources (datasets, models, etc.) should be downloaded in the example code
  • Do not track external files in the repository
  • Include code to download and prepare resources when the example is run

Best Practices

  1. Code Organization

    • Keep the code modular and well-documented
    • Use clear, descriptive variable and function names
    • Include comments explaining complex operations
  2. Error Handling

    • Implement proper error handling
    • Provide meaningful error messages
    • Handle resource download failures gracefully
  3. Performance

    • Optimize for reasonable execution time
    • Include progress indicators for long-running operations
    • Consider adding configuration options for different hardware capabilities
  4. Documentation

    • Document all public APIs
    • Include inline comments for complex logic
    • Explain any non-obvious implementation details

Sharing Code with the Book

When a book chapter shows code from a maintained example, use an mdBook include so changes to the example also update the chapter. Include the whole file when it is short, or select a named region to preserve a tutorial's step-by-step presentation. Prefer named regions over line numbers, which shift as the source changes.

Mark a region in the Rust source with ordinary comments:

// ANCHOR: example_model
#[derive(Module, Debug)]
pub struct Model {
    linear: Linear,
}
// ANCHOR_END: example_model

Inside the chapter's Rust code fence, reference the source path and region. For example, the basic-workflow model chapter uses:

{{#include ../../../examples/guide/src/model.rs:model}}

Paths are relative to the Markdown file containing the include. The example above is relative to burn-book/src/basic-workflow/model.md. mdBook removes the anchor comments from rendered snippets. Keep conceptual or deliberately abbreviated snippets inline when the maintained implementation would obscure the explanation. Update the surrounding prose when an included API changes.

Build both books from the repository root to check includes:

cargo xtask books burn build
cargo xtask books contributor build

These commands install mdBook if needed. To preview either book in your browser, replace build with open.

An include keeps the displayed code synchronized with its source; compiling and testing the example remains a separate check.

Submitting Your Example

  1. Ensure your example follows all the guidelines above
  2. Test your example thoroughly
  3. Create a pull request with:
    • A clear description of what the example demonstrates
    • Any relevant issue numbers
    • Screenshots or output examples (if applicable)

Feel free to ask questions in the pull request if you need clarification or guidance.

Frequently Encountered Issues

This is a collection of issues people have encountered and asked about on the Discord server. This section is separated from the guides since it can involve lots of details that are only relevant to a small subset of contributors.

Issues encountered while adding ops

Below are some of the issues that were encountered while adding ops to the project. If you encounter an issue while adding an op that isn't listed here, and it's not obvious how to fix it, you can add it to this list or reach out on the Discord server if you need help.

Off by .000001 errors

---- fusion::base::tests::maxmin::tests::test_mean_dim_2d stdout ---- thread 'fusion::base::tests::maxmin::tests::test_mean_dim_2d' panicked at burn-wgpu/src/fusion/base.rs:185:5: assertion `left == right` failed left: Data { value: [1.0, 4.0], shape: Shape { dims: [2, 1] } } right: Data { value: [0.99999994, 3.9999998], shape: Shape { dims: [2, 1] } } ----

tests::maxmin::tests::test_mean_dim_2d stdout ---- thread 'tests::maxmin::tests::test_mean_dim_2d' panicked at burn-wgpu/src/lib.rs:49:5: assertion `left == right` failed left: Data { value: [1.0, 4.0], shape: Shape { dims: [2, 1] } } right: Data { value: [0.99999994, 3.9999998], shape: Shape { dims: [2, 1] } }

If you encounter this, swap out the assert_eq! in the failing test for tensor1.to_data().assert_approx_eq with 3 as the second argument. The second arguments specifies the level of precision: 3 is equivalent to a less than 10-3 (0.001) difference between the elements of the two tensors.