惯性聚合 高效追踪和阅读你感兴趣的博客、新闻、科技资讯
阅读原文 在惯性聚合中打开

推荐订阅源

D
DataBreaches.Net
Engineering at Meta
Engineering at Meta
AI
AI
Threat Intelligence Blog | Flashpoint
Threat Intelligence Blog | Flashpoint
C
CXSECURITY Database RSS Feed - CXSecurity.com
S
Schneier on Security
H
Hackread – Cybersecurity News, Data Breaches, AI and More
F
Fortinet All Blogs
T
Threat Research - Cisco Blogs
B
Blog
K
Kaspersky official blog
Cisco Talos Blog
Cisco Talos Blog
T
The Exploit Database - CXSecurity.com
U
Unit 42
NISL@THU
NISL@THU
D
Docker
Vercel News
Vercel News
C
Check Point Blog
Blog — PlanetScale
Blog — PlanetScale
GbyAI
GbyAI
C
CERT Recently Published Vulnerability Notes
J
Java Code Geeks
Hugging Face - Blog
Hugging Face - Blog
Latest news
Latest news
Martin Fowler
Martin Fowler
Microsoft Azure Blog
Microsoft Azure Blog
I
InfoQ
Know Your Adversary
Know Your Adversary
A
Arctic Wolf
L
LINUX DO - 热门话题
IT之家
IT之家
SecWiki News
SecWiki News
博客园 - 【当耐特】
Schneier on Security
Schneier on Security
C
Cybersecurity and Infrastructure Security Agency CISA
The Last Watchdog
The Last Watchdog
S
Secure Thoughts
P
Proofpoint News Feed
N
News and Events Feed by Topic
S
Security @ Cisco Blogs
Google DeepMind News
Google DeepMind News
钛媒体:引领未来商业与生活新知
钛媒体:引领未来商业与生活新知
爱范儿
爱范儿
罗磊的独立博客
C
Cyber Attacks, Cyber Crime and Cyber Security
D
Darknet – Hacking Tools, Hacker News & Cyber Security
Cloudbric
Cloudbric
O
OpenAI News
V
V2EX
Cyber Security Advisories - MS-ISAC
Cyber Security Advisories - MS-ISAC

博客园 - PKICA

rust类型系统标记 编译配置解答 git实用命令 rust底层设计理念值得注意的几个地方总结 rust可变引用作为函数参数的机理详解 Rust内存重解释transmute C与Rust类型映射 Rust FFI 安全抽象范式 rust延迟初始化原语 rust重借用机制与原理 rust参数传递模型 汇编语言语法详解 gdb汇编调试 gdb-pwndbg的安装与使用指南 gdb调试插件gef C语言thread_local linux系统readelf命令使用指南 gcore转储进程内存 gdb查看命令 RGB与YUV颜色编码的区别 Rust原子类型 C++ STL求两个集合交集差集 gdb调试集锦 ubuntu24.0.4使用root用户登录 ubuntu24.0.4输入密码后跳回登录界面 AI内存压缩技术TurboQuant及存疑 ubuntu切换到指定内核版本 在没有顶级科技大佬直接背书的情况下deepseek为啥能够异军突起? HuggingFace和deepseek的关系 当前主流AI大模型 Rust写时克隆Cow系列2
告别大显存依赖!用 Rust 新一代深度学习框架 Burn 打造纯 CPU 文本分类推理引擎
PKICA · 2026-07-21 · via 博客园 - PKICA

这是一个基于 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);
}