utils.py — FSDP 工具函数¶
文件概述¶
提供 FSDP 引擎所需的设备网格创建和分片策略选择工具。
核心函数¶
create_device_mesh¶
def create_device_mesh(world_size, fsdp_size):
"""创建 FSDP 设备网格
- fsdp_size < 0 或 >= world_size: 纯 FSDP (FULL_SHARD)
mesh_shape = (world_size,)
- 0 < fsdp_size < world_size: 混合分片 (HYBRID_SHARD)
mesh_shape = (world_size // fsdp_size, fsdp_size)
即: (DDP组数, 每组FSDP大小)
"""
get_sharding_strategy¶
def get_sharding_strategy(device_mesh):
"""根据设备网格维度选择分片策略
1D mesh → FULL_SHARD(完全分片)
2D mesh → HYBRID_SHARD(混合分片:组内 FSDP,组间 DDP)
"""
apply_npu_fsdp_patches¶
为华为 NPU 设备应用 FSDP 兼容性补丁。
小结¶
这些工具函数简化了 FSDP 的分布式配置,自动根据并行度选择最优的分片策略。