mindspore-ai / mindspore-ai/hyper-parallel

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

Open
#680 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
53
Forks
63
Avg merge
23h 45m
Merged PRs (30d)
63

Description

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 泛化能力。

schema_version: 1
source: gitcode
gitcode_repo: mindspore/hyper-parallel
gitcode_issue: 239
source_url: https://gitcode.com/mindspore/hyper-parallel/issues/239

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

No repository files, tests, or concrete entry points are named. Start by locating the existing DCP save/load implementation and documenting which DP, TP, EP, PP, and HSDP cases it already covers. Done should be a maintainer-approved, phased scope with explicit compatibility and validation requirements.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
distributed-systems
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Active
Clarity
Needs clarification
Newbie friendliness
25/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.