





















这是一个基于 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);
}
此内容由惯性聚合(RSS阅读器)自动聚合整理,仅供阅读参考。 原文来自 — 版权归原作者所有。