ARTICLE DETAIL

资讯详情

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

NVIDIA 让 KV Cache 跨模型搬家:一次矩阵求解,32K 预填充从 6975 毫秒压到 278 毫秒

NVIDIA 让 KV Cache 跨模型搬家:一次矩阵求解,32K 预填充从 6975 毫秒压到 278 毫秒 一次模型切换要付两遍钱这是很多团队上线多模型编排之后才发现的账。用户的对话已经跑了三十轮上下文攒到两万多个 token路由层决定把这一轮交给更大的模型处理。大模型接手之后的第一件事不是生成回答而是把这两万多个 token 从头到尾再读一遍。这一遍读取叫预填充它的唯一产物就是 KV Cache而这份缓存小模型刚刚才算过一次。NVIDIA 的 Taekyung Heo、Rasoul Shafipour 和 Bita Darvish Rouhani 等人在 8 月 4 日挂到 arXiv 上的论文把这笔重复开销当成了一个可以直接省掉的目标。编号是 2608.03893标题叫跨模型 KV Cache 迁移副标题点明了方法的性质一个用于预填充复用的闭式线性映射。闭式这个词是整篇论文最值得注意的地方它意味着不需要反向传播不需要训练循环一次矩阵求解就拿到映射器。他们在 Qwen3 的 14B 到 32B 这一对上把 32K 上下文的预填充从6975 毫秒压到了 278 毫秒。这个结论听起来太便宜了所以值得先把它的边界说清楚。论文只做族内迁移也就是同一个模型家族里不同尺寸之间的搬运跨家族的情况明确列进了未来工作。六对模型里只有四对拿到了可用的精度另外两对垮得非常彻底。速度收益和精度损失是同一件事的两面谁都不能单独拿出来说。把这两面一起读完才能判断这套方法能不能进你的推理栈。预填充这笔钱为什么一直在重复付要理解论文省掉的是什么先要看清预填充在推理链路里的位置。模型生成回答分成两个阶段预填充负责把输入的每一个 token 过一遍全部网络层算出每层每个注意力头的键和值堆起来就是 KV Cache。解码阶段每生成一个新 token只需要读这份缓存再追加一行不必重算历史。所以预填充是一次性的重活它的成本随模型规模和提示长度一起涨。业界对这笔成本早有对策叫前缀缓存。同一个模型收到共享前缀的多个请求时把前缀那段 KV 存下来复用后面的请求直接从断点接着算。这套机制在单模型内部工作得很好vLLM 和 SGLang 都做了成熟实现系统提示词和长文档这类固定前缀能省下大量重复计算。它的前提是接收方和生产方是同一个模型缓存的数值格式完全对得上。一旦换了模型这个前提立刻失效前缀缓存就退化成一份没人能读的字节。问题恰好出在这个前提上。生产部署已经不是单模型系统了论文点出三种常见形态成本质量级联、对话中途切换、以及请求路由。这三种形态的共同点是同一个会话会在不同尺寸的模型之间来回跳而每一次跳跃都让接收方从第零个 token 开始重算。长会话让提示越来越长多模型编排让切换越来越频繁两个趋势叠在一起预填充的账就滚了起来。直觉上这件事似乎无解因为源模型和目标模型的层数不同、隐藏维度不同、注意力头配置也可能不同。一份 14B 模型算出来的缓存直接倒给 32B 模型形状根本对不上张量维度在第一步就会报错。既然对不上就只能重算这是过去默认的结论也是各家推理框架当前的实际行为。论文的切入点正是这里它没有接受这个结论而是把问题重新定义了一次。既然预填充的唯一产物就是 KV Cache那么跳过预填充就不是一个计算问题而是一个表示问题。手上已经有一份缓存需要的是把它翻译成接收方期待的格式。翻译的可行性取决于两份缓存之间有没有可利用的结构而不取决于两个模型的层数是否相等。这个转向让整件事从不可能变成了可测量。已有的路都要训练这条不用跨模型复用 KV 并不是全新的想法论文在相关工作里列了四条已有路线。C2C 为每一对模型训练神经融合器LatentAlign 学习把每个模型映射进共享潜空间的适配器IAM 替换的是小模型的注意力模式而不是 KV 数值DroidSpeak 则要求两个模型架构完全一致。四条路各有各的适用面但它们共享同一个成本。这个共同成本就是梯度训练或者强架构假设。要么你得为每一对模型跑一轮反向传播要么你得接受两个模型必须长得几乎一样。前者意味着上线一个新尺寸就要重训一次映射训练数据、显卡和调参时间都得重新排期后者意味着方法在真实的家族内部就用不上因为同一家族的不同尺寸本来就有架构差异。论文明确说据他们所知没有工作研究过跨模型 KV 的关系是否简单到可以用闭式映射解决。这句话既是他们的贡献声明也是一个被长期忽略的检查项。这个问题问得很朴素也正是全文的支点。如果关系足够线性那么最小二乘就够了不需要神经网络不需要优化器也不需要学习率。线性方法有一个工程上的巨大好处就是可以用一次矩阵求解拿到解析解成本可预测行为可解释。于是论文先不设计模型而是去测量这个关系到底有多线性。还有一条容易混淆的边界需要提前划清。模型内部的跨层 KV 共享研究的是同一个模型里不同层之间的冗余前缀缓存研究的是同一个模型不同请求之间的复用。这篇论文做的是把缓存数值从一个模型搬到另一个模型三者是正交的可以叠加使用。搞清这一点才不会把它误当成又一个层间共享的变体。先量线性再谈方法论文没有一上手就设计映射器而是先做了一件更基础的事用最笨的办法测量两个模型的 KV 之间有多少线性关系。做法是取源模型的某一层、目标模型的某一层、某一个注意力头在 token 级别拟合一个最普通的最小二乘回归看能解释掉多少方差。这个探针不带任何技巧只有一个源层没有正则化就是教科书上的普通线性回归。先测量再设计这个顺序让后面每一个组件都有实测依据而不是凭直觉堆上去的。结果比预期的强。在 Qwen3 的 14B 到 32B 上单个源层就解释了目标键的 56% 方差值的 32%。这已经不是噪声水平了说明两个模型对同一段文本算出的中间表示共享着相当大的一块结构。当他们把多个源层的信息一起用上时键的解释度升到79%值升到65%。最好的单个源层与目标层组合在去掉位置编码之后达到了 0.81。把所有层组合画成热力图之后四个规律很清楚。第一线性拟合度沿对角线明显高于零也就是源模型的浅层对应目标模型的浅层深层对应深层这个直觉是对的。第二两个模型越接近对角线越锐利架构和深度差距越大这条亮带就越弥散。第三旋转位置编码会污染拟合把它剥掉之后对角线普遍变得更清晰。第四键比值更好预测两者的解释度通常差着约 0.2。第四条规律其实有直观解释。键参与的是注意力打分它的作用是决定看哪里这个决策在同一家族的模型之间高度一致值承载的是看到了什么内容随着模型容量增加内容表示的丰富度差异更大所以更难线性预测。这个差异后面还会再出现一次在最终的精度结果里键的映射质量对下游表现的影响远大于值。换句话说这套方法的成败主要押在键上值的误差有更大的容忍空间。还有一个关键的测量结论决定了方法的形状。既然单个源层只解释了 56%那多少个源层才够论文用贪心前向选择做了实验每一步加入能让解释度提升最多的那个源层。答案是信息确实分散在多个源层里只用一个源层键只能拿到全部源层版本的 66%值只有 42%。从一层加到四层收益最大加到六层左右就基本饱和了。三个组件一次矩阵求解基于上面的测量论文的映射器由三个组件拼成每一个都直接对应一条实测规律。第一个是逐头岭回归为目标模型的每一个层、每一个注意力头单独拟合一个独立的线性映射。这样做绕开了源和目标在层数、头维度、头数量上的全部不匹配问题因为映射是按目标侧的结构来组织的。目标要多少层多少头就拟合多少个独立的小回归源侧的形状差异被吸收进了输入维度里。第二个组件是跨层源选择对应信息分散在多层这条发现。对每个目标层按解释度挑出最有预测力的前 k 个源层把它们的键值特征拼接起来当作输入。同一个目标层内的所有头共享同一组选中的源层这允许跨头的信息流动。k 是唯一需要按模型对扫描的超参数论文在 1 到全部之间扫了十一个取值。第三个组件是内容空间映射处理位置编码的污染问题。旋转位置编码给键施加了一个依赖位置的旋转缓存里存的是旋转之后的键。论文的做法是先用逆旋转把位置信息剥掉在无位置的内容空间里做映射映射完再用接收方的位置编码重新旋转回去。因为旋转矩阵是正交的逆运算精确且几乎不花成本。这一步的价值不在短文本上的精度而在长度泛化。论文坦白说直接在带位置编码的键上拟合在短上下文基准上的表现落在噪声范围内看不出差别。但那样拟合出来的权重被绑死在校准时见过的 1024 token 位置分布上换到 32K 就不成立了。解耦之后的公式天然支持任意长度这才是服务长提示的前提。拟合本身简单到有点反直觉。校准数据只用 500 条 FineWeb-Edu 序列每条 1024 个 token按步长 4 下采样每个目标头拿到约 12.8 万个 token 的观测。然后解一个带 Tikhonov 正则的正规方程正则系数取 0.01。用正则而不是纯最小二乘的原因很实际k 较大时特征维度能到几万而被选中的源层本来就相关矩阵接近奇异加一点对角项能稳定求逆而几乎不引入偏差。整个拟合在单个八卡 H100 节点上耗时约 47 到 87 分钟全程没有梯度训练。有意思的是拟合时间随目标头数的增长是次线性的因为主导时间的那个矩阵乘积只需按目标层算一次该层的所有头共享。最终映射器的参数量在 10.1 亿到 33.6 亿之间存储 4 到 12 GB。这个体积不算小但相对于它服务的模型规模是可以接受的。把这套流程跑一遍论文的三个组件描述得很清楚但读公式和跑代码是两件事。下面这段程序把完整链路实现了一遍构造一个人造的两模型家族其中目标模型的 KV 是若干源层的线性读出加噪声然后用论文的三个组件去恢复它并和缺少组件的朴素版本对照。它只依赖 NumPy可以直接运行。我先把它跑通再写进文章紧接着给出的输出是本机真实运行结果。#!/usr/bin/env python3 Cross-model KV transfer: closed-form ridge mapper (arXiv 2608.03893 core path). import numpy as np D_H, N_KV, T_CAL, LAM 64, 4, 4096, 0.01 rng np.random.default_rng(0) def rope_matrix(pos, d_h, base10000.0): Build per-position rotation angles for half the head dim. inv base ** (-np.arange(0, d_h // 2) / (d_h // 2)) ang np.outer(pos, inv) return np.cos(ang), np.sin(ang) def rope_apply(k, cos, sin, inverseFalse): Rotate (or un-rotate) keys in place of the models own RoPE. a, b k[..., ::2], k[..., 1::2] s -sin if inverse else sin out np.empty_like(k) out[..., ::2] a * cos - b * s out[..., 1::2] a * s b * cos return out def fit_ridge(x, y, lamLAM): Closed-form W (XtX lam I)^-1 XtY on centered features. xm, ym x.mean(0), y.mean(0) xc, yc x - xm, y - ym g xc.T xc lam * np.eye(xc.shape[1]) w np.linalg.solve(g, xc.T yc) return w, ym - xm w def r2(y, yh): ss_res ((y - yh) ** 2).sum() ss_tot ((y - y.mean(0)) ** 2).sum() return 1.0 - ss_res / ss_tot def select_top_k(src_layers, tgt, k): Rank source layers by single-layer R2, keep the k most predictive. scored [] for i, s in enumerate(src_layers): w, b fit_ridge(s, tgt) scored.append((r2(tgt, s w b), i)) scored.sort(reverseTrue) return [i for _, i in scored[:k]] def attn_output(q, k, v): Single-head attention output, used as the fidelity diagnostic. logits q k.T / np.sqrt(q.shape[-1]) logits - logits.max(-1, keepdimsTrue) p np.exp(logits) return (p / p.sum(-1, keepdimsTrue)) v # --- synthetic two-model family: target KV is a linearnoise read of source layers --- n_src 6 src [rng.normal(size(T_CAL, D_H)) for _ in range(n_src)] truth {i: rng.normal(size(D_H, D_H)) * 0.3 for i in (1, 3, 4)} tgt_content sum(src[i] w for i, w in truth.items()) 0.05 * rng.normal(size(T_CAL, D_H)) pos np.arange(T_CAL) cos, sin rope_matrix(pos, D_H) src_roped [rope_apply(s, cos, sin) for s in src] # what the cache actually stores tgt_roped rope_apply(tgt_content, cos, sin) # 1) strip RoPE so the fit is position-free src_stripped [rope_apply(s, cos, sin, inverseTrue) for s in src_roped] tgt_stripped rope_apply(tgt_roped, cos, sin, inverseTrue) # 2) cross-layer selection, then one ridge solve on the concatenated features picked select_top_k(src_stripped, tgt_stripped, k3) X np.concatenate([src_stripped[i] for i in picked], axis1) W, B fit_ridge(X, tgt_stripped) # 3) map, then re-encode with the receivers RoPE pred_stripped X W B pred_roped rope_apply(pred_stripped, cos, sin) # baseline: no RoPE stripping, single source layer w1, b1 fit_ridge(src_roped[picked[0]], tgt_roped) naive src_roped[picked[0]] w1 b1 q rng.normal(size(256, D_H)) cs lambda a, b: float((a * b).sum() / (np.linalg.norm(a) * np.linalg.norm(b))) gt_out attn_output(q, tgt_roped, tgt_content) print(selected source layers:, sorted(picked), | ground truth:, sorted(truth)) print(fR2 full pipeline : {r2(tgt_roped, pred_roped):.4f}) print(fR2 naive (k1,RoPE): {r2(tgt_roped, naive):.4f}) print(fattn-output cosine full : {cs(gt_out, attn_output(q, pred_roped, tgt_content)):.4f}) print(fattn-output cosine naive: {cs(gt_out, attn_output(q, naive, tgt_content)):.4f})在我本机跑出来的结果是这样的跨层选择准确挑回了真实的源层 1、3、4完整流程的解释度 0.9999注意力输出余弦 0.9997而砍掉位置解耦和跨层选择的朴素版本解释度掉到 0.0341余弦只有 0.1835。人造数据当然比真实模型友好这里要看的不是绝对数值而是三个组件缺失时的塌陷幅度它和论文消融实验的方向完全一致。论文在真实模型上的对应数字是键解释度从 0.79 掉到 0.56量级比这里温和得多但方向一样。差别在于真实模型的 KV 关系只是接近线性而这段代码里的关系本来就是线性构造出来的。这段代码里有两处细节值得单独指出。一个是 rope_apply 用同一个函数处理正向和逆向旋转只翻转正弦项的符号这利用了旋转矩阵正交的性质也是论文说逆运算精确且几乎免费的原因。另一个是 fit_ridge 先对特征和响应做中心化再求解截距通过均值反推回来这样正则项只作用在斜率上不会把截距一起压向零。这两处都不是性能优化而是保证数学上成立的必要步骤去掉任何一个结果都会偏。四对能用两对报废论文在三个家族的六对模型上做了评测全部是键值匹配的对也就是源和目标共享注意力头数和每头维度。评测基准有五个ARC-Challenge、HellaSwag、WinoGrande、五样本 MMLU、以及八样本带思维链的 GSM8K另外用 WikiText-2 困惑度和 CoQA 多轮对话做补充。保留率的定义很直接用迁移之后的准确率除以目标模型自己预填充的准确率。这个比值等于一百说明映射缓存和真实缓存对下游任务是等价的。所有六对都在各自扫出来的最优源层数下评测用的是完整流程没有为某一对单独调整方法。结果的分化程度超出了论文自己的预期六对里的平均保留率从 42% 一直铺到 98%。这不是一条平滑的曲线而是清晰的两档。四对落在可用区间两对彻底垮掉中间几乎没有过渡带。下面这张表是论文的头条数字我按平均保留率从高到低排列。第四列的地板归一化是把随机猜测校正掉之后的结果它比原始保留率更能反映真实水平后面会单独解释这个校正为什么必须做。家族模型对平均保留率地板归一化HellaSwagMMLUGSM8KQwen314B→32B97.6%96.3%97.6%95.0%95.6%Qwen38B→32B87.5%80.7%95.2%88.5%68.8%Ministral 33B→8B76.2%65.9%93.3%69.4%36.6%Llama 3.18B→70B72.8%62.9%94.4%73.3%18.2%Ministral 33B→14B44.2%14.7%68.0%32.0%3.2%Ministral 38B→14B41.6%11.1%58.7%32.7%1.6%先看最好的那一行。Qwen3 从 14B 迁到 32B平均保留 97.6%ARC-Challenge 甚至到了 101%就是说迁移之后比目标模型自己算的还高了一点点这属于评测噪声范围内的正常波动。这一对的两个模型架构最接近深度差距最小正好对应前面热力图里对角线最锐利的情况。它是这套方法的最佳案例也是论文用来做全部消融实验的对象。地板归一化这一列必须一起读否则会高估效果。多选题基准都有随机猜测的基线四选一是 25%二选一的 WinoGrande 是 50%。如果不减掉这个地板一个完全失效的映射器在 WinoGrande 上也能拿到接近 70% 的保留率看起来还行实际上什么都没保住。归一化之后把随机猜测放到 0%、目标模型自己放到 100%数字才诚实。做了这个校正Llama 3.1 的 8B 到 70B 就从 72.8% 掉到 62.9%WinoGrande 单项从 87.1% 直接掉到 58.5%。这一对是全部评测里参数比例最悬殊的从 80 亿搬到 700 亿能保住六成多已经算不错。真正刺眼的是它的 GSM8K只有 18.2%而目标模型自己能做到 81.12%。也就是说这个映射器保住了模型的常识判断却几乎废掉了它的数学推理能力。这两种能力在同一份缓存上的存活率差了四倍多这个落差本身就值得注意。GSM8K 是这张表里最残酷的一列因为它的地板本来就接近零归一化前后数字不变没有虚高的空间。它考的是八样本带思维链的数学推理要求模型沿着多步链条一路推下去中间任何一步的表示失真都会让整条链断掉。所以这一列可以当成一个压力测试多选题看的是模型能不能选对数学题看的是模型的内部状态还够不够干净。四对可用的模型里只有 Qwen3 14B 到 32B 的 GSM8K 站住了 95.6%其余三对分别是 68.8%、36.6% 和 18.2%衰减非常陡。两个垮掉的对都来自 Ministral 3目标都是 14B。8B 到 14B 归一化后只剩 11.1%3B 到 14B 是 14.7%GSM8K 分别是 1.6% 和 3.2%基本等于完全不会做题了。这两对的键值配置是匹配的走的是同一套流程同样的校准数据同样的超参扫描结果就是不行。论文没有掩饰这个结果反而把它当成最重要的线索键值匹配与迁移成功相关但不构成保证。为什么同样的拟合质量会有不同结局到这里出现了一个真正有意思的谜题。工程上最自然的做法是拿拟合的解释度当筛选指标如果它能预测下游表现那么部署前只要看一眼拟合质量就能决定这对模型能不能用。论文测了这个想法答案是不行。这个否定结论比一个正面指标更有价值因为它拦掉了一条看起来顺理成章、实际会误导部署决策的捷径。如果按拟合质量来筛选你会把一对废掉的组合放进生产环境。反例给得很干净。Llama 3.1 的 8B 到 70B校准集上键的解释度是 0.84小到大方向 HellaSwag 保住 94%可反过来大到小只剩 37%。Ministral 3B 到 8B 拟合出来的解释度一模一样也是 0.84两个方向都稳定保住 93%。同样的拟合质量下游结局差了一倍多。解释度在单对模型内部仍然有用比如用来挑源层但跨对比较时它给不出答案。同一个方向上的两对模型可以有相同的拟合质量和完全不同的结局这说明缺失的信息不在回归本身里。论文找到的答案在于误差落在哪里而不是误差有多大。解释度衡量的是每个通道的重建精度所有维度一视同仁但注意力不是这样工作的它拿键去和查询打分再用得到的注意力权重给值加权。真正决定下游行为的是目标模型最终会算出来的那个注意力输出。于是他们直接测量这个量比较用映射缓存和用真实缓存算出的注意力输出之间的余弦相似度。这个指标的预测力明显更好。在三个家族的十二次配对评测里注意力输出余弦与 HellaSwag 保留率的皮尔逊相关系数是0.57而校准集上键的解释度相关系数是-0.20等于没有关系甚至方向相反。这个结论对工程实践的意义很实在要判断一对模型能不能迁移别看回归拟合得多漂亮去看注意力输出还剩多少相似度。论文也诚实地指出了这个指标的局限它是事后诊断必须先把映射器拟合出来才能测所以还不能用来在拟合之前预筛模型对。找一个拟合前就能算的可迁移性信号被列进了未来工作。再往下追一层论文提出了误差集中度的概念来解释为什么余弦会不同。做法是把映射器的键误差投影到目标查询矩阵的右奇异向量上按对应奇异值的平方加权再除以全部分量的平均误差。集中度大于一说明误差恰好落在注意力会读取的方向上小于一说明误差落在注意力忽略的地方。同样大小的误差藏对了位置就无害藏错了位置就致命。非线性能救回来但救的是特定的病既然线性在两对上失效自然的追问是换成非线性能不能救。论文训练了一个多层感知机作为替代两个 1024 单元的 ReLU 隐藏层用和岭回归相同的均方误差损失在推理时直接替换掉线性映射。除了映射器的函数形式其他一切保持不变这样对照才干净。评测覆盖了四对模型从最成功的一直到最失败的都包含在内。值得注意的是换成多层感知机就放弃了闭式求解的全部好处重新回到需要梯度训练的路上。结果分成两半很值得玩味。在岭回归本来就成功的对上多层感知机反而略微落后Qwen3 14B 到 32B 从 97.6% 变成 97.3%Ministral 3B 到 8B 从 93.3% 掉到 91.8%。在岭回归失败的两对上它把 HellaSwag 保留率抬高了24.3 到 36.8 个百分点8B 到 14B 从 58.7% 直接拉到 95.5%四对模型全部越过了 90%。3B 到 14B 那一对也从 68.0% 抬到了 92.3%两个原本报废的组合都被救了回来。这个反差说明问题不在数据也不在流程而在映射函数的表达能力上。换句话说非线性不是普遍更强它只在特定情况下有用。哪里的跨模型 KV 关系本来就是线性的线性映射就够了加复杂度反而引入了额外的拟合噪声。哪里的关系不够线性才需要非线性来补。这个判断避免了一个很容易犯的错误就是默认更复杂的映射器总能带来更好的结果。失败对上究竟出了什么问题误差集中度给了答案。在用于评测的 HellaSwag token 上岭回归的键解释度是深度负值3B 到 14B 是 -7.818B 到 14B 是 -3.22。负值意味着映射器的预测比直接用均值还差也就是在校准集上拟合出来的线性关系根本没有外推到评测数据上。多层感知机把这个数字分别拉高了 7.62 和 3.08虽然仍在零以下但差距缩小了大半。同时发生的是误差位置的重新分配。在这两对上多层感知机把键的误差集中度平均降了约 2.5注意力输出余弦平均提升约 0.45。误差总量并没有消失它只是被挪到了注意力读不到的方向上去。这就是 24 到 37 个百分点收益的真正来源不是拟合得更准而是错得更不要紧。这个视角对做推理优化的人很有用评估一个近似方法的时候误差的分布位置可能比误差的总量更值得关心。同样的道理也适用于量化和稀疏化它们本质上都是在往模型里注入可控的误差。反过来的证据同样重要。Ministral 3B 到 8B 这一对换成多层感知机之后集中度和余弦两项都改善了HellaSwag 却还是掉了 1.5 个百分点。所以重新分配误差本身不是充分条件它只在被放错位置的误差大到足以造成影响时才起作用。这条边界让整个机制解释站得住脚而不是变成一个万能的事后归因。速度这一面的真实数字精度的账算完了该看省下来的时间到底有多少。测量环境是单个八卡 H100 节点带 NVLinkbf16 精度每个测量点跑 50 次预热和 30 次计时。对照组是目标模型跑一遍完整的预填充用的是 flash_attention_2不含语言模型头。实验组是映射器把源缓存翻译成目标格式包含跨卡搬运缓存所需的传输时间。这个对照口径是公平的因为它把迁移方案自己的额外开销也算进去了。下面这张表取 Qwen3 14B 和 32B 这一对两个方向都列出来。序列长度映射器 14B→32B重新预填充 32B加速比映射器 32B→14B重新预填充 14B加速比6414.0 ms61.7 ms4×11.6 ms39.2 ms3×8K67.8 ms1154.8 ms17×101.9 ms501.0 ms5×32K277.6 ms6975.3 ms25×427.1 ms2952.7 ms7×规律很清晰序列越长收益越大。64 个 token 时只有 4 倍因为这时候映射器的墙上时间被一个固定开销占满了Python 调度和跨卡传输加起来 14 毫秒跟实际计算量没什么关系。到 32K 时加速比冲到 25 倍因为重新预填充的成本随长度和模型规模一起涨而映射器只是一堆按层批处理的矩阵乘法增长慢得多。这个趋势正好对上了长会话智能体的使用形态。两个方向的收益不对称值得单独说一下。小到大方向在 32K 上是 25 倍大到小只有 7 倍。原因在分母上接收方是 14B 的时候它自己重新预填充本来就只要 2952 毫秒比 32B 的 6975 毫秒便宜得多所以省下来的绝对值和相对值都更小。反过来看把缓存往大模型上搬省掉的正是最贵的那一次计算。七对模型的完整数据也给了小到大方向从 2.7 倍到 25.1 倍Llama 3.1 的 8B 到 70B 是 4.5 到 14.9 倍把 11562 毫秒压到 777 毫秒。大到小方向 Qwen3 的 32B 到 14B 是 3.3 到 6.9 倍Llama 70B 到 8B 是 2.8 到 7.6 倍。所有配置都在正收益区间没有出现映射比重算还慢的情况。Ministral 3 的两对倍数最低3B 到 8B 只有 2.7 到 4.0 倍因为目标模型本身就小重新预填充的基数不高。倍数最高的始终是目标模型最大的那些配置。论文对这组数字的诚实程度值得肯定它自己列了三条限制。两组测量都用合成输入隔离出了纯计算成本端到端的迁移还要把映射好的缓存发给目标进程这部分没有测。映射器跑在 eager 模式下没用 torch.compile 也没用 CUDA 图意味着还有优化空间。Ministral 3 的重新预填充只算了语言模型解码器主体不含视觉塔。多轮切换会不会越漂越远对话中途切换是论文列的三个应用场景之一这个场景有一个特有的风险如果每次切换都引入一点误差那么来回切几十轮之后误差会不会累积成灾。这个担心是合理的因为映射之后的缓存会成为下一轮的历史误差有机会自我放大。论文用 CoQA 做了测量100 段对话每段约 15 轮覆盖五个领域在 Qwen3 14B 和 32B 之间来回切。选这一对是因为它保留率最高如果连它都漂其他对就不用谈了。衡量方式是漂移定义为目标模型自己的 F1 分数与映射器在同一轮的 F1 分数之差。结果是两个方向的漂移都很小。小到大方向从第一轮到第十轮差距只扩大了 1.7 个百分点而且这个扩大主要来自 32B 的天花板在上升映射器本身保持稳定。大到小方向的漂移呈线性增长每轮 0.33 个百分点。论文给的判断很克制两个斜率都太小不足以在十轮之内造成级联失败但大到小方向的线性漂移在非常长的会话里仍然会累积。0.33 个百分点每轮听起来微不足道跑到五十轮就是十六个百分点这对一个长期运行的智能体会话来说不能忽略。工程上的对策是设一个漂移预算超过阈值就强制做一次真实的预填充来重置状态。这样既拿到了大部分切换的加速又给累积误差设了上限。还有一个细节说明这个测量的可靠性。源层数这个超参数是在多选题基准上选定的GSM8K、CoQA 多轮和预填充延迟三项从选择过程中排除是构造上的留出集。论文还做了留一交叉验证去掉一个选择基准再重新选超参看被去掉的那个基准的分数变化多少。四对可用模型上最大变化 1.45 个百分点全部六对上最大 2.49 个百分点。另外三个完全没参与选择的基准也测了PIQA、BoolQ 和 ARC-Easy。四对可用模型的平均保留率都在 96% 以上两对失败的分别是 63.7% 和 59.3%分档结构完整重现。这说明前面的分档不是超参过拟合出来的假象而是真实的能力边界。这类留出验证在工程论文里经常被省掉做了就值得指出来。保留率甚至比样本内还略高一点因为 PIQA 和 ARC-Easy 本身比选择基准更简单。这套方法进你的推理栈之前把论文的全部结论收在一起能得到几条可以直接用的判断。第一条是适用范围只做族内迁移源和目标必须共享注意力头数和每头维度也就是论文说的键值匹配。跨家族的情况完全没测论文把它列进了未来工作因为不同家族可能共享足够的表示结构也可能需要完全不同的机制。另外三个家族用的都是稠密全注意力混合架构和线性注意力的情况同样没有覆盖。第二条是筛选方法。键值匹配是必要条件但不是充分条件六对里有两对匹配却失败了。真正的筛选指标是注意力输出余弦它和保留率的相关系数是 0.57而拟合解释度是 -0.20。代价是这个指标必须先拟合映射器才能算所以上线一对新模型的流程是拟合一次测余弦和一个代表性下游基准通过了再进生产。第三条是成本结构。拟合一次约 47 到 87 分钟单节点八卡是一次性的离线开销映射器本体 10.1 亿到 33.6 亿参数占 4 到 12 GB 存储是常驻成本。这笔存储要算进显存预算尤其是当你有多对模型需要各自的映射器时数量是按模型对而不是按模型算的。校准数据方面有个好消息样本量超过 200 条就基本饱和50 条也只差 1.6 个百分点但校准语料的领域会影响结果代码语料在 HellaSwag 上掉了 5.24 个百分点。第四条是收益预期别按 25 倍去算账。25 倍是 32K 上下文小到大方向的最好情况短提示只有 3 到 4 倍大到小方向普遍在 3 到 7 倍。真正能吃到高倍数的是长上下文加频繁上切的场景也就是长会话智能体和成本质量级联。如果你的会话都很短或者切换不频繁这套方法的收益会被固定开销吃掉大半。第五条关于精度取舍。四对可用的模型里多选题类任务普遍保住九成以上但数学推理衰减剧烈最差的只有 18.2%。所以决策不该是这对模型能不能迁移而应该是这对模型在我的任务上能不能迁移。如果你的业务主要是检索问答和分类风险可控如果涉及多步推理、代码生成或者数学计算必须用自己的任务重新测一遍。最后一条是关于线性这件事本身的判断。这篇论文最有意思的地方不是 25 倍加速而是它证明了同一家族不同尺寸的模型对同一段文本算出的中间表示之间存在相当强的线性关系强到用 500 条序列拟合的岭回归就能捕捉大半。这个事实的含义超出了缓存复用它说明模型规模增长带来的表示变化有很大一部分是可以用线性变换刻画的。至于这个结构从哪里来是共享的预训练数据、相似的架构选择还是别的什么论文把它留成了开放问题。
返回列表