数据截至 (上游 commit dd417662e5bd)
04 · 模型后端接口
这一章讲什么: 评测框架怎么跟「模型」这个变量解耦——
LM契约长什么样,TemplateLM收编了哪些公共样板,以及最常用的HFLM怎么把几万条 loglikelihood 请求跑得又快又稳。要接新后端或抠性能,看这一章。
1. 它要解决的小问题
框架要支持 30+ 种模型:本地 HF 权重、vLLM 服务、OpenAI API、llama.cpp、NeMo…… 它们的分词器、batch 能力、并行方式完全不同。
但第 3 章说过,评测只需要四种请求。所以契约就是:把三种打分方法实现掉,剩下的(排序、分批、缓存、分布式)要么由模板提供,要么由你自由选择。 新后端接入的最小成本 = 实现三个方法。
2. LM 契约:三个方法 + 一套分布式原语
抽象基类 LM(lm_eval/api/model.py:25)的核心是三个抽象方法:
| 方法 | 输入(Instance.args) | 输出 | 语义 |
|---|---|---|---|
loglikelihood | (context, continuation) | (logprob, is_greedy) | continuation 的 log 概率 + 贪心是否会 产出它(lm_eval/api/model.py:39-55) |
loglikelihood_rolling | (string,) | float | 整篇文本的 log 概率,文档各自独立算(lm_eval/api/model.py:57-96) |
generate_until | (context, gen_kwargs) | str | 生成到停止符(lm_eval/api/model.py:99-111) |
注意契约里的两个刻意设计:
- 输入输出都是字符串,tokenization 无关(类 docstring,
lm_eval/api/model.py:28-30)——框架从不碰 token,后端想怎么分词怎么分。 is_greedy要后端给:「逐 token argmax 是否等于 continuation」只有拿着 logits 的人才知道,框架算不了。
另外还有一套分布式原语,默认实现是单机 no-op,多卡后端覆写:rank / world_size / all_gather / gather_object / barrier(lm_eval/api/model.py:171-205)。evaluate() 靠它们切数据、收结果(第 5 章)。
3. TemplateLM:把 token 化样板收进基类
TemplateLM(lm_eval/api/model.py:331)是为「本地 tokenizer + causal/seq2seq 模型」准备的中间层,实现了 loglikelihood 的公共前半段,只留下 _loglikelihood_tokens 给子类。
3.1 _encode_pair:空格归位
context 与 continuation 的边界是分数的敏感点。_encode_pair(lm_eval/api/model.py:368-406)先做一件事:把 context 尾巴上的空格挪给 continuation:
# 摘自 lm_eval/api/model.py:390-393
n_spaces = len(context) - len(context.rstrip())
if n_spaces > 0:
continuation = context[-n_spaces:] + continuation
context = context[:-n_spaces]
为什么:BPE 分词里 " world"(带前导空格)和 "world" 是不同的 token 序列。词边界空格属于「下一个词」一边,先归位再分词,context_enc + continuation_enc 才等于整句分词的无损切分。causal 模型接着用「整句编码 − 前半段编码」求 continuation 的 token(lm_eval/api/model.py:395-400)。
3.2 空上下文:用 BOS 垫一个 token
loglikelihood(lm_eval/api/model.py:408-446)处理 context == "" 的情形:continuation 不加特殊符号单独编码,然后拿 prefix_token_id(通常 BOS/EOS)充当一个 token 的「上下文」(lm_eval/api/model.py:430-440)。这样下游 _loglikelihood_tokens 永远能假设 context 非空——互信息任务那批 ("", choice) 请求(第 3 章)就走这条路。
4. HFLM 的打分机器:_loglikelihood_tokens
HFLM(lm_eval/models/huggingface.py:62,注册名 hf/hf-auto/huggingface)是最常用的后端。它的 _loglikelihood_tokens(lm_eval/models/huggingface.py:1329)把「几万条 (context, continuation) 打分」组织成一条高效的批处理流水线。
4.1 总览图
请求列表
│
▼ Collator:长度降序排序(+可选按 context 分组)
▼ 按 batch_size 切 chunk
每个 chunk:
① 拼 context+continuation,超长左截断
② 右 padding 对齐,一次前向
③ log_softmax
④ 逐条:切出 continuation 段的 logits
⑤ argmax → is_greedy;gather 目标 token 概率 → logprob
▼
get_original:还原回原始请求顺序
4.2 长度降序:一个排序解决三件事
_collate(lm_eval/models/huggingface.py:1338-1348)的排序键是 (-len(toks), tuple(toks))——最长优先。源码注释自己列了理由:时间预估只会高估不会低估;每个 chunk 的第一个 元素就决定了 padding 长度,自适应 batch 好实现;OOM 会在一开始就炸,而不是跑到 90% 才炸。
配套的 _batch_scheduler(lm_eval/models/huggingface.py:1312)支持 batch_size=auto:N:每 1/N 进度重新探测一次当前能放下的最大 batch。探测缓存每轮请求重置(lm_eval/models/huggingface.py:1390-1391)——因为上一轮 ARC(短句)测出的 batch size 用到 MMLU(长句)上会 OOM,注释里挂着 issue #1678。
4.3 一次前向,同时算出 logprob 和 is_greedy
每个 chunk 的核心十来行(lm_eval/models/huggingface.py:1485-1557):
# 摘自 lm_eval/models/huggingface.py:1485-1557(精简,保留骨架)
multi_logits = F.log_softmax(self._model_call(batched_inps, **call_kwargs), dim=-1, ...)
...
logits = self._select_cont_toks(logits, contlen=contlen, inplen=ctx_len) # 只留 continuation 段
greedy_tokens = logits.argmax(dim=-1)
...
max_equal = (greedy_tokens[:, -cont_toks.shape[1]:] == cont_toks).all() # is_greedy
logits = torch.gather(logits, 2, cont_toks.unsqueeze(-1)).squeeze(-1) # 目标 token 的 logprob
answer = (float(logits.sum()), bool(max_equal))
读法:
_select_cont_toks(lm_eval/models/huggingface.py:1191)把 context 段和右 padding 切掉,只留 continuation 对应的 logits 窗口。is_greedy= 「每个位置上 argmax 是否恰好等于 continuation 的 token」——一次前向,判对错和算概率同时完成,这是loglikelihood便宜的关键。- 太长的请求从左边截断(
lm_eval/models/huggingface.py:1424-1436并打警告):保住 continuation,丢 context 的开头。
4.4 单 token 续写缓存:选择题的倍增器
MMLU 的四个选项是 " A"/" B"/" C"/" D"——前缀完全一样,只差最后一个 token。HFLM 默认开 logits_cache(lm_eval/models/huggingface.py:84):Collator 按 context + cont[:-1] 分组(_lookup_one_token_cont,lm_eval/models/huggingface.py:1350-1357),同组只跑一条请求的前向,get_cache(lm_eval/models/utils.py:327)把这条的 logits 广播给组内其他请求(截不同长度窗口)。
效果:K 个单 token 选项,前向次数约除以 K。源码注释的原话是「speeds up some multiple-choice tasks proportionally to the number of choices」(lm_eval/models/huggingface.py:1352-1355)——MMLU 这种字母选项的任务是它的最佳主场;选项是完整句子时省不了多少,因为缓存键不含最后 token,只惠及「只差结尾」的请求组。
4.5 Collator:排序与还原的通用件
Collator(lm_eval/models/utils.py:238)是上述排序/分组/分批/缓存的载体,三个方法:get_batched(:283)产出重排后的批次,get_cache(:327)做单 token 缓存命中,get_original(:402)把结果还原回请求原顺序。评测器全程不知道请求被重排过——顺序还原是后端内部的事。
5. generate_until:生成请求的三板斧
HFLM.generate_until(lm_eval/models/huggingface.py:1574)结构相似,但有三处生成特有的处理:
- 按 gen_kwargs 分组批处理:贪心(temperature=0)和采样(temperature=0.8)不能混在一个 batch,
Collator(group_by="gen_kwargs")把同参数请求凑一起(lm_eval/models/huggingface.py:1617-1622)。 - 停止序列归一:
until统一成 list 并追加 EOS(handle_stop_sequences,lm_eval/models/utils.py:641);normalize_gen_kwargs(lm_eval/models/utils.py:657)把max_new_tokens等别名统一成max_gen_toks,do_sample=False时强制 temperature=0。 - 后处理在字符串层面再做一遍:
postprocess_generated_text(lm_eval/models/utils.py:943)按停止符截断(取split(term)[0]),并处理think_end_token——带思考链的模型只取</think>之后的部分再匹配停止符,因为推理过程里常含\n\n这类「假停止符」(lm_eval/models/utils.py:957-960注释)。
为什么停止符要在 token 级(early stop)和字符串级(截断)各做一遍:token 级早停省算力,字符串级保证输出干净——两者对「停止符跨越 token 边界」的情形行为不同,字符串级是兜底。
6. loglikelihood_rolling:滑窗算整文概率
perplexity 任务的麻烦在于:文档比上下文窗口长。HFLM.loglikelihood_rolling(lm_eval/models/huggingface.py:1226)的做法:
- 整篇分词后,
get_rolling_token_windows(lm_eval/utils.py:335)按max_seq_len切出重叠滑窗,make_disjoint_window(lm_eval/utils.py:378)把每窗裁成「上下文 + 互不重叠的预测段」。 - 每个 token 恰好被预测一次,且每段预测都带满前文——这是与「把多篇文档拼起来算」的朴素做法的本质区别(契约 docstring 专门强调,
lm_eval/api/model.py:62-66)。 - 所有窗过
_loglikelihood_tokens打分,按请求下标聚回、求和(lm_eval/models/huggingface.py:1292-1300)。