拓冰建站拓冰建站
首页 / 资讯中心 / 正文

FreeToken:边缘侧Transformer动态Token稀疏化推理优化实战

最近在啃一篇做边缘侧推理的论文方法名叫FreeToken。第一眼看到这个名字还以为是某个tokenizer的库读进去才发现它是一套专门给边缘设备设计的推理框架核心思路很直接把Transformer里面的token当成一种可以动态分配的计算资源该省的省该放的放。我在边缘侧推理框架这个方向上摸爬滚打了几年以前做优化基本靠剪枝、量化、蒸馏三板斧看到这个方案的时候心里咯噔一下终于有人把“动态跳过无用token”这件事做成了一个工程上能落地的框架。这篇文章不是什么官方教程就是我啃完论文之后的工作笔记再加上我在实际移植和部署过程中踩过的一些坑。如果你是做边缘端AI部署、模型加速或者单纯对Transformer推理性能优化感兴趣的这篇内容应该能帮你少走不少弯路。我尽量用大白话讲清楚FreeToken做了什么、为什么这么做、以及部署时真正要命的那些细节。1. 边缘侧推理的真实痛点与FreeToken要解决的事1.1 边缘侧推理为什么越来越绕不开以前做AI推理大家都习惯把数据传到云端让服务器算完再把结果拉回来。但这两年边缘侧推理的需求明显压过来了核心原因有三个时延、隐私、带宽。拿工业质检来说产线上的相机每秒要拍几十张图如果每张图都传到云端识别等待时间就成了瓶颈而且一旦网络抖动整条产线就得停摆。再比如仓储机器人、车端感知、医疗影像辅助诊断这些场景要么对延迟极其敏感要么数据根本不允许出设备。边缘盒子这几年卖得越来越好本质上就是这些需求在驱动。但问题在于边缘设备的算力增长远远跟不上模型规模的增长。我手头常用的几款边缘盒子算力从几TOPS到几十TOPS不等跑个MobileNet、ResNet系列还算轻松一旦换成Vision Transformer、Swin Transformer这种架构延迟立刻飙升。Transformer对算力的需求主要体现在自注意力机制上序列越长token之间的两两交互就越昂贵这种计算量在服务器上都烧钱放到边缘设备上更是折磨。1.2 现有方案卡在哪既然模型跑不动传统思路就是压缩模型。这个方向我试过不少招说说实际感受。结构化剪枝是最先上手的方法直接把不重要的通道或者头剪掉好处是模型变小了推理确实变快。但问题也很明显剪枝是静态的一旦剪完所有输入都走同样的计算路径。而实际场景里大部分输入里面其实有大量冗余信息静态剪枝没办法针对单张图的内容动态调整计算量。量化是另一个常用手段从FP32压到INT8在边缘盒子上收益非常明显。但量化对动态结构的支持不太好如果模型计算路径本身是变化的量化后的精度掉点就很难控制。蒸馏也试过用小模型学大模型效果不错但训练成本高而且压缩比例到了一定程度就上不去了毕竟模型容量在那里摆着。这些方案的本质都是“把模型变小”但FreeToken的思路不太一样它让模型学会“看情况少算”。同样是跑一张图简单的内容少算一点复杂的内容多算一点计算路径是动态的而不是一次性把模型压死。1.3 FreeToken的出发点把token当成一种资源读论文的时候FreeToken标题里这个“Free”我琢磨了一会儿后来想通了它指的是“释放”——把那些不重要的token释放掉不参与后续计算这样计算资源就free出来了。这个出发点非常贴合Transformer的实际情况。一张输入图片切成一堆patch每个patch对应一个token但并不是所有token都同等重要。比如一张商品检测图背景区域占了很大面积那些背景token从头到尾都没啥贡献又比如一张纹理重复的布料图相邻好几个token长得几乎一样全算一遍纯属浪费更别提很多pipeline里为了对齐batch加进去的padding token完全就是负担。FreeToken的思路就是在每个Transformer阶段用一个轻量的门控机制给token打分保留得分高的、丢掉得分低的。随着网络加深token数量逐层减少形成一个从密到疏的结构。这样一来模型的计算量不再是一个固定值而是跟着输入内容动态变化。我读到这里的时候觉得挺妙因为它避开了剪枝“一刀切”的疼点又能实实在在省算力。2. FreeToken核心设计拆解2.1 整体架构与计算流程FreeToken整体还是以Vision Transformer为主干没有刻意去改动attention的计算公式而是在Transformer的阶段之间插入了token筛选模块。如果用一句话概括它的结构逐步筛选、渐进稀疏。流程大致是这样的图片输入后先切成patch经过stem映射成初始token序列然后进入第一个Transformer阶段正常计算阶段结束后门控模块根据每个token的特征输出一个重要性分数按分数排序保留比例较高的token其余token直接丢掉不进入下一阶段剩下的token继续走后面的阶段重复这个过程。我在笔记里画了个简化的数据流输入序列 → Block1~Block4 → 门控筛选 → 保留部分token → Block5~Block8 → 门控筛选 → ... → 最终分类。实际上FreeToken在每个阶段后面都接了一个门控token数量是阶梯式下降的。这里要注意一个关键设计被丢弃的token是直接“物理删除”而不是像一些稀疏注意力方法那样把不重要的token置零、但还保留在序列里。物理删除的好处是后续阶段的序列长度真的变短了计算量和内存占用都能同步降下来缺点是序列长度动态变化给部署带来了一些麻烦这个后面再说。2.2 token稀疏化的具体实现门控模块的实现是整篇论文里我最关心的部分。它不能太复杂否则筛选操作本身的计算量就能抵消掉省下来的时间也不能太简陋否则重要性分数不准会把关键token扔了。FreeToken选的是轻量方案对每个token的特征做一次线性变换过一层激活函数再过一个线性层输出一个标量分数。这个结构很轻参数量大概只占整个模型的不到2%。打分之后对所有token的分数做排序取前k个保留k由预设的保留比例决定。有两个细节值得细品。第一这里用的是top-k硬选择而不是设一个固定阈值。因为不同输入的特征分布差异很大固定阈值很容易在某种输入上失效要么漏掉重要token要么保留一堆没用的。top-k至少能保证每一层都只保留信息量最大的那部分。第二是梯度怎么传。top-k选择这个操作本身是不可导的直接反向传播会断掉。论文里用的是类似Gumbel-Softmax或者straight-through estimator的思路让门控模块在训练时能接收到梯度信号学会“什么样的token应该被留下”。还有一个有意思的细节门控模块并不直接预测“这个token重不重要”而是预测“这个token对后续任务的贡献有多大”这个贡献是跟任务相关的。同样是背景token如果任务是分类可能果断丢掉没问题但如果任务是细粒度分割背景也是有用的上下文信息。所以FreeToken的门控不是单独训的而是跟主任务一起联合训练让任务损失来定义“重要性”。2.3 训练策略与精度补偿动态稀疏结构有一个通病训练的时候如果直接端到端硬训模型很容易不稳定甚至出现“门控模块越权”的情况比如把所有token都丢了靠一个残差分支瞎猜结果。FreeToken在训练时做了几个约束这个我认为是论文里很关键的贡献。首先是联合训练时的辅助损失。光靠最后的分类损失去约束门控信号太稀疏了。FreeToken在中间阶段也加了辅助分类头让每个阶段的保留token都能直接感受到分类压力这样门控模块的优化会平稳很多。其次是蒸馏损失。这个在论文里没有特别强调但我复现的时候发现它非常关键。用完整训练的teacher模型指导学生模型不仅让学生的输出向teacher靠拢更重要的是一起蒸馏中间特征这样学生模型在token被筛选之后依然能维持跟teacher相似的表达能力。最后是dropout策略。训练的时候保留比例会随机扰动而不是固定一个值。这样做的好处是让模型适应不同计算量的情况避免训练时用0.7的保留比例、部署时换到0.5就崩掉。我在实际做的时候还有一个体会保留比例的设计不能头尾一样大越靠前的层越接近输入像素冗余信息多可以多丢一点越靠后的层token已经比较“精炼”了再丢就容易伤筋动骨。所以我的配置通常是前面几层保留0.3~0.4中间0.5最后一层0.8左右整体下来计算量能省将近一半。2.4 和注意力机制的关系读FreeToken的时候我一直在想一个问题它的门控跟Transformer自注意力里天然存在的attention score是什么关系attention score本身就能反映token之间的关联强度能不能直接用它来做筛选说实话这两种思路各有利弊。用attention score做筛选的优点是几乎零额外参数量缺点是attention模块本身已经被大量任务损失压着再让它同时承担token重要性的职责往往会干扰它对长距离依赖的建模。而且attention score是对所有token两两计算的单看某一行或者某一列怎么聚合才合适本身也是个问题。FreeToken选择用独立门控模块来承担这个职责等于是把“内容理解”和“结构决策”解耦了。attention模块专心建模token之间的关系门控模块专心判断哪些token可以不要。这个设计在工程上是很便利的两个模块可以分别调试门控模块出问题了也不会影响主干特征提取的质量。我在实际部署中还尝试过另一种方案不额外加门控直接把注意力权重的峰值位置作为重要token的依据。效果确实一般主要问题是注意力矩阵要用softmax把概率算出来这一步在边缘设备上本身就是不小的开销而FreeToken的门控只是一个线性层几乎可以忽略不计。3. 部署实操从权重到边缘设备跑起来3.1 部署前的模型准备论文读得再热闹最后得跑起来才算数。我主要是在Jetson Orin和瑞芯微RK3588两个平台上做了验证下面这些步骤是通用的。第一步先把模型权重拿到手。现在网上关于FreeToken的信息挺杂搜索栏里经常蹦出“FreeToken下载”“FreeToken官网”这些词我个人的建议是不要随便去第三方下载站找现成的包源码直接clone下来自己编译权重也尽量用官方仓库或者你信任的模型库。这年头供应链上投毒的事情不是没有能自己编就自己编能校验哈希就校验哈希。第二步把PyTorch模型转成ONNX。这一步有几个注意事项。输入尺寸要固定虽然FreeToken的token数量是动态的但初始输入尺寸最好固定为训练时的分辨率比如224x224或者256x256。转ONNX的时候要把token筛选模块中的top-k、gather这类动态算子显式保留下来不要试图用静态优化去绕否则后面几乎一定会出问题。第三步转成平台专用的格式。Jetson上我用TensorRTRK3588上走RKNN。这一步真正恶心的地方是动态shape支持。FreeToken的序列长度是逐层变化的TensorRT还好一些RKNN对动态shape的支持就相当保守后面单独讲。3.2 关键参数配置参考部署FreeToken的时候最核心的参数就是每个阶段的token保留比例。以ViT-Small为例我整理了一份配置参考你可以根据自己的设备算力来微调。参数项推荐范围说明阶段1保留比例0.3 ~ 0.4越靠近输入冗余越多可以激进一些阶段2保留比例0.5 ~ 0.6过渡阶段保守一点阶段3保留比例0.7 ~ 0.8特征已经精炼少动为妙最终平均保留比例0.45 ~ 0.6整体计算量缩减40%左右batch size1边缘端推理几乎都是batch1动态shape最省心数据精度FP16优先边缘设备FP16收益明显INT8需要充分校准batch size这里多说一句。很多做服务端推理的同学习惯开大batch但边缘端场景绝大多数是单路视频流或者单张图片请求batch1是最实际的配置。FreeToken在batch1下收益最大因为动态token筛选不需要给不同样本对齐序列长度。3.3 性能对比怎么测才可信部署完之后怎么测性能是有讲究的。你要是只盯着FPS这一个指标很容易被误导。我一般会同时看四个维度均值延迟、p95延迟、功耗和精度。均值延迟反映整体吞吐能力p95延迟反映稳定性。FreeToken这种动态结构有个特点不同输入的计算量差异大简单图片可能跑得飞快复杂图片会慢一些。如果只看均值很容易忽略掉那些“倒霉的复杂图片”带来的卡顿。精度的验证用测试集的平均指标还不够我会额外留一批分布比较极端的样本专门看那些难样本的表现。因为动态稀疏最怕的情况就是门控在常规样本上学得很好遇到没见过的分布就开始乱丢token。实测下来FreeToken在Jetson Orin上跑ViT-Small整体延迟相比原始模型大概能降40%~50%精度掉点可以控制在1个百分点以内。在RK3588上效果略差一些主要是因为动态shape对算子融合不友好但也比完全不优化强得多。4. 常见问题与排查技巧实录4.1 推理精度掉得厉害这是很多人上来就会遇到的问题我在第一版测试时也踩了。明明训练的时候指标还行一部署到边缘设备上精度就崩了。先排查保留比例是不是调得太激进了。训练时保留比例是0.6部署时为了追求速度改成0.3精度稳掉。建议先对齐训练时的配置确认能复现出训练精度再一步步收紧。再排查量化。FP16一般影响不大但INT8就不好说了。FreeToken的动态结构对量化误差更敏感尤其gather和top-k操作选出来的token一旦因为量化噪声导致排序变化后续所有层都会受影响。解决思路是增加校准集规模同时考虑对门控模块单独保持FP16只对主干做量化。最后要检查训练和推理的一致性。训练时的随机扰动到推理时要彻底关掉这个我犯过糊涂训练时保留比例随机扰动没关部署后门控看到的分布跟训练完全不一样掉点特别明显。4.2 延迟不降反升这个坑更隐蔽。理论上token保留比例从0.7降到0.4计算量应该大幅下降结果一测延迟反而涨了。如果你的延迟掉到这份上先别急着骂框架大概率是下面几个原因。最典型的是mask计算和gather操作落在了CPU上。有些转换工具遇到动态算子会保守地把链路切开一部分跑CPU一部分跑NPU/GPU中间的拷贝开销非常吓人。我遇到过gather操作本身只有几百微秒但因为CPU和GPU之间的数据来回搬运整个延迟多了几十毫秒。排查方法很直接用profiler看每一层的耗时分布。如果发现token筛选模块的耗时占比在总耗时里超过10%那就是这个模块没有真正融合进主计算图里。解决方式是调整转换工具的算子支持表或者手动重写一部分图结构。还有一个容易忽略的问题内存碎片化。动态shape导致每个batch分配的显存/内存大小都不一样频繁申请释放会产生碎片。边缘盒子内存本来就紧张碎片一多系统开始换页延迟自然暴涨。建议提前分配好最大序列长度的缓存区或者用内存池来管理。4.3 算子兼容性与动态形状的坑最后是兼容性问题。边缘设备上跑动态结构最痛苦的就是NPU/DSP对动态shape的支持普遍不完善。top-k和gather这类操作在服务器级GPU上是常规操作但在边缘设备上有可能压根没有对应算子或者性能极其拉胯。我之前在RK3588上调模型RKNN对gather的支持就非常有限最后只能把token筛选部分放在CPU上执行但通过优化算子实现、减少CPU和NPU的交换次数才把延迟压下来。内存布局问题也要留意。Transformer在GPU上通常用NHWC或者NCHW布局但NPU对layout的要求五花八门。FreeToken动态改变token数量后内存布幅经常要重新排布这个转换开销很容易被忽略。我的建议是尽可能让token筛选发生在比较靠前的阶段这样后面所有层的序列长度都短了反而能掩盖掉一部分转换开销。再有一个就是模型导出时遇到的不确定性问题。同一个ONNX文件在不同版本转换工具下的结果可能不一样动态shape的解析尤其脆弱。我遇到过导出时提示成功但实际推理时某个阶段的token数量变成了负数排查半天发现是工具对top-k输入shape的推断有问题。遇到这种情况别硬调试着把保序操作换成等价的固定shape表达式往往能绕过。5. 一些个人体会和后续打算FreeToken这篇论文让我印象最深的一点是它把“动态计算”和“工程落地”之间的距离拉近了一步。以前我总觉得动态稀疏结构是学术界用来刷点的玩具工业落地还是得靠静态压缩但FreeToken在部署层面做了很多扎实的工作让我改变了这个看法。写这篇文章的时候我注意到网上已经有不少人在讨论FreeToken的部署教程甚至有人用Codex这类工具来辅助安装和调试。我的观点是工具永远是辅助真正让你解决问题的是你对模型结构的理解和对部署平台脾气的把握。Codex能帮你写脚本、查文档但如果你不知道FreeToken的门控是top-k硬选择、不知道序列长度动态变化会带来内存碎片化问题再智能的工具也救不了你。后续我准备在更多模型上跑FreeToken的验证包括检测类模型和一些多模态模型。token稀疏化在分类任务上效果不错但在检测任务上很多“看起来不重要”的背景其实是上下文关键怎么平衡就值得再琢磨。有新的结果我再来更新这篇笔记。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门