告别大显存依赖!用 Rust 新一代深度学习框架 Burn 打造纯 CPU 文本分类推理引擎

在当今大模型横行的时代,似乎不搞块英伟达显卡就没法玩深度学习了。但如果你只是想在边缘设备、轻量化服务器或者 CPU 环境下部署一个高效的文本分类模型,杀鸡何必用牛刀?
今天,我们将基于 Rust 全新一代、多后端驱动的深度学习框架 Burn(目前在 GitHub 上备受瞩目),手把手搭建一个完全独立、零显卡依赖(纯 CPU)的文本分类推理 Demo。
环境:rust 1.86.0

💡 为什么选择 Burn 框架?

  1. 极速的 CPU 后端:通过 burn-ndarray 后端,Burn 可以在没有 GPU 的情况下,利用 CPU 矩阵运算跑出极高的推理性能。
  2. 现代化的无畏并发:得益于 Rust 的所有权和生命周期系统,Burn 在编译期就能帮你死守张量安全,杜绝内存泄漏。
  3. 零成本抽象:一套代码,通过切换 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 包含三个核心组件:
 
  1. Embedding 层:将离散的单词 ID 转换为高维连续向量。
  2. Transformer 编码器层:捕捉句子中单词之间的上下文长距离关联。
  3. 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",
}
...

🔍 核心技术点解析

  1. TransformerEncoderInput 的必要性:
    在最新版的 Burn 中,TransformerEncoder 的前向传播要求将 Tensor 包装进 TransformerEncoderInput 中。这是官方为了方便开发者同时传入文本张量与填充掩码(Padding Mask)而做的优秀设计,能完美处理变长文本。
  2. 解除 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 库(如 tokenizers crate)来将中英文句子转成代码里的 ID。
  • 探讨如何将这段 Rust 推理逻辑打包成 C 语言动态链接库(FFI),供其他语言调用。

参考资料:

https://burn.dev/

posted @ 2026-07-21 14:57  PKICA  阅读(1)  评论(0)    收藏  举报