【Bug已解决】Vllm_importance_sampling_correction sequence-level mode aggregates per-token log-ratios with sum instead of mean 解决方案
原始报错:Vllm_importance_sampling_correction sequence-level mode aggregates per-token log-ratios with sum instead of mean 场景:vLLM 的重要性采样校正(importance sampling correction)用来修正"训练策略"与"生成时用的旧策略"之间的分布偏差。在 token 级模式下,逐 token 的 log-ratio 求和是对的(因为 exp(sum) = 连乘 = 序列的 IS 权重)。但切到"序列级(sequence-level)"模式时,代码仍用
sum把逐 token 的 log-ratio 加起来,导致长序列的校正因子被长度放大——序列越长,sum 越大,IS 权重被人为放大,训练被长度绑架。序列级模式的本意是"整条序列一个校正因子",应当用mean(长度归一化)而非sum。 关键词:重要性采样校正、IS correction、log-ratio、sequence-level、token-level、sum vs mean、长度归一化、vLLM、离线策略、RL。
一、现象长什么样
序列越长,校正越离谱:
- token 级模式:校正 =
exp(Σ log-ratio),即各 token 比率连乘,正确; - 序列级模式:代码简单地把逐 token log-ratio 也
sum了,然后exp; - 但序列级本应代表"整条序列一个校正因子",长度应被归一;
- 用
sum时,长序列的 log-ratio 累加值更大,exp 后权重被指数放大,短序列则被压小; - 后果:训练被序列长度偏见主导,长序列的梯度被不当地放大,短序列被忽视;
- 表现:loss/优势随长度系统性偏移,且只在序列级模式(而非 token 级)出现。
核心问题:序列级模式的聚合方式错用 token 级的sum,没做长度归一(mean)。
二、背景:token 级 sum 对,但序列级 mean 才对
重要性采样校正的核心是策略比π_new(a|s) / π_old(a|s),取 log 得到log-ratio。对一条序列:
- token 级:序列的 IS 权重 = 各 token 比率的连乘=
exp(Σ log-ratio)。这里sum是正确的,因为 sum of logs = log of product。 - 序列级:把整条序列当作"一个动作"来校正,目标是得到一个与长度无关的、代表整条序列偏差的标量因子。若仍用
sum,这个因子会随长度线性增长(在 log 空间),exp 后指数增长——长短序列完全不可比。正确做法是mean(或sum / 长度),让校正因子反映"平均每个 token 的偏差",长度归一。
一句话:token 级要的是"乘积"(sum of logs),序列级要的是"平均偏差"(mean of logs)。两种模式的聚合语义不同,sum只适用于 token 级。
三、根因:序列级模式复用了 token 级的 sum 聚合
根因拆解:
- 模式混用:序列级直接调用 token 级的
sum聚合,没区分语义; - 无长度归一:序列级没除以 token 数,log-ratio 随长度累积;
- 指数放大:
exp(sum)让长序列权重被指数放大,短序列被压; - 长度偏见:训练目标被序列长度绑架,偏离真实策略偏差;
- 仅序列级暴露:token 级用 sum 是对的,所以只在切序列级时出错;
- 缺模式分支:聚合函数没按
mode="token"/"sequence"分支处理。
下面用最小模型复现"序列级用 sum 被长度放大",再给"序列级用 mean"的修复。
四、最小可运行复现
import math def is_correction(log_ratios, mode="token"): if mode == "token": return math.exp(sum(log_ratios)) # token 级:sum 正确 # 序列级错误写法:仍 sum return math.exp(sum(log_ratios)) # 被长度放大 if __name__ == "__main__": short = [0.1, 0.1] # 2 个 token long = [0.1] * 20 # 20 个 token,平均偏差一样 print("token 级 short:", is_correction(short, "token")) print("token 级 long :", is_correction(long, "token")) # 序列级:用 sum 时,长序列权重被放大 10 倍(exp(2.0) vs exp(0.2)) print("序列级(错,sum) short:", is_correction(short, "sequence")) print("序列级(错,sum) long :", is_correction(long, "sequence"))运行可见序列级用 sum 时,长短序列权重差 10 倍(尽管平均偏差相同)——长度偏见现场。
五、方案:序列级用 mean 做长度归一
第一层:序列级模式把逐 token log-ratio取均值再 exp,得到长度无关的校正因子:
def is_correction_fixed(log_ratios, mode="token"): if not log_ratios: return 1.0 if mode == "token": return math.exp(sum(log_ratios)) # token 级:sum(连乘) # 序列级:mean(长度归一) return math.exp(sum(log_ratios) / len(log_ratios)) if __name__ == "__main__": short = [0.1, 0.1] long = [0.1] * 20 print("序列级(对,mean) short:", is_correction_fixed(short, "sequence")) print("序列级(对,mean) long :", is_correction_fixed(long, "sequence")) # 现在长短一致(平均偏差相同 -> 校正因子相同)序列级用 mean 后,长短序列得到相同的校正因子,长度偏见消除。
六、方案:按模式分支聚合,统一入口
第二层:把聚合收成按模式分支的单一函数,token 级 sum、序列级 mean,调用方只传 mode:
def aggregate_log_ratios(log_ratios, mode): """按模式聚合逐 token log-ratio。""" if mode == "token": return sum(log_ratios) # 用于连乘(exp 后) if mode == "sequence": if not log_ratios: return 0.0 return sum(log_ratios) / len(log_ratios) # 长度归一 raise ValueError(mode) def is_correction_v2(log_ratios, mode): agg = aggregate_log_ratios(log_ratios, mode) return math.exp(agg) if __name__ == "__main__": print("统一入口 token:", is_correction_v2([0.1, 0.1], "token")) print("统一入口 seq :", is_correction_v2([0.1]*20, "sequence"))统一入口保证两种模式用各自正确的聚合,调用方不踩坑。
七、方案:空序列与长度归一边界守卫
第三层:序列级聚合要处理空序列(长度为 0)和极短序列,避免除零或无效校正:
def aggregate_log_ratios_safe(log_ratios, mode): if not log_ratios: # 空序列:返回中性值(log-ratio=0 -> 校正=1) return 0.0 if mode == "token": return sum(log_ratios) # 序列级:长度归一(此处 length 已保证 > 0) return sum(log_ratios) / len(log_ratios) def is_correction_safe(log_ratios, mode): agg = aggregate_log_ratios_safe(log_ratios, mode) return math.exp(agg) if __name__ == "__main__": print("空序列 token:", is_correction_safe([], "token")) # 1.0 print("空序列 seq :", is_correction_safe([], "sequence")) # 1.0(中性)空序列返回中性校正(1.0),不除零、不崩,边界安全。
八、验证:把"序列级用 mean、token 级用 sum"锁进测试
def test_token_uses_sum(): # token 级 exp(sum) = 连乘 assert abs(is_correction_v2([0.1, 0.2], "token") - math.exp(0.3)) < 1e-9 def test_sequence_uses_mean(): short = [0.1, 0.1] long = [0.1] * 20 # 序列级平均偏差相同 -> 校正因子相同 assert abs(is_correction_v2(short, "sequence") - is_correction_v2(long, "sequence")) < 1e-9 def test_empty_neutral(): assert is_correction_safe([], "token") == 1.0 assert is_correction_safe([], "sequence") == 1.0 if __name__ == "__main__": test_token_uses_sum() test_sequence_uses_mean() test_empty_neutral() print("IS 校正聚合测试通过。")九、排查清单("序列级 IS 校正被长度放大"按顺序查)
- 模式分支:聚合是否按 token/sequence 模式分支?没有则序列级误用 sum。
- 长度归一:序列级是否除以 token 数(mean)?没除则被长度放大。
- 指数放大:是否 exp(sum) 让长序列权重指数增长?是则长度偏见。
- token 级正确:token 级用 sum(连乘)是否正确?是,勿误改成 mean。
- 仅序列级暴露:是否 token 级正常、切序列级才错?则聚合语义混用。
- 空序列:序列级聚合是否处理空序列(除零)?需返回中性值。
- 统一入口:是否单一函数按 mode 分支?有则调用方不踩坑。
十、小结
"序列级 IS 校正用 sum 而非 mean"是聚合语义混用:token 级要"连乘"(sum of log-ratios 正确),序列级要"平均偏差"(mean of log-ratios,长度归一),但代码在序列级仍复用 token 级的 sum,使长序列的校正因子被长度指数放大,训练被长度偏见绑架。
修复三层:
- 序列级 mean:逐 token log-ratio 取均值再 exp,长度归一,长短序列可比;
- 模式分支:单一聚合函数按
mode分支,token 级 sum、序列级 mean; - 空序列守卫:空序列返回中性校正(1.0),避免除零。
核心原则:逐 token 量聚合到序列级时,token 级要"连乘"(sum of logs),序列级要"平均"(mean of logs)。凡是"序列级模式仍用 sum 聚合 per-token log-ratio"的写法,都应改为 mean 做长度归一——否则序列越长,校正越强,训练被长度而非策略偏差主导。