mindspore-ai / mindspore-ai/hyper-parallel

HP质量防护分析

Open
#320 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

HyperParallel PyTorch 后端质量防护网设计

本文聚焦 HyperParallel 代码仓中基于 PyTorch 后端构建的分布式训练能力,目标是为测试方案建立一套面向训练主链路的分层质量防护体系。

本文约束如下:

  • 仅覆盖 PyTorch 后端
  • 仅讨论训练相关能力
  • 暂不考虑 HSDP
  • 暂不考虑 torch.compile

1. HyperParallel 分布式训练能力总览

1.1 纳入本方案的能力范围

按当前仓库实现与测试现状,PyTorch 后端纳入本方案的能力域如下:

  • 基础通信与运行时底座
  • DTensor / DeviceMesh / Layout
  • Shard / 声明式切分注入
  • Tensor Parallel
  • Fully Shard / FSDP
  • Pipeline Parallel / VPP(含 P2P 通信重叠)
  • Context Parallel(含 Async CP)
  • Activation Checkpoint
  • Activation Swap
  • Distributed Checkpoint
  • 训练恢复与集成层
  • Clip Grad
  • Init Weights / Meta Init
  • Expert Parallel / MoE
  • Symmetric Memory(单边通信)
  • 分布式随机数 / Seed 管理
  • Mixed Precision Policy
1.2 能力域总表

“依赖的 PyTorch 核心特性”以 /Users/liuchongming/Documents/Obsidian Vault/HyperParallel/PyTorch运行机制总结.md 为主基线,并补充纳入以下工程相关机制:

  • torch.distributed 的 process group、collective、P2P 语义
  • DeviceMesh / DTensor / layout / redistribute 语义
  • 参数初始化、meta tensor、materialization 生命周期
  • optimizer / scheduler 的状态保存恢复与 step 语义
功能域 当前在 HyperParallel 的对象/边界 依赖的 PyTorch 核心特性 对应源码目录
基础通信与运行时底座 process group、device mesh 绑定、collective/P2P 调用、平台抽象 torch.distributed.init_process_group()torch.distributed.new_group()torch.distributed.destroy_process_group()torch.distributed.all_reduce() / all_gather() / reduce_scatter()torch.distributed.send() / recv() hyper_parallel/collectives hyper_parallel/platform/torch
DTensor / DeviceMesh / Layout 分布式张量表示、layout 描述、redistribute、local/global 语义转换 torch.distributed.device_mesh.DeviceMesh;DTensor distribute/redistribute 语义;torch.nn.Module.register_forward_pre_hook()torch.nn.Module.register_forward_hook()torch.nn.Module.__call__() hyper_parallel/core/dtensor hyper_parallel/platform/torch/dtensor.py
Shard / 声明式切分注入 shard()ShardingPlan、模块输入输出 layout 注入、参数布局注入 torch.nn.Module.register_forward_pre_hook(with_kwargs=True)torch.nn.Module.register_forward_hook(with_kwargs=True);kwargs/positional args 透传语义;模块边界输入输出改写 hyper_parallel/core/shard
Tensor Parallel 1D TP、算子级分片与聚合、parallelize_module / distribute_module 路径 torch.nn.Module.register_forward_pre_hook()torch.nn.Module.register_forward_hook();DTensor dispatch / redistribute;torch.Tensor.backward() / torch.autograd.backward();参数 .grad 累积语义;collective 通信 hyper_parallel/core/tensor_parallel hyper_parallel/core/shard hyper_parallel/core/dtensor
Fully Shard / FSDP fully_shard()、参数 all-gather/unshard、backward reduce-scatter、reshard、mixed precision policy torch.nn.Module.register_forward_pre_hook(with_kwargs=True)torch.nn.Module.register_forward_hook();输出张量 torch.Tensor.register_hook()torch.autograd.backward()torch.autograd.Variable._execution_engine.queue_callback(...)torch.nn.Module.state_dict()torch.nn.Module.load_state_dict()torch.nn.Module.register_load_state_dict_post_hook() hyper_parallel/core/fully_shard hyper_parallel/platform/torch/fully_shard
Pipeline Parallel / VPP(含 P2P 通信重叠) PipelineStage、GPipe/1F1B/VPP 调度、stage 间 send/recv、micro-batch 调度、overlap_p2p 通信与计算重叠 显式 forward() / backward() 调度;torch.autograd.backward()torch.distributed.send() / recv();micro-batch 执行时序;shared parameter 梯度同步 hyper_parallel/core/pipeline_parallel hyper_parallel/platform/torch/pipeline_parallel
Context Parallel(含 Async CP) attention 上下文切分、QKV 重排、sync/async CP、Async CP 下 projection GEMM 与 AllToAll 通信重叠 torch.nn.Module.register_forward_pre_hook(with_kwargs=True)torch.nn.Module.register_forward_hook();异步路径上的 torch.nn.Module.register_full_backward_pre_hook()torch.autograd.backward();A2A 类 collective 语义 hyper_parallel/core/context_parallel
Activation Checkpoint checkpoint wrapper、选择性重算、函数/模块级 checkpoint torch.utils.checkpoint.checkpoint();autograd 图重放语义;forward 重算与 backward 路径差异;saved tensor 生命周期 hyper_parallel/core/activation_checkpoint hyper_parallel/platform/torch/activation_checkpoint
Activation Swap saved tensor swap/offload、pack/unpack、层间预取与换回编排 torch.autograd.graph.saved_tensors_hooks()torch.nn.Module.register_forward_pre_hook(prepend=True)torch.nn.Module.register_forward_hook()torch.nn.Module.register_full_backward_pre_hook(prepend=True)torch.nn.Module.register_full_backward_hook() hyper_parallel/core/activation_checkpoint/swap.py hyper_parallel/platform/torch/activation_checkpoint
Distributed Checkpoint DCP planner、metadata、save/load、跨 mesh restore torch.nn.Module.state_dict()torch.nn.Module.load_state_dict()torch.optim.Optimizer.state_dict()torch.optim.Optimizer.load_state_dict();DTensor/full tensor/scalar 持久化与重建语义 hyper_parallel/core/distributed_checkpoint
训练恢复与集成层 trainer 集成、resume、optimizer/scheduler state 恢复、导出 torch.optim.Optimizer.step()Optimizer.state_dict() / load_state_dict()torch.optim.lr_scheduler.LRScheduler.state_dict() / load_state_dict()torch.nn.Module.load_state_dict() 后继续训练的参数-优化器状态映射关系 hyper_parallel/integration/llamafactory
Clip Grad 分布式梯度 norm 聚合与裁剪、main_grad 兼容 torch.Tensor.grad 可见语义;Tensor.register_post_accumulate_grad_hook() 对应的累积后时机基线;torch.nn.utils.clip_grad_norm_() 的插入点语义;optimizer step 前梯度后处理 hyper_parallel/core/utils/clip_grad.py hyper_parallel/platform/torch/clip_grad.py
Init Weights / Meta Init 延迟初始化、分片初始化、meta materialization 参数初始化与 meta device 生命周期;materialization 语义;torch.nn.Module.load_state_dict()torch.nn.Module.register_load_state_dict_post_hook();分布式初始化与 layout 对齐 hyper_parallel/core/dtensor/init_weights.py hyper_parallel/core/dtensor/parameter_init.py hyper_parallel/platform/torch/init_weights.py
Expert Parallel / MoE MoE FFN mega kernel(AllToAll-Dispatch→GMM→SwiGLU→AllToAll-Combine)、AIV/AIC 核心并行、事件驱动的通信-计算重叠 torch.distributed.all_to_all();GEMM 算子;SwiGLU 激活函数;事件同步语义 hyper_parallel/core/multicore hyper_parallel/platform/torch/multicore
Symmetric Memory(单边通信) 对称内存分配、单边操作 shmem_put/get、原子信号操作、集合通信 shmem_allgather/alltoall、融合操作 fused_all_gather_matmul/fused_matmul_reduce_scatter 底层共享内存语义;信号量同步;集合通信原语 hyper_parallel/core/symmetric_memory
分布式随机数 / Seed 管理 OffsetBasedRNGTracker、分布式随机数生成、跨 rank seed 同步 torch.random;随机数状态管理与同步 hyper_parallel/core/dtensor/random.py
Mixed Precision Policy FSDP 混合精度策略、dtype 管理与转换 torch.amp;dtype 自动转换语义 hyper_parallel/core/fully_shard hyper_parallel/platform/torch/fully_shard
1.3 当前工程视角下的重点能力

从训练主链路风险看,本方案优先关注四类能力:

  • 改写张量语义的能力:DTensor、Shard、TP、CP
  • 改写参数生命周期的能力:Fully Shard / FSDP、Init Weights / Meta Init
  • 改写执行时序的能力:PP/VPP、Activation Checkpoint、Activation Swap
  • 改写训练闭环状态的能力:Distributed Checkpoint、Resume、Clip Grad

这些能力覆盖了 HyperParallel 在 PyTorch 后端上最主要的训练语义变更点,也是测试方案应优先设防的对象。

2. Megatron/TorchTitan 训练场景对标

本章只基于本地代码仓中已核实的测试、实现与文档,提炼 Megatron-LM 与 TorchTitan 目前实际防护了哪些训练场景,以及这些场景对应的代码路径。

2.1 Megatron-LM 已核实的训练场景与用例路径

本节代码依据:

  • Megatron-LM/tests/unit_tests
  • Megatron-LM/tests/functional_tests
已核实场景 防护内容 主要用例路径 对 HyperParallel 的参考价值
并行状态初始化与 group 构造 校验 TP、PP、CP、EP、DP 等并行组的初始化、rank、world size、group rank 集合与不同初始化顺序的一致性 tests/unit_tests/test_parallel_state.py HyperParallel 的 process group / mesh / rank 语义应单独设防,不能只靠后续训练跑通来间接证明
Tensor Parallel 基础语义 校验 TP 初始化、layer 切分、mapping、cross entropy、随机数与工具函数 tests/unit_tests/tensor_parallel/test_initialization.py tests/unit_tests/tensor_parallel/test_layers.py tests/unit_tests/tensor_parallel/test_mappings.py tests/unit_tests/tensor_parallel/test_cross_entropy.py tests/unit_tests/tensor_parallel/test_random.py HyperParallel TP 不应只测 end-to-end,还应有对切分规则、映射与随机数语义的基础防护
Pipeline Parallel 调度与通信 校验 PP forward/backward 调度选择、schedule table、micro-batch 顺序、communicator、多模块 schedule 与 pipeline layout tests/unit_tests/pipeline_parallel/test_schedules.py tests/unit_tests/pipeline_parallel/test_bridge_communicator.py tests/unit_tests/pipeline_parallel/test_multimodule_schedules.py tests/unit_tests/pipeline_parallel/test_pipeline_layout.py HyperParallel PP/VPP 应把调度、通信、layout 从单纯训练场景中拆出来单测
激活 offload 与 PP 交互 校验 fine-grained activation offloading 与 pipeline path 的组合行为 tests/unit_tests/pipeline_parallel/test_fine_grained_activation_offloading.py HyperParallel 的 AC/swap 与 PP 组合场景应单列,而不是只挂在 AC 或 PP 任一侧
Dist Checkpoint 基础序列化 校验单进程/多进程保存加载、metadata、strict/non-strict load、partition change save/load、ShardedObject 序列化 tests/unit_tests/dist_checkpointing/test_serialization.py HyperParallel DCP 需要覆盖 save/load 基础语义、strictness 与分片变化,而不是只测一步恢复
Dist Checkpoint 与 PP 布局重配置 校验 PP/VPP layout 改变后的 checkpoint 保存、加载和并行重配置 tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py HyperParallel 的 PP checkpoint 不能只测同构恢复,应覆盖 stage layout 改变后的恢复路径
Dist Checkpoint 与 Optimizer 重分片 校验 optimizer state、DP/TP/PP 变化下的 save/load、resharding、fully-reshardable / dp-reshardable 路径 tests/unit_tests/dist_checkpointing/test_optimizer.py HyperParallel resume 测试必须纳入 optimizer state 与并行重配置,不应只比较 model state_dict
训练主闭环基线 对预训练 pipeline 的指标进行 deterministic / approximate 检查 tests/functional_tests/python_test_utils/test_pretraining_regular_pipeline.py HyperParallel 第 4 层需要引入基于 loss/metric 曲线的训练闭环验证,而不是只比较单步输出
Resume checkpoint 闭环 对恢复后训练的后半段指标与连续训练后半段进行对齐检查 tests/functional_tests/python_test_utils/test_pretraining_resume_checkpoint_pipeline.py HyperParallel 的 resume 防护应以”恢复后继续训练曲线是否连续”为核心断言
MoE / EP 模型与算子测试 校验 MoE layer 实现、expert 切分、token 路由、AllToAll dispatch、shared expert 支持 tests/unit_tests/transformer/moe/ HyperParallel 已有完整 MoE 实现(core/multicore/),应建立对等的 EP 单元测试体系
EP 梯度同步测试 校验 expert parallel 配置下的梯度同步正确性 tests/unit_tests/distributed/test_grad_sync_with_expert_parallel.py HyperParallel EP 训练场景下梯度同步应独立设防
A2A 通信重叠测试 校验 AllToAll 通信与计算重叠的正确性 tests/unit_tests/a2a_overlap/ HyperParallel Symmetric Memory 与 Async CP 均涉及通信-计算重叠,此场景直接相关
Optimizer 梯度匹配测试 校验 optimizer 梯度计算的数值正确性 tests/functional_tests/python_test_utils/test_optimizer_grads_match.py HyperParallel 分布式 optimizer 梯度计算应有独立的数值匹配验证
2.2 TorchTitan 已核实的训练场景与用例路径

本节代码依据:

  • torchtitan/tests/unit_tests
  • torchtitan/tests/integration_tests
已核实场景 防护内容 主要用例路径 对 HyperParallel 的参考价值
并行维度校验与 mesh 操作 校验 ParallelDims 构造、auto-calculate、enabled properties、mesh/get_mesh/get_optional_mesh、world_size 约束,以及单卡/8 卡 mesh 操作 tests/unit_tests/test_parallel_dims.py HyperParallel 需要把并行维度和 mesh 语义单独测透,而不是等 TP/FSDP/CP 场景来被动暴露
FSDP 混合精度与 mesh 基线 test_parallel_dims.py 中包含 TestSingleGPUMixedPrecisionFSDP,用于对齐 composable FSDP mixed precision 的关键行为 tests/unit_tests/test_parallel_dims.py HyperParallel 的 FSDP mixed precision 也应有独立基线测试,不应只放在组合场景中出现
Checkpoint manager 单元防护 校验 save/load 恢复、保留最近 K 个 checkpoint、latest checkpoint 发现、interval、生效 rank、model-only 保存加载、异步保存与 pinned memory staging tests/unit_tests/test_checkpoint.py HyperParallel DCP / resume 需要细分为 manager 语义测试,而不是把所有风险都压到 E2E 恢复场景
Activation Checkpoint 单元防护 校验 no-AC、selective AC、full AC 的 FLOPs、显存、数值正确性,以及 per-op selective recompute 配置行为 tests/unit_tests/test_activation_checkpoint.py HyperParallel AC 不应只测开关前后 loss,一并要测重算开销、显存收益和 per-op 策略行为
确定性与种子分配 校验不同 mesh 维度上的 seed 唯一性与共享规则,例如 PP 维不同 seed、TP 维共享 seed tests/unit_tests/test_set_determinism.py HyperParallel 的 determinism 防护应与并行维度绑定,而不是只做全局固定种子
集成测试场景定义 在统一特性表中显式定义 1D/2D compile、checkpoint integration、PP schedule、PP+DP、PP+TP、FSDP+CP、FSDP+TP+CP、gradient accumulation、validation 等集成场景 tests/integration_tests/features.py HyperParallel 第 3/4 层可以借鉴这种“统一场景注册表”做法,把关键组合场景显式枚举出来
Checkpoint 集成闭环 在集成场景中显式覆盖 full checkpoint、model-only HF checkpoint、optional checkpoint、save/load resume tests/integration_tests/features.py HyperParallel 应把 checkpoint/resume 从单测延伸到集成层,验证训练脚本级的真实恢复路径
PP 调度集成场景 在集成场景中显式覆盖 GPipe、1F1B、Interleaved1F1B、InterleavedZeroBubble、ZBV、自定义 CSV schedule tests/integration_tests/features.py tests/assets/custom_schedule.csv HyperParallel PP/VPP 建议采用“调度类型显式枚举”的场景组织方式,而不是只保留一个通用 PP 冒烟
并行组合集成场景 在集成场景中显式覆盖 PP+TPPP+DP+TPFSDP+CPFSDP+TP+CP、validation with tp+cp+pp tests/integration_tests/features.py HyperParallel 第 3 层的组合场景表可以直接借鉴这种”组合名 -> CLI 覆盖项”的组织方法
Expert Parallel 单元与多后端测试 校验 ExpertParallelExpertTensorParallelDeepEPExpertParallelTorchAOExpertParallel 四种 EP 实现的正确性 tests/unit_tests/test_expert_parallel.py HyperParallel MoE 实现需要覆盖不同 EP 策略路径的单元测试
FSDP + MoE Sharding 交互测试 校验 FSDP 分片与 MoE 专家分片的组合行为 tests/unit_tests/test_fsdp_moe_sharding.py HyperParallel 需要覆盖 FSDP × EP 组合场景
MoE Compile 测试 校验 MoE 模型在 compile 模式下的正确性 tests/unit_tests/test_compile_moe.py HyperParallel 未来如支持 compile,需要覆盖 MoE compile 路径
TP KV Heads 校验测试 校验 TP 配置下 KV heads 数量的合法性与分片正确性 tests/unit_tests/test_tp_kv_heads_validation.py HyperParallel TP 应包含参数合法性校验的独立测试
FP8 / Float8 仿真训练 校验 Float8 精度下的训练正确性 tests/integration_tests/features.pyfloat8_emulation HyperParallel 如支持 FP8 训练,需要建立对等的精度验证场景
Gradient Accumulation 集成 校验梯度累积在分布式配置下的正确性 tests/integration_tests/features.pygradient_accumulation HyperParallel 长稳训练场景应覆盖梯度累积
3D 并行 + compile 组合场景 校验 torchcomms_3d_dp+cp+pp+compiletorchcomms_3d_dp+tp+pp+compile 等三维并行 + compile 组合 tests/integration_tests/features.pytorchcomms_3d_*3d_compile HyperParallel 第 3 层组合场景表应参考这种三维并行组合覆盖方式
FlexAttention / VarLen Attention + CP 校验 FlexAttention 和变长 Attention 在 CP 配置下的行为 tests/integration_tests/features.pyfsdp+flex_attnfsdp+varlen_attn+per_op_sac HyperParallel CP 路径应验证对不同 Attention 实现的兼容性
Seed Checkpoint 与跨维度 Determinism 校验 seed checkpoint 的保存恢复与跨并行维度的确定性 tests/integration_tests/features.pyseed_checkpoint HyperParallel 分布式随机数管理需要建立跨并行维度的确定性验证
SFT 训练场景 校验 SFT 训练模式在分布式配置下的正确性 tests/integration_tests/features.pysft HyperParallel 训练闭环测试应覆盖 SFT 场景
Model-only HF Checkpoint 导出 校验仅保存模型权重并导出为 HuggingFace 格式的正确性 tests/integration_tests/features.pymodel_only_hf_checkpointlast_save_model_only_fp32last_save_model_only_bf16 HyperParallel DCP 应覆盖 model-only 保存和格式转换场景
2.3 对 HyperParallel 首版方案的直接启发

从两边已核实的测试代码看,对 HyperParallel 最有价值的不是“支持了哪些并行功能”,而是它们如何把功能拆成可防护的场景。

可以直接借鉴的设计原则有四条:

  • 并行维度、process group、mesh、rank 语义必须独立设防,不能完全依赖训练场景兜底
  • checkpoint / resume 必须拆成 manager 语义、save/load 语义、并行重配置语义、训练闭环语义四层
  • PP/TP/CP/FSDP 等组合场景应显式枚举,而不是笼统写“混合并行 E2E”
  • 训练闭环测试应以 loss/metric 轨迹连续性为断言,而不只比较单步数值

因此,HyperParallel 首版方案中第二章不应再只列“Megatron/TorchTitan 支持哪些能力”,而应以“哪些场景已经被它们的测试体系明确防护”为参照,反推:

  • 我们缺哪些单元防护
  • 我们缺哪些组合场景
  • 我们缺哪些训练闭环恢复场景
2.4 HyperParallel / Megatron / TorchTitan 能力矩阵

下表中的分布式训练能力为HyperParallel+Megatron+TorchTitan能力全集,下方能力矩阵用于评估:

  1. 哪些能力已经是业界主线能力,应该优先补齐同等级测试防护
  2. 哪些能力是 HyperParallel 的差异化形态,需要围绕其特有机制单独设防
分布式训练能力 HyperParallel Megatron-LM TorchTitan
Process Group / Mesh / Rank 运行时底座 支持。具备 DeviceMesh、process group、collective/P2P 运行时底座。 支持。parallel_state 明确管理 TP/PP/CP/EP/DP group。 支持。ParallelDims/DeviceMesh 是所有并行维度入口。
DTensor / DeviceMesh 作为一等抽象 支持。hyper_parallel/core/dtensor 是 TP、FSDP、checkpoint 等主线基础。 部分支持。DTensor/DeviceMesh 主要出现在 Megatron-FSDP 路径与相关测试中,TP/PP/CP 主线仍以 Megatron 自有并行抽象为主。 支持。TP、FSDP2、CP、checkpoint、seed/determinism 都直接建立在 DTensor / DeviceMesh 之上。
声明式切分 / 模块边界 hook 注入 支持。存在 shard() / ShardingPlan / module hook 注入,是 HyperParallel 的明显特征。 不以内建通用能力提供。更偏显式并行配置与模块实现,不提供对等的通用声明式切分层。 不以内建通用能力提供。存在并行化 helper 和 style,但没有对等的通用 ShardingPlan 抽象。
纯数据并行 / DDP 未见独立 DDP 训练主路径。当前更偏 fully_shard / HSDP / DTensor 体系。 支持。README 与训练栈均包含 distributed / DDP 路径。 支持。features.py 中有显式 ddp 集成场景。
FSDP 支持。hyper_parallel/core/fully_shard 已形成 PyTorch 主路径。 支持。既有 Megatron-FSDP,也有 Torch FSDP2 训练路径与单测。 支持。以 PyTorch composable fully_shard / FSDP2 为主。
HSDP 支持。已有独立 HSDP API、scheduler、state 与测试;但不在本文当前测试方案范围内。 支持(实验性)。torch_fully_sharded_data_parallel.py 提供了 HSDP 路径,相关单测已存在。 支持。features.py 中显式覆盖 hsdphsdp+tphsdp+cp
Tensor Parallel 支持,但当前以 1D TP 为主。READMEtest_api.py 均体现 1D 约束。 支持。是 Megatron 的核心主线能力。 支持。集成测试覆盖 2D eager/compile、TP+PP、FSDP+TP+CP。
Sequence Parallel 部分支持。仓库内存在 sequence-parallel 相关算子/注意力约束,但未见独立、完整的通用 SP 能力面。 支持。配置、训练参数和 functional case 均明确覆盖 sequence parallel。 支持。enable_sequence_parallel 是 TP 常规配置,且有 2d_eager_no_sp 对照场景。
Pipeline Parallel 支持。存在 ScheduleGPipeSchedule1F1B 与 stage/runtime 实现。 支持。PP 是核心主线能力。 支持。features.py 显式覆盖 GPipe、1F1B、多种 PP 组合。
Virtual Pipeline / Interleaved PP 支持。已有 ScheduleInterleaved1F1B,对应本文中的 VPP。 支持。functional case 中大量 vp* / interleaved 场景。 支持。Interleaved1F1BInterleavedZeroBubble 已进入集成场景。
ZeroBubble Pipeline 未见明确支持。当前 PP 主线为 GPipe / 1F1B / Interleaved1F1B。 未见明确现成能力证据。当前已核实主线集中在常规 PP / VPP。 支持。features.py 显式覆盖 InterleavedZeroBubbleZBVZeroBubble
Context Parallel 部分支持。存在 ContextParallel / AsyncContextParallel,但 README 明确显示 Ulysses、Ring Attention 尚未完成。 支持。README、parallel state 和 functional case 均覆盖 CP;支持 hierarchical CP 和 hybrid CP scheduling。 支持。集成测试显式覆盖 cp_allgathercp_alltoallfsdp+cp
Expert Parallel / MoE 并行 支持。core/multicore/ 提供完整 MoE FFN mega kernel 实现(AllToAll-Dispatch→GMM→SwiGLU→AllToAll-Combine),平台层有 torch/multicore/mindspore/multicore/ 适配。 支持。README、parallel state、functional case 和 MoE case 均覆盖 EP/MoE。 支持。提供多种 EP 实现(ExpertParallelExpertTensorParallelDeepEPExpertParallelTorchAOExpertParallel),含 DeepEP/HybridEP 后端和 MXFP8 支持。
Activation Checkpoint 支持。已有 checkpoint wrapper、SAC context 与平台适配。 支持。存在 full recompute、uniform recompute、activation checkpoint 相关 case。 支持。单测显式覆盖 none / selective / full,集成场景也直接引用。
Activation Swap / Activation Offload 部分支持。代码中已有 activation swap / saved tensor hook 实现,但 README 中 selective swap 协同与自动策略仍未完成。 部分支持。已核实 fine-grained activation offloading 与 pipeline 交互测试,但不是 TorchTitan/HyperParallel 这种通用 saved-tensor swap 形态。 未见对等的 activation swap 能力。当前主线是 activation checkpoint,CPU offload 相关场景未形成稳定主表。
Distributed Checkpoint 支持。已有自研 distributed_checkpoint 模块与 planner/storage/api。 支持。dist_checkpointing 是主线能力,测试覆盖 serialization、optimizer、PP layout 等。 支持。DCP、async checkpoint、HF 互转、seed checkpoint 都是正式能力。
Resume + 并行重分片恢复 部分支持。存在 trainer resume、DCP load 与 load_state_dict 修复路径,但当前工程测试覆盖仍需要补强。 支持。functional case 和 unit test 已覆盖 ckpt-resume、reshard、optimizer state 恢复。 支持。full_checkpointoptional_checkpointseed_checkpoint、PP/TP/CP 组合恢复均已显式建场景。
Distributed Optimizer 未见独立成型能力。当前更偏依赖底层 optimizer state 与 DCP 恢复。 支持。存在 distributed optimizer 与 overlap/reshard/resume 相关场景。 部分支持。优化器状态会随 FSDP2/DCP 正常分布式保存恢复,但未形成 Megatron 式独立 distributed optimizer 产品面。
Meta Init / 延迟 materialization 支持。init_weights.pyparameter_init.pyfully_shard 路径都覆盖 meta 生命周期。 支持。训练主路径和 FSDP/Megatron-FSDP 测试都覆盖 meta 初始化。 支持。README、docs/fsdp.md、trainer 与 model parallelize 路径都将 meta init 作为正式能力。
分布式梯度裁剪 支持。已有 DTensor/FSDP/HSDP 感知的 clip_grad_norm_ 未见同等独立能力面作为框架特征突出提供,更多依附于训练栈实现。 支持。文档与实现明确说明 optimizer/grad clip 面向 DTensor/FSDP2 工作。
FP8 Training 未见独立 FP8 训练能力。 支持。megatron/core/fp8_utils.py 提供完整 FP8 精度支持与 Transformer Engine 集成。 部分支持。集成测试包含 float8_emulation 场景;Expert Parallel 路径支持 MXFP8。
Overlap 计算通信重叠 部分支持。Pipeline overlap_p2p、Async CP、Symmetric Memory 均提供通信-计算重叠能力。 支持。配置层面支持 P2P 通信与计算重叠。 部分支持。通过 async TP、async checkpointing 实现部分重叠。
torch.compile 集成 未见明确支持。当前方案暂不考虑 torch.compile 支持。代码中多处使用 torch.compile 支持。集成测试显式覆盖 1d_compile2d_compile3d_compile 等场景。
Gradient Accumulation 未见独立内置能力。 支持。通过 microbatch 训练实现梯度累积。 支持。features.py 显式定义 gradient_accumulation 集成场景。
Async TP 未见独立能力。 未见明确独立能力。 支持。enable_async_tensor_parallel 配置项提供异步 TP 通信与计算重叠。
Custom Pipeline Schedule 未见独立能力。当前 PP 主线为 GPipe/1F1B/Interleaved1F1B。 未见独立能力。 支持。支持通过 CSV 文件定义自定义 PP 调度(pp_custom_csv)。
Per-op Selective AC 未见独立能力。 未见明确独立能力。 支持。集成测试覆盖 fsdp+flex_attn+per_op_sacfsdp+varlen_attn+per_op_sac 等场景。

3. HyperParallel 分层防护体系设计

3.1 为什么需要分层防护

HyperParallel 的主要风险并不集中在单一算子或单一 API,而是来自以下三类系统性问题:

  • 功能逻辑正确,但依赖的 PyTorch 机制假设被破坏
  • 单能力测试通过,但多能力组合后时序、生命周期或状态恢复出错
  • 单步可运行,但多 step、resume 或长稳训练中出现数值漂移、挂死或状态不一致

因此,测试体系不能只按“目录”或“功能模块”组织,而必须分层组织,使不同类型的问题在不同层级被尽早发现。

3.2 五层防护的职责边界
防护层 目标 主要发现的问题
第 0 层:纯逻辑与静态单测 快速验证纯逻辑、静态推导、轻量接口行为 layout 推导错误、planner 错误、调度逻辑错误、参数映射错误
第 1 层:PyTorch 核心机制防护 验证 HyperParallel 对 PyTorch 机制的依赖假设没有被破坏 hook 顺序错误、autograd 边界错误、state_dict 生命周期错误、materialization 时机错误
第 2 层:单并行能力正确性 验证单个能力域在训练语义上正确 loss/grad/参数更新不一致、单能力数值错误
第 3 层:组合能力交互防护 验证多个能力叠加后仍保持正确训练语义 生命周期错位、通信与调度冲突、resume 与并行语义冲突
第 4 层:长稳、恢复与回归防护 验证训练闭环长期稳定,且关键指标不回退 训练卡死、NaN/Inf、显存泄漏、resume 曲线断裂、性能明显退化
3.3 能力域与防护层的映射原则

不同能力域不需要在每一层平均用力,而应按风险映射:

  • DTensor、Shard、Init Weights 更依赖第 0 层和第 1 层
  • TP、FSDP、PP、CP 更依赖第 2 层和第 3 层
  • DCP、resume、clip_grad 更依赖第 1 层、第 3 层和第 4 层
  • AC / swap 横跨第 1 层到第 4 层,因为它们同时改变机制、数值和长稳行为
3.4 优先级原则

测试建设优先级应遵循以下顺序:

  1. 先补 PyTorch 机制防护,再补功能矩阵穷举
  2. 先补单能力训练正确性,再补大规模组合
  3. 先补训练闭环与 resume,一般性能回归放后
  4. 先保证能尽早暴露严重错误,再追求覆盖面最大化

4. 五层防护展开

4.1 第 0 层:纯逻辑与静态单测

这一层的目标是快速发现不依赖真实多卡执行的逻辑错误,运行应足够轻量、适合高频触发。

应重点覆盖的能力与特性
能力域 应测特性
DTensor / DeviceMesh / Layout mesh 构造、placement 表达、layout 推导、redistribute 推导、local/global 语义映射
Shard ShardingPlan 解析、输入输出布局注入规则、参数布局声明、kwargs/positional args 映射
Tensor Parallel style / plan 解析、模块切分规则、参数切分后的静态 shape 规则
Fully Shard / FSDP state / param group / scheduler 的静态规则、mixed precision policy 解析、state_dict key 规则
Pipeline Parallel / VPP stage 切分、schedule 生成、micro-batch 顺序、send/recv 任务编排
Context Parallel attention layout 规则、QKV 维度转换、mesh 维映射
Distributed Checkpoint planner、metadata、layout 映射、save/load 计划生成
Clip Grad norm 聚合逻辑、裁剪系数计算、main_gradgrad 选择逻辑
Init Weights / Meta Init 延迟初始化规则、参数 materialize 规则、初始化布局对齐
这一层不解决的问题
  • 不证明真实多卡数值正确
  • 不证明通信顺序正确
  • 不证明训练闭环稳定
4.2 第 1 层:PyTorch 核心机制防护

这一层的目标是保护 HyperParallel 赖以成立的 PyTorch 执行假设,重点不是“训练结果像不像”,而是“挂钩点和生命周期对不对”。

应重点覆盖的机制型测试
能力域 应测机制
DTensor / Shard torch.nn.Module.__call__() 是否按预期触发 register_forward_pre_hook() / register_forward_hook();直接 forward() 调用是否会绕过框架假设
Tensor Parallel forward hook 注入后的输入输出布局是否稳定;torch.Tensor.backward() 后参数 .grad 累积边界是否符合预期
Fully Shard / FSDP register_forward_pre_hook(with_kwargs=True) / register_forward_hook() 触发时机;输出张量 torch.Tensor.register_hook() 是否准确进入 backward 入口;queue_callback 收尾时机;register_load_state_dict_post_hook() 后参数状态是否可继续训练
Pipeline Parallel / VPP torch.autograd.backward() 与 stage 调度边界是否一致;send/recv 的调用次序是否与调度计划匹配;shared parameter 梯度同步是否落在正确边界
Context Parallel register_forward_pre_hook(with_kwargs=True) / register_forward_hook() 触发时 attention 张量布局是否符合假设;async 路径上 register_full_backward_pre_hook() 的时机是否稳定
Activation Checkpoint torch.utils.checkpoint.checkpoint() 后 backward 是否走重算路径;重算路径下 hook 与保存张量生存期是否符合预期
Activation Swap torch.autograd.graph.saved_tensors_hooks() 的 pack/unpack 次序;register_full_backward_pre_hook() / register_full_backward_hook() 的预取/释放边界是否正确
Distributed Checkpoint state_dict() / load_state_dict() 的对象生命周期;模型状态和优化器状态加载后是否仍与并行对象绑定正确
Clip Grad .grad / main_grad 可见性;clip 是否发生在 optimizer step 前;冻结参数与无 grad 参数是否被正确跳过
Init Weights / Meta Init meta 参数 materialize 时机;load_state_dict() 后参数对象是否完成修复;延迟初始化是否不会破坏后续 hook 假设
这一层的典型失败信号
  • hook 没触发、触发顺序错、触发边界错
  • backward 回调在错误时间点运行
  • load_state_dict() 后对象能加载但不能继续训练
  • checkpoint / swap / materialization 改写了预期的生命周期
4.3 第 2 层:单并行能力正确性

这一层的目标是验证单个能力域在训练语义上的正确性,核心断言是数值、梯度和参数更新与基线一致。

应重点覆盖的能力测试
能力域 应测特性
Tensor Parallel 单卡 vs TP 的 forward output / loss / grad / step 后参数 对齐;关键模块如 linear、embedding、attention 相关路径对齐
Fully Shard / FSDP 单卡或非分片基线 vs FSDP 的 loss / grad shard / step 后参数 对齐;mixed precision 开关组合;state_dict 保存恢复后继续训练
Pipeline Parallel / VPP GPipe、1F1B、VPP 独立正确性;最后 stage loss 与基线一致;各 stage 梯度闭环正确;无训练卡死
Context Parallel CP off/on 的 attention 输出、loss、grad 对齐;不同 layout 与 causal/non-causal 场景正确
Activation Checkpoint none / recompute 模式下 loss、grad、step 后参数对齐
Activation Swap swap off/on 的 loss 与 grad 对齐;保存张量换回后 backward 正确
Distributed Checkpoint DCP save/load 后模型、优化器、标量状态正确重建
Clip Grad clip off/on、不同并行配置下 global norm 与裁剪后梯度与基线一致
Init Weights / Meta Init 普通初始化 vs meta init / 延迟初始化 后训练结果一致
这一层的断言重点
  • 输出一致
  • loss 一致
  • grad 一致
  • 参数更新后一致
  • 保存恢复后单能力继续训练一致
4.4 第 3 层:组合能力交互防护

这一层的目标是验证真实训练中更高风险的能力叠加场景,因为大量系统问题只会在组合场景中暴露。

首版关键组合场景
组合场景 覆盖能力 覆盖的交互机制组合 主要风险 核心断言
Shard + TP 输入输出布局注入场景 Shard / Tensor Parallel torch.nn.Module.register_forward_pre_hook(with_kwargs=True) × torch.nn.Module.register_forward_hook(with_kwargs=True)torch.nn.Module.__call__() × kwargs/positional args 透传 输入 layout 改写与输出 layout 回写顺序错误;kwargs 丢失或错位;local/global 语义漂移 输入输出 layout 与预期一致;不同输入形式行为一致;forward output 与基线对齐
TP + FSDP 单 step 训练场景 Tensor Parallel / Fully Shard torch.nn.Module.register_forward_pre_hook() × torch.nn.Module.register_forward_hook();输出张量 torch.Tensor.register_hook() × torch.autograd.Variable._execution_engine.queue_callback(...) forward 侧张量改写与参数生命周期冲突;backward 入口与 backward 收尾之间状态未闭环;梯度同步后参数状态错误 loss / grad / step 后参数 与基线对齐;backward 后参数状态正确;无训练卡死
TP + FSDP state_dict / resume 场景 Tensor Parallel / Fully Shard / Distributed Checkpoint / Resume torch.nn.Module.state_dict() × torch.nn.Module.load_state_dict()torch.nn.Module.register_load_state_dict_post_hook() × FSDP 参数修复路径 load 后参数对象、分片状态或布局状态不一致;恢复后训练一步跳变 save -> load -> train 与连续训练一致;参数与 optimizer state 映射正确;恢复后可稳定继续训练
TP + PP 微批调度场景 Tensor Parallel / Pipeline Parallel torch.nn.Module.register_forward_pre_hook() × torch.nn.Module.register_forward_hook()torch.autograd.backward() × torch.distributed.send() / torch.distributed.recv() TP 改写后的张量进入 stage 边界时 shape/layout 不一致;send/recv 次序与 micro-batch 调度错位 无训练卡死;micro-batch 顺序正确;最终 loss 与基线对齐
CP + TP attention 布局协同场景 Context Parallel / Tensor Parallel torch.nn.Module.register_forward_pre_hook(with_kwargs=True) × torch.nn.Module.register_forward_hook();async 路径下 torch.nn.Module.register_full_backward_pre_hook() × backward 回流 attention 前后张量重排与 TP 切分冲突;async 路径等待点错误 attention 输出对齐;grad 对齐;async 路径无训练卡死
FSDP + CP attention 训练场景 Fully Shard / Context Parallel torch.nn.Module.register_forward_pre_hook(with_kwargs=True) × torch.nn.Module.register_forward_hook()torch.nn.Module.register_full_backward_pre_hook() × 输出张量 torch.Tensor.register_hook();CP 通信 × FSDP 参数 all-gather/reduce-scatter 生命周期 attention 张量重排与参数 unshard/reshard 冲突;backward 前 CP 等待点与 FSDP backward 入口错位;通信次序不一致导致训练卡死或静默错误 attention 输出、loss、grad 与基线对齐;backward 后参数状态正确;无训练卡死
FSDP + AC 重算训练场景 Fully Shard / Activation Checkpoint 输出张量 torch.Tensor.register_hook() × torch.autograd.Variable._execution_engine.queue_callback(...)torch.utils.checkpoint.checkpoint() × backward 重算路径 backward 重算与参数 all-gather / reduce-scatter / reshard 生命周期冲突;回调边界错位 recompute 开关前后 loss / grad / step 后参数 对齐;backward 后参数状态正确;无训练卡死
PP + AC 微批重算场景 Pipeline Parallel / Activation Checkpoint torch.utils.checkpoint.checkpoint() × torch.autograd.backward();micro-batch schedule × torch.distributed.send() / recv() 重算窗口与 stage 调度错位;微批回传顺序错误;通信等待失配 无训练卡死;micro-batch 回传顺序正确;loss 与基线对齐
Swap + AC 保存张量生命周期场景 Activation Swap / Activation Checkpoint torch.autograd.graph.saved_tensors_hooks() × torch.utils.checkpoint.checkpoint()torch.nn.Module.register_full_backward_pre_hook() × torch.nn.Module.register_full_backward_hook() pack/unpack、预取、重算路径互相干扰;保存张量生命周期异常;释放时机错误 backward 正确;saved tensor 生命周期符合预期;无训练卡死、无异常增长
Resume + FSDP + Meta Init 恢复场景 Distributed Checkpoint / Fully Shard / Init Weights / Meta Init torch.nn.Module.load_state_dict() × torch.nn.Module.register_load_state_dict_post_hook();meta materialization × FSDP 参数修复 load 后参数对象未正确 materialize;分片状态与真实参数不一致;恢复后训练立即出错或结果跳变 load 后可继续训练;resume 一步与连续训练一致;参数状态、分片状态、optimizer state 一致
Clip Grad + FSDP 梯度后处理场景 Clip Grad / Fully Shard torch.Tensor.register_post_accumulate_grad_hook() 对应的累积后边界 × torch.nn.utils.clip_grad_norm_();梯度同步完成边界 × optimizer step 裁剪发生在错误时机;裁剪前梯度尚未完成同步;main_grad / grad 读取错误 global norm 与基线一致;裁剪后梯度一致;optimizer step 后参数一致
这一层的关注点
  • 不仅看能否运行,还要看组合后生命周期是否闭环
  • 不仅看 loss,还要看 resume 后是否出现训练轨迹跳变
  • 不仅防数值错误,还要防训练卡死和静默错误结果
4.5 第 4 层:长稳、恢复与回归防护

这一层的目标是把测试从“单次正确”推进到“持续稳定”,用于拦截长稳问题、恢复问题和重大回归。

应重点覆盖的场景
场景 应测特性
多 step 长稳训练 无训练卡死、无 NaN/Inf、无显存持续增长、loss 曲线稳定
中断恢复 save -> stop -> load -> continue train 与连续训练在 loss、grad、参数轨迹上保持一致
Cross-mesh restore save mesh != load mesh 时模型与优化器状态仍能正确恢复
关键性能护栏 TP、FSDP、PP、AC、swap 开关下吞吐与峰值显存无明显非预期回退
关键 nightly / weekly canary 用少量代表性组合长期守护主链路,例如 TP × FSDPTP × PPFSDP × ACresume × FSDP
这一层的判定重点
  • 训练是否持续可运行
  • 状态恢复是否真正连续
  • 是否存在慢性泄漏或长周期不稳定
  • 是否出现明显性能退化

结论

HyperParallel PyTorch 后端的测试方案,不应再停留在“功能存在性验证”层面,而应围绕以下主线构建质量防护网:

  • 先从能力域完整盘点训练主链路对象
  • 再以 Megatron-Core 与 TorchTitan 的成熟训练场景作为对标参照
  • 用五层防护体系把逻辑错误、机制错误、单能力错误、组合错误和长稳错误分层拦截

在当前范围内,首版最优先补强的不是更多算子测试,而是:

  • 第 1 层的 PyTorch 机制防护
  • 第 3 层的关键组合场景
  • 第 4 层的 resume 与长稳防护

这三类测试最能提升 HyperParallel PyTorch 后端在真实训练闭环中的可靠性。

参考来源

  • PyTorch 运行机制基线:/Users/liuchongming/Documents/Obsidian Vault/HyperParallel/PyTorch运行机制总结.md
  • HyperParallel 仓库现状:README.mdtests/torch/*hyper_parallel/platform/torch/*hyper_parallel/core/*
  • Megatron-Core 官方仓库:[Megatron-LM / Megatron-Core](https://github.com/NVIDI

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

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

Start by reading the referenced Megatron-LM and TorchTitan unit and integration test paths, especially parallel-state, checkpoint, activation-checkpoint, and pretraining pipeline tests. Compare those coverage areas with the HyperParallel directories listed in the design, then define a prioritized quality-protection plan with explicit test entry points and completion criteria for the covered PyTorch training capabilities.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
distributed-systems, machine-learning, testing-qa
Issue type
Documentation
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.