告别大显存依赖!用 Rust 新一代深度学习框架 Burn 打造纯 CPU 文本分类推理引擎
在当今大模型横行的时代,似乎不搞块英伟达显卡就没法玩深度学习了。但如果你只是想在边缘设备、轻量化服务器或者 CPU 环境下部署一个高效的文本分类模型,杀鸡何必用牛刀?
今天,我们将基于 Rust 全新一代、多后端驱动的深度学习框架 Burn(目前在 GitHub 上备受瞩目),手把手搭建一个完全独立、零显卡依赖(纯 CPU)的文本分类推理 Demo。
环境:rust 1.86.0
💡 为什么选择 Burn 框架?
- 极速的 CPU 后端:通过
burn-ndarray后端,Burn 可以在没有 GPU 的情况下,利用 CPU 矩阵运算跑出极高的推理性能。 - 现代化的无畏并发:得益于 Rust 的所有权和生命周期系统,Burn 在编译期就能帮你死守张量安全,杜绝内存泄漏。
- 零成本抽象:一套代码,通过切换
Backend泛型,就能无缝在 CPU(NdArray)、GPU(WGPU/CUDA)之间切换。
🛠️ 项目环境准备
首先,创建一个全新的 Rust 独立项目:
cargo new burn_cpu_demo
cd burn_cpu_demo
在
Cargo.toml 中引入 Burn 核心库 与 NdArray CPU 后端。由于我们要使用纯 CPU 运算,无需开启复杂的硬件加速特性:[package]
name = "burn_cpu_demo"
version = "0.1.0"
edition = "2021"
[dependencies]
# 引入 Burn 框架及训练/推理核心组件
burn = { version = "0.16.0", features = ["train"] }
# 引入纯 CPU 运行的 ndarray 后端
burn-ndarray = { version = "0.16.0" }
🧠 核心架构:Transformer 编码器分类网络
现代文本分类(如经典的 AG News 新闻分类)通常采用 Transformer 架构。我们的 Demo 包含三个核心组件:
- Embedding 层:将离散的单词 ID 转换为高维连续向量。
- Transformer 编码器层:捕捉句子中单词之间的上下文长距离关联。
- Linear 输出层:将提取的特征映射到最终的类别得分(Logits)。
同时,针对 Transformer 提取完特征后的数据聚合,我们实现了解析文本常用的两种经典前向传播策略:
- 📈 平均池化(Mean Pooling):全句特征取平均,代表整句话的全局语义。
- ✂️ 切片裁剪(Slice Pooling):类似 BERT 模型的
[CLS]标记,只提取第一个单词的特征进行分类。
📝 完整源码展示
代码
src/main.rs:这是一个基于 Burn 框架的文本分类模型完整 Demo,它包含了带有详尽注释的网络架构定义、配置初始化以及模拟文本 ID 的推理过程。该代码演示了使用 Transformer 编码器处理输入序列,并对比了切片(Slice)与平均(Mean)两种池化策略的输出结果。
// 引入 Burn 框架的配置宏与模块宏组件
use burn::config::Config;
use burn::module::Module;
// 引入神经网络核心层:全连接层(Linear)、Transformer 编码器(TransformerEncoder)、词嵌入层(Embedding)
use burn::nn::{
Linear, LinearConfig,
transformer::{TransformerEncoder, TransformerEncoderConfig},
Embedding, EmbeddingConfig,
};
// 引入最新版 Transformer 编码器所需的前向传播输入包装结构体
use burn::nn::transformer::TransformerEncoderInput;
// 引入张量后端特质(Backend),使模型具备多硬件平台(CPU/GPU)的可扩展性
use burn::tensor::backend::Backend;
// 引入 Burn 的核心张量结构体(Tensor)以及整型数据标记(Int)
use burn::tensor::{Tensor, Int};
/// 1. 显式定义纯 CPU 计算后端
/// 使用标准的 32 位浮点数(f32)作为基本算术单元。
/// 这对应了在不配置 `f16` 特性时,项目所默认采用的底层矩阵运算类型。
type CpuBackend = burn_ndarray::NdArray<f32>;
/// 2. 定义文本分类神经网络的拓扑结构
/// 通过 `#[derive(Module)]` 宏,Burn 会自动追踪和管理其内部所有子层(网络层)的权重参数。
/// `#[derive(Debug)]` 允许使用标准格式化占位符 `{:?}` 打印出整个模型的内部架构。
#[derive(Module, Debug)]
pub struct TextClassificationModel<B: Backend> {
transformer: TransformerEncoder<B>, // 核心特征提取器:Transformer 编码器层
embedding: Embedding<B>, // 词特征映射层:将离散的单词 ID 转换为高维连续向量
output_linear: Linear<B>, // 分类输出层:将高维特征映射到具体的类别得分空间
}
/// 3. 定义模型的超参数配置结构体
/// `#[derive(Config)]` 宏会自动为该结构体生成编译期动态方法,如 `EmbeddingConfig::new`。
/// 同时也允许将此配置轻松序列化为 JSON 等文件保存。
#[derive(Config)]
pub struct TextClassificationModelConfig {
pub n_classes: usize, // 最终的分类目标数量(例如 AG News 数据集是 4 分类任务)
pub n_features: usize, // 词嵌入维度 / 隐藏层特征维度(Embedding Dimension,如 128 维)
pub vocab_size: usize, // 词汇表的最大容量大小(决定了模型能够识别多少个不同的单词)
pub n_heads: usize, // 多头注意力机制(Multi-Head Attention)中“头”的数量
pub n_layers: usize, // Transformer 编码器堆叠的层数
}
impl TextClassificationModelConfig {
/// 核心初始化函数:基于当前配置,在指定硬件设备上创建模型并赋予随机的初始权重。
pub fn init<B: Backend>(&self, device: &B::Device) -> TextClassificationModel<B> {
// 初始化词嵌入层:输入为 (词表大小, 向量维度),并在指定设备(如 CPU)上分配内存
let embedding = EmbeddingConfig::new(self.vocab_size, self.n_features).init(device);
// 初始化 Transformer 编码器:
// 参数依次为:(特征维度, 前馈神经网络隐层维度, 注意力头数, 堆叠层数)
// 通常前馈网络的隐层维度设定为特征维度的 4 倍(即 self.n_features * 4)
let transformer = TransformerEncoderConfig::new(
self.n_features,
self.n_features * 4,
self.n_heads,
self.n_layers,
)
.init(device);
// 初始化全连接输出层:将特征维度平滑降维映射到分类的类别数量(从 128 维映射到 4 维)
let output_linear = LinearConfig::new(self.n_features, self.n_classes).init(device);
// 将所有初始化完毕的子层组装进自定义的模型结构体中并返回
TextClassificationModel {
transformer,
embedding,
output_linear,
}
}
}
/// 4. 前向传播策略 A:平均池化(Mean Pooling)
impl<B: Backend> TextClassificationModel<B> {
/// 接收输入的单词 ID 序列张量,通过全序列特征取平均的方式,输出最终的分类概率对数几率(Logits)
pub fn forward_mean(&self, tokens: Tensor<B, 2, Int>) -> Tensor<B, 2> {
// [步骤 1] 将离散的 Token ID 映射为稠密的语义向量
// 输入张量形状 (Shape): [Batch_Size, Seq_Length] -> 模拟数据为 [2, 10]
// 输出张量形状 (Shape): [Batch_Size, Seq_Length, n_features] -> 变为 [2, 10, 128]
let x = self.embedding.forward(tokens);
// [步骤 2] 将标准张量包装进新版 Burn 要求的 Transformer 输入专属结构体中
let input = TransformerEncoderInput::new(x);
// [步骤 3] 送入 Transformer 编码器进行长距离上下文关联计算
// 输出张量形状 (Shape) 保持不变: [Batch_Size, Seq_Length, n_features] -> 依然是 [2, 10, 128]
let x = self.transformer.forward(input);
// [步骤 4] 核心聚合操作(平均池化):对维度 1(即时间步 / 单词序列维度)求平均值
// 这一步会将一句话中所有单词的特征融合成一个平均特征,代表整句话的全局语义
// 形状变换: [2, 10, 128] -> 聚合后变为 [2, 1, 128]
let x = x.mean_dim(1);
// [步骤 5] 降维消除孤立维度:将大小正好为 1 的维度 1 强行挤压抹去
// 形状变换: [2, 1, 128] -> 平坦化为标准的二维矩阵 [2, 128]
let x = x.squeeze(1);
// [步骤 6] 全连接层分类映射
// 形状变换: [2, 128] 与全连接权重矩阵 [128, 4] 相乘 -> 最终输出类别得分 [2, 4]
self.output_linear.forward(x)
}
}
/// 5. 前向传播策略 B:切片裁剪法(Slice Pooling)
impl<B: Backend> TextClassificationModel<B> {
/// 类似于 BERT 模型,忽略后续单词,仅仅抽取每句话的第 0 个单词(通常作为 [CLS] 标记)的特征来进行全句分类
pub fn forward_slice(&self, tokens: Tensor<B, 2, Int>) -> Tensor<B, 2> {
// [步骤 1] 词嵌入映射转换
// 形状变换: [2, 10] -> 转换为 [2, 10, 128]
let x = self.embedding.forward(tokens);
// [步骤 2] 构建 Transformer 的标准流输入
let input = TransformerEncoderInput::new(x);
// [步骤 3] 经过 Transformer 多头注意力层的特征洗礼
// 形状保持: 依然是 [2, 10, 128]
let x = self.transformer.forward(input);
// [步骤 4] 获取当前张量各个维度的动态精确尺寸(得到一个数组,例如)
let dims = x.dims();
// [步骤 5] 核心聚合操作(切片法):精准截取第 0 个时间步位置的矩阵切片
// 具体的范围指定规则为:
// 维度 0(Batch 维度) : 保持全选 -> 取 0..dims[0] (即 0..2)
// 维度 1(Seq 维度) : 只要第一个单词 -> 取 0..1 (包含第 0 项,不含第 1 项)
// 维度 2(Feature 维度) : 保持全选 -> 取 0..dims[2] (即 0..128)
// 形状变换: [2, 10, 128] -> 截取后缩减为 [2, 1, 128]
let x = x.slice([0..dims[0], 0..1, 0..dims[2]]);
// [步骤 6] 降维消除孤立维度:因为维度 1 的尺寸现在变成了 1,可以安全地通过 squeeze 将其抹除
// 形状变换: [2, 1, 128] -> 平坦化为标准的二维矩阵 [2, 128]
let x = x.squeeze(1);
// [步骤 7] 全连接层映射得出结果
// 形状变换: [2, 128] -> 通过层映射最终转换为 [2, 4]
self.output_linear.forward(x)
}
}
/// 6. 主程序入口
fn main() {
// 实例化纯 CPU 的运算设备
let device = burn_ndarray::NdArrayDevice::Cpu;
println!("🚀 正在使用纯 CPU 后端初始化文本分类模型...");
// 初始化超参数配置实例:设定为 4 分类、特征 128 维、词表包含 1000 词、2 个注意力头、2 layer
let config = TextClassificationModelConfig {
n_classes: 4,
n_features: 128,
vocab_size: 1000,
n_heads: 2,
n_layers: 2,
};
// 驱动配置实例化具体的模型,并将模型所有的初始权重直接绑定并加载到 CPU 内存上
let model = config.init::<CpuBackend>(&device);
println!("📝 正在模拟输入文本数据 (Batch Size: 2, Sequence Length: 10)...");
// ✨【数据修复位置】:模拟构造一个批次(Batch)的文本数字信号:
// 包含 2 句话(Batch Size = 2),每句话由 10 个词的内部 ID 组成(Sequence Length = 10)。
// 该数组在内存中完美契合一个二维形状的数学矩阵:形状为 [2, 10]
let token_data = [
[1, 2, 3, 4, 5, 6, 7, 8, 9, 10], // 第一句话:由 10 个离散词 ID 构成
[11, 12, 13, 14, 15, 0, 0, 0, 0, 0], // 第二句话:较短,末尾 5 位用 0 进行了标准的对齐填充(Padding)
];
// 使用 `from_data` 静态方法,将 Rust 原生的二维常规数组包装成 Burn 系统的高级 `Tensor`
// 显式指定泛型类型为:<CPU后端, 2维张量, 整型数据类型>,并将其绑定到 CPU 设备上
let tokens = Tensor::<CpuBackend, 2, Int>::from_data(token_data, &device);
println!("⚡ 正在纯 CPU 上执行前向推理...");
// 执行前向传播运算。因为 Tensor 默认采用移动语义(Move),
// 为了防止 tokens 张量在第一次调用后被提前销毁,我们使用 `.clone()` 显式复制一份描述符传入。
// 这两个函数输出的结果,均为未经过 Softmax 归一化的分类对数几率(即原始的 Logits 分数)
let output_slice = model.forward_slice(tokens.clone()); // 运行切片池化网络流程
let output_mean = model.forward_mean(tokens.clone()); // 运行平均池化网络流程
// 格式化输出最终得到的两个结果张量,其最终的 Shape 均为标准的 [2, 4] 二维矩阵
println!("\n✅ 推理成功完成!输出的分类张量结果:");
println!("--- [切片法 (Slice Pooling) 输出得分] ---");
println!("{}", output_slice);
println!("--- [平均池化法 (Mean Pooling) 输出得分] ---");
println!("{}", output_mean);
}
⚡ 性能避坑指南:必须使用 --release!
在深度学习和复杂的矩阵运算中,Rust 的默认 Debug 编译模式性能极差(由于未开启编译器优化,速度可能慢上百倍)。
在运行此 Demo 时,请务必在终端中加上
--release 标志:cargo run --release
📊 运行结果预览:
🚀 正在使用纯 CPU 后端初始化文本分类模型...
📝 正在模拟输入文本数据 (Batch Size: 2, Sequence Length: 10)...
⚡ 正在纯 CPU 上执行前向推理...
✅ 推理成功完成!输出的分类张量结果:
--- [切片法 (Slice Pooling) 输出得分] ---
Tensor {
data:
[[-0.2142, 0.5123, -0.1942, 0.0124],
[ 0.3512, -0.1412, 0.8141, -0.6102]],
shape:,
device: Cpu,
backend: "ndarray",
kind: "Float",
}
...
🔍 核心技术点解析
TransformerEncoderInput的必要性:
在最新版的 Burn 中,TransformerEncoder的前向传播要求将 Tensor 包装进TransformerEncoderInput中。这是官方为了方便开发者同时传入文本张量与填充掩码(Padding Mask)而做的优秀设计,能完美处理变长文本。- 解除 Squeeze 报错陷阱:
直接对形状为[2, 10, 128]的张量调用squeeze(1)会引发 CPU 的 ❌ Panic。我们必须通过.mean_dim(1)或是.slice(...)将序列长度(维度 1)强行压缩为 1 之后,才能安全地剥离它。
🎯 结语
使用 Rust 的 Burn 框架搭建深度学习模型是一种极为丝滑的体验。它彻底打破了“Python 垄断深度学习”的固有印象。对于需要高并发、轻量化部署的工业级线上项目,这种纯 CPU、低内存开销的 Rust 推理引擎无疑极具吸引力。
下一步,我们将探讨如何把实际训练好的大模型参数(
.mpk 文件)加载到这个 CPU Demo 中,并加入中文分词器(Tokenizer)来实现对真实中文网页文本的动态分类!你对 Rust 玩深度学习怎么看?欢迎在评论区留下你的看法!🦀
本文代码基于 Burn v0.16.0 编写,完全适配最新的语言规范。
互动引导:如果这篇博客对您有启发,接下来您可以:
- 尝试引入真实的
Tokenizer库(如tokenizerscrate)来将中英文句子转成代码里的 ID。 - 探讨如何将这段 Rust 推理逻辑打包成 C 语言动态链接库(FFI),供其他语言调用。
参考资料:
https://burn.dev/
浙公网安备 33010602011771号