Expand description
§Burn
Burn is a deep learning framework written in Rust. It covers tensors, automatic differentiation, neural network modules, optimizers, training and model storage, and runs the same model code on GPUs (CUDA, ROCm, Metal, Vulkan, WebGPU), CPUs and WebAssembly.
§Quick start
Burn ships no execution backend by default. Enable one or more with Cargo features:
[dependencies]
burn = { version = "0.22", features = ["wgpu"] }Models are plain structs that derive Module. Tensors carry their rank in
the type and their backend in their Device, so model code has no backend
type parameter:
use burn::nn::{Linear, LinearConfig, Relu};
use burn::prelude::*;
#[derive(Module, Debug)]
struct Mlp {
hidden: Linear,
activation: Relu,
output: Linear,
}
impl Mlp {
fn new(device: &Device) -> Self {
Self {
hidden: LinearConfig::new(784, 128).init(device),
activation: Relu::new(),
output: LinearConfig::new(128, 10).init(device),
}
}
fn forward(&self, input: Tensor<2>) -> Tensor<2> {
let x = self.activation.forward(self.hidden.forward(input));
self.output.forward(x)
}
}
// An enabled backend in priority order (GPUs before CPUs), unless `BURN_DEVICE` names one.
// `Device::wgpu(..)`, `Device::cuda(0)`, ... pick one explicitly.
let device = Device::default();
let model = Mlp::new(&device);
let logits = model.forward(Tensor::zeros([32, 784], &device));The Burn Book walks through a full training workflow.
§Crate map
tensor:Tensor,Device, dtypes and tensor operations.moduleandnn: theModuletrait and neural network layers.config: serializable configuration structs with#[derive(Config)].optim,lr_scheduler,grad_clipping: optimizers and training utilities.data: datasets, transformations and data loaders.store: saving and loading weights in burnpack, SafeTensors and PyTorch formats.train: theLearner, metrics and the training dashboard (trainfeature).vision,signal,linalg: domain-specific tensor operations (features of the same names).remoteandserver: run tensors on devices hosted by another machine (remoteandremote-serverfeatures).prelude: the types most programs import.
§Backends
Every enabled backend is available at runtime through a Device constructor, and several can
be used side by side:
| Backend | Feature | Device |
|---|---|---|
| CUDA | cuda | Device::cuda(0) |
| ROCm | rocm | Device::rocm(0) |
| wgpu (any graphics API) | wgpu | Device::wgpu(Default::default()) |
| Metal | metal | Device::metal(Default::default()) |
| Vulkan | vulkan | Device::vulkan(Default::default()) |
| WebGPU | webgpu | Device::webgpu(Default::default()) |
| CubeCL CPU | cpu | Device::cpu() |
| Flex (pure Rust CPU) | flex | Device::flex() |
Autodiff and kernel fusion are decorators over these backends: device.autodiff() enables
gradients for tensors created on a device, and the CubeCL backends fuse operations by default.
NdArray (ndarray) and LibTorch (tch) are deprecated.
§Quantization
Burn supports post-training quantization of weights and activations, per tensor or per block, to 8, 4 and 2-bit integers and to FP8 and FP4 formats on supported backends. Quantization-aware training is not supported yet. See the quantization chapter.
§Feature Flags
The following feature flags are available.
Default features include std and optim (and therefore autodiff), but no execution backend.
Select a backend explicitly, for example features = ["wgpu"] or ["flex"].
Specialized operations are also opt-in, for example features = ["flex", "signal"].
Backend-free builds can define tensor/model APIs without installing an execution backend.
Device::default() panics if no execution backend is available; graph capture remains
available through Device::capture() with the capture feature.
- Training
train: Enables featuresdatasetandoptimand provides a training environmentoptim: Enables optimizers and learning rate schedulers (impliesautodiff)rl: Enables reinforcement learning utilitiestui: Includes Text UI with progress bar and plots (requirestrain)metrics: Includes system info metrics (CPU/GPU usage, etc.) (requirestrain)
- Dataset
dataset: Includes a datasets libraryaudio: Enables audio datasets (SpeechCommandsDataset)sqlite: Stores datasets in an SQLite database, backed by Tursosqlite-bundled: Deprecated alias forsqlitevision: Enables vision datasets (MnistDataset) and theburn-visionops module
- Backends
wgpu: Makes available the WGPU backend, on whichever graphics API the platform provideswebgpu: AddsDevice::webgpu, pinned to WebGPU (implieswgpu)vulkan: AddsDevice::vulkan, pinned to Vulkan (implieswgpu)metal: AddsDevice::metal, pinned to Metal with native MSL (implieswgpu)cuda: Makes available the CUDA backendrocm: Makes available the ROCm backendcpu: Makes available the CubeCL CPU backendtch: Makes available the LibTorch backend (deprecated - use a CubeCL backend instead)flex: Makes available the Flex backend (pure-Rust CPU, std/no_std/WASM)ndarray: Makes available the NdArray backend (deprecated - useflexinstead)
- Backend specifications
simd: Enable SIMD kernels in the Flex and NdArray backendsrayon: Enable multi-threaded execution in the Flex and NdArray backendsaccelerate,blas-netlib,openblas,openblas-system: BLAS providers for the NdArray backendautotune: Enable running benchmarks to select the best kernel in backends that support it.autotune-checks: Check that every autotune candidate produces the same output (debugging).x86-v4: Enable AVX-512 matmul kernels in the Flex backend.apple-amx: Enable the experimental Apple AMX matmul kernels in the Flex backend.template: Enable hand-written, non-JIT custom kernels in the CubeCL backends.fusion: Enable operation fusion in backends that support it.tracing: Enable diagnostic tracing in the selected backends (disabled by default).
- Backend decorators
autodiff: Makes available the Autodiff backend
- Model Storage
store: Enables theburn-storesnapshot tooling and burnpack stores; withstd, this also includes SafeTensorssafetensors: Enables SafeTensors import and export inno_stdbuilds (impliesstore)pytorch: Enables PyTorch checkpoint import (impliesstore)
- Others:
std: Activates the standard library (deactivate for no_std)linalg: Enables linear algebra operationscapture: Makes the non-executing graph capture backend available.ir: Makes Burn’s operation intermediate representation available.cubecl: Re-exports CubeCL asburn::cubeclfor writing custom kernels.signal: Enables signal processing operations fromburn-signal.extension: Enables the backend extension API, includingTensor::from_primitive.remote: Enables remote devices over Iroh;remote-websocketadds the WebSocket transport.remote-server: Enables the remote server (impliesremote).network: Enables network utilities (currently, only a file downloader with progress bar)
You can also check the details in sub-crates burn-core and burn-train.
§Backend tracing
Add "tracing" to the features of your burn dependency to compile backend instrumentation,
including autodiff and fusion spans. When depending directly on burn-autodiff or
burn-fusion, enable their tracing feature instead. These spans are opt-in: configuring a
tracing subscriber alone does not enable them. Configure your subscriber to include the
trace level to observe tensor operation spans.
The feature propagates to enabled backends without selecting an additional backend. Normal training logs remain available without this feature.
Modules§
- backend
- Backend module.
- config
- The configuration module.
- data
- Data module.
- grad_
clipping - Gradient clipping module.
- linalg
- Linear algebra operations.
- lr_
scheduler - Learning rate scheduler module.
- module
- Core module infrastructure and neural-network initializers.
- nn
- Neural network module.
- optim
- Optimizers module.
- prelude
- Structs and macros used by most projects. Add
use burn::prelude::*to your code to quickly get started with Burn. - rl
- Module for reinforcement learning.
- serde
- Serde
- signal
- Signal processing module.
- store
- Model storage and serialization: the non-generic record system (always available), plus,
with the
storefeature, the snapshot tooling and burnpack stores. Thesafetensorsandpytorchfeatures add those importers. - tensor
- Tensor types and compatibility re-exports.
- train
- Train module
- vision
- Vision module.
Macros§
- empty
- Constant macro.
Structs§
- Burn
Config - Represents the global configuration for Burn.
- Tensor
- A tensor with a given backend, shape and data type.
Functions§
- runtime_
config - Returns the current
BurnConfig, cached in thread-local storage on native targets.