跳到主要内容

数据截至 (上游 commit 090253dac668)

01 · 模型定义:配置驱动的 transformer

这一章讲什么: olmo/model.py(约 1900 行)如何用一个配置 dataclass + 两种 block 实现覆盖从 GPT-2 到 OLMo-2 的全部架构变体;以及那些只在大规模实战里才会学到的细节(BufferCache、embedding padding、初始化策略)。


1. 它要解决的小问题

OLMo 项目从 2023 到 2024 出了好几代模型,架构一直在改:OLMo-1 是 GPT-2 式的 learned position + LayerNorm + 有偏置;OLMo-2 换成 RoPE + RMSNorm + QK-norm + 无偏置 + post-norm。

如果把每代模型写成一个单独的类,仓库会变成一坨互相抄改的代码。OLMo 的选择是:结构差异全部变成配置字段,模型代码只有一份ModelConfig 的默认值精确等于 GPT-2 base(olmo/config.py:240 的注释明写),而每代模型的真实架构由 configs/official-*/ 下的 YAML 决定。

教学价值就在这里:读一份 YAML 就能看出「现代 LLM 相对 GPT-2 到底改了哪几处」


2. 顶层结构:一次前向的数据流

OLMo.forwardolmo/model.py:1253)的形状很简单——嵌入、N 个 block、最终 norm、logits:

input_ids (B, T)


① wte 词嵌入 ──(可选 emb_norm)── (可选 + wpe 位置嵌入)── emb_drop
│ ↑ 仅当 rope 和 alibi 都关掉时才有 wpe

② N × OLMoBlock 每个 block: x = x + Attn(LN(x)); x = x + MLP(LN(x))
│ (norm_after=True 时 norm 挪到子层之后)

③ ln_f 最终 layer norm


④ logits = ff_out(x) 或 weight_tying 时 x @ wte.weightᵀ
(可选 scale_logits: 再除 sqrt(d_model))

几个值得记住的形状细节:

  • 权重绑定(weight tying) 是默认开的(olmo/config.py:418),输出层直接复用词嵌入矩阵(olmo/model.py:1454-1455)。OLMo-2 关掉了它(weight_tying: false),改用独立的 ff_outolmo/model.py:1129-1139)。
  • 词表 pad 到 128 的倍数。 embedding_size 默认 50304(GPT-2 词表 50257 向上对齐);不是 128 的倍数会警告「可能伤吞吐」(olmo/model.py:1086-1091)。OLMo-2 词表 100278 pad 到 100352。
  • block_group_size 把若干 block 包进一个 OLMoBlockGroupolmo/model.py:1116-1123),不改变参数量,纯粹为 FSDP 包裹粒度服务(详见第 3 章)。

3. 核心机制一:OLMo-2 相对 GPT-2 改了哪几处

小问题: 「现代 LLM 架构」到底比 GPT-2 多了什么?论文里散落各处,而这里两份配置一对比就是答案。

ModelConfig 默认值 = GPT-2 base;configs/official-1124/OLMo2-7B-stage1.yaml:5-36 = OLMo-2 7B。逐项 diff:

字段GPT-2 默认(代码)OLMo-2 7B(配置)为什么
归一化layer_norm_type: defaultolmo/config.py:360rmsRMSNorm 少算均值,更快更稳
归一化位置norm_after: false(pre-norm)true(post-norm)OLMo-2 论文称 post-norm 配合 QK-norm 更稳
位置编码无(用 learned wperope: truerope_theta: 500000旋转位置编码,外推性好
QK 归一化attention_layer_norm: falsetrue压住 attention logit 增长,防 loss 尖峰
偏置include_bias: truefalse大模型偏置趋近 0,省掉无损失
激活activation_type: swiglu(默认已是)swiglu门控线性单元
Dropout0.10.0万亿 token 单轮训练不过拟合,不需要
权重绑定weight_tying: truefalse独立输出头表达力更强
精度字段init_device: meta + layer_norm_eps: 1e-6meta 初始化配合 FSDP 懒初始化

直觉一句话: 这份 diff 就是 2022→2024 年 LLM 架构演化的浓缩清单。每一处改动背后都有论文,但在这里它就是一个配置字段。


4. 核心机制二:两种 block,一个接口

小问题: 怎么既保留自己的高效实现,又能逐比特对齐 Llama 做对照实验?

OLMoBlock.buildolmo/model.py:668-675)按 block_type 二选一:

block投影方式norm 位置注意力实现
OLMoSequentialBlock(默认)融合 QKV:一个 att_proj 线性层,输出切三段支持 pre/post 两种flash-attn / SDPA 分发
OLMoLlamaBlock分离的 q_proj / k_proj / v_proj仅 pre-norm手写 matmul + softmax(仿 Llama)

融合 QKV 长什么样

OLMoSequentialBlock.__init__olmo/model.py:685-701)先算好三段维度再建一个线性层:

# 摘自 olmo/model.py:689-697
head_dim = config.d_model // config.n_heads
self.fused_dims = (
config.d_model, # q 段
config.effective_n_kv_heads * head_dim, # k 段
config.effective_n_kv_heads * head_dim, # v 段
)
self.att_proj = nn.Linear(config.d_model, sum(self.fused_dims), ...)

前向时一次矩阵乘、一次 split(olmo/model.py:754-759):

qkv = self.att_proj(h)
if self.config.clip_qkv is not None:
qkv.clamp_(min=-self.config.clip_qkv, max=self.config.clip_qkv)
q, k, v = qkv.split(self.fused_dims, dim=-1)

一次 GEMM 代替三次,GPU 利用率高一截;effective_n_kv_headsolmo/config.py:477)让 MQA/GQA 也能复用同一个融合投影——q 段始终占满 d_model,k/v 段按 kv 头数收窄。

原理演示:post-norm 开关

norm_after 这一个布尔值改变 block 的计算序(示意,非源码):

# 示意,非源码
if not norm_after: # pre-norm:先归一化再进子层
h = attn_norm(x)
att = attention(proj(h))
x = x + dropout(att)
else: # post-norm:子层输出之后再归一化
att = attention(proj(x))
x = x + dropout(attn_norm(att))

真实代码里这个开关在 OLMoSequentialBlock.forwardolmo/model.py:746-790)里出现四次(attn 前后、FFN 前后各一次)。OLMo-2 用 post-norm + QK-norm 的组合,这是它稳定性配方的一部分。


5. 核心机制三:attention 内核里的实战细节

两种 block 共用基类 OLMoBlock.attentionolmo/model.py:588)。进去之后的处理顺序是固定的:

q, k, v
│ ① QK-norm(可选):q_norm(q), k_norm(k) model.py:602-605
│ ② reshape 成 (B, n_heads, T, head_dim) model.py:609-613
│ ③ 拼接 KV cache(推理时) model.py:615-618
│ ④ RoPE 旋转 q, k model.py:623-625
│ ⑤ 合并 causal bias / padding mask model.py:627-635
▼ ⑥ 分发到具体注意力实现 model.py:639-648
attn_out(att)

第 ⑥ 步的分发在 _scaled_dot_product_attentionolmo/model.py:533-586),优先级从高到低:

  1. 带文档掩码 → flash_attn_varlen_func:548-563):把 batch 拉平成 (B*T, ...),用 cu_doc_lens 告诉 flash-attn 每条序列里有哪些文档边界,注意力不跨文档泄漏。这是预训练 packing 的关键设施。
  2. 有 flash-attn 且无 mask → flash_attn_func:564-568)。
  3. 兜底 PyTorch SDPA:569-586):注意这里有个补丁——SDPA 不支持 GQA,所以先把 k/v 用 repeat_interleave 复制到 q 的头数再算(:570-577)。

一个精度细节: RoPE 默认全精度施加(rope_full_precision: true)。RotaryEmbedding.forwardolmo/model.py:307-324)先把 q/k 转 float、关掉 autocast,旋转完再转回原 dtype——在低精度训练里,旋转矩阵的 sin/cos 若按 bf16 算会引入可见误差。

另一个 -inf 细节: 把 padding mask 加进 attention bias 时可能出现 -inf + -inf,SDPA 对 -inf 的处理会产生 NaN。代码显式调 ensure_finite_-inf 钳回 dtype.minolmo/model.py:1370-1376),注释里写明了原因。


6. 核心机制四:BufferCache——为 FSDP 让路的缓存

小问题: causal mask、RoPE 的 sin/cos 表这些「算一次、到处用」的张量,放哪?

按 PyTorch 习惯应该注册成 buffer,但 OLMo 不敢。BufferCache 的 docstring 是本章最有价值的一段注释(olmo/model.py:108-116):

We avoid using buffers because we've run into various issues doing so with FSDP. ... it does synchronize them across processes, which we want to avoid since (A) it isn't necessary, and (B) we sometimes have -inf in these biases which might get turned into NaNs when they're synchronized due to casting or some other issue.

翻译过来:FSDP 不分片 buffer,但会跨进程同步它们;而 causal bias 里的 -inf 在同步的精度转换里可能变成 NaN。于是 OLMo 用一个普通 dict 当缓存:

class BufferCache(dict, MutableMapping[str, torch.Tensor]):
"""Cache for attention biases and other things ..."""

get_causal_attention_biasolmo/model.py:384-393)和 RotaryEmbedding.get_rotary_embeddingolmo/model.py:270-296)都是「查缓存 → 没有或不够长才算 → 按 device 搬移」的模式,跨 rank 各自算各自的,永不同步。

这条经验可以带走: 在 FSDP 下,纯函数式的本地缓存比 buffer 更安全。


7. 核心机制五:初始化与显存策略

三种初始化方案

OLMo.reset_parametersolmo/model.py:1173)按 init_fn 分发,block 级在 OLMoBlock.reset_parametersolmo/model.py:483-506):

init_fn规则出处
normal全部截断正态,init_std=0.02GPT-2 传统,OLMo-2 用它
mitchell输出层 std 随层深衰减:1/sqrt(2·d·(layer_id+1))olmo/model.py:493-496,OLMo-1 用
full_megatron输出层 init_std / sqrt(2·n_layers):498-500,Llama 2 同款

「输出投影按深度缩小初始方差」是深网稳定的经典技巧,两种进阶方案只是衰减律不同。

激活检查点:按层粒度省显存

should_checkpoint_blockolmo/model.py:91-105)提供 whole_layer / one_in_two / ... / fine_grained 八种策略,用一个取模判断决定哪些层重算。activation_checkpoint_function:78-88)有个细节:只有 dropout 非零才需要保存 RNG 状态(没 dropout 时重算是确定性的,省一步)。


8. 坑与边界

  • ALiBi 和 RoPE 互斥,且 ALiBi 不支持 flash-attn——OLMo.__init__ 里直接抛 OLMoConfigurationErrorolmo/model.py:1077-1081)。
  • OLMoLlamaBlock 不支持文档掩码_scaled_dot_product_attention 里遇到 max_doc_len 直接 NotImplementedErrorolmo/model.py:899-902)。要 packing 必须用 sequential block。
  • init_device: meta 时不在 __init__ 里初始化权重olmo/model.py:1143-1145),由 FSDP 的 param_init_fnreset_parameters() 延迟做——脚本侧的对应逻辑在 scripts/train.py:176-230
  • torch.backends.cuda.enable_mem_efficient_sdp(False)olmo/model.py:1102-1103):作者认为 mem-efficient 后端「super slow」,直接禁用,逼 PyTorch 选 flash 后端。
  • _make_state_dict_compatibleolmo/model.py:1787)承担历史包袱:去 _fsdp_wrapped_module. 前缀、把远古单 norm checkpoint 拆成 attn_norm/ff_norm 两份、按不同 block_group_size 重新分组。长期项目的 checkpoint 兼容逻辑都堆在这里。

9. 代码地图

主题文件路径符号名
模型入口olmo/model.py:1070OLMo
前向olmo/model.py:1253OLMo.forward
block 基类与构建olmo/model.py:411:668OLMoBlockOLMoBlock.build
默认 blockolmo/model.py:678OLMoSequentialBlock
Llama 对照 blockolmo/model.py:826OLMoLlamaBlock
attention 内核olmo/model.py:588:533OLMoBlock.attention_scaled_dot_product_attention
RoPEolmo/model.py:258RotaryEmbedding
归一化族olmo/model.py:139:228LayerNormBaseRMSLayerNorm
FSDP 避坑缓存olmo/model.py:108BufferCache
初始化olmo/model.py:1173:483OLMo.reset_parametersOLMoBlock.reset_parameters
FSDP 包裹策略olmo/model.py:1467get_fsdp_wrap_policy
配置字段全集olmo/config.py:234ModelConfig
真实架构参数configs/official-1124/OLMo2-7B-stage1.yaml(YAML model: 段)