跳到主要内容

数据截至 (上游 commit bd2a0fc7c314)

02 · 模型组件与数据管线

这一章讲什么: 配置里那行 _component_: torchtune.models.llama3_1.llama3_1_8b 到底造出了什么。torchtune 自己实现了全部模型(不依赖 transformers),本章讲它的积木体系、注意力与 RoPE 的实现要点,以及数据从「对话消息」到「训练 batch」的管道。


1. 它要解决的小问题

一个微调库需要模型代码,有两条常见路:依赖 transformers 的模型实现(省事但牵进一整个生态),或者自己写(可控但要维护十几个模型家族)。torchtune 选了后者,而且要同时满足两个矛盾的要求:

  • 积木要通用:LoRA 版的 q_proj 和普通版的 q_proj 应该能互换,不用改注意力代码。
  • 型号要多:Llama2/3/3.1/4、Qwen2/2.5/3、Gemma、Phi、Mistral……每族还有若干尺寸。

2. 思路:积木 + 两级 builder

torchtune 的答案写在 torchtune/models/llama3_1/_component_builders.py:29-36 的模块 docstring 里:积木本身做得尽量灵活(比如 MultiHeadAttentionq_projnn.LinearLoRALinear 都行),再让 builder 函数负责把积木缝成整体,这样积木的构造函数能保持简单。

具体是三层:

┌─────────────────────────────────────────────┐
│ 型号 builder(_model_builders.py) │ llama3_1_8b() —— 填 8B 的参数表
│ vocab=128256, 32 层, 32 头, 8 kv 头 … │
└──────────────────┬──────────────────────────┘

┌─────────────────────────────────────────────┐
│ 组件 builder(_component_builders.py) │ llama3_1(...) —— for 循环拼 32 个
│ MultiHeadAttention + FeedForward + RMSNorm │ TransformerSelfAttentionLayer
└──────────────────┬──────────────────────────┘

┌─────────────────────────────────────────────┐
│ 积木(torchtune/modules/) │ 通用 nn.Module,与具体模型无关
└─────────────────────────────────────────────┘

真实代码对照:llama3_1_8b()(torchtune/models/llama3_1/_model_builders.py:19)只是一张参数表(num_layers=32, num_heads=32, num_kv_heads=8, embed_dim=4096, rope_base=500_000……),全部工作委托给组件 builder llama3_1()(torchtune/models/llama3_1/_component_builders.py:41)。后者用一段 for _ in range(num_layers) 循环逐层组装,最后返回一个 TransformerDecoder(torchtune/models/llama3_1/_component_builders.py:115-122 附近)。

这套结构的可换性在 LoRA 上体现得最清楚:lora_llama3_1(torchtune/models/llama3_1/_component_builders.py:138)与 llama3_1 结构逐行平行,只是把指定的投影换成 LoRALinear——注意力、归一化、decoder 外壳一个字符都不用改。LoRA 侧的细节见第 5 章。

3. TransformerDecoder:组装出来的那个东西

所有文本模型最终都是 TransformerDecoder(torchtune/modules/transformer.py:331)的实例。它就是教科书式的 decoder:

  • tok_embeddingslayers(TransformerSelfAttentionLayer,torchtune/modules/transformer.py:17,即 pre-norm 的 attn + MLP 残差块,forward 在 :88)→ normoutput 投影。
  • forward(torchtune/modules/transformer.py:574)按层循环,支持 input_pos(packing 时的相对位置)、mask(bool 或 flex 的 BlockMask)、output_hidden_states(抽中间层给蒸馏用)。

两个非常教科书的口子值得注意:

口子干什么位置
skip_output_layer置 True 后 unembed 只返回 hidden states,跳过 output 投影——配合 LinearCrossEntropyLoss 把 lm_head 挪进 loss 里分块算(第 4 章)torchtune/modules/transformer.py:401:678-688(unembed)
num_output_chunks旧版 chunked CE 的分块数,已标注 deprecated,指向 LinearCrossEntropyLosstorchtune/modules/transformer.py:409-413

KV cache 是推理用的:每层 TransformerSelfAttentionLayer.setup_caches 会建 KVCache(torchtune/modules/kv_cache.py:11)。训练 recipe 不碰它。

4. MultiHeadAttention:刻意直白的实现

MultiHeadAttention(torchtune/modules/attention.py:18)的 forward(:181)几乎是论文伪代码的直译,读它就能复习 GQA:

  1. q = self.q_proj(x),view 成 [b, s, n_kv*q_per_kv, head_dim],过 RoPE。
  2. k, v 同理(k 也过 RoPE),有 kv_cache 就先写缓存。
  3. GQA 的 kv 扩展是显式 expand 出来的:k.unsqueeze(2).expand(expand_shape).flatten(1, 2)(torchtune/modules/attention.py:288-290)把 n_kv 个头复制成 n_h 个,再走普通 MHA。没有用分组查询的专用 kernel——可读性优先,性能交给底层的 SDPA。
  4. 注意力本体委托给 self._attention_call(torchtune/modules/attention.py:140),它由 _sdpa_or_flex_attention()(torchtune/modules/attention_utils.py:185)在模块初始化时选定:mask 是 BlockMask 且 torch≥2.5 且算力 ≥7.5 时走 flex attention,否则走 SDPA(torchtune/modules/attention_utils.py:186-193 的 docstring 列了全部条件)。
  5. 输出 transpose 回来过 output_proj

可选的 q_norm/k_norm(Qwen3 用)就是两行 if(torchtune/modules/attention.py:246-248:278-280 附近)——积木化让这种变体不需要子类。

5. RoPE:缓存、不进 state dict、meta device 友好

RotaryPositionalEmbeddings(torchtune/modules/position_embeddings.py:13)的实现有三个教学点:

① 预计算缓存。 rope_init(:41)先算 theta = 1 / base^(2i/d),再 build_rope_cache(:49)把 [max_seq_len, dim/2] 个位置的 cos/sin 外积算好,torch.stack[max_seq_len, dim/2, 2] 的缓存。forward(:69)里按 input_pos 索引缓存(:95-97),把输入 view 成复数对做旋转。推理/训练共享同一份缓存。

persistent=False theta 和 cache 都注册为不持久化 buffer(torchtune/modules/position_embeddings.py:51:67)——它们能从超参重算,不该进 checkpoint,也不该占 state_dict 的键。这是 RoPE 权重转换能「跳过 rotary_emb.inv_freq」的前提(第 3 章)。

③ meta device 延迟初始化。 分布式 recipe 在 meta device 上实例化模型(不分配真显存),此时算不了真缓存。所以 Llama3ScaledRoPE.rope_init(torchtune/models/llama3_1/_position_embeddings.py:62)开头检查 freqs.is_meta 就提前返回(:72-74);recipe 在加载权重之前统一补调:if hasattr(m, "rope_init"): m.rope_init()(recipes/full_finetune_distributed.py:694-698)。一句话:谁需要真值,谁负责触发

Llama3.1 特有的频率缩放(apply_scaling,torchtune/models/llama3_1/_position_embeddings.py:103)按波长分三段:短波不动、长波除以 scale_factor、中间段平滑过渡——这就是 8K 外推到 128K 上下文的那套官方公式,照论文实现,没有额外发明。

6. 数据管线:从消息到 batch

数据侧是一条直线,四个站:

原始样本 ──► message_transform ──► model_transform(tokenizer) ──► [pack] ──► collate
(dict) (变成 messages 列表) (tokens + mask + labels) 拼包 成 batch tensor

第 1、2 站:SFTTransform SFTDataset.__getitem__(torchtune/datasets/_sft.py:128)对每个样本调 SFTTransform.__call__(torchtune/datasets/_sft.py:146):先用 message_transform 把原始 dict 变成 Message 列表,再交给 tokenizer(model_transform)。

第 2 站的掩码语义。Llama3Tokenizer(torchtune/models/llama3/_tokenizer.py:47)为例,tokenize_messages(:270)逐条消息 tokenize,并产出等长的 bool mask:message.masked=True 的轮次(通常是 user/system)整段标 True,BOS/EOS 恒为 True。SFTTransform 随后把 mask 翻译成 labels——np.where(mask[1:], CROSS_ENTROPY_IGNORE_IDX, tokens[1:]) 再左移一位对齐 next-token 预测(torchtune/datasets/_sft.py:165-176)。一句话:mask 为真的位置不进 loss

第 3 站:packing(可选)。 PackedDataset(torchtune/datasets/_packed.py:17)在初始化时一次性把数据集 tokenize 并拼接成 max_seq_len 长的包(_pack,:102),每个包记录 seq_lens(包内各样本长度)和 input_pos(包内相对位置)。训练时零 tokenization 开销、零 padding 浪费。

第 4 站:collate。 没 packing 用 padded_collate_sft(torchtune/data/_collate.py:190)右 padding;packing 用 padded_collate_packed(torchtune/data/_collate.py:562),它把 seq_lens 交给 packed_block_causal_mask(torchtune/modules/attention_utils.py:133)生成块对角因果 mask——同一样本内因果、跨样本不可见:

包内两个样本 [a a a b b],mask(1=可见):

a₁ a₂ a₃ b₁ b₂
a₁ [1 0 0 0 0]
a₂ [1 1 0 0 0]
a₃ [1 1 1 0 0]
b₁ [0 0 0 1 0]
b₂ [0 0 0 1 1]

支持 flex attention 时,这个 mask 以 mask_mod 函数形式交给 create_block_mask,只存压缩表示(torchtune/modules/attention_utils.py:160-179),注意力 kernel 内部按文档边界截断——长序列下比物化 [b, s, s] 的 bool mask 省一个数量级的显存。

7. 关键细节 / 坑

  • 积木直译的代价是扩展性要自己缝。 想加一个 torchtune 没有的模型,要手写一族 builder + 权重转换(第 3 章),工作量比在 transformers 里注册一个 config 大。
  • tokenizer 同时是 Transform Llama3Tokenizer 继承 ModelTokenizerTransform(基类 Protocol 在 torchtune/modules/transforms/tokenizers/_utils.py:52),所以它既能出现在配置 tokenizer: 块,又能当 dataset 的 model_transform 用。
  • packing 与 CP/compile 的交互。 分布式 recipe 的 dataloader 固定 drop_last=True,注释写明是为了避免 compile + flex attention 的 shape 问题(recipes/full_finetune_distributed.py:843-844)。
  • torchtune.modules.tokenizers 已废弃。 旧路径只是个转发 shim,真实基类在 torchtune/modules/transforms/tokenizers/(torchtune/modules/tokenizers/__init__.py:8-22 的 NOTE)。

8. 代码地图

主题文件路径符号名
型号参数表torchtune/models/llama3_1/_model_builders.pyllama3_1_8blora_llama3_1_8b
组件拼装torchtune/models/llama3_1/_component_builders.pyllama3_1lora_llama3_1lora_llama3_attention
decoder 外壳torchtune/modules/transformer.pyTransformerDecoderTransformerSelfAttentionLayerTransformerDecoder.unembed
注意力torchtune/modules/attention.pyMultiHeadAttention.forward
SDPA/flex 选择torchtune/modules/attention_utils.py_sdpa_or_flex_attentionpacked_block_causal_mask
RoPEtorchtune/modules/position_embeddings.pyRotaryPositionalEmbeddings.rope_initbuild_rope_cache
Llama3.1 频率缩放torchtune/models/llama3_1/_position_embeddings.pyLlama3ScaledRoPE.apply_scaling
tokenizer 与掩码torchtune/models/llama3/_tokenizer.pyLlama3Tokenizer.tokenize_messages
SFT 样本变换torchtune/datasets/_sft.pySFTDatasetSFTTransform.__call__
packingtorchtune/datasets/_packed.pyPackedDataset._pack
collatetorchtune/data/_collate.pypadded_collate_sftpadded_collate_packed