本文聚焦 HyperParallel 代码仓中基于 PyTorch 后端构建的分布式训练能力,目标是为测试方案建立一套面向训练主链路的分层质量防护体系。
本文约束如下:
torch.compile
按当前仓库实现与测试现状,PyTorch 后端纳入本方案的能力域如下:
“依赖的 PyTorch 核心特性”以 /Users/liuchongming/Documents/Obsidian Vault/HyperParallel/PyTorch运行机制总结.md 为主基线,并补充纳入以下工程相关机制:
/Users/liuchongming/Documents/Obsidian Vault/HyperParallel/PyTorch运行机制总结.md
torch.distributed
DeviceMesh
meta
step
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
torch.distributed.device_mesh.DeviceMesh
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()
ShardingPlan
torch.nn.Module.register_forward_pre_hook(with_kwargs=True)
torch.nn.Module.register_forward_hook(with_kwargs=True)
hyper_parallel/core/shard
parallelize_module
distribute_module
torch.Tensor.backward()
torch.autograd.backward()
.grad
hyper_parallel/core/tensor_parallel
fully_shard()
torch.Tensor.register_hook()
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
PipelineStage
overlap_p2p
forward()
backward()
hyper_parallel/core/pipeline_parallel
hyper_parallel/platform/torch/pipeline_parallel
torch.nn.Module.register_full_backward_pre_hook()
hyper_parallel/core/context_parallel
torch.utils.checkpoint.checkpoint()
hyper_parallel/core/activation_checkpoint
hyper_parallel/platform/torch/activation_checkpoint
torch.autograd.graph.saved_tensors_hooks()
torch.nn.Module.register_forward_pre_hook(prepend=True)
torch.nn.Module.register_full_backward_pre_hook(prepend=True)
torch.nn.Module.register_full_backward_hook()
hyper_parallel/core/activation_checkpoint/swap.py
torch.optim.Optimizer.state_dict()
torch.optim.Optimizer.load_state_dict()
hyper_parallel/core/distributed_checkpoint
torch.optim.Optimizer.step()
Optimizer.state_dict()
load_state_dict()
torch.optim.lr_scheduler.LRScheduler.state_dict()
hyper_parallel/integration/llamafactory
main_grad
torch.Tensor.grad
Tensor.register_post_accumulate_grad_hook()
torch.nn.utils.clip_grad_norm_()
hyper_parallel/core/utils/clip_grad.py
hyper_parallel/platform/torch/clip_grad.py
hyper_parallel/core/dtensor/init_weights.py
hyper_parallel/core/dtensor/parameter_init.py
hyper_parallel/platform/torch/init_weights.py
torch.distributed.all_to_all()
hyper_parallel/core/multicore
hyper_parallel/platform/torch/multicore
shmem_put/get
shmem_allgather/alltoall
fused_all_gather_matmul/fused_matmul_reduce_scatter
hyper_parallel/core/symmetric_memory
OffsetBasedRNGTracker
torch.random
hyper_parallel/core/dtensor/random.py
torch.amp
从训练主链路风险看,本方案优先关注四类能力:
这些能力覆盖了 HyperParallel 在 PyTorch 后端上最主要的训练语义变更点,也是测试方案应优先设防的对象。
本章只基于本地代码仓中已核实的测试、实现与文档,提炼 Megatron-LM 与 TorchTitan 目前实际防护了哪些训练场景,以及这些场景对应的代码路径。
本节代码依据:
Megatron-LM/tests/unit_tests
Megatron-LM/tests/functional_tests
tests/unit_tests/test_parallel_state.py
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
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
tests/unit_tests/pipeline_parallel/test_fine_grained_activation_offloading.py
tests/unit_tests/dist_checkpointing/test_serialization.py
tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py
tests/unit_tests/dist_checkpointing/test_optimizer.py
tests/functional_tests/python_test_utils/test_pretraining_regular_pipeline.py
tests/functional_tests/python_test_utils/test_pretraining_resume_checkpoint_pipeline.py
tests/unit_tests/transformer/moe/
core/multicore/
tests/unit_tests/distributed/test_grad_sync_with_expert_parallel.py
tests/unit_tests/a2a_overlap/
tests/functional_tests/python_test_utils/test_optimizer_grads_match.py
torchtitan/tests/unit_tests
torchtitan/tests/integration_tests
ParallelDims
tests/unit_tests/test_parallel_dims.py
test_parallel_dims.py
TestSingleGPUMixedPrecisionFSDP
tests/unit_tests/test_checkpoint.py
tests/unit_tests/test_activation_checkpoint.py
tests/unit_tests/test_set_determinism.py
tests/integration_tests/features.py
tests/assets/custom_schedule.csv
PP+TP
PP+DP+TP
FSDP+CP
FSDP+TP+CP
tp+cp+pp
ExpertParallel
ExpertTensorParallel
DeepEPExpertParallel
TorchAOExpertParallel
tests/unit_tests/test_expert_parallel.py
tests/unit_tests/test_fsdp_moe_sharding.py
tests/unit_tests/test_compile_moe.py
tests/unit_tests/test_tp_kv_heads_validation.py
float8_emulation
gradient_accumulation
torchcomms_3d_dp+cp+pp+compile
torchcomms_3d_dp+tp+pp+compile
torchcomms_3d_*
3d_compile
fsdp+flex_attn
fsdp+varlen_attn+per_op_sac
seed_checkpoint
sft
model_only_hf_checkpoint
last_save_model_only_fp32
last_save_model_only_bf16
从两边已核实的测试代码看,对 HyperParallel 最有价值的不是“支持了哪些并行功能”,而是它们如何把功能拆成可防护的场景。
可以直接借鉴的设计原则有四条:
因此,HyperParallel 首版方案中第二章不应再只列“Megatron/TorchTitan 支持哪些能力”,而应以“哪些场景已经被它们的测试体系明确防护”为参照,反推:
下表中的分布式训练能力为HyperParallel+Megatron+TorchTitan能力全集,下方能力矩阵用于评估:
parallel_state
fully_shard
features.py
ddp
torch_fully_sharded_data_parallel.py
hsdp
hsdp+tp
hsdp+cp
README
test_api.py
enable_sequence_parallel
2d_eager_no_sp
ScheduleGPipe
Schedule1F1B
ScheduleInterleaved1F1B
vp*
Interleaved1F1B
InterleavedZeroBubble
ZBVZeroBubble
ContextParallel
AsyncContextParallel
cp_allgather
cp_alltoall
fsdp+cp
torch/multicore/
mindspore/multicore/
none / selective / full
distributed_checkpoint
dist_checkpointing
load_state_dict
full_checkpoint
optional_checkpoint
init_weights.py
parameter_init.py
docs/fsdp.md
clip_grad_norm_
megatron/core/fp8_utils.py
1d_compile
2d_compile
enable_async_tensor_parallel
pp_custom_csv
fsdp+flex_attn+per_op_sac
HyperParallel 的主要风险并不集中在单一算子或单一 API,而是来自以下三类系统性问题:
因此,测试体系不能只按“目录”或“功能模块”组织,而必须分层组织,使不同类型的问题在不同层级被尽早发现。
不同能力域不需要在每一层平均用力,而应按风险映射:
测试建设优先级应遵循以下顺序:
这一层的目标是快速发现不依赖真实多卡执行的逻辑错误,运行应足够轻量、适合高频触发。
grad
这一层的目标是保护 HyperParallel 赖以成立的 PyTorch 执行假设,重点不是“训练结果像不像”,而是“挂钩点和生命周期对不对”。
register_forward_pre_hook()
register_forward_hook()
register_forward_pre_hook(with_kwargs=True)
queue_callback
register_load_state_dict_post_hook()
register_full_backward_pre_hook()
register_full_backward_hook()
state_dict()
这一层的目标是验证单个能力域在训练语义上的正确性,核心断言是数值、梯度和参数更新与基线一致。
forward output / loss / grad / step 后参数
loss / grad shard / step 后参数
none / recompute
这一层的目标是验证真实训练中更高风险的能力叠加场景,因为大量系统问题只会在组合场景中暴露。
Shard + TP
Shard / Tensor Parallel
TP + FSDP
Tensor Parallel / Fully Shard
loss / grad / step 后参数
Tensor Parallel / Fully Shard / Distributed Checkpoint / Resume
save -> load -> train
TP + PP
Tensor Parallel / Pipeline Parallel
torch.distributed.recv()
CP + TP
Context Parallel / Tensor Parallel
FSDP + CP
Fully Shard / Context Parallel
FSDP + AC
Fully Shard / Activation Checkpoint
PP + AC
Pipeline Parallel / Activation Checkpoint
Swap + AC
Activation Swap / Activation Checkpoint
Resume + FSDP + Meta Init
Distributed Checkpoint / Fully Shard / Init Weights / Meta Init
Clip Grad + FSDP
Clip Grad / Fully Shard
torch.Tensor.register_post_accumulate_grad_hook()
这一层的目标是把测试从“单次正确”推进到“持续稳定”,用于拦截长稳问题、恢复问题和重大回归。
save -> stop -> load -> continue train
save mesh != load mesh
TP × FSDP
TP × PP
FSDP × AC
resume × FSDP
HyperParallel PyTorch 后端的测试方案,不应再停留在“功能存在性验证”层面,而应围绕以下主线构建质量防护网:
在当前范围内,首版最优先补强的不是更多算子测试,而是:
这三类测试最能提升 HyperParallel PyTorch 后端在真实训练闭环中的可靠性。
README.md
tests/torch/*
hyper_parallel/platform/torch/*
hyper_parallel/core/*
HyperParallel PyTorch 后端质量防护网设计
本文聚焦 HyperParallel 代码仓中基于 PyTorch 后端构建的分布式训练能力,目标是为测试方案建立一套面向训练主链路的分层质量防护体系。
本文约束如下:
torch.compile1. HyperParallel 分布式训练能力总览
1.1 纳入本方案的能力范围
按当前仓库实现与测试现状,PyTorch 后端纳入本方案的能力域如下:
1.2 能力域总表
“依赖的 PyTorch 核心特性”以
/Users/liuchongming/Documents/Obsidian Vault/HyperParallel/PyTorch运行机制总结.md为主基线,并补充纳入以下工程相关机制:torch.distributed的 process group、collective、P2P 语义DeviceMesh/ DTensor / layout / redistribute 语义metatensor、materialization 生命周期step语义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/collectiveshyper_parallel/platform/torchtorch.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/dtensorhyper_parallel/platform/torch/dtensor.pyshard()、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/shardparallelize_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_parallelhyper_parallel/core/shardhyper_parallel/core/dtensorfully_shard()、参数 all-gather/unshard、backward reduce-scatter、reshard、mixed precision policytorch.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_shardhyper_parallel/platform/torch/fully_shardPipelineStage、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_parallelhyper_parallel/platform/torch/pipeline_paralleltorch.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_paralleltorch.utils.checkpoint.checkpoint();autograd 图重放语义;forward 重算与 backward 路径差异;saved tensor 生命周期hyper_parallel/core/activation_checkpointhyper_parallel/platform/torch/activation_checkpointtorch.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.pyhyper_parallel/platform/torch/activation_checkpointtorch.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_checkpointtorch.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/llamafactorymain_grad兼容torch.Tensor.grad可见语义;Tensor.register_post_accumulate_grad_hook()对应的累积后时机基线;torch.nn.utils.clip_grad_norm_()的插入点语义;optimizer step 前梯度后处理hyper_parallel/core/utils/clip_grad.pyhyper_parallel/platform/torch/clip_grad.pymetadevice 生命周期;materialization 语义;torch.nn.Module.load_state_dict();torch.nn.Module.register_load_state_dict_post_hook();分布式初始化与 layout 对齐hyper_parallel/core/dtensor/init_weights.pyhyper_parallel/core/dtensor/parameter_init.pyhyper_parallel/platform/torch/init_weights.pytorch.distributed.all_to_all();GEMM 算子;SwiGLU 激活函数;事件同步语义hyper_parallel/core/multicorehyper_parallel/platform/torch/multicoreshmem_put/get、原子信号操作、集合通信shmem_allgather/alltoall、融合操作fused_all_gather_matmul/fused_matmul_reduce_scatterhyper_parallel/core/symmetric_memoryOffsetBasedRNGTracker、分布式随机数生成、跨 rank seed 同步torch.random;随机数状态管理与同步hyper_parallel/core/dtensor/random.pytorch.amp;dtype 自动转换语义hyper_parallel/core/fully_shardhyper_parallel/platform/torch/fully_shard1.3 当前工程视角下的重点能力
从训练主链路风险看,本方案优先关注四类能力:
这些能力覆盖了 HyperParallel 在 PyTorch 后端上最主要的训练语义变更点,也是测试方案应优先设防的对象。
2. Megatron/TorchTitan 训练场景对标
本章只基于本地代码仓中已核实的测试、实现与文档,提炼 Megatron-LM 与 TorchTitan 目前实际防护了哪些训练场景,以及这些场景对应的代码路径。
2.1 Megatron-LM 已核实的训练场景与用例路径
本节代码依据:
Megatron-LM/tests/unit_testsMegatron-LM/tests/functional_teststests/unit_tests/test_parallel_state.pytests/unit_tests/tensor_parallel/test_initialization.pytests/unit_tests/tensor_parallel/test_layers.pytests/unit_tests/tensor_parallel/test_mappings.pytests/unit_tests/tensor_parallel/test_cross_entropy.pytests/unit_tests/tensor_parallel/test_random.pytests/unit_tests/pipeline_parallel/test_schedules.pytests/unit_tests/pipeline_parallel/test_bridge_communicator.pytests/unit_tests/pipeline_parallel/test_multimodule_schedules.pytests/unit_tests/pipeline_parallel/test_pipeline_layout.pytests/unit_tests/pipeline_parallel/test_fine_grained_activation_offloading.pytests/unit_tests/dist_checkpointing/test_serialization.pytests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.pytests/unit_tests/dist_checkpointing/test_optimizer.pytests/functional_tests/python_test_utils/test_pretraining_regular_pipeline.pytests/functional_tests/python_test_utils/test_pretraining_resume_checkpoint_pipeline.pytests/unit_tests/transformer/moe/core/multicore/),应建立对等的 EP 单元测试体系tests/unit_tests/distributed/test_grad_sync_with_expert_parallel.pytests/unit_tests/a2a_overlap/tests/functional_tests/python_test_utils/test_optimizer_grads_match.py2.2 TorchTitan 已核实的训练场景与用例路径
本节代码依据:
torchtitan/tests/unit_teststorchtitan/tests/integration_testsParallelDims构造、auto-calculate、enabled properties、mesh/get_mesh/get_optional_mesh、world_size 约束,以及单卡/8 卡 mesh 操作tests/unit_tests/test_parallel_dims.pytest_parallel_dims.py中包含TestSingleGPUMixedPrecisionFSDP,用于对齐 composable FSDP mixed precision 的关键行为tests/unit_tests/test_parallel_dims.pytests/unit_tests/test_checkpoint.pytests/unit_tests/test_activation_checkpoint.pytests/unit_tests/test_set_determinism.pytests/integration_tests/features.pytests/integration_tests/features.pytests/integration_tests/features.pytests/assets/custom_schedule.csvPP+TP、PP+DP+TP、FSDP+CP、FSDP+TP+CP、validation withtp+cp+pptests/integration_tests/features.pyExpertParallel、ExpertTensorParallel、DeepEPExpertParallel、TorchAOExpertParallel四种 EP 实现的正确性tests/unit_tests/test_expert_parallel.pytests/unit_tests/test_fsdp_moe_sharding.pytests/unit_tests/test_compile_moe.pytests/unit_tests/test_tp_kv_heads_validation.pytests/integration_tests/features.py(float8_emulation)tests/integration_tests/features.py(gradient_accumulation)torchcomms_3d_dp+cp+pp+compile、torchcomms_3d_dp+tp+pp+compile等三维并行 + compile 组合tests/integration_tests/features.py(torchcomms_3d_*、3d_compile)tests/integration_tests/features.py(fsdp+flex_attn、fsdp+varlen_attn+per_op_sac)tests/integration_tests/features.py(seed_checkpoint)tests/integration_tests/features.py(sft)tests/integration_tests/features.py(model_only_hf_checkpoint、last_save_model_only_fp32、last_save_model_only_bf16)2.3 对 HyperParallel 首版方案的直接启发
从两边已核实的测试代码看,对 HyperParallel 最有价值的不是“支持了哪些并行功能”,而是它们如何把功能拆成可防护的场景。
可以直接借鉴的设计原则有四条:
因此,HyperParallel 首版方案中第二章不应再只列“Megatron/TorchTitan 支持哪些能力”,而应以“哪些场景已经被它们的测试体系明确防护”为参照,反推:
2.4 HyperParallel / Megatron / TorchTitan 能力矩阵
下表中的分布式训练能力为HyperParallel+Megatron+TorchTitan能力全集,下方能力矩阵用于评估:
DeviceMesh、process group、collective/P2P 运行时底座。parallel_state明确管理 TP/PP/CP/EP/DP group。ParallelDims/DeviceMesh是所有并行维度入口。hyper_parallel/core/dtensor是 TP、FSDP、checkpoint 等主线基础。shard()/ShardingPlan/ module hook 注入,是 HyperParallel 的明显特征。ShardingPlan抽象。fully_shard/ HSDP / DTensor 体系。features.py中有显式ddp集成场景。hyper_parallel/core/fully_shard已形成 PyTorch 主路径。fully_shard/ FSDP2 为主。torch_fully_sharded_data_parallel.py提供了 HSDP 路径,相关单测已存在。features.py中显式覆盖hsdp、hsdp+tp、hsdp+cp。README与test_api.py均体现 1D 约束。enable_sequence_parallel是 TP 常规配置,且有2d_eager_no_sp对照场景。ScheduleGPipe、Schedule1F1B与 stage/runtime 实现。features.py显式覆盖 GPipe、1F1B、多种 PP 组合。ScheduleInterleaved1F1B,对应本文中的 VPP。vp*/ interleaved 场景。Interleaved1F1B、InterleavedZeroBubble已进入集成场景。features.py显式覆盖InterleavedZeroBubble与ZBVZeroBubble。ContextParallel/AsyncContextParallel,但 README 明确显示 Ulysses、Ring Attention 尚未完成。cp_allgather、cp_alltoall、fsdp+cp。core/multicore/提供完整 MoE FFN mega kernel 实现(AllToAll-Dispatch→GMM→SwiGLU→AllToAll-Combine),平台层有torch/multicore/和mindspore/multicore/适配。ExpertParallel、ExpertTensorParallel、DeepEPExpertParallel、TorchAOExpertParallel),含 DeepEP/HybridEP 后端和 MXFP8 支持。none / selective / full,集成场景也直接引用。distributed_checkpoint模块与 planner/storage/api。dist_checkpointing是主线能力,测试覆盖 serialization、optimizer、PP layout 等。load_state_dict修复路径,但当前工程测试覆盖仍需要补强。full_checkpoint、optional_checkpoint、seed_checkpoint、PP/TP/CP 组合恢复均已显式建场景。init_weights.py、parameter_init.py与fully_shard路径都覆盖 meta 生命周期。meta初始化。docs/fsdp.md、trainer 与 model parallelize 路径都将 meta init 作为正式能力。clip_grad_norm_。megatron/core/fp8_utils.py提供完整 FP8 精度支持与 Transformer Engine 集成。float8_emulation场景;Expert Parallel 路径支持 MXFP8。overlap_p2p、Async CP、Symmetric Memory 均提供通信-计算重叠能力。torch.compile。torch.compile。1d_compile、2d_compile、3d_compile等场景。features.py显式定义gradient_accumulation集成场景。enable_async_tensor_parallel配置项提供异步 TP 通信与计算重叠。pp_custom_csv)。fsdp+flex_attn+per_op_sac、fsdp+varlen_attn+per_op_sac等场景。3. HyperParallel 分层防护体系设计
3.1 为什么需要分层防护
HyperParallel 的主要风险并不集中在单一算子或单一 API,而是来自以下三类系统性问题:
因此,测试体系不能只按“目录”或“功能模块”组织,而必须分层组织,使不同类型的问题在不同层级被尽早发现。
3.2 五层防护的职责边界
3.3 能力域与防护层的映射原则
不同能力域不需要在每一层平均用力,而应按风险映射:
3.4 优先级原则
测试建设优先级应遵循以下顺序:
4. 五层防护展开
4.1 第 0 层:纯逻辑与静态单测
这一层的目标是快速发现不依赖真实多卡执行的逻辑错误,运行应足够轻量、适合高频触发。
应重点覆盖的能力与特性
ShardingPlan解析、输入输出布局注入规则、参数布局声明、kwargs/positional args 映射main_grad与grad选择逻辑这一层不解决的问题
4.2 第 1 层:PyTorch 核心机制防护
这一层的目标是保护 HyperParallel 赖以成立的 PyTorch 执行假设,重点不是“训练结果像不像”,而是“挂钩点和生命周期对不对”。
应重点覆盖的机制型测试
torch.nn.Module.__call__()是否按预期触发register_forward_pre_hook()/register_forward_hook();直接forward()调用是否会绕过框架假设torch.Tensor.backward()后参数.grad累积边界是否符合预期register_forward_pre_hook(with_kwargs=True)/register_forward_hook()触发时机;输出张量torch.Tensor.register_hook()是否准确进入 backward 入口;queue_callback收尾时机;register_load_state_dict_post_hook()后参数状态是否可继续训练torch.autograd.backward()与 stage 调度边界是否一致;send/recv 的调用次序是否与调度计划匹配;shared parameter 梯度同步是否落在正确边界register_forward_pre_hook(with_kwargs=True)/register_forward_hook()触发时 attention 张量布局是否符合假设;async 路径上register_full_backward_pre_hook()的时机是否稳定torch.utils.checkpoint.checkpoint()后 backward 是否走重算路径;重算路径下 hook 与保存张量生存期是否符合预期torch.autograd.graph.saved_tensors_hooks()的 pack/unpack 次序;register_full_backward_pre_hook()/register_full_backward_hook()的预取/释放边界是否正确state_dict()/load_state_dict()的对象生命周期;模型状态和优化器状态加载后是否仍与并行对象绑定正确.grad/main_grad可见性;clip 是否发生在 optimizer step 前;冻结参数与无 grad 参数是否被正确跳过meta参数 materialize 时机;load_state_dict()后参数对象是否完成修复;延迟初始化是否不会破坏后续 hook 假设这一层的典型失败信号
load_state_dict()后对象能加载但不能继续训练4.3 第 2 层:单并行能力正确性
这一层的目标是验证单个能力域在训练语义上的正确性,核心断言是数值、梯度和参数更新与基线一致。
应重点覆盖的能力测试
forward output / loss / grad / step 后参数对齐;关键模块如 linear、embedding、attention 相关路径对齐loss / grad shard / step 后参数对齐;mixed precision 开关组合;state_dict 保存恢复后继续训练none / recompute模式下 loss、grad、step 后参数对齐这一层的断言重点
4.4 第 3 层:组合能力交互防护
这一层的目标是验证真实训练中更高风险的能力叠加场景,因为大量系统问题只会在组合场景中暴露。
首版关键组合场景
Shard + TP输入输出布局注入场景Shard / Tensor Paralleltorch.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 透传TP + FSDP单 step 训练场景Tensor Parallel / Fully Shardtorch.nn.Module.register_forward_pre_hook()×torch.nn.Module.register_forward_hook();输出张量torch.Tensor.register_hook()×torch.autograd.Variable._execution_engine.queue_callback(...)loss / grad / step 后参数与基线对齐;backward 后参数状态正确;无训练卡死TP + FSDPstate_dict / resume 场景Tensor Parallel / Fully Shard / Distributed Checkpoint / Resumetorch.nn.Module.state_dict()×torch.nn.Module.load_state_dict();torch.nn.Module.register_load_state_dict_post_hook()× FSDP 参数修复路径save -> load -> train与连续训练一致;参数与 optimizer state 映射正确;恢复后可稳定继续训练TP + PP微批调度场景Tensor Parallel / Pipeline Paralleltorch.nn.Module.register_forward_pre_hook()×torch.nn.Module.register_forward_hook();torch.autograd.backward()×torch.distributed.send()/torch.distributed.recv()CP + TPattention 布局协同场景Context Parallel / Tensor Paralleltorch.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 回流FSDP + CPattention 训练场景Fully Shard / Context Paralleltorch.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 生命周期FSDP + AC重算训练场景Fully Shard / Activation Checkpointtorch.Tensor.register_hook()×torch.autograd.Variable._execution_engine.queue_callback(...);torch.utils.checkpoint.checkpoint()× backward 重算路径loss / grad / step 后参数对齐;backward 后参数状态正确;无训练卡死PP + AC微批重算场景Pipeline Parallel / Activation Checkpointtorch.utils.checkpoint.checkpoint()×torch.autograd.backward();micro-batch schedule ×torch.distributed.send()/recv()Swap + AC保存张量生命周期场景Activation Swap / Activation Checkpointtorch.autograd.graph.saved_tensors_hooks()×torch.utils.checkpoint.checkpoint();torch.nn.Module.register_full_backward_pre_hook()×torch.nn.Module.register_full_backward_hook()Resume + FSDP + Meta Init恢复场景Distributed Checkpoint / Fully Shard / Init Weights / Meta Inittorch.nn.Module.load_state_dict()×torch.nn.Module.register_load_state_dict_post_hook();meta materialization × FSDP 参数修复Clip Grad + FSDP梯度后处理场景Clip Grad / Fully Shardtorch.Tensor.register_post_accumulate_grad_hook()对应的累积后边界 ×torch.nn.utils.clip_grad_norm_();梯度同步完成边界 × optimizer stepmain_grad/grad读取错误这一层的关注点
4.5 第 4 层:长稳、恢复与回归防护
这一层的目标是把测试从“单次正确”推进到“持续稳定”,用于拦截长稳问题、恢复问题和重大回归。
应重点覆盖的场景
save -> stop -> load -> continue train与连续训练在 loss、grad、参数轨迹上保持一致save mesh != load mesh时模型与优化器状态仍能正确恢复TP × FSDP、TP × PP、FSDP × AC、resume × FSDP这一层的判定重点
结论
HyperParallel PyTorch 后端的测试方案,不应再停留在“功能存在性验证”层面,而应围绕以下主线构建质量防护网:
在当前范围内,首版最优先补强的不是更多算子测试,而是:
这三类测试最能提升 HyperParallel PyTorch 后端在真实训练闭环中的可靠性。
参考来源
/Users/liuchongming/Documents/Obsidian Vault/HyperParallel/PyTorch运行机制总结.mdREADME.md、tests/torch/*、hyper_parallel/platform/torch/*、hyper_parallel/core/*