跳到主要内容

数据截至 (上游 commit 090253dac668)

05 · 生成与输出形态:beam search、HF 桥、checkpoint 加载

这一章讲什么: 模型训完之后怎么用它——OLMo.generate 自带的 beam search 怎么工作、olmo/beam_search.py 的三层可插拔设计、以及训练格式到 Hugging Face 格式的转换。


1. 它要解决的小问题

一个预训练框架为什么还要自己写生成?两个原因:

  1. 评测和冒烟测试需要它。 训出 checkpoint 要立刻能「说句话」验证(in-loop 的下游评测虽然只用 loglikelihood 打分,但 sanity check 需要真的生成)。
  2. checkpoint 格式是训练内部的。 sharded checkpoint 是几百个分片文件,外部世界要的是 HF 格式——中间需要一座桥。

所以生成侧拆成三块:OLMo.generate(模型方法)、BeamSearch(通用解码器)、hf_olmo(生态桥)。


2. OLMo.generate:把模型塞进 beam search 的 step 函数

OLMo.generateolmo/model.py:1603)本体很薄——真正的逻辑在一个闭包 step:1676-1714)里。beam search 框架每走一步调一次 step(last_predictions, state),它做三件事:

① 组织输入: 第 0 步用完整 prompt;之后每步只喂上一步的 1 个新 token
+ 从 state 还原 past_key_values(KV cache)
② 前向: self(input_ids, past_key_values=..., use_cache=True,
last_logits_only=True) ← 只算最后一个位置的 logits
③ 打包: log_softmax 得 log 概率;新的 KV cache 拍平回 state

两个实现细节:

  • KV cache 在 dict 和 list-of-tuples 之间来回转换flatten_past_key_values / unflatten_past_key_values:1657-1674),因为 beam search 框架要求 state 是张量字典。
  • last_logits_only=True 让前向在最终 norm 前把序列砍到只剩最后一个位置(olmo/model.py:1441-1443),逐 token 解码时省掉整段序列的 logits 计算。
  • beam_size=1 时 beam search 退化成贪心解码——docstring 明说这就是默认用法(:1619)。

3. BeamSearch:三层可插拔

olmo/beam_search.py 是一个从 AllenNLP 移植来的通用 beam search,与模型完全解耦(它只见 step 函数)。可插拔的三层:

抽象内置实现
采样器 Sampler从 log 概率里挑候选节点/挑 beamDeterministicSamplerMultinomialSamplerTopKSamplerTopPSamplerGumbelSampler
约束 Constraint每步修改 log 概率RepeatedNGramBlockingConstraint
终局打分 FinalSequenceScorer给完成的序列重排序SequenceLogProbabilityScorerLengthNormalizedSequenceLogProbabilityScorer

类定义坐标:采样器们在 olmo/beam_search.py:44-421,约束在 :494-641BeamSearch 本体在 :649

3.1 主循环的形状

BeamSearch._search:813 起)维护两个列表:predictions(每步的 token)和 backpointers(每个 beam 上一步来自哪个 beam)。每步的核心是把「beam 内累计分 + 新 token 分」拉平成 beam_size × vocab 的候选池,取 top-beam_size,backpointer 记录来源。结束后沿 backpointers 回溯出完整序列(_reconstruct_sequences:722)。

一个框架级约定值得注意_search 开头构造 log_probs_after_end:894-898)——一个除 EOS 外全为 -inf 的分布。已经结束的 beam 下一步被强制「继续输出 EOS、得分不变」,这样不同长度的序列能在同一个张量结构里公平比分,结束的 beam 不会再膨胀。

3.2 约束示例:n-gram 阻断

RepeatedNGramBlockingConstraint:593)把 beam search 变成防复读机:每条 beam 维护 seen_ngrams(前缀 → 已被禁止的下一 token 列表)和滑动窗口 current_prefixapply 时把会构成重复 n-gram 的 token 的 log 概率置为 dtype 最小值(:608-623),_update_state 更新窗口与登记表(:625-641)。约束逐 beam 跑 Python 循环——慢,但只在生成时用到,训练不碰。

3.3 采样器的分工

采样器行为用途
DeterministicSampler:102纯 top-k默认,贪心/beam search
TopKSampler:148top-k 内按概率采样多样生成
TopPSampler:205累积概率 ≤ p 的核内采样多样生成
GumbelSampler:291Gumbel-top-k 无放回采样随机 beam search(论文同款)

4. from_checkpoint:训练格式怎么变回模型

OLMo.from_checkpointolmo/model.py:1730)是给「拿着 checkpoint 目录想用模型」的人的入口,流程:

checkpoint_dir
│ ① 猜类型: 有 model.pt → unsharded;否则 sharded (:1740-1747)
│ ② 读 config.yaml → ModelConfig (:1750-1751)

unsharded 分支 sharded 分支
init_device=cpu 建模型 按 sharded_checkpointer 类型分流:
torch.load(model.pt) - olmo_core → load_model_and_optim_state
load_state_dict(...) - torch_new → checkpoint.load_model_state
.to(device) (就地填进已建好的模型, :1763-1783)

中间过一道 _make_state_dict_compatible:1787)处理历史 key 变迁(去 FSDP 前缀、旧 norm 拆分、block group 重组——见第 1 章 §8)。checkpoint 目录里的 config.yaml 是格式契约的一部分:模型结构不从代码默认值来,而是从当时存盘的配置来。


5. hf_olmo:通往 HF 生态的薄桥

hf_olmo/ 是独立的子包(有自己的 pyproject.toml),做两件事:

一是运行时封装。 OLMoForCausalLMhf_olmo/modeling_olmo.py:41)继承 PreTrainedModel + GenerationMixin,内部直接持有一个原生 OLMo 实例;create_model_config_from_pretrained_config:19-38)把 HF 的 OLMoConfig 逐字段翻译成 ModelConfig(连 _attn_implementation == "flash_attention_2" 这种 HF 侧开关都映射回 flash_attention 布尔值)。有了这个封装,.generate()from_pretrained 等 HF 全家桶方法直接可用——README 里的推理示例走的就是这条路。

二是离线转换。 hf_olmo/convert_olmo_to_hf.pywrite_model:67)把 unsharded checkpoint 写成 HF 目录结构(write_configconfig.json:52);maybe_unshard:237)意味着它还能先合并分片。发布流水线(scripts/release.shscripts/s3_unshard_to_hf.py)就是把训练产物批量转成 HF 格式上传。

inference/ 目录则是更早期的零散脚本(量化、MMLU benchmark 外壳),属于历史遗留,不是主线。


6. 坑与边界

  • generate 不 batch 解码友好性有限:beam search 状态逐 beam 复制 KV cache,显存随 batch × beam_size 涨;严肃的高吞吐推理应该走 vLLM 之类的外部引擎(本仓库不含)。
  • 约束是 Python 循环,只在 debug/小规模生成时划算(inferred:apply 里逐 batch 逐 beam 的 dict 操作无法向量化)。
  • from_checkpoint 的 sharded 分支只支持两种 checkpointerolmo_core / torch_new);torch_legacy 格式的 sharded checkpoint 要先转 unsharded(scripts/unshard.py)再加载。
  • HF 封装是「barebones」(作者自称,modeling_olmo.py:43 docstring):不支持 flash-attn 之外的高级特性映射,init_params=False 的默认说明它假定你马上 from_pretrained 灌权重。
  • beam search 对 -inf 序列不免疫search 的 docstring 明确警告——如果 beam_size 小于有限概率动作数,返回的「最优」序列可能本身就是 -inf,调用方要自己检查(olmo/beam_search.py:763-772)。

7. 代码地图

主题文件路径符号名
生成入口olmo/model.py:1603OLMo.generate
beam search 主循环olmo/beam_search.py:649:813BeamSearchBeamSearch._search
采样器族olmo/beam_search.py:44Sampler 及各子类
n-gram 阻断olmo/beam_search.py:593RepeatedNGramBlockingConstraint
终局打分olmo/beam_search.py:424:462FinalSequenceScorerLengthNormalizedSequenceLogProbabilityScorer
checkpoint 加载olmo/model.py:1730OLMo.from_checkpoint
分片加载olmo/checkpoint.py:320load_model_state
HF 封装hf_olmo/modeling_olmo.py:41OLMoForCausalLM
HF 配置翻译hf_olmo/modeling_olmo.py:19hf_olmo/configuration_olmo.py:14create_model_config_from_pretrained_configOLMoConfig
格式转换hf_olmo/convert_olmo_to_hf.py:67:105write_modelconvert_checkpoint
分片合并脚本scripts/unshard.py(sharded → unsharded)