burn
A flexible, backend-agnostic deep-learning framework - train and infer in Rust with pluggable compute backends.
Search across all documentation pages
A flexible, backend-agnostic deep-learning framework - train and infer in Rust with pluggable compute backends.
use burn::backend::NdArray;
use burn::module::Module;
use burn::nn::Linear;
use burn::tensor::Tensor;
type B = NdArray<f32>;
#[derive(Module, Debug)]
struct Model<B: burn::backend::Backend> {
linear: Linear<B>,
}
fn forward<B: burn::backend::Backend>(m: &Model<B>, x: Tensor<B, 2>) -> Tensor<B, 2> {
m.linear.forward(x)
}When to reach for this:
use burn::backend::NdArray;
use burn::module::Module;
use burn::nn::{Linear, LinearConfig};
use burn::tensor::Tensor;
type B = NdArray<f32>;
#[derive(Module, Debug)]
struct Mlp<B: burn::backend::Backend> {
fc: Linear<B>,
}
fn main() {
let device = <B as burn::backend::Backend>::Device::default();
let mlp = Mlp::<B> {
fc: LinearConfig::new(4, 2).init(&device),
};
let input = Tensor::<B, 2>::from_data([[1.0, 2.0, 3.0, 4.0]], &device);
let out = mlp.fc.forward(input);
println!("{:?}", out.into_data());
}What this demonstrates:
Module derive for parameter registrationLinearConfig builder pattern for layer setupBackend.Learner API wires datasets, optimizers, metrics, and checkpointing.| Backend | Target |
|---|---|
| NdArray | CPU dev, CI |
| WGPU | Cross-platform GPU |
| CUDA | NVIDIA servers |
| Alternative | Use When | Don't Use When |
|---|---|---|
| candle | HF transformer weights day one | You need generic backend swapping |
| ort | Production ONNX inference only | Training inside Rust |
| tch | Must reuse arbitrary PyTorch code | Pure Rust dependency policy |
| JAX/Python | Large-scale TPU training | Edge deploy without Python |
Inference paths mature faster than large-scale training. Validate accuracy and perf on your model before cutting Python.
Change the type alias type B = ... and rebuild with appropriate feature flags - no source rewrite for module code.
Yes via burn-import and custom modules - transformer support is active but check examples for your architecture.
Implement Dataset trait or adapt CSV/Parquet readers from the data section into tensor batches.
Experimental paths exist; expect size and SIMD limits - profile before betting on browser training.
Use LearnerBuilder with optimizer config, metric trackers, and learner.fit(dataloader) - see burn book for full snippet.
ONNX via ort is great for frozen graphs; native burn modules allow Rust-only hot paths.
Run golden tensor comparisons on a fixed batch after import/export - do not trust loss curves alone.
WGPU/CUDA backends cache allocations; drop unused tensors and limit batch size when OOM.
candle is HF-focused; burn is general - see candle for transformer zoo shortcuts.
Stack versions: This page was written for Rust 1.97.0 (edition 2024), Tokio 1.x, Axum 0.8, serde 1.0, sqlx 0.8, clap 4, and Polars 0.46+.
Reviewed by Chris St. John·Last updated Jul 16, 2026