Skip to main content

Crate burn

Crate burn 

Source
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.
  • module and nn: the Module trait 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: the Learner, metrics and the training dashboard (train feature).
  • vision, signal, linalg: domain-specific tensor operations (features of the same names).
  • remote and server: run tensors on devices hosted by another machine (remote and remote-server features).
  • 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:

BackendFeatureDevice
CUDAcudaDevice::cuda(0)
ROCmrocmDevice::rocm(0)
wgpu (any graphics API)wgpuDevice::wgpu(Default::default())
MetalmetalDevice::metal(Default::default())
VulkanvulkanDevice::vulkan(Default::default())
WebGPUwebgpuDevice::webgpu(Default::default())
CubeCL CPUcpuDevice::cpu()
Flex (pure Rust CPU)flexDevice::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 features dataset and optim and provides a training environment
    • optim: Enables optimizers and learning rate schedulers (implies autodiff)
    • rl: Enables reinforcement learning utilities
    • tui: Includes Text UI with progress bar and plots (requires train)
    • metrics: Includes system info metrics (CPU/GPU usage, etc.) (requires train)
  • Dataset
    • dataset: Includes a datasets library
    • audio: Enables audio datasets (SpeechCommandsDataset)
    • sqlite: Stores datasets in an SQLite database, backed by Turso
    • sqlite-bundled: Deprecated alias for sqlite
    • vision: Enables vision datasets (MnistDataset) and the burn-vision ops module
  • Backends
    • wgpu: Makes available the WGPU backend, on whichever graphics API the platform provides
    • webgpu: Adds Device::webgpu, pinned to WebGPU (implies wgpu)
    • vulkan: Adds Device::vulkan, pinned to Vulkan (implies wgpu)
    • metal: Adds Device::metal, pinned to Metal with native MSL (implies wgpu)
    • cuda: Makes available the CUDA backend
    • rocm: Makes available the ROCm backend
    • cpu: Makes available the CubeCL CPU backend
    • tch: 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 - use flex instead)
  • Backend specifications
    • simd: Enable SIMD kernels in the Flex and NdArray backends
    • rayon: Enable multi-threaded execution in the Flex and NdArray backends
    • accelerate, blas-netlib, openblas, openblas-system: BLAS providers for the NdArray backend
    • autotune: 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 the burn-store snapshot tooling and burnpack stores; with std, this also includes SafeTensors
    • safetensors: Enables SafeTensors import and export in no_std builds (implies store)
    • pytorch: Enables PyTorch checkpoint import (implies store)
  • Others:
    • std: Activates the standard library (deactivate for no_std)
    • linalg: Enables linear algebra operations
    • capture: Makes the non-executing graph capture backend available.
    • ir: Makes Burn’s operation intermediate representation available.
    • cubecl: Re-exports CubeCL as burn::cubecl for writing custom kernels.
    • signal: Enables signal processing operations from burn-signal.
    • extension: Enables the backend extension API, including Tensor::from_primitive.
    • remote: Enables remote devices over Iroh; remote-websocket adds the WebSocket transport.
    • remote-server: Enables the remote server (implies remote).
    • 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 store feature, the snapshot tooling and burnpack stores. The safetensors and pytorch features add those importers.
tensor
Tensor types and compatibility re-exports.
train
Train module
vision
Vision module.

Macros§

empty
Constant macro.

Structs§

BurnConfig
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.