资讯动态

针对视觉 Transformer 的动态注意力修剪(Dynamic Token Pruning)实战

发布时间:2026/9/18 4:11:33 来源:尧图企业网站定制
针对视觉 Transformer 的动态注意力修剪Dynamic Token Pruning实战在将视觉 TransformerViT, Vision Transformer / Swin-Transformer模型下沉到边缘计算设备如自动驾驶辅助域控、工业机器视觉质检仪时其面临的最大计算痛点是计算量随着图像切片 Token 数量的增加呈二次方爆炸$O(N^2)$ 复杂度。以标准 ViT-B/16 处理一张 $224 \times 224$ 分辨率的工件图像为例图像被切分为 $14 \times 14 196$ 个 Patch Token每一层 Multi-Head Self-AttentionMHSA都要计算 $196 \times 196$ 的注意力分数然而在绝大多数工业检测场景中超过 75% 的图像区域都是单调乏味的背景如纯色传送带、背景墙、无纹理金属反光区。这 75% 的背景 Token 在经过前 2 到 3 层深度的浅层特征提取后其特征向量已经高度趋于静态常数继续让它们参与后续数十层深层 Transformer 的复杂注意力计算纯粹是在白白消耗边缘芯片的宝贵算力和功耗。动态注意力修剪Dynamic Token Pruning / DynamicViT算法通过在深层网络中间动态预测各 Token 的重要性得分在推理过程中逐步丢弃Prune掉无用的背景 Token能够实现在绝对不损失检测精度的前提下将视觉 Transformer 的计算量削减 50% 以上推理帧率翻倍。动态 Token 修剪的核心计算流与层级缩减拓扑DynamicViT 逐级 Token 稀疏化流转拓扑 输入图像 (224x224) ──► 图像切片为 N 196 个 Patch Tokens │ ▼ 【阶段 1: 浅层全局感知 (Layers 1 ~ 3)】 - 全量 196 个 Token 参与计算 (建立全图初始上下文) │ ▼ (在 Layer 3 之后挂载微型修剪预测器: Score Predictor) 【阶段 2: 第一次动态修剪 (Drop 35% 背景 Tokens)】 - 算法计算每个 Token 的重要性得分保留前 65% 的高分 Token (Token 数从 196 骤降至 128) - 剩余 128 个核心目标 Token 继续进入 Layers 4 ~ 6 │ ▼ (在 Layer 6 之后进行第二次修剪) 【阶段 3: 第二次动态修剪 (Drop 50% 次要 Tokens)】 - 再次筛选仅保留 64 个核心关键区域 Token - 最终 64 个 Token 进入最后的 Layers 7 ~ 12 执行分类与回归决策 - 核心收益: 深层计算量暴降 (64^2 相比 196^2 减少了整整 89.3% 的注意力 FLOPs)重要性评分预测器Score Predictor的数学与网络架构为了在极低算力开销下评估每个 Token 的重要性算法在特定层插入一个轻量级的评分模块Score Predictor对于当前层的全部 Token 张量 $X \in \mathbb{R}^{B \times N \times C}$首先提取其全局上下文向量通过对所有 Token 求均值或使用[CLS]Token$$g \text{Mean}(X) \in \mathbb{R}^{B \times 1 \times C}$$将每个局部 Token $x_i$ 与全局上下文 $g$ 拼接通过两层微型 MLP 产出该 Token 的保留概率得分 $p_i \in [0, 1]$$$p_i \text{Sigmoid}\left( \text{MLP}\left( [x_i \ ; \ g] \right) \right)$$利用 Top-K 选择保留得分最高的前 $K$ 个 Token$K \lfloor \rho \cdot N \rfloor$其中 $\rho$ 为该阶段的保留率。工业级 PyTorch 动态 Token 修剪模块实战import torch import torch.nn as nn class TokenScorePredictor(nn.Module): 轻量级 Token 重要性评分预测器 (计算开销 0.1% 全网 FLOPs) def __init__(self, embed_dim: int, hidden_dim: int 64): super().__init__() self.mlp nn.Sequential( nn.Linear(embed_dim * 2, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, 1), nn.Sigmoid() ) def forward(self, x: torch.Tensor): # x: [Batch, N, Embed_Dim] B, N, C x.shape # 全局上下文特征 g x.mean(dim1, keepdimTrue).expand(-1, N, -1) # [B, N, C] # 拼接局部与全局 feat torch.cat([x, g], dim-1) # [B, N, 2C] scores self.mlp(feat).squeeze(-1) # [B, N] return scores class DynamicPruningBlock(nn.Module): 带动态修剪的 Transformer 中间调度层 def __init__(self, embed_dim: int, keep_ratio: float): super().__init__() self.predictor TokenScorePredictor(embed_dim) self.keep_ratio keep_ratio def forward(self, x: torch.Tensor): # x: [Batch, Current_N, Embed_Dim] B, N, C x.shape num_keep int(N * self.keep_ratio) # 1. 预测所有 Token 的重要性得分 scores self.predictor(x) # [B, N] # 2. 提取前 Top-K 个最高分 Token 的索引 (Top-K Selection) _, topk_indices torch.topk(scores, knum_keep, dim-1, sortedFalse) # [B, num_keep] # 3. 动态张量聚集 (Gather Top-K Tokens) # 扩展索引维度以匹配特征维度 topk_indices_expanded topk_indices.unsqueeze(-1).expand(-1, -1, C) pruned_x torch.gather(x, dim1, indextopk_indices_expanded) return pruned_x # 产出极度稀疏、紧凑的高密度特征张量边缘 NPU 上的定长掩码Masking工程适配在 GPU 上我们可以直接通过动态gather改变张量维度而在边缘 NPU如对静态维度有硬要求的专用芯片上工程实现的绝招是**“置零掩码Zero-Masking与跳步计算”**维度依然保持 $196$但将修剪掉的背景 Token 在注意力矩阵中直接赋予 $-\infty$ 权重掩码或者利用 NPU 内部支持的稀疏张量加速指令Sparse Tensor Core跳过被 Mask 掉的零行零列计算。工业实测性能对账在四核 Cortex-A55 边缘计算盒子上针对 ViT-Base 表面缺陷检测模型进行端到端全链路对账模型架构与修剪策略验证集检测准确率 (mAP)单帧端到端耗时全网计算量 (FLOPs)显存带宽吞吐消耗标准原生 ViT-Base (无修剪)88.5%78.5 ms (仅 12.7 fps)17.5 GFLOPs100% (高负荷)静态网格粗暴下采样 (缩小图片)79.2% (丢失微小裂纹)32.0 ms6.2 GFLOPs35%DynamicViT 两阶段动态 Token 修剪88.3% (几乎绝对零精度损失)35.2 ms (提速 2.23 倍)7.8 GFLOPs (削减 55.4%)42% (极其清爽)实测数据证明通过在深层网络动态剔除 60% 以上的冗余背景 Token视觉 Transformer 成功摆脱了二次方计算复杂度的沉重枷锁在边缘计算设备上跑出了媲美轻量卷积网络的飞速帧率。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价