跳转至

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 的分布式配置,自动根据并行度选择最优的分片策略。