You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

FSDP中init_device_mesh的mesh_dim_names参数未生效问题

FSDP设备网格维度配置问题

问题代码

device_mesh = init_device_mesh(
                "cuda",
                mesh_shape=(1, 8),
                mesh_dim_names=("replicate","shard"),
            )
model = FSDP(
    model,
    auto_wrap_policy=size_based_auto_wrap_policy,
    device_id=device,
    # sharding_strategy=ShardingStrategy.FULL_SHARD,
    mixed_precision=MixedPrecision(
        param_dtype=torch.float,
        reduce_dtype=torch.float,
        buffer_dtype=torch.float,
    ),
    sync_module_states=True,
    limit_all_gathers=True,
    use_orig_params=True,
    device_mesh=device_mesh,
)

问题描述

按照预期,将shard维度设为8、replicate维度设为1时模型应完全分片,但实际是no_shard的FSDP配置;交换mesh_dim_names顺序后问题依旧,说明FSDP并未按维度名称识别分片/复制维度,而是默认将第一个维度视为分片维度、第二个视为复制维度。需要明确如何可靠控制各维度用途,尤其是涉及张量并行(TP)、序列并行(SP)等多维度场景时。


解决方案

1. 显式指定FSDP的分片维度

FSDP不会自动通过mesh_dim_names识别分片维度,必须显式通过shard_dim参数指定目标维度(可传维度名称或索引),同时建议显式开启FULL_SHARD策略避免默认行为歧义。

修改后的FSDP初始化代码:

model = FSDP(
    model,
    auto_wrap_policy=size_based_auto_wrap_policy,
    device_id=device,
    sharding_strategy=ShardingStrategy.FULL_SHARD,
    mixed_precision=MixedPrecision(
        param_dtype=torch.float,
        reduce_dtype=torch.float,
        buffer_dtype=torch.float,
    ),
    sync_module_states=True,
    limit_all_gathers=True,
    use_orig_params=True,
    device_mesh=device_mesh,
    shard_dim="shard"  # 或传索引1,对应mesh_shape的第二个维度
)

2. 多并行维度(TP/SP)的处理

当同时使用FSDP+TP/SP时,需要先构建包含所有并行维度的设备网格,再为每个并行组件显式指定对应的维度:

  • 例如,若张量并行度为2、FSDP分片度为8,设备网格可设为mesh_shape=(2, 8),mesh_dim_names=("tp", "fsdp")
  • FSDP初始化时指定shard_dim="fsdp"
  • 张量并行组件需配置为使用tp维度,确保各组件的并行逻辑绑定到对应设备网格维度

内容的提问来源于stack exchange,提问作者李瑞斌

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.12 10:42:43