跳到主要内容

数据截至 (上游 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.weightsafetensors 分片 + index
Meta 原生(official repo)layers.{i}.attention.wq.weightconsolidated.*.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 投影做转置(下一节)。

3. _permute:RoPE 交错方式的历史差异

这是整个转换层里唯一「动数值」的地方,也是最妙的细节。

背景: RoPE 对 head_dim 维度的配对方式有两种约定——Meta/llama 原始实现把相邻两维配成一对(interleaved:(x₀,x₁),(x₂,x₃)…),HF transformers 的实现把前半维与后半维配对(rotate-half:(x₀,x_{d/2})…)。数学等价,但 q/k 投影输出行的排列顺序不同。

torchtune 的选择: 内部统一用 Meta 的 interleaved 约定;加载 HF 权重时把 q/k 投影矩阵的行重排一次,之后训练、保存都不用再管两套约定。这个重排就是 hf_to_tune 里的 _permute(torchtune/models/convert_weights.py:151-157):

def _permute(t, n_heads):
# 行视图 [n_heads, 2, head_dim//2, dim] → 交换 2 与 head_dim//2 → 摊平
return (
t.view(n_heads, 2, head_dim // 2, dim)
.transpose(1, 2)
.reshape((head_dim * n_heads), dim)
)

调用点就在查表之后:凡是键里含 q_proj 就按 num_heads 转,含 k_proj 就按 num_kv_heads 转(torchtune/models/convert_weights.py:160-165)。保存时 tune_to_hf(:166)里的 _permute 是逆操作(:189-195,view 的维度顺序对调)。

重点看: 这个修复集中在转换函数内部,注意力代码本身完全不用知道另一种约定的存在——脏数据在边界洗干净,内部保持单一约定。这是值得抄的架构卫生。

4. Checkpointer 家族:一个模型类型一个翻译策略

所有 checkpointer 实现同一个 Protocol(_CheckpointerInterface,torchtune/training/checkpointing/_checkpointer.py:54),按「权重的来源格式 × 读写方式」分成四类:

Checkpointer场景位置
FullModelTorchTuneCheckpointer读写 torchtune 自家格式(不转换)torchtune/training/checkpointing/_checkpointer.py:122
FullModelHFCheckpointer读写 HF 格式(最常用的入口)同上 :369
FullModelMetaCheckpointer读写 Meta 原生格式同上 :1039
DistributedCheckpointer分布式中间断点(DCP 格式,含 optimizer 分片)同上 :1305

按模型分流的翻译策略由配置里的 model_type 决定(ModelType 枚举,torchtune/training/checkpointing/_utils.py:81):大多数家族走通用的 hf_to_tune/tune_to_hf,但有特殊结构的(Phi3 的 qkv 融合、Qwen2 的 tied embeddings、Llama3.2 Vision 的双塔)各自有专属转换函数,在 load_checkpoint/save_checkpoint 里 if-elif 分流(torchtune/training/checkpointing/_checkpointer.py:604-650:790-875)。

4.1 加载:记住每个键来自哪个分片

FullModelHFCheckpointer.load_checkpoint(torchtune/training/checkpointing/_checkpointer.py:548)的非 DCP 路径做三件事:

  1. 逐个读 safetensors 分片(safe_torch_load,torchtune/training/checkpointing/_utils.py:229——safetensors 用 safe_open 逐键读,.pthtorch.load(mmap=True),避免把整份权重同时驻留 CPU 内存)。
  2. 顺手记录 self._weight_map[key] = 分片号(torchtune/training/checkpointing/_checkpointer.py:590-592)。
  3. 合并后按 model_type 翻译成 torchtune 格式,放进 state_dict[MODEL_KEY](MODEL_KEY = "model",常量集中在 torchtune/training/checkpointing/_utils.py:57-71)。

4.2 保存:逆向翻译 + 按原布局分片

save_checkpoint(torchtune/training/checkpointing/_checkpointer.py:734)先做逆向翻译(tune_to_hf 等,见 §3),然后写盘分两种模式:

  • DCP 模式:HuggingFaceStorageWriter + _weight_map 推出的 fqn_to_index_mapping,让保存的分片布局和加载时一致(torchtune/training/checkpointing/_checkpointer.py:882-895)。
  • 普通模式:save_torch_state_dict(..., max_shard_size="5GB") 重新分片(:897-903)。

adapter(LoRA)权重单独处理:原样存一份 .pt(torchtune 格式,续训用),再翻译一份成 PEFT 的 adapter_model.safetensors(生态互通用),外加 adapter_config.json(torchtune/training/checkpointing/_checkpointer.py:907-986)。翻译规则见 §6。

5. CheckpointClient:recipe 的统一入口

recipe 不直接摸 checkpointer,而是通过 CheckpointClient(torchtune/training/checkpointing/_checkpoint_client.py:63)。它的 save_checkpoint(:403)按两个维度分流:

full_tensors?
┌───────┴───────┐
True(最终) False(中间断点)
│ │
│ enable_async_checkpointing?
│ ┌────┴────┐
│ True False
▼ ▼ ▼
_save_checkpoint_sync _save_checkpoint_async
(HF 格式,可分发) (DCP 异步) (同步,含 optimizer/dataloader 态)

两个设计点:

  • 「中间 vs 最终」决定带不带训练状态。 中间断点要支持续训,所以除了模型还要存 optimizer 状态、dataloader 状态和 TrainingProgress(seed / epochs_run / steps_run / max_steps_per_epoch,torchtune/training/checkpointing/_checkpoint_client.py:36);最终产物只要纯模型权重,能直接给推理引擎用。判断逻辑在 save_checkpoint 开头(:437-449)。
  • 分布式下同步保存要先聚合。 _save_checkpoint_sync(:239)在非 DCP 且非单机时,用 training.gather_cpu_state_dict(torchtune/training/_distributed.py:486)把 FSDP 分片在 rank 0 的 CPU 上拼回完整 state dict 再写盘;若还有异步 checkpoint 在飞,先等它落地,避免两条写盘路径抢 CPU(:277-292)。

6. 导出 PEFT:让 torchtune 的 LoRA 能被 PEFT 加载

torchtune 自己实现的 LoRA(第 5 章)和 HuggingFace PEFT 的命名不同(lora_a vs lora_A)。导出函数 tune_to_peft_adapter_weights(torchtune/models/convert_weights.py:263)的做法很经济:

  • 复用主映射表派生 adapter 映射:遍历 _TO_PEFT_KEYS(:218,lora_a→lora_A 等)与 _FROM_HF,用字符串替换生成 attn.q_proj.lora_a.weight → base_model.model.model.layers.{}.self_attn.q_proj.lora_A.weight 这类映射(:275-288)。
  • B 矩阵也要 permute:q/k 的 lora_B 输出维就是 q/k 的行维,同样受 RoPE 交错差异影响,所以导出时对 lora_B 做一次 _permute_lora_matrix(:290-296:305-309)。
  • tune_to_peft_adapter_config(:240)把 torchtune 的模块名翻成 PEFT 的 target_modules(w1→gate_proj 等,_TO_PEFT_TARGET_MODULES:225),并补上 base_model_name_or_path

7. 关键细节 / 坑

  • QLoRA 保存时的反量化 hook。 NF4 基座权重在 state_dict() 时由 reparametrize_as_dtype_state_dict_post_hook(torchtune/modules/common_utils.py:24)还原成 bf16 并可 offload 到 CPU,避免保存时 GPU 显存峰值;LoRA builder 在组装模型时统一注册(torchtune/models/llama3_1/_component_builders.py:271_register_reparametrize_state_dict_hooks)。
  • 部分模型的 adapter 不能存 PEFT 格式。 Phi3/Phi4、Llama3.2 Vision、Llama4 走 torchtune 格式存 adapter,保存时只打 warning(torchtune/training/checkpointing/_checkpointer.py:922-933)。
  • optimizer-in-backward 的断点不能跨模式恢复。 它的 optimizer 状态是「每参数一个 optimizer」,恢复时若原 run 没开 optimizer_in_bwd 会直接抛错(recipes/full_finetune_distributed.py:767-780)。
  • 续训的前置条件。 resume_from_checkpoint: True 之外,config 里的 checkpointer.recipe_checkpoint 必须指向中间断点文件;恢复出的 recipe state 若与当前 config 冲突(seed、max_steps_per_epoch),以 checkpoint 为准并打 warning(recipes/full_finetune_distributed.py:284-323_update_recipe_state)。

8. 代码地图

主题文件路径符号名
键名映射表torchtune/models/convert_weights.py_FROM_HF_FROM_METAget_mapped_key
RoPE 交错转置torchtune/models/convert_weights.pyhf_to_tune(_permute)、tune_to_hf
PEFT 导出torchtune/models/convert_weights.pytune_to_peft_adapter_weightstune_to_peft_adapter_config_TO_PEFT_KEYS
HF checkpointertorchtune/training/checkpointing/_checkpointer.pyFullModelHFCheckpointer.load_checkpointsave_checkpoint
分布式断点torchtune/training/checkpointing/_checkpointer.pyDistributedCheckpointer
recipe 侧入口torchtune/training/checkpointing/_checkpoint_client.pyCheckpointClient.save_checkpoint_save_checkpoint_sync_save_checkpoint_asyncTrainingProgress
安全加载(mmap)torchtune/training/checkpointing/_utils.pysafe_torch_loadModelType、各 *_KEY 常量
FSDP 聚合回 CPUtorchtune/training/_distributed.pygather_cpu_state_dict
QLoRA 反量化 hooktorchtune/modules/common_utils.pyreparametrize_as_dtype_state_dict_post_hook_register_reparametrize_state_dict_hooks