mindspore-ai / mindspore-ai/hyper-parallel
【RFC】多维混合并行场景下,DCP保存加载功能泛化覆盖和完善,叠加DP + TP +EP +PP + HSDP
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]:
"""
支持跨拓扑自动迁移的加载接口
"""
...
四、实施建议
- 第一阶段:夯实基础,优先支持 DP + TP + PP 的静态拓扑保存加载
- 第二阶段:引入
TopologyMapper,实现跨 TP/PP size 的动态加载 - 第三阶段:叠加 HSDP 去重存储 + EP (MoE) 专家并行支持
- 第四阶段:完善版本迁移体系 + 增量保存优化
该方案的核心思想是保存逻辑分片、加载时物理重分片,通过统一的 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
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- 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