Burn Framework (Phần 2): Thiết Kế Mô Hình và Đa Backend
Hướng dẫn định nghĩa kiến trúc neural network với trait Module trong Burn và cách khai thác hệ thống multi-backend.
Một trong những điểm sáng nhất của Burn ↗ là khả năng tách biệt hoàn toàn giữa định nghĩa kiến trúc mô hình và backend tính toán.
Định nghĩa Neural Network với Module Derive Macro#
Trong Burn, mọi mô hình đều implement trait Module. Nhờ có proc-macro #[derive(Module)], việc khai báo lớp mạng cực kỳ rõ ràng và type-safe.
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)
}
}rustSức mạnh của kiến trúc Multi-Backend#
Burn cho phép đổi backend linh hoạt bằng cách đổi Generic type B:
- WGPU Backend: Hoàn hảo cho GPU cross-platform (Vulkan, Metal, DirectX).
- LibTorch Backend: Tối ưu khi cần dùng lại các PyTorch C++ bind.
- Candle Backend: Tối ưu nhẹ cho Hugging Face Candle.
- NdArray Backend: Dành cho môi trường CPU không có GPU.
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!("Khởi tạo mô hình thành công!");
}rust