Feat/checkpoint resharding - #196
JYMiracle305 wants to merge 4 commits into
Conversation
5153a9e to
a3f6181
Compare
57ab888 to
7689011
Compare
e0540b5 to
c85e813
Compare
c85e813 to
7028203
Compare
5514a0b to
cf3cce5
Compare
7028203 to
8e4012b
Compare
8e4012b to
83a73ee
Compare
59a9674 to
9d99035
Compare
83a73ee to
57bc905
Compare
9d99035 to
862a378
Compare
862a378 to
f6f4715
Compare
57bc905 to
9928ff2
Compare
f6f4715 to
9d30401
Compare
845f617 to
8ba9890
Compare
9d30401 to
6799ce4
Compare
d64951e to
770fc02
Compare
6799ce4 to
9d406c9
Compare
770fc02 to
8a3ddc2
Compare
86aefcd to
48ca7e4
Compare
8a3ddc2 to
340c458
Compare
340c458 to
32b6d48
Compare
32b6d48 to
fb9fc5f
Compare
8323452 to
b66c25c
Compare
b66c25c to
82a301e
Compare
fb9fc5f to
31e9655
Compare
31e9655 to
40e28ec
Compare
7f25f0a to
3328359
Compare
4f610e7 to
df1c7e3
Compare
| nn::parallel::function::AllReduce(token, nn::parallel::function::ReduceOpType::kSum, nullptr, true)->Synchronize(); | ||
| } | ||
|
|
||
| void WaitForWriterManifests(const std::filesystem::path &staging_root, int tp_size, int pp_size, |
There was a problem hiding this comment.
这里在一直轮询所有 metadata.json,但进程同步语义不应该绑定到文件是否出现上。上面已经有SynchronizeCheckpointRanks 函数使用 AllReduce 实现同步,应该可以直接调用,通过 AllReduce 来同步。WaitForWriterManifests 和 WaitForGlobalMetadata 应该可以直接删掉。
| obj_pos = obj_end + 1; | ||
| } | ||
|
|
||
| LOG(INFO) << "[CKPT] Loaded metadata.json: " << meta.tensors.size() << " tensors, iteration=" << meta.iteration; |
There was a problem hiding this comment.
这里如果继续用轮询检测 metadata file 的话可能有刷屏风险。
There was a problem hiding this comment.
不再使用轮训的方式没有这个风险了
| } | ||
| } | ||
|
|
||
| Checkpoint::Load(resume_dir, *args.model, args.optimizer.get(), args.state, args.lr_scheduler.get()); |
There was a problem hiding this comment.
关于 Checkpoint::Save/Load 与 LoadDistributedCheckpoint 的分层,建议调整一下。
现在 Checkpoint::Save/Load 已经没有调用方了(只剩 test_trainer_state.cc、test_lr_scheduler_state.cc、test_checkpoint_serialization.cc 在调),实际入口变成了 Checkpoint::SaveSharded 和 LoadDistributedCheckpoint。这样不太规范,建议改成:
- Checkpoint::Save/Load —— 格式层,唯一入口。负责分片布局、metadata、rank 协同。TP=1 是 world_size=1 的退化情形,不是另一条代码路径。这也正是本 PR 的目的:TP=1 和 TP=8 的 ckpt 可以互读,那么入口自然也该是同一个,不需要 LoadDistributedCheckpoint 这个名字。
- SaveCheckpoint/ResumeFromCheckpoint —— 策略层。负责 iter_xxxxxxx 目录、架构校验等。
相应地 reshard.h/reshard.cc 可以整体并入 Checkpoint::Load —— Load 本身就 reshard-aware 之后,不存在第二个 reshard 入口。
| WriteItem item; | ||
| item.key = key; | ||
| item.filename = is_optimizer ? "optimizer.ckpt" : "model.ckpt"; | ||
| item.offset = offset; |
There was a problem hiding this comment.
这个字段不需要吧,真正用的是 storage.data_offset
| } | ||
|
|
||
| // Compute one rank's balanced interval, including non-divisible dimensions. | ||
| inline std::pair<int64_t, int64_t> GetRankSliceRange(int64_t global_size, int world_size, int rank) { |
| } | ||
| } | ||
| if (!filtered_sd.empty()) { | ||
| model_file_index = SaveStateDict(checkpoint_dir / "model.ckpt", filtered_sd); |
There was a problem hiding this comment.
建议把 model.ckpt、optimizer.ckpt、metadata.json 等字符串都提出为常量,不要在各个地方硬编码。
| for (auto &[name, parameter] : parameters) { | ||
| name = RemapLayerKey(name, local_layers, global_layers); | ||
| if (!sharded_state.tensors.contains(name)) { | ||
| continue; |
There was a problem hiding this comment.
这里应该 CHECK 或 LOG(WARNING),tensors 不在 ShardedStateDict 里是不正常的。
| #include <unordered_set> | ||
| #include <vector> | ||
|
|
||
| #include "infini_train/include/checkpoint/shard_spec.h" |
There was a problem hiding this comment.
nn 是基础设施,checkpoint 是上层功能,nn 层不应该依赖 checkpoint 层。PyTorch 里也是 DCP 依赖 nn.Module,不是反过来。建议把 shard 描述类型下沉到基础层(infini_train/include/shard_spec.h,namespace infini_train),或让 nn 层只暴露一个不含 checkpoint 语义的轻量结构。
|
|
||
| #include "infini_train/include/nn/functional.h" | ||
| #include "infini_train/include/nn/init.h" | ||
| #include "infini_train/include/nn/lora/lora_parallel_linear.h" |
There was a problem hiding this comment.
Transformer 层不应该 include lora 层,这里引用是为了拿 kParamLoraBName ,但也不能挪到其他地方,因为 lora 本身不知道 qkv 是不是 packed 的。这里要不加个 fixme,之后把 packed QKV 本身抽成专门的 sharding abstraction。
| emitted_items.push_back(&item); | ||
| } | ||
| } | ||
| int dp_rank = 0, tp_rank = 0, pp_rank = 0; |
There was a problem hiding this comment.
我们现在都是硬编码dp_rank = 0,但 Megatron / PyTorch 是 replica_id 挂在 ShardedTensor 上,由数据自己决定谁写,这里可以加个TODO。
|
|
||
| checkpoint::ShardedStateDict global_state; | ||
| for (auto &[local_key, tensor] : local_state.tensors) { | ||
| const auto global_key = RemapLayerKey(local_key, local_layers, global_layers); |
There was a problem hiding this comment.
这里能不能构造时就把 global layer_number 传进来,TransformerModel::ShardedStateDict() 中直接按 layer 生成,而不是 Remap 一遍。
| return state; | ||
| } | ||
|
|
||
| void Checkpoint::SaveTrainerStateFile(const std::filesystem::path &path, const TrainerState &state) { |
There was a problem hiding this comment.
这个函数其实是为了把 private 的 SaveTrainerState 方法暴露出来。
如果分布式 checkpoint 变成默认的 Checkpoint::Save/Load,那这些 wrapper 基本都可以删掉。顶层只有Checkpoint::Save/Load,内部直接 SaveTrainerState 即可。下同。
| std::unordered_map<std::string, std::shared_ptr<Tensor>> filtered_sd; | ||
| for (const auto &[key, info] : sharded_sd.tensors) { | ||
| // Optimizer tensors are serialized separately. | ||
| if (key.starts_with("adam.")) { |
| ofs << " \"n_embd\": " << state.n_embd << ",\n"; | ||
| ofs << " \"vocab_size\": " << state.vocab_size << "\n"; | ||
| ofs << " },\n"; | ||
| ofs << " \"tensors\": [\n"; |
There was a problem hiding this comment.
TrainerState 和 metadata.json 的部分信息是重复的。建议
// trainer_state.json:
// 训练语义 / common state
global_step
consumed_train_samples
model config:
n_layer
n_head
n_kv_head
n_embd
original_vocab_size
padded_vocab_size
runtime config:
tp_size
pp_size
dp_size
sp_size
vpp_size
// metadata.json:
// checkpoint format/version
tensor A:
key
dtype
global_shape
shard 0:
global_offset
local_shape
segments
writer/file
file_offset
byte_size
shard 1:
...也就是说 metadata 存
{
"version": 3,
"format": "infinitrain_sharded",
"tensors": [...]
}
就够了。
|
|
||
| namespace infini_train::checkpoint { | ||
|
|
||
| void LoadDistributedCheckpoint(const std::filesystem::path &checkpoint_dir, nn::Module &model, Optimizer *optimizer, |
| state = Checkpoint::LoadTrainerStateFile(checkpoint_dir / "trainer_state.json"); | ||
| const int current_tp = nn::parallel::global::GetTensorParallelSize(); | ||
| const int current_pp = nn::parallel::global::GetPipelineParallelSize(); | ||
| const bool topology_changed |
There was a problem hiding this comment.
不应该 由 TP/PP 是否变化来控制逻辑,应该 load 时先根据当前运行拓扑重新构造 target sharded state dict,由 saved layout + target layout 驱动是否 load,参考 Megatron https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/training/checkpointing.py?#L1838-L1908 。
另外这里如果暂时不支持的 VPP 话可以加一个抛出错误,同时说明 TODO。
背景
现有 distributed checkpoint 与保存时的 TP/PP 拓扑绑定,恢复训练时要求使用相同的并行配置。本 PR 引入基于全局张量坐标的 checkpoint resharding,使 checkpoint 可以在不同 TP/PP 配置之间恢复。
设计文档:Checkpoint Resharding 设计
主要修改
ShardedTensor/ShardedStateDict,描述张量的全局形状、本地分片、全局偏移和切分方式。SavePlanner,统一规划模型参数和 Adam optimizer state 的本地写入布局,并生成可用于 reshard 的全局 metadata。LoadPlanner,根据源 checkpoint 与当前目标拓扑的分片坐标计算重叠区间。IndexedRegionLoadStrategy,按 metadata 中的文件和 offset 直接读取所需区域,在加载阶段完成重组,无需预先生成中间 checkpoint。m/v共用参数分片信息,训练状态和 LR scheduler 状态随 checkpoint 一并恢复。dp_rank=0的 TP/PP ranks 写入 shard,并在所有 rank metadata 就绪后原子发布全局 metadata。当前限制
segments表达。DistributedOptimizer/ ZeRO optimizer state 的保存和恢复。测试