已开启
【RFC】多维混合并行场景下,DCP保存加载功能泛化覆盖和完善,叠加DP + TP +EP +PP + HSDP #239
zhangbuxue创建于 6月23日
7月27日 关联了pull request:feat: generalize and improve DCP save and load for hybrid parallel topologies
7月27日 关联了pull request:feat: generalize and improve DCP save and load for hybrid parallel topologies
8月20日 关联了pull request:feat: generalize and improve DCP save and load
User: 维混合并行场景下,叠加DP + TP +EP +PP + HSDP,DCP保存加载功能泛化覆盖和完善
针对该特性写一下背景介绍和简单的方案设计
Kimi: 我来为您设计这个大规模分布式训练场景下的 Checkpoint 保存加载功能泛化方案。
一、背景介绍
1.1 业务背景
在大模型训练领域,随着模型规模从百亿向千亿、万亿参数演进,单一并行策略已无法满足显存和算力需求。多维混合并行(如数据并行 DP + 张量并行 TP + 序列并行 EP + 流水线并行 PP + 分层数据并行 HSDP)已成为大模型训练的标配架构。在此复杂拓扑下,Checkpoint(DCP,Distributed Checkpoint)的保存与加载面临严峻挑战:
1.2 核心痛点
二、方案设计
2.1 整体架构
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: """根据目标拓扑计算所需分片位置""" ...核心能力:
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_spec2.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_path2.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 manifest2.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_dict2.5 关键特性覆盖矩阵
FlatParameter保存Shard(dim)标记,支持跨 TP size 加载global_offset定位,支持 stage 数变更TopologyMapper自动计算重分片策略async_write+ 双缓冲,训练与 IO 流水化三、接口设计
# 用户层 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]: """ 支持跨拓扑自动迁移的加载接口 """ ...四、实施建议
TopologyMapper,实现跨 TP/PP size 的动态加载该方案的核心思想是保存逻辑分片、加载时物理重分片,通过统一的
ShardingSpec协议解耦存储拓扑与计算拓扑,从而实现复杂混合并行场景下的 Checkpoint 泛化能力。