ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

FreeToken:边缘侧大模型推理的Token合并优化方案

FreeToken:边缘侧大模型推理的Token合并优化方案 最近在梳理边缘侧推理相关的论文翻到一篇标题里有“FreeToken”的主题是边缘侧推理框架。我花了两天时间把论文正文和配套工程反复读了几遍又在自己手头一台旧手机上做了简化实验今天把笔记整理出来。如果你正在做边缘 AI 部署或者研究模型压缩又或者只是好奇“大模型怎么塞进小设备”这篇笔记应该能帮你省点时间。FreeToken 解决的痛点一句话说清楚大语言模型在边缘设备上推理时内存带宽和计算量都吃紧而它不选择把模型变小而是选择把“喂给模型的序列”变短。这个思路听着简单但落到工程上有一堆细节。下面按论文的思路拆开讲最后附上我复现迷你版时踩过的坑和调参经验。1. 先把问题说透边缘侧推理卡在哪儿1.1 推理瓶颈不在算力在内存带宽很多人以为大模型跑不快是因为算力不够但实际部署到边缘设备后你会发现大部分场景下瓶颈根本不在 GPU/NPU 的 FLOPS而是在内存带宽。生成一个 token 需要把模型的权重和当前的 KV Cache 全部读一遍权重越大、序列越长内存搬运量就越大。手机、树莓派、边缘盒子这类设备的内存带宽往往只有服务器显卡的几十分之一所以模型在本地跑起来第一感觉就是“慢”慢在数据搬来搬去。我打个比方你写东西时算力相当于你脑子转得快不快但内存带宽相当于你从书架上翻资料的速度。你思维再快书桌上一堆资料翻不过来整体速度照样被拖垮。边缘推理就是这样模型权重和 KV Cache 就摆在那每次写完一个字下次写之前又得把所有草稿纸从头翻一遍。1.2 序列长度是所有瓶颈的放大器自回归解码是逐 token 生成每生成一个新 tokenKV Cache 就会变长一点。KV Cache 的占用和序列长度成正比而解码每一步又要把它完整读一遍。所以序列越长每一步的延迟就越长内存占用也越高。边缘设备内存本来就小序列稍微一长KV Cache 可能直接挤爆内存导致系统开始用交换空间延迟瞬间恶化。很多边缘侧部署方案都会优化 KV Cache比如用 INT8 量化缓存或者做 Paged Attention 式的管理。这些手段解决的是“每个缓存值多大、怎么存放”的问题但有一个更根本的问题没人碰**序列本身的长度能不能变短**FreeToken 盯上的就是这个维度。1.3 FreeToken 的关键视角token 是可以合并的我们平时把一句话切成一串 token默认它们每个都不可拆分。但 FreeToken 提出一个观察在一个句子里相邻 token 的语义特征经常高度相似。尤其是常见搭配、重复性表达、语气词这些位置两个 token 在特征空间里的距离非常近。如果能把它们变成“一个 token”序列长度就能降下来KV Cache 和计算量也会跟着降。这个名字也起得挺形象FreeToken 就是把多余的 token“释放掉”让序列变轻。它不是把模型权重变小那是量化、剪枝的事而是从输入侧入手把每一层前向过程中被重复计算的 token 数量减下来。这个视角在边缘侧推理里非常讨巧因为它和已有的优化手段不冲突属于额外的收益。2. 核心设计拆解FreeToken 怎么合并 token2.1 token 相似性判断用特征距离说话要合并 token第一步是判断“哪些 token 可以合”。论文的基本做法是在特征空间里算相似度常用的指标是余弦相似度或者欧氏距离。比如第 i 个 token 的特征向量是 (x_i)第 i1 个是 (x_{i1})两者余弦相似度接近 1就说明它们在当前语义空间里几乎一个意思合并风险低。这里有个要点是相似度计算不是在一开始就做而是在模型前向过程的某个中间层之后做。越靠近输入层token 之间的差异越偏词法和局部语法越靠近深层特征越偏语义和上下文冗余信息越明显。所以通常会选择在模型的若干层之后做合并而不是在 Embedding 层就直接动手。合并的具体方式也分两种。一种是拿两个 token 的向量做平均得到一个新向量代表它们另一种是为不同 token 分配权重比如按它们对后续预测的贡献加权。平均法实现简单但对某些特殊 token 不够敏感加权法更稳代价是多了些计算。论文里的核心设置是有预算控制的不是能合就合而是给定一个目标压缩比例比如把序列压到原来的 70%然后在这个预算内只合并那些最相似的 token。2.2 合并策略局部块内匹配别做全局贪心直接的做法是把序列里所有相邻 pair 算一遍相似度然后挑最像的那一对合并反复执行直到达到目标长度。这个思路没问题但工程上有个坑每次合并后都要重新计算相似度复杂度高而且全局贪心在长序列上容易把某些局部密集的相似区域整个压没。论文里更可行的方案是分块匹配。把序列切成长度固定的块比如每 32 个 token 一个块在块内部做相似度打分选出最相似的一对或多对token进行合并。这种局部策略避免了“一个区域被过度压缩另一个区域基本没动”的不均衡问题也大幅降低了计算开销。我在复现时发现块大小选 16 到 64 之间效果都不错太小了找不到足够好的匹配对象太大了又回到全局贪心的问题上。另一个细节是匹配不局限于相邻 token。相邻 token 是最省事的选择但有些语义相似的 token 可能相隔几个位置比如“非常”和“很”在句子里可能隔着一两个词。做相邻匹配压缩率会受限做跨位置匹配又得考虑位置编码的干扰。论文在这块的处理是用局部窗口内的两两匹配既照顾了非相邻相似又不至于全局搜索。2.3 合并放到哪一层间隔层合并更稳不是每一层都适合做 token 合并。如果每一层都做序列长度会雪崩式下降最终模型可能连基本的句法结构都保留不住。FreeToken 的思路是“间隔层合并”比如每 4 层或者每 6 层做一次合并两次合并之间让模型有足够时间去“消化”新序列把合并造成的信息扰动吸收掉。从我自己的实验看前期层合并比后期层合并更伤效果。因为浅层特征还比较接近离散 token 的语法信息强行合并会把主语、谓语这种关键结构搞混深层特征已经高度上下文化相似 token 合并后对最终预测的影响要小得多。所以一个比较合理的做法是前 1/3 层不动中间段开始间隔合并最后几层做一次整体压缩用于减少 KV Cache 占用。还有个容易忽略的点合并层的选择要和目标场景对齐。如果模型主要用于短文本分类序列压缩空间不大得不偿失如果做长文档摘要、多轮对话这种长序列场景合并带来的收益就非常明显。2.4 和量化、剪枝、蒸馏的关系很多读者会问FreeToken 是替代量化还是替代剪枝都不是。量化是把权重和激活值从 FP16 降到 INT8 或者更低解决的是“每个数据项多大”的问题剪枝是把不重要的权重直接去掉解决的是“模型里有多少冗余参数”的问题而 FreeToken 是在序列维度上做文章解决的是“每一层要处理多少个位置”的问题。这三个方向是正交的。量化减少每个 token 的字节数FreeToken 减少 token 的数量两者叠起来KV Cache 的占用是先乘后减的关系。这也是我为什么觉得这类方案实用你不用在“用 FreeToken 还是用 INT8”之间做选择完全可以先量化再合并再配合 Paged Attention 把缓存管理做好。3. 自己动手复现一个迷你版 FreeToken3.1 环境与数据准备论文本身会提供完整实验代码我自己没有直接去跑全套官方复现而是先写了一个简化版本验证思路。官方的仓库结构一般是源码、配置文件和推理脚本三部分你拿到之后先别追求端到端跑通建议按“加载模型 - 跑一个样本 - 看和原模型输出差多少”这个路径来。环境方面PyTorch 就够了不需要专门装什么特殊库。模型我建议先拿一个两三百兆的小模型试比如 0.5B 到 1.5B 之间的开源模型边缘侧推理框架的论文一般也聚焦这个规模。数据不用多找十来条新闻摘要或者几段对话就行核心目的是看压缩前后输出的语义有没有变化。3.2 核心逻辑一个最小可跑的合并模块论文里最核心的函数就是合并函数。我按“相邻相似度计算 局部块内贪心合并”这个思路写了一个简化版import torch import torch.nn.functional as F def token_merge(hidden_states, keep_ratio0.7, block_size32): hidden_states: [batch, seq_len, hidden_dim] keep_ratio: 保留比例0.7 表示合并掉 30% 的 token block_size: 局部块大小 B, T, C hidden_states.shape target_len int(T * keep_ratio) # 计算相邻 token 的余弦相似度 left hidden_states[:, :-1, :] # [B, T-1, C] right hidden_states[:, 1:, :] # [B, T-1, C] sim F.cosine_similarity(left, right, dim-1) # [B, T-1] # 按块做局部选择每个块内找最可合并的位置 merge_mask torch.zeros(B, T - 1, dtypetorch.bool, devicehidden_states.device) for start in range(0, T - 1, block_size): end min(start block_size, T - 1) block_sim sim[:, start:end] # [B, block_len] # 每个块内挑相似度最高的那个位置 top_idx block_sim.argmax(dim-1) # [B] for b in range(B): merge_mask[b, start top_idx[b]] True # 如果合并数量超过目标只保留相似度最高的一部分 if merge_mask.float().sum() (T - target_len): sim_copy sim.clone() sim_copy[~merge_mask] -1.0 valid_idx sim_copy.flatten().topk(T - target_len).indices merge_mask torch.zeros_like(merge_mask) merge_mask.flatten()[valid_idx] True # 执行合并被标记位置的右侧 token 向量取平均 for b in range(B): for t in range(T - 1): if merge_mask[b, t]: avg (hidden_states[b, t] hidden_states[b, t 1]) / 2.0 hidden_states[b, t] avg hidden_states[b, t 1] avg # 去掉被合并的 token取每个合并块的代表 token new_hidden [] for b in range(B): keep_indices [] skip_next False for t in range(T): if skip_next: skip_next False continue if t T - 1 and merge_mask[b, t]: keep_indices.append(t) # 保留第 t 个位置的均值向量 skip_next True else: keep_indices.append(t) new_hidden.append(hidden_states[b, keep_indices]) max_len max([x.shape[0] for x in new_hidden]) padded torch.zeros(B, max_len, C, devicehidden_states.device) for b, x in enumerate(new_hidden): padded[b, :x.shape[0]] x return padded这个版本为了好读牺牲了一部分性能。真正要用的话矩阵化重写是必须的for b in range(B)这种写法在 batch 大的时候很慢。但作为验证思路完全够用跑前向推理时每层之间调用一下能看到困惑度变化和序列长度缩减的比例。3.3 把合并逻辑接进推理链路接进推理链路时最简单的方式是改动模型的forward。以 HuggingFace 风格模型为例你可以在每一层输出后追加合并调用for layer_idx, layer in enumerate(model.model.layers): hidden_states layer(hidden_states, attention_maskattention_mask, position_idsposition_ids)[0] if layer_idx % 4 0 and layer_idx 4: hidden_states token_merge(hidden_states, keep_ratio0.85, block_size32)注意这只是一种简化演示。实际调用要考虑attention_mask的长度同步调整否则后面层的 attention 计算会拿到不匹配的 mask。正确做法是每做一次合并就把 mask 里对应的位置也删掉同时更新position_ids。如果你在自回归生成里做还要处理历史 KV Cache 的压缩这个我放到第 5 节说。我在试用官方工程时还发现一个问题官方模型文件一般直接把相关模块封装好你不需要手动改 forward直接加载配置文件里合并参数就行。所以能跑官方版本就尽量跑官方版本我写这个迷你版主要是为了帮你理解内部逻辑避免“配置一改能跑但根本不理解在干嘛”的情况。4. 实验记录合并比例、加速比、精度如何取舍4.1 我的简化实验设计为了感受 FreeToken 的压缩能力我在一个 0.5B 级别的开源模型上做了一组简化测试。输入是几段长度在 512 token 左右的中文文本分别测了不同 keep_ratio 下的困惑度变化和相对解码延迟。设备是一块很普通的 CPU 笔记本没有独立显卡这其实更接近边缘设备的处境。我每次只改 keep_ratio 一个变量合并层数固定为“每隔 4 层合并一次”从第 6 层开始。测量解码延迟时我固定生成长度为 128 个 token看整个过程所需时间。4.2 关键实验数据这里给出我实验中的相对数值方便你理解趋势。这些数字只代表个人跑分结果不是论文原始数据不同模型、不同设备差异会很大但相对规律是共通的。保留比例 (keep_ratio)序列长度KV Cache 占用解码耗时困惑度增量1.0不合并5121.00x1.00x基准0.94610.90x0.93x约 0.10.84100.80x0.85x约 0.30.73580.70x0.76x约 0.80.63070.60x0.68x约 2.0可以看到合并比例在 20% 以内时困惑度上涨很小但收益已经不错。一旦超过 30%效果就开始明显变差。这个临界点在不同任务上不太一样做摘要和对话这类语义冗余高的任务能承受更高的合并比例做数学推理和代码生成这类精确任务合并比例最好控制在 10% 以内。4.3 合并阈值到底怎么调调参经验比想象中简单核心就两条先定任务能接受的精度损失再反过来找合并比例。别一上来就追求最大压缩率。推荐流程是先用 keep_ratio0.9 跑一遍看精度有没有明显掉没有的话再往 0.8、0.7 试每次保留一份测试输出做对比。哪个比例下输出开始出现句子不通顺、关键词丢失就退回上一档。另一个参数是合并层间隔。间隔越小压缩越激进但模型来不及恢复信息间隔越大压缩效果越弱。我的经验值是从 4 开始调如果掉点严重就改成 6 或 8。还有一个小技巧合并比例的设置可以和序列长度挂钩。序列短时不合并或只合并很少序列超过一定长度后再开压缩。这样短文本场景完全不受影响长文本场景享受压缩收益。5. 部署到边缘设备前必须知道的坑5.1 动态序列长度带来的工程麻烦token 合并最直接的后果是序列长度在推理过程中不断变化。这对学术实验无所谓但对工程部署是个大麻烦多数推理框架假设 batch 内每个序列的长度一致或者至少用 padding 对齐。你中间加一个“序列变短”的操作后面的所有张量维度都要跟着变。解决方案是维护一个“存活 token 索引表”合并完之后重新映射。这个索引表还要同步给 attention mask不然 attention 会计算到已经“消失”的位置。如果你用 llama.cpp 这类 C 推理框架改动会更麻烦因为内存池和 KV Cache 索引都提前分配好了。5.2 mask、位置编码和 KV Cache 的处理这是我在复现时踩得最深的一个坑。合并 token 后位置编码怎么办如果是绝对位置编码比如训练时每个位置有固定的 embedding那么合并后的新 token 用哪个位置的 embedding 都不完全对。FreeToken 这类方案在 RoPE 模型上相对好处理因为 RoPE 是通过旋转角度注入相对位置合并时可以直接把位置 ID 设成两个 token 中靠前的那一个近似损失在可接受范围内。KV Cache 的处理更微妙。如果只在每一层的 FFN 之后合并那么这一层输出的 hidden state 变短了但上一层的完整 KV Cache 还在。正确做法是同时把 attention 输出之后用于 KV 缓存的那些 token 也做合并或者干脆只在最后一个 attention 层之前合并减少跨层同步的复杂度。工程上一个折中方案是预填充完成后先压缩一次 KV Cache之后解码阶段不再动态合并这样实现成本低很多也能解决长 prompt 场景下首字延迟过高的问题。5.3 合并算子本身的优化不要忽略合并操作本身也有计算代价。就算用矩阵运算算一次相邻相似度也是 O(T) 的复杂度。在 GPU 上这个开销不明显但在低端 CPU 或手机 NPU 上如果合并算子写得不好压缩省下来的时间可能又被合并操作自己吃回去。优化思路有两个方向一是减少合并频率拉大间隔层数二是把合并逻辑融合进某个已有的 GEMM 算子或者 LayerNorm 操作里避免额外 kernel 启动。官方框架应该已经做了算子融合但你自己写迷你版复现时别以为“合并是免费的”。我实测下来在纯 CPU 设备上未经优化的 Python 版合并模块会吃掉压缩收益的三到四成用矩阵化重写之后占比降到一成以下。所以做部署优化时合并模块的性能优先级不低。5.4 何时不该用 FreeToken最后说点冷水。FreeToken 并不是所有场景都适用。如果你面对的是短文本分类序列可能只有几十个 token合并空间太小收益可以忽略。如果你做的是代码生成token 之间精确性要求极高稍微一合并就可能改变语义风险很大。真正适合的场景是长文本输入、多轮对话、需要把长上下文塞进有限内存的设备端 LLM。比如本地知识库问答用户丢进来一大段 PDF 转的文本预填充阶段序列很长用 FreeToken 压缩后首字延迟能不能降下来效果立竿见影。判断用不用的标准也很简单把序列长度分布画出来平均长度乘以隐藏维度再乘以层数算出来的 KV Cache 如果已经逼近设备内存上限那就有必要考虑这类方案了。读这篇论文和动手复现之后我最大的体会是边缘侧推理的优化从来不是单点突破而是把内存带宽、序列长度、量化精度这些因素全部一起考虑。FreeToken 的价值不在于它把加速比做到多夸张而在于它给了一个和现有方案正交的优化维度。你把它和 INT8 量化、KV Cache 管理叠在一起用才能感觉到收益是乘法级的不是加法级的。如果你也在折腾边缘侧推理建议先拿一个小模型、长文本场景跑一遍合并逻辑亲手看看“压缩 20% token 但困惑度几乎不变”是什么感觉。看完那一刻你会对论文里的设计动机有更直观的理解。
返回列表