Skip to content

Feat/checkpoint resharding - #196

Open
JYMiracle305 wants to merge 4 commits into
masterfrom
feat/checkpoint_resharding
Open

JYMiracle305 wants to merge 4 commits into
masterfrom
feat/checkpoint_resharding

Conversation

@JYMiracle305

@JYMiracle305 JYMiracle305 commented Jul 31, 2026 •

Copy link
Copy Markdown
Contributor

背景

现有 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。
  • 支持 TP 扩大、缩小以及 PP stage 变化;支持 QKV 非连续分段布局、词表尾部 padding 和保存/目标 dtype 转换。
  • 模型与 Adam m/v 共用参数分片信息,训练状态和 LR scheduler 状态随 checkpoint 一并恢复。
  • 保存阶段仅由 dp_rank=0 的 TP/PP ranks 写入 shard,并在所有 rank metadata 就绪后原子发布全局 metadata。

当前限制

  • 仅支持 checkpoint version 3。
  • 每个张量目前只支持一个 fragmented axis;QKV 等非连续布局通过 segments 表达。
  • 暂不支持 DistributedOptimizer / ZeRO optimizer state 的保存和恢复。
  • optimizer resharding 依赖稳定的 named parameters。

测试

  • TP 2 -> 4、TP 4 -> 2 的 overlap planning。
  • 非均匀 shard 的显式 global offset。
  • QKV segmented layout 在 TP 变化后的重组。
  • PP stage 变化时只加载目标 stage 所需参数。
  • vocabulary padding、BF16 -> FP32 dtype 转换。
  • model/optimizer state metadata 和 checkpoint serialization round trip。

@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch 3 times, most recently from 5153a9e to a3f6181 Compare August 3, 2026 10:17
@JYMiracle305
JYMiracle305 changed the base branch from master to feat/checkpoint-consumed-micro-batches August 3, 2026 10:17
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint-consumed-micro-batches branch from 57ab888 to 7689011 Compare August 4, 2026 07:43
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch 2 times, most recently from e0540b5 to c85e813 Compare August 6, 2026 07:13
@JYMiracle305
JYMiracle305 changed the base branch from feat/checkpoint-consumed-micro-batches to feat/checkpoint-optimizer-state-control August 6, 2026 07:16
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch from c85e813 to 7028203 Compare August 6, 2026 07:19
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint-optimizer-state-control branch from 5514a0b to cf3cce5 Compare August 6, 2026 08:12
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch from 7028203 to 8e4012b Compare August 6, 2026 08:12
@JYMiracle305
JYMiracle305 changed the base branch from feat/checkpoint-optimizer-state-control to feat/checkpoint-consumed-micro-batches August 6, 2026 08:14
@JYMiracle305
JYMiracle305 changed the base branch from feat/checkpoint-consumed-micro-batches to feat/checkpoint-optimizer-state-control August 6, 2026 08:16
@JYMiracle305
JYMiracle305 changed the base branch from feat/checkpoint-optimizer-state-control to feat/checkpoint-consumed-micro-batches August 6, 2026 08:17
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch from 8e4012b to 83a73ee Compare August 6, 2026 08:44
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint-consumed-micro-batches branch from 59a9674 to 9d99035 Compare August 6, 2026 08:44
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch from 83a73ee to 57bc905 Compare August 6, 2026 09:47
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint-consumed-micro-batches branch from 9d99035 to 862a378 Compare August 6, 2026 09:47
@kilinchange
kilinchange force-pushed the feat/checkpoint-consumed-micro-batches branch from 862a378 to f6f4715 Compare August 7, 2026 02:18
@kilinchange
kilinchange force-pushed the feat/checkpoint_resharding branch from 57bc905 to 9928ff2 Compare August 7, 2026 02:18
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint-consumed-micro-batches branch from f6f4715 to 9d30401 Compare August 12, 2026 01:37
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch 5 times, most recently from 845f617 to 8ba9890 Compare August 12, 2026 08:27
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint-consumed-micro-batches branch from 9d30401 to 6799ce4 Compare August 12, 2026 14:55
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch 2 times, most recently from d64951e to 770fc02 Compare August 14, 2026 09:19
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint-consumed-micro-batches branch from 6799ce4 to 9d406c9 Compare August 14, 2026 09:20
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch from 770fc02 to 8a3ddc2 Compare August 14, 2026 09:35
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint-consumed-micro-batches branch from 86aefcd to 48ca7e4 Compare August 20, 2026 08:56
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch from 8a3ddc2 to 340c458 Compare August 20, 2026 08:56
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch from 340c458 to 32b6d48 Compare August 25, 2026 03:24
@JYMiracle305
JYMiracle305 changed the base branch from feat/checkpoint-consumed-micro-batches to fix/adam-fp32-state August 25, 2026 03:34
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch from 32b6d48 to fb9fc5f Compare August 26, 2026 08:09
@kilinchange
kilinchange force-pushed the feat/checkpoint_resharding branch from fb9fc5f to 31e9655 Compare August 26, 2026 08:15
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch from 31e9655 to 40e28ec Compare August 27, 2026 01:33
@JYMiracle305
JYMiracle305 changed the base branch from fix/adam-fp32-state to master August 27, 2026 01:37
@JYMiracle305 JYMiracle305 changed the title [WIP] Feat/checkpoint resharding Feat/checkpoint resharding Aug 27, 2026
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch 2 times, most recently from 7f25f0a to 3328359 Compare August 27, 2026 06:14
@JYMiracle305
JYMiracle305 force-pushed the feat/checkpoint_resharding branch 2 times, most recently from 4f610e7 to df1c7e3 Compare September 4, 2026 07:38
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,

@chen2021673 chen2021673 Sep 20, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里在一直轮询所有 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;

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里如果继续用轮询检测 metadata file 的话可能有刷屏风险。

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

不再使用轮训的方式没有这个风险了

}
}

Checkpoint::Load(resume_dir, *args.model, args.optimizer.get(), args.state, args.lr_scheduler.get());

@chen2021673 chen2021673 Sep 20, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

关于 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;

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这个字段不需要吧,真正用的是 storage.data_offset

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

已修改

}

// 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) {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这个函数没有被使用,可以删掉。

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

done

}
}
if (!filtered_sd.empty()) {
model_file_index = SaveStateDict(checkpoint_dir / "model.ckpt", filtered_sd);

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

建议把 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;

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里应该 CHECK 或 LOG(WARNING),tensors 不在 ShardedStateDict 里是不正常的。

#include <unordered_set>
#include <vector>

#include "infini_train/include/checkpoint/shard_spec.h"

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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"

@chen2021673 chen2021673 Sep 20, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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;

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

我们现在都是硬编码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);

@chen2021673 chen2021673 Sep 20, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里能不能构造时就把 global layer_number 传进来,TransformerModel::ShardedStateDict() 中直接按 layer 生成,而不是 Remap 一遍。

return state;
}

void Checkpoint::SaveTrainerStateFile(const std::filesystem::path &path, const TrainerState &state) {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这个函数其实是为了把 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.")) {

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

同上,不要硬编码。

ofs << " \"n_embd\": " << state.n_embd << ",\n";
ofs << " \"vocab_size\": " << state.vocab_size << "\n";
ofs << " },\n";
ofs << " \"tensors\": [\n";

@chen2021673 chen2021673 Sep 21, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

同上,这个函数需要修改。

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

@chen2021673 chen2021673 Sep 21, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

不应该 由 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。

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants