1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
//! # 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:
//!
//! ```toml
//! [dependencies]
//! burn = { version = "0.22", features = ["wgpu"] }
//! ```
//!
//! Models are plain structs that derive [`Module`](module::Module). Tensors carry their rank in
//! the type and their backend in their [`Device`](tensor::Device), so model code has no backend
//! type parameter:
//!
//! ```rust,no_run
//! 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](https://burn.dev/books/burn/) walks through a full training workflow.
//!
//! ## Crate map
//!
//! - [`tensor`]: [`Tensor`], [`Device`](tensor::Device), dtypes and tensor
//! operations.
//! - [`module`] and [`nn`]: the [`Module`](module::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:
//!
//! | 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](https://burn.dev/books/burn/performance/quantization.html).
//!
//! ## 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](https://turso.tech/)
//! - `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`](https://docs.rs/burn-core) and [`burn-train`](https://docs.rs/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.
pub use *;
/// Linear algebra operations.
/// Core module infrastructure and neural-network initializers.
/// Tensor types and compatibility re-exports.
/// Train module
/// Module for reinforcement learning.
pub use remote;
pub use server;
/// 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.
/// Neural network module.
pub use ;
/// Optimizers module.
// For backward compat, `burn::lr_scheduler::*`
/// Learning rate scheduler module.
// For backward compat, `burn::grad_clipping::*`
/// Gradient clipping module.
/// CubeCL module re-export.
/// Vision module.
/// Signal processing module.