已开启
【RFC】多维混合并行场景下,DCP保存加载功能泛化覆盖和完善,叠加DP + TP +EP +PP + HSDP #239
zhangbuxue创建于  6月23日
zhangbuxue成员
6月23日 创建

User: 维混合并行场景下,叠加DP + TP +EP +PP + HSDP,DCP保存加载功能泛化覆盖和完善

针对该特性写一下背景介绍和简单的方案设计
Kimi: 我来为您设计这个大规模分布式训练场景下的 Checkpoint 保存加载功能泛化方案。

一、背景介绍

1.1 业务背景

在大模型训练领域,随着模型规模从百亿向千亿、万亿参数演进,单一并行策略已无法满足显存和算力需求。多维混合并行(如数据并行 DP + 张量并行 TP + 序列并行 EP + 流水线并行 PP + 分层数据并行 HSDP)已成为大模型训练的标配架构。在此复杂拓扑下,Checkpoint(DCP,Distributed Checkpoint)的保存与加载面临严峻挑战:

  • 拓扑复杂性:不同并行维度对参数/优化器状态的分片方式各异,Checkpoint 需精准捕获每个 rank 持有的分片信息
  • 动态扩缩容:训练过程中可能调整并行策略(如 TP 从 4 扩到 8),要求 Checkpoint 具备跨拓扑恢复能力
  • 存储效率:全量保存万亿参数会产生 TB 级文件,需支持差异化保存(如仅保存模型权重 vs 完整训练状态)
  • 容错可靠性:大规模集群下节点故障常态化,需保证 Checkpoint 的一致性和可恢复性

1.2 核心痛点

痛点 具体表现
分片元信息缺失 TP/EP 分片后,各 rank 仅持有部分参数,缺乏全局 sharding spec 无法重建
跨拓扑加载失败 从 8 卡 TP 保存的 ckpt,无法在 16 卡 TP 环境下直接加载
冗余存储 DP 组内各 rank 重复保存相同优化器状态,浪费 50%+ 存储
版本兼容 模型结构迭代后,旧版 ckpt 无法适配新版代码

二、方案设计

2.1 整体架构

┌─────────────────────────────────────────────────────────────┐
│                    DCP Save/Load Controller                  │
├─────────────┬─────────────┬─────────────┬─────────────────┤
│  Sharding   │  Topology   │  Storage    │   Versioning    │
│  Metadata   │  Mapping    │  Optimizer  │   & Migration   │
│   Manager   │   Engine    │             │                 │
├─────────────┴─────────────┴─────────────┴─────────────────┤
│              Unified Tensor Storage Format                   │
│         (支持 FSDP ShardedTensor / TP DTensor)              │
├─────────────────────────────────────────────────────────────┤
│  DP Group  │  TP Group   │  EP Group   │  PP Stage       │
│  (HSDP DP) │  (Intra-TP) │ (Intra-EP)  │ (Inter-Stage)   │
└─────────────────────────────────────────────────────────────┘

2.2 关键模块设计

2.2.1 分片元数据管理(Sharding Metadata)

# 统一分片描述协议
@dataclass
class ShardingSpec:
    tensor_name: str
    global_shape: Tuple[int, ...]
    global_offset: Tuple[int, ...]      # 该分片在全局张量中的偏移
    shard_shape: Tuple[int, ...]        # 本分片的实际形状
    placement: List[Placement]          # [Shard(dim), Replicate(), ...]
    # 多维并行叠加信息
    dp_axis: Optional[int] = None       # HSDP 分片维度
    tp_axis: Optional[int] = None       # TP 分片维度
    ep_axis: Optional[int] = None       # EP 分片维度
    pp_stage: Optional[int] = None      # PP 流水线阶段

class GlobalTensorIndex:
    """全局张量索引:支持从任意拓扑定位到具体分片"""
    def __init__(self):
        self.index: Dict[str, List[ShardingSpec]] = {}
    
    def locate_shard(self, tensor_name: str, 
                     target_topology: ParallelConfig) -> ShardingSpec:
        """根据目标拓扑计算所需分片位置"""
        ...

核心能力:

  • 拓扑无关性:保存时记录逻辑分片信息,而非物理 rank 映射
  • 自动重分片:加载时根据当前拓扑重新计算 all_gather / slice 策略

2.2.2 拓扑映射引擎(Topology Mapping)

class TopologyMapper:
    """支持跨拓扑的 Checkpoint 转换"""
    
    def save(self, state_dict: Dict[str, torch.Tensor], 
             current_topology: ParallelConfig) -> DCPManifest:
        """保存时生成拓扑无关的 manifest"""
        manifest = DCPManifest()
        for name, tensor in state_dict.items():
            # 识别该 tensor 的并行属性
            sharding = self._infer_sharding(tensor, current_topology)
            # 仅保存必要的 rank 数据(去重)
            if self._is_unique_shard(name, sharding):
                manifest.add_shard(name, sharding, tensor)
        return manifest
    
    def load(self, manifest: DCPManifest, 
             target_topology: ParallelConfig) -> Dict[str, torch.Tensor]:
        """加载时根据新拓扑重组张量"""
        result = {}
        for name, spec in manifest.iter_shards():
            # 计算当前 rank 需要持有的分片
            local_spec = self._remap_sharding(spec, target_topology)
            if local_spec.is_local:
                # 从存储读取并 possibly all-gather / slice
                tensor = self._materialize_tensor(manifest, name, local_spec)
                result[name] = tensor
        return result

关键算法:跨 TP 加载

def remap_tp_sharding(saved_spec: ShardingSpec, 
                      new_tp_size: int) -> ShardingSpec:
    """
    示例:从 TP=4 保存的 ckpt,加载到 TP=8 环境
    - 原逻辑:将 param 按 dim 0 切分为 4 份,每份 shape=[N/4, M]
    - 新逻辑:需将每份再细分为 2 份,通过 all-gather 原 2 个 rank数据后再 slice
    """
    if saved_spec.tp_axis is not None:
        old_tp_size = saved_spec.tp_size
        ratio = new_tp_size // old_tp_size  # 扩展倍数
        
        # 计算当前 rank 在旧拓扑中对应的数据
        local_rank = get_rank()
        old_rank = local_rank // ratio
        inner_offset = local_rank % ratio
        
        # 更新分片描述
        new_spec = copy.deepcopy(saved_spec)
        new_spec.shard_shape = (
            saved_spec.shard_shape[0] // ratio,
            saved_spec.shard_shape[1]
        )
        new_spec.global_offset = (
            saved_spec.global_offset[0] + inner_offset * new_spec.shard_shape[0],
            saved_spec.global_offset[1]
        )
        return new_spec

2.2.3 存储优化层(Storage Optimizer)

class DedupStorageManager:
    """消除 DP/TP 组内冗余存储"""
    
    def __init__(self, topology: ParallelConfig):
        self.dedup_groups = self._build_dedup_groups(topology)
        # HSDP: 仅 intra-node 保存,inter-node 复用
        # TP: 同 TP 组内仅保存一份
    
    def _build_dedup_groups(self, topology) -> List[List[int]]:
        """构建冗余消除组"""
        groups = []
        # HSDP: 按 node 分组,同 node 内 DP 冗余
        if topology.hsdp_enabled:
            for node in topology.nodes:
                groups.append(node.dp_ranks)
        # TP: 同 TP 组仅保存 rank0
        for tp_group in topology.tp_groups:
            groups.append(tp_group)
        return groups
    
    def save_shard(self, name: str, tensor: torch.Tensor, 
                   rank: int) -> Optional[str]:
        """仅唯一代表 rank 执行实际写入"""
        for group in self.dedup_groups:
            if rank in group and rank != group[0]:
                return None  # 跳过,由 group[0] 保存
        # 执行写入...
        return storage_path

2.2.4 版本兼容与迁移(Versioning)

@dataclass
class CheckpointVersion:
    format_version: str = "2.0"           # DCP 格式版本
    model_version: str                   # 模型结构版本
    torch_version: str                   # PyTorch 版本
    topology_signature: str              # 拓扑指纹(用于快速匹配)

class CheckpointMigrator:
    """支持旧版 ckpt 向新版迁移"""
    
    MIGRATION_RULES: Dict[str, Callable] = {
        "v1->v2": self._migrate_v1_to_v2,
        "transformer_block_rename": self._rename_blocks,
    }
    
    def load_with_migration(self, path: str, 
                           target_version: str) -> Dict[str, Any]:
        manifest = self._read_manifest(path)
        
        # 检测版本差异
        if manifest.version.model_version != target_version:
            migrator = self._get_migrator(
                manifest.version.model_version, 
                target_version
            )
            manifest = migrator(manifest)
        
        return manifest

2.3 保存流程(Save Pipeline)

async def distributed_save(state_dict, path, topology):
    # 1. 生成分片元数据(各 rank 并行)
    local_manifest = {}
    for name, tensor in state_dict.items():
        spec = ShardingSpec.from_tensor(tensor, topology)
        local_manifest[name] = spec
    
    # 2. 全局聚合 manifest(仅 rank0 保存完整索引)
    global_manifest = all_gather_manifest(local_manifest)
    
    # 3. 冗余消除:确定实际需保存的 rank
    save_plan = DedupStorageManager(topology).get_save_plan(global_manifest)
    
    # 4. 异步写入(非阻塞训练)
    if save_plan.should_save(local_rank):
        for name, spec in save_plan.local_shards:
            await async_write(spec.serialize(), f"{path}/{name}")
    
    # 5. 写入全局 manifest 文件
    if is_master():
        global_manifest.save(f"{path}/manifest.json")

2.4 加载流程(Load Pipeline)

async def distributed_load(path, topology, strict=True):
    # 1. 读取全局 manifest
    manifest = DCPManifest.load(f"{path}/manifest.json")
    
    # 2. 版本检查与迁移
    if manifest.version != CURRENT_VERSION:
        manifest = CheckpointMigrator().migrate(manifest)
    
    # 3. 计算当前 rank 所需分片
    mapper = TopologyMapper()
    required_shards = mapper.compute_required_shards(manifest, topology)
    
    # 4. 并行读取(利用拓扑局部性优化读取顺序)
    local_tensors = {}
    for shard_spec in required_shards:
        # 优先从同节点缓存读取,其次从存储读取
        data = await read_with_locality_optimization(shard_spec)
        local_tensors[shard_spec.name] = data
    
    # 5. 张量重组(all-gather / slice / concat)
    state_dict = mapper.assemble_tensors(local_tensors, manifest, topology)
    
    # 6. 严格模式校验
    if strict:
        assert validate_consistency(state_dict, manifest)
    
    return state_dict

2.5 关键特性覆盖矩阵

特性 支持方案 关键实现
DP (Data Parallel) ✅ 全支持 FSDP 参数分片 + 优化器状态分片,按 FlatParameter 保存
TP (Tensor Parallel) ✅ 全支持 列/行并行参数按 Shard(dim) 标记,支持跨 TP size 加载
EP (Expert Parallel, MoE) ✅ 全支持 Expert 路由表 + 专家权重分片,支持 EP 组动态调整
PP (Pipeline Parallel) ✅ 全支持 Stage 间通过 global_offset 定位,支持 stage 数变更
HSDP (Hybrid Sharding) ✅ 全支持 分层分片:intra-node 全副本 + inter-node 分片,存储去重
跨拓扑加载 ✅ 核心能力 TopologyMapper 自动计算重分片策略
异步保存 ✅ 非阻塞 async_write + 双缓冲,训练与 IO 流水化
增量保存 ⚠️ 部分支持 仅保存变更的 optimizer state(需配合追踪器)

三、接口设计

# 用户层 API
class DistributedCheckpoint:
    @staticmethod
    def save(
        state_dict: Dict[str, Any],
        path: str,
        parallel_config: ParallelConfig,
        options: SaveOptions = SaveOptions()
    ) -> SaveResult:
        """
        多维混合并行场景下的统一保存接口
        """
        ...
    
    @staticmethod
    def load(
        path: str,
        parallel_config: ParallelConfig,
        options: LoadOptions = LoadOptions(),
        strict: bool = True
    ) -> Dict[str, Any]:
        """
        支持跨拓扑自动迁移的加载接口
        """
        ...

四、实施建议

  1. 第一阶段:夯实基础,优先支持 DP + TP + PP 的静态拓扑保存加载
  2. 第二阶段:引入 TopologyMapper,实现跨 TP/PP size 的动态加载
  3. 第三阶段:叠加 HSDP 去重存储 + EP (MoE) 专家并行支持
  4. 第四阶段:完善版本迁移体系 + 增量保存优化

该方案的核心思想是保存逻辑分片、加载时物理重分片,通过统一的 ShardingSpec 协议解耦存储拓扑与计算拓扑,从而实现复杂混合并行场景下的 Checkpoint 泛化能力。

likedislike
YYzzYYzz
7月27日 关联了pull request:feat: generalize and improve DCP save and load for hybrid parallel topologies
YYzzYYzz
8月20日 关联了pull request:feat: generalize and improve DCP save and load