版本:v1.1 | 范围:fully_shard API + HSDPModule 接口
fully_shard
原则:覆盖 fully_shard() 函数入参 + HSDPModule 方法的覆盖 gap;语义相关接口合并到一个 case;2–4 卡;简单 MLP;目标是接口行为正确性,不做完整精度对标。
fully_shard()
mesh=None
test_fully_shard_none_mesh
mesh=DeviceMesh
reshard_after_forward=True/False
shard_placement_fn
_test_hsdp_param
mp_policy
fp16/fp32/fp32
offload_policy
pin_memory=True
ignored_params
replicate_params
comm_fusion
module=List[Module]
set_requires_gradient_sync
test_fully_shard_module
zero_grad
set_modules_to_forward_prefetch
set_modules_to_backward_prefetch
reshard()
unshard()
set_is_last_backward
set_requires_all_reduce
set_reshard_after_forward
set_reshard_after_backward
set_reduce_op_type
test_dynamic_reshard_control
_test_fully_shard_module.py
set_reshard_after_forward(False)
param.data.numel()==full_numel
shard_numel
test_set_requires_all_reduce
atol=1e-6
recurse=False
test_set_reduce_op_type
"sum"
"avg"
test_zero_grad
zero_grad()
param.grad is None
param.data
test_manual_reshard_unshard_lifecycle
async_op=True
test_ignored_params
_test_fully_shard_precision.py
numel==full_numel
test_custom_shard_placement_fn
_test_hsdp_param.py
param.shape[1]==full_shape[1]//world_size
atol=1e-5
Replicate()
test_comm_fusion_numerical_correctness
abs(loss_fusion - loss_nofusion) < 1e-6
test_fully_shard_list_input
MixedPrecisionPolicy 有 5 个字段:param_dtype、reduce_dtype、output_dtype、cast_forward_inputs、apply_grad_on_fp32_main_grad。 当前仅覆盖了 fp16/fp32/fp32 一种组合,以下补强:
MixedPrecisionPolicy
param_dtype
reduce_dtype
output_dtype
cast_forward_inputs
apply_grad_on_fp32_main_grad
test_mp_bf16_param_fp32_reduce
param_dtype=bf16, reduce_dtype=fp32, output_dtype=fp32
atol=1e-3
test_mp_bf16_all_bf16
param_dtype=bf16, reduce_dtype=bf16, output_dtype=bf16
test_mp_cast_forward_inputs_false
param_dtype=fp16, reduce_dtype=fp32, cast_forward_inputs=False
test_mp_apply_grad_on_fp32_main_grad
param_dtype=fp16, reduce_dtype=fp32, apply_grad_on_fp32_main_grad=True
param.grad
混合精度覆盖矩阵(补强后):
目标文件:L1-10~13 均归入 tests/torch/fully_shard/_test_fully_shard_precision.py
tests/torch/fully_shard/_test_fully_shard_precision.py
原则:简单网络(≤6 层);精度统一对标单卡/DDP;按特性命名归档文件。
test_gradient_accumulation
test_frozen_params_mixed_training
grad is None
test_reentrant_ac_fsdp
_test_fully_shard_with_ac.py
test_nonreentrant_ac_fsdp
test_selective_ac_fsdp
test_ac_activation_swap_fsdp
test_grad_clip_with_dtensor_grads
test_dynamic_sequence_length
test_pp_fsdp_coupling
_test_fully_shard_with_pp.py
test_ep_fsdp_coupling
_test_fully_shard_with_ep.py
atol=1e-4
L2 目标文件命名说明:
原则:
test_llm_decoder_hsdp_vs_ddp
_test_fully_shard_with_llm.py
(replicate=2, shard=2)
test_llm_decoder_fsdp_1d_vs_ddp
test_moe_llm_ep_fsdp_vs_ddp
(replicate=1, fsdp=2, ep=2)
test_vit_fsdp_vs_ddp
_test_fully_shard_with_vit.py
test_chunked_ce_fsdp_vs_ddp
test_multimodal_deepstack_fsdp
L3-01 / L3-02(LLM 1D/2D-mesh vs DDP):
L3-03(MoE EP+FSDP vs 单卡):
L3-04(ViT FSDP vs DDP):
L3-05(Chunked CE FSDP vs DDP):
sum(per_chunk_losses) == total_loss
L3-06(多模态 Deep Stack):
L3 目标文件命名说明:
精度自洽规则:
1e-5
1e-4
hyper_parallel/core/fully_shard/api.py
hyper_parallel/core/fully_shard/utils.py
hyper_parallel/platform/torch/fully_shard/
tests/torch/fully_shard/_test_fully_shard_module.py
tests/torch/fully_shard/_test_hsdp_param.py
tests/torch/fully_shard/_test_fully_shard_with_ac.py
tests/torch/fully_shard/_test_fully_shard_with_pp.py
tests/torch/fully_shard/_test_fully_shard_with_ep.py
tests/torch/fully_shard/_test_fully_shard_with_llm.py
tests/torch/fully_shard/_test_fully_shard_with_vit.py
HyperParallel FSDP 测试方案
Level 1:HSDPModule 接口单点测试
原则:覆盖
fully_shard()函数入参 + HSDPModule 方法的覆盖 gap;语义相关接口合并到一个 case;2–4 卡;简单 MLP;目标是接口行为正确性,不做完整精度对标。fully_shard() 入参 gap 分析
mesh=Nonetest_fully_shard_none_mesh)mesh=DeviceMeshreshard_after_forward=True/Falseshard_placement_fn_test_hsdp_param中有单元,但fully_shardAPI 层未独立测试mp_policy全字段组合fp16/fp32/fp32一种组合offload_policypin_memory=True一种ignored_paramsreplicate_paramscomm_fusionmodule=List[Module]HSDPModule 方法 gap 分析
set_requires_gradient_synctest_fully_shard_module)zero_gradset_modules_to_forward_prefetchset_modules_to_backward_prefetchreshard()/unshard()set_is_last_backwardset_requires_all_reduceset_reshard_after_forward(动态)set_reshard_after_backward(动态)set_reduce_op_typeL1 用例列表
test_dynamic_reshard_control_test_fully_shard_module.pyset_reshard_after_forward(False)→ forward 后param.data.numel()==full_numel;toggle 回 True →shard_numel;set_reshard_after_backward同理;动态切换不抛异常test_set_requires_all_reduce_test_fully_shard_module.pyatol=1e-6一致;recurse=False只影响顶层 moduletest_set_reduce_op_type_test_fully_shard_module.py"sum"→ grad == ref_grad × world_size;"avg"→ grad == ref_grad;mid-train 切换不污染 optimizer statetest_zero_grad_test_fully_shard_module.pyzero_grad():param.grad is None;param.data与 snapshot 一致;后续 backward 产生新鲜 gradtest_manual_reshard_unshard_lifecycle_test_fully_shard_module.pyasync_op=True时 handle wait 后方可读取test_ignored_params_test_fully_shard_precision.pynumel==full_numel;grad 不跨 rank 同步;其余参数正常 reduce-scattertest_custom_shard_placement_fn_test_hsdp_param.pyparam.shape[1]==full_shape[1]//world_size;loss + gradatol=1e-5vs 单卡;Replicate()返回时 shape == full_shapetest_comm_fusion_numerical_correctness_test_fully_shard_precision.pyabs(loss_fusion - loss_nofusion) < 1e-6;per-param gradatol=1e-6(comm_fusion 开/关数值等价性)test_fully_shard_list_input_test_fully_shard_module.pyL1-10~13:混合精度策略充分校验(重点补强)
MixedPrecisionPolicy有 5 个字段:param_dtype、reduce_dtype、output_dtype、cast_forward_inputs、apply_grad_on_fp32_main_grad。当前仅覆盖了
fp16/fp32/fp32一种组合,以下补强:mp_policy配置test_mp_bf16_param_fp32_reduceparam_dtype=bf16, reduce_dtype=fp32, output_dtype=fp32atol=1e-3vs 单卡 bf16 参考test_mp_bf16_all_bf16param_dtype=bf16, reduce_dtype=bf16, output_dtype=bf16atol=1e-3test_mp_cast_forward_inputs_falseparam_dtype=fp16, reduce_dtype=fp32, cast_forward_inputs=Falsetest_mp_apply_grad_on_fp32_main_gradparam_dtype=fp16, reduce_dtype=fp32, apply_grad_on_fp32_main_grad=Trueparam.graddtype 为 fp32;精度atol=1e-5vs fp32 单卡参考混合精度覆盖矩阵(补强后):
param_dtypereduce_dtypecast_forward_inputsapply_grad_on_fp32_main_gradLevel 2:特性耦合测试
原则:简单网络(≤6 层);精度统一对标单卡/DDP;按特性命名归档文件。
L2 用例列表
test_gradient_accumulation_test_fully_shard_precision.pyatol=1e-5;中间步 rank 间 grad 不等(确认未同步)test_frozen_params_mixed_training_test_fully_shard_precision.pygrad is None;data 不变;非 frozen 层 gradatol=1e-5vs 单卡;FSDP 不为 frozen 参数发起 reduce-scattertest_reentrant_ac_fsdp_test_fully_shard_with_ac.pyatol=1e-5;per-layer grad normatol=1e-5;首/末层 grad tensoratol=1e-5vs 单卡test_nonreentrant_ac_fsdp_test_fully_shard_with_ac.pytest_selective_ac_fsdp_test_fully_shard_with_ac.pytest_ac_activation_swap_fsdp_test_fully_shard_with_ac.pyatol=1e-5;grad normatol=1e-5;swap 回 GPU 后无 device mismatch;recompute 时激活在 GPU 上test_grad_clip_with_dtensor_grads_test_fully_shard_precision.pyatol=1e-5vs 单卡;全局 norm 而非 per-shard normtest_dynamic_sequence_length_test_fully_shard_precision.pytest_pp_fsdp_coupling_test_fully_shard_with_pp.pyatol=1e-5vs 4卡 DDP;stage boundary gradatol=1e-5;无死锁(60s timeout)test_ep_fsdp_coupling_test_fully_shard_with_ep.pyatol=1e-4;per-expert grad normatol=1e-4;无 NCCL 报错L2 目标文件命名说明:
_test_fully_shard_precision.py(已有)_test_fully_shard_with_ac.py(新建)_test_fully_shard_with_pp.py(新建)_test_fully_shard_with_ep.py(新建)Level 3:真实网络结构测试
原则:
L3 用例列表
test_llm_decoder_hsdp_vs_ddp_test_fully_shard_with_llm.py(replicate=2, shard=2)vs DDPtest_llm_decoder_fsdp_1d_vs_ddp_test_fully_shard_with_llm.pytest_moe_llm_ep_fsdp_vs_ddp_test_fully_shard_with_llm.py(replicate=1, fsdp=2, ep=2)vs 单卡test_vit_fsdp_vs_ddp_test_fully_shard_with_vit.pytest_chunked_ce_fsdp_vs_ddp_test_fully_shard_with_llm.pytest_multimodal_deepstack_fsdp_test_fully_shard_with_vit.pyfully_shard各 L3 用例断言规格
L3-01 / L3-02(LLM 1D/2D-mesh vs DDP):
atol=1e-5;step1 loss(optimizer 更新后)atol=1e-5atol=1e-5atol=1e-5atol=1e-5L3-03(MoE EP+FSDP vs 单卡):
atol=1e-4;dense 层 grad normatol=1e-5;per-expert grad normatol=1e-4L3-04(ViT FSDP vs DDP):
atol=1e-5;grad normatol=1e-5atol=1e-5L3-05(Chunked CE FSDP vs DDP):
atol=1e-5;per-chunk lossatol=1e-5;sum(per_chunk_losses) == total_loss(atol=1e-6,一致性验证)atol=1e-5;输出层 weight grad normatol=1e-5L3-06(多模态 Deep Stack):
atol=1e-5;ViT grad normatol=1e-5;Decoder grad normatol=1e-5atol=1e-5L3 目标文件命名说明:
_test_fully_shard_with_llm.py(新建)_test_fully_shard_with_vit.py(新建)精度自洽矩阵
精度自洽规则:
1e-5失败但1e-4通过,必须 triage 根因,说明存在数值累积 bug,不允许静默放宽关键文件映射
hyper_parallel/core/fully_shard/api.pyhyper_parallel/core/fully_shard/utils.pyhyper_parallel/platform/torch/fully_shard/tests/torch/fully_shard/_test_fully_shard_module.pytests/torch/fully_shard/_test_fully_shard_precision.pytests/torch/fully_shard/_test_hsdp_param.pytests/torch/fully_shard/_test_fully_shard_with_ac.pytests/torch/fully_shard/_test_fully_shard_with_pp.pytests/torch/fully_shard/_test_fully_shard_with_ep.pytests/torch/fully_shard/_test_fully_shard_with_llm.pytests/torch/fully_shard/_test_fully_shard_with_vit.py