数据截至 (上游 commit bd2a0fc7c314)
03 · checkpoint 与权重格式转换
这一章讲什么: torchtune 的模型命名体系是自创的(
layers.0.attn.q_proj.weight),而世界上流通的权重是 HF 格式(model.layers.0.self_attn.q_proj.weight)和 Meta 原生格式(layers.0.attention.wq.weight)。本章讲这套三向翻译怎么做、为什么 q/k 投影还要多转一次置,以及训练中的 checkpoint 读写如何组织。
1. 它要解决的小问题
加载一个 Llama-3.1-8B,你可能拿到三种形态:
| 来源 | 键名风格 | 文件形态 |
|---|---|---|
HF Hub(*-hf 仓库) | model.layers.{i}.self_attn.q_proj.weight | safetensors 分片 + index |
| Meta 原生(official repo) | layers.{i}.attention.wq.weight | consolidated.*.pth |
| torchtune 自己的中间产物 | layers.{i}.attn.q_proj.weight | .pt / safetensors |
光改键名还不够——q/k 投影的行排列顺序在 HF 与 Meta/torchtune 之间也不一样(RoPE 实现约定的差异,见 §3)。而保存时还要能逆向翻回去,让产物直接被 transformers / vLLM 使用。
2. 思路:双向映射表 + 一次转置
整个翻译层就两个文件的事:
- 映射表:
torchtune/models/convert_weights.py:13(_FROM_META)与:35(_FROM_HF),{}占位符表示层号。 - 查表函数:
get_mapped_key(torchtune/models/convert_weights.py:55)先把具体层号替换成{}查表,再把层号填回去。查不到就抛带键名的异常——加载错格式的 checkpoint 会在这里第一时间炸,而不是等 shape 对不上。
逐键转换时顺手做两件事:跳过位置编码缓存(rotary_emb.inv_freq / rope.freqs 映射到 None 或被显式跳过,因为它能从超参重算,见第 2 章 §5),以及对 q/k 投影做转置(下一节)。