
先说结论如果你手里只有一颗算力不高的边缘芯片却想跑带Transformer结构的模型FreeToken这篇论文提供了一条和传统量化、静态剪枝都不太一样的技术路线——它不追求把模型本身变小而是让模型在推理时动态“偷懒”把计算集中到真正关键的token上。这篇文章我是当成一篇论文笔记来整理的重点讲FreeToken的核心方法、框架设计里的工程细节以及我在复现和部署过程中踩过的坑。适合正在做边缘侧推理、端侧模型加速的算法工程师和嵌入式开发同学参考。凡是论文原文没详细展开、属于我实践经验补充的部分我会在文中明确标注出来。1. 论文在解决什么问题1.1 边缘侧推理的真正瓶颈不是算力是内存和带宽很多刚接触边缘部署的同学容易陷入一个误区以为只要芯片算力够模型就能跑得动。实际上等我真在开发板上跑过一轮就发现边缘侧推理最卡脖子的往往不是TOPS算力而是内存容量和带宽。这类设备的内存通常只有几个GB甚至几百MB内存带宽也就几十GB/s级别而Transformer模型的计算特征恰恰是“中间结果特别多”。每层都要读写KV Cache、注意力矩阵、中间激活值序列越长KV Cache线性膨胀注意力计算量干脆按序列长度的平方往上涨。举个例子我在RK3588上跑一个轻量级ViT模型模型本身参数量不大但输入分辨率一提高中间激活值立刻把内存顶满推理速度从几十毫秒直接掉到几百毫秒。这时候乘加运算不是瓶颈数据在内存和计算单元之间搬来搬去的耗时才是。换句话说边缘侧做Transformer推理本质上是和内存带宽做斗争。还有一个容易被忽略的点并非所有token都值得同等计算。视觉任务里大块平滑背景对应的patch token信息量很低语音任务里静音帧和重复音节的token同样没什么区分度文本里也有一批停用词。对这些token做完整的注意力计算和FFN变换属于实打实的资源浪费。FreeToken这篇论文就是针对这两个痛点来的——既要控制内存占用又要避免低信息量token浪费算力。1.2 FreeToken的定位让模型学会“按需计算”FreeToken本质上是一个边缘侧推理框架核心思想是运行时动态稀疏计算。它不再要求模型对每个输入都走同一条计算路径而是让模型根据输入内容自己决定哪些token需要保留、哪些token可以被剪掉、哪些样本在浅层就可以输出结果。它涉及三个相互配合的机制token剪枝、级联退出、稀疏注意力。token剪枝负责砍掉不重要的token级联退出负责让简单样本提前结束推理稀疏注意力负责让保留下来参与计算的token减少无效的注意力计算。三者合在一起效果是计算量能跟输入难度动态匹配简单样本跑浅层、少量token复杂样本才跑深层、全量token。我花了一周左右时间复现论文思路分别在不同任务上做了验证。结论是在分类和语音指令词这类任务上这套方法确实能把中间层浮点运算量砍掉一半左右精度损失控制在1个点以内但如果任务本身对每个token都极度依赖收益就没那么大。这个差异后面我会详细展开。先看框架的核心设计。2. FreeToken框架的设计思路与方法拆解2.1 整体设计把“每层全算”改成“每层选算”传统Transformer的推理流程很固定输入的所有token在每一层都完整参与多头注意力和FFN计算层与层之间全量传递。FreeToken改成了另一个流程每一层开始前先对当前所有token做一次重要性评估决定保留Top-K个token参与本层计算同时判断当前层的输出置信度是否足够高如果足够高就直接在这里输出结果不再继续往深层走。这个设计跟常见的静态剪枝、蒸馏最大的区别在于静态剪枝是在训练完成后把权重矩阵整体砍掉一部分不管输入是什么计算路径都一样而FreeToken是在推理时根据输入动态选择计算路径不同样本走的层数和token集合都不同。打个比方以前是所有人不管题难不难都必须做完12道大题现在是先扫一眼题目简单的做3道就能交卷难的就做满12道。省下来的时间自然就变成了推理提速。从框架角度看FreeToken可以理解为一个带自适应能力的推理引擎。它内部包含一个轻量的token打分模块、一组可配置的退出分支、以及支持动态稀疏计算的算子库。模型结构本身不一定要改动多大关键是训练阶段要让模型具备“知道自己哪些token重要、哪些样本已经能输出”的能力。2.2 两级Token选择粗筛和精筛分开做论文里最值得一提的细节是token选择策略它没有直接用一个注意力分数做Top-K。如果直接算完整注意力矩阵再做筛选虽然理论上可行但那个完整的O(N²)注意力矩阵恰恰是我们想省掉的东西先算再剪等于白忙活。FreeToken的做法是把token选择拆成两级。第一级是粗筛用一个非常轻量的打分器完成。这个打分器把每个token的embedding做一次全局池化再接一个很小的全连接层输出一个标量重要性分数。这个打分器的计算成本很低我实测下来大概只占一个Transformer层计算量的1%左右。根据这个分数先按比例或阈值砍掉一批明显不重要的token比如背景patch、静音帧、重复片段。第二级是精筛只对第一轮保留下来的候选token计算真正的多头注意力然后用注意力得分再做一次筛选决定最终进入FFN的token集合。这样做既避免了全量注意力带来的平方级开销又保证了剪枝不是随便乱剪。我在复现时专门做过对比实验只保留第一级粗筛、去掉第二级精筛精度损失明显变大把两级都加上后效果接近完整模型。原因也容易理解粗筛打分器毕竟是个轻量模块只能抓住“明显不重要”的token而注意力得分能反映出token之间的上下文关系这个信息对保留重要token很有价值。2.3 级联退出浅层能搞定就不必深算级联退出是FreeToken里另一个减计算量的大头。它的逻辑很朴素很多任务的简单样本在浅层就已经有很高的置信度了没必要非得跑完所有层。论文在每个Block出口处都接了一个小的分类头或回归头计算代价很小。推理时如果当前层的输出置信度超过设定阈值τ就直接返回结果否则继续往下层走。这里有个我在实践中的体会不是所有任务都适合级联退出。分类、检索、指令词识别这类语义比较直接的任务浅层特征往往已经足够退出收益很大但检测、分割这类对空间细节要求高的任务太早退出容易丢边界、丢小目标必须把阈值调得非常高提速效果就有限了。如果业务场景允许建议把级联退出模块单独训练成“轻量分类器”比直接在原始任务头上接退出分支更稳。另外要注意级联退出对batch推理不友好。如果端侧一次只处理一条样本退出策略可以灵活生效但如果用大的batch并行推理提前退出的样本和继续深算的样本会互相拖累处理起来非常麻烦。FreeToken本身是按单样本动态路径设计的这点和边缘侧常见的小batch、低延迟场景刚好契合。2.4 训练阶段的Loss怎么设计动态剪枝和级联退出不是训练完就可以直接加的模型必须在一开始就学会“什么时候该保留token、什么时候该提前退出”。FreeToken在训练阶段把所有选择操作都做成了可微的。具体来说token保留操作被实现成带Gumbel噪声的离散采样这样既保留了梯度回传能力推理时又能把采样换成确定性的Top-K。损失函数方面论文的思路是联合多个目标一起优化主任务loss保证精度token保留率正则loss控制剪枝比例计算量正则loss跟模型总FLOPs挂钩。这三个loss相加之后相当于训练过程中自动在“精度”和“速度”之间找平衡点。想更快就把计算量正则系数调大想更准就调小一个参数就能控制整体倾向。我的理解是FreeToken本质上做了一次轻量级的“结构搜索”只是它搜索的不是网络宽度和深度而是每一层的token保留比例和退出层号而且这个搜索过程完全融合在训练里不需要额外步骤。对做工程的人来说这种设计有个额外好处训练完成后你不需要在端侧跑一个复杂的选择网络所有决策都固化成了阈值和Top-K操作部署时非常干净。3. 关键实现细节与工程化要点3.1 稀疏注意力算子的正确姿势动态token剪枝有个很容易犯的错先在代码里把完整注意力矩阵算出来再套一个mask把不需要的位置置成负数。这样做的代码确实好写但内存一点没省甚至比原来更慢完全违背了做token剪枝的初衷。正确做法是把稀疏注意力当作一个真正的稀疏计算来实现。对于单query的场景可以用top-k先选出每个query最关心的k个key位置然后通过索引直接从key和value里gather出对应部分再在选出来的子矩阵上做attention。下面是我在复现时用PyTorch写的一个简化版def sparse_attention(query, key, value, keep_k32): # query: [B, H, T, D], key/value: [B, H, T, D] scores torch.matmul(query, key.transpose(-1, -2)) / (D ** 0.5) # 对每个query只保留最关心的keep_k个位置 topk_scores, topk_idx torch.topk(scores, keep_k, dim-1) # 按索引从key/value中gather出子集 key_select torch.gather( key, dim-2, indextopk_idx.unsqueeze(-1).expand(-1, -1, -1, D) ) value_select torch.gather( value, dim-2, indextopk_idx.unsqueeze(-1).expand(-1, -1, -1, D) ) attn_weights torch.softmax(topk_scores, dim-1) output torch.matmul(attn_weights.unsqueeze(-2), value_select) return output.squeeze(-2), topk_idx如果你的场景是先做全局token剪枝再在剪枝后的token集合内部计算注意力流程会反过来先选出query集合再在这些query之间做稀疏注意力。两种模式对应不同业务需求前者适合保留全局感知、按需看重点后者适合直接缩小计算域速度提升更明显。我建议两种模式都实现业务上灵活切换。3.2 KV Cache优化与算子融合token剪枝之后KV Cache的长度不是固定的了这对推理引擎的适配是个考验。很多边缘推理框架绑定的是静态shape图动态长度意味着要么用动态shape模式要么把最长序列固定下来再padding。padding到最长的做法会牺牲一部分剪枝收益但胜在兼容性好。如果对性能要求高我建议优先选支持动态shape的引擎。内存优化还有几个细节值得注意。第一KV Cache可以尝试用INT8存储精度损失通常很小但对内存占用的改善非常直接。第二QKV三个矩阵的乘法尽量融合成一个算子减少kernel launch次数。第三LayerNorm、残差连接和激活函数这几个逐元素操作也可以在编译期融合。这些优化单独看每一项收益都不算大但叠加起来在内存带宽受限的边缘设备上效果很明显。论文里给的FLOPs数据是纯计算量的对比实际工程落地时内存访问量的减少往往比FLOPs减少更能带来端到端提速。3.3 边缘设备适配动态shape是最大的坎FreeToken这种动态计算特性在PyTorch里跑起来很舒服但一旦要部署到边缘设备第一个拦路虎就是动态shape。大多数边缘推理引擎比如RKNN、TensorRT、CoreML虽然都宣称支持动态shape但实际用起来限制不少有的只支持动态batch不支持动态序列长度有的支持但性能会下降。我自己的做法是先导出两个ONNX静态图一个是完整模型一个是剪枝后的模型。在端侧先用轻量打分器判断当前输入的复杂度根据阈值决定调用哪一个子图。这种方式虽然粗暴丢掉了逐层动态调整的能力但胜在稳定而且在输入复杂度两级分化明显的业务里可用的加速比也很可观。如果你要保留完整的逐层动态能力建议直接把FreeToken的token选择逻辑在端侧引擎里写成自定义算子工程量会大不少但上限也高。4. 复现与部署实操记录4.1 环境准备和代码获取论文如果公开了代码一般会放在GitHub或论文项目主页。拿下来之后先用git clone把仓库拉到本地然后创建一个干净的conda环境再装依赖。这里提醒一句别一上来就装一堆最新版依赖先看README里锁定的版本。PyTorch版本变动很容易导致CUDA算子重编浪费一两个小时很常见。我推荐的环境组合是Python 3.9或3.10PyTorch 2.0以上ONNX和ONNX Runtime负责导出验证再根据目标设备选装RKNN Toolkit或TensorRT。如果只是先看效果不需要一步到位部署到板子先用小模型在服务器上复现就够了。小模型推荐DeiT-T、TinyViT这类参数量在5M到20M之间的backbone训练速度快调试起来也方便。板上验证放在最后一步确认模型效果没问题再折腾跨平台导出。4.2 从训练到导出的完整流程复现FreeToken的完整流程我建议按照下面这几个步骤走每一步都比较可控。第一步选好backbone并加载预训练权重。我试过从零训练效果明显不如在预训练权重上做增量训练。因为token选择模块和退出分支都是新增结构但backbone的底层特征提取能力还是预训练权重更靠谱。第二步准备任务数据。分类任务直接复用ImageNet或自己的业务数据集检测、分割任务需要额外处理。第三步配置动态训练参数开始训练。训练时重点关注token保留率的变化曲线如果保留率一直掉到很低说明计算量正则loss系数调太大了模型在牺牲精度换速度要及时调回去。训练结束后把模型导出成ONNX格式。导出时要特别注意动态轴的设置把序列长度维度标成动态。导出之后用ONNX Runtime跑一遍和PyTorch的推理结果做数值对比。我遇到过导出后精度正常但输出张量shape对不上的情况基本都是动态轴设置和实际推理数据不匹配导致的排查起来并不难。最后再根据板子的SDK把ONNX转成对应格式在端侧做benchmark和精度验证。4.3 几个核心参数的调法FreeToken里最影响效果的参数有三个初始token保留率r、级联退出的置信度阈值τ、计算量正则系数λ。它们对结果的影响不是互相独立的我调参时的经验是先固定两个只动一个逐个找规律。token保留率r建议初始值设在0.5到0.75之间。太高了剪枝效果不明显太低了精度崩得快。阈值τ建议从0.9起步这个值比较保守先确保精度不跌再慢慢调低提速。计算量正则系数λ建议从0.01开始如果训练过程中token保留率下不去可以适当加大如果精度掉了就先减小。整体来看这三个参数共同决定模型的“懒惰程度”核心原则是每次只动一个参数记录好精度和FLOPs的变化再决定下一步方向。4.4 常见问题与排查速查表复现过程中我整理了下面这个常见问题速查表都是自己实际碰到过的按现象排查比较快。问题现象可能原因解决思路训练时精度崩了token保留率过低或计算量正则系数太大调大保留率调小λ先恢复精度再谈提速token保留率训练中不下降gate没有收敛学习率过高导致训练不稳定降低学习率增加训练轮数检查gate输入特征是否归一化剪枝后推理速度不升反降注意力mask实现方式不对先算全量再加mask改成top-k gather的稀疏实现导出ONNX时动态轴报错动态shape设置和实际数据不一致检查sequence维度的dynamic_axes配置确保导出时输入shape正确边缘设备上延迟比预期高动态shape在端侧被padding成固定长度或算子没融合尝试双子图切换方案或使用自定义算子实现动态选择级联退出后单个样本延迟不稳定简单样本和困难样本计算量差异大缓存命中率不稳定设置最小计算量约束防止退出条件过激进统计P95延迟来评估这些坑里面最容易被忽视的是“推理速度不升反降”这一类。很多人在PyTorch里用mask方式实现了稀疏注意力看FLOPs确实降了但实际推理时间没变甚至更慢原因就是全量注意力矩阵照样算了内存访问一点没省。这个我在前面讲稀疏注意力实现时特别强调过工程上一定要看端到端的延迟不要只看FLOPs。5. 效果表现与适用场景分析5.1 论文报告的数据与我的复现观察理论部分讲完说点更实际的数据。论文里报告的一组典型结果是在ImageNet分类任务上使用DeiT-T作为backbone通过FreeToken动态剪枝和级联退出FLOPs减少了大约45%Top-1精度损失控制在0.8个百分点左右。我在自己数据集上复现时得到的结果比较接近一个四分类工业缺陷检测模型FLOPs减少约40%准确率从98.2%降到97.6%损失不到1个点但端到端延迟在RK3588上缩短了约27%。另一个用语音指令词识别的实验里FreeToken带来了约30%的延迟下降准确率几乎没掉。原因很容易理解语音指令词里大量帧是静音和过渡段这些token的信息量很低被剪掉之后几乎不影响语义判断。但要提醒一句这些数据有很强的任务和硬件相关性。分类任务天然适合动态剪枝因为决定分类结果的往往只是少数几个判别性区域检测任务里背景token也很多但对小目标来说稍有不慎就会把关键区域剪掉。所以不要看到论文数据就无脑上一定要在自己的数据集上实测。我建议先拿一个基准模型把FLOPs减少比例和精度损失画成一条曲线看看trade-off是否符合业务预期再决定是否往工程方向投入。5.2 什么业务适合用什么业务别硬上基于我的实际测试FreeToken特别适合下面这几类业务第一类是视觉分类和粗筛查场景比如工业质检的良品/不良品初筛、安防场景的图片粗分类、无人机航拍图的场景判断。这类任务输入内容变化大简单样本占比高级联退出收益明显。第二类是语音命令词和唤醒词识别静音帧和过渡帧天然适合token剪枝而且对低延迟要求极高FreeToken的两大机制都能发挥作用。第三类是手机端或平板端的多模态助手输入是图像和文本混合内容其中一部分模态的信息量远大于另一部分动态计算的优势能直接转化为功耗和发热的降低。反过来有些业务我不建议上FreeToken。第一类是像素级任务比如语义分割、关键点检测、图像超分。这类任务几乎要求每个输出位置都要感知全图信息剪掉一个token都可能直接影响边缘或小目标。第二类是强时序依赖任务比如视频动作识别里需要精准捕捉连续几帧变化的任务token剪枝会打断时序上下文串联。第三类是模型已经非常小的情况比如只有两三层、几十万参数的极轻量模型剪枝空间本来就小再加上动态分支的开销很可能得不偿失。从框架选型角度讲如果你的业务满足“速度快、精度准、板子内存小”这三个要求中的至少两个FreeToken这套思路就值得试。这也是我复现之后最大的感受边缘侧推理的优化不能光盯着模型参数和算力把计算策略本身做成动态的往往能带来更大的提升空间。最后分享一点个人体会。我在复现FreeToken的过程中最有价值的收获不是某个算子的实现而是确认了“动态稀疏计算”这一整套思路的工程可行性。以前做端侧优化思路基本固定在剪枝、量化、蒸馏这三件套上FreeToken让我看到了另一条路模型不一定要变笨只要学会在简单问题上偷懒就行。如果你手里正好有边缘侧Transformer推理的优化任务建议先从分类或指令词识别这类任务切入用比较小的成本把论文跑通再逐步扩展到更复杂的业务。真遇到问题也别慌动态shape、延迟抖动、算子融合这些坑都是可以逐个解决的关键是先把端到端流程跑起来再谈优化。