Burn 框架指南 (第二部分):模型构建与多后端架构
讲解如何利用 Module 派生宏构建神经网络,以及如何驾驭 Burn 的多计算后端系统。
Burn ↗ 最突出的设计亮点之一在于将 模型架构定义 与 底层硬件计算后端 彻底解耦。
使用 Module 派生宏构建神经网络#
在 Burn 中,所有神经网络层都需要实现 Module Trait。借助 #[derive(Module)] 过程宏,层声明变得清晰且类型安全。
src/model.rs
use burn::nn::{Linear, LinearConfig, Relu};
use burn::module::Module;
use burn::tensor::backend::Backend;
use burn::tensor::Tensor;
#[derive(Module, Debug)]
pub struct Model<B: Backend> {
linear1: Linear<B>,
linear2: Linear<B>,
activation: Relu,
}
impl<B: Backend> Model<B> {
pub fn new(input_dim: usize, hidden_dim: usize, output_dim: usize, device: &B::Device) -> Self {
let linear1 = LinearConfig::new(input_dim, hidden_dim).init(device);
let linear2 = LinearConfig::new(hidden_dim, output_dim).init(device);
Self {
linear1,
linear2,
activation: Relu::new(),
}
}
pub fn forward(&self, input: Tensor<B, 2>) -> Tensor<B, 2> {
let x = self.linear1.forward(input);
let x = self.activation.forward(x);
// [!code focus]
self.linear2.forward(x)
}
}rust多后端架构的强大之处#
只需更改泛型参数 B,即可自由无缝地切换计算后端:
- WGPU Backend: 跨 Vulkan、Metal 和 DirectX 的高性能 GPU 后端。
- LibTorch Backend: 与 PyTorch C++ 绑定兼容。
- Candle Backend: 基于 Hugging Face Candle 的轻量级后端。
- NdArray Backend: 适用于纯 CPU 环境。
src/main.rs
use burn::backend::wgpu::{Wgpu, WgpuDevice};
fn main() {
let device = WgpuDevice::default();
let model = Model::<Wgpu>::new(784, 128, 10, &device);
println!("模型初始化成功!");
}rust