作者​:昇腾实战派​知识地图​:【昇腾实战派】综合指导知识地图-CSDN博客

背景概述

随着大模型训练的普及,分布式训练技术成为提升训练效率、降低显存占用的关键手段。PyTorch 提供的 Fully Sharded Data Parallel(FSDP)是一种模型分片并行策略,能够将模型参数、梯度和优化器状态分布到多个设备上,从而支持训练超出单设备显存容量的大模型。

FSDP2 是 FSDP 的下一代实现,基于 DTensor 进行参数管理,相比 FSDP1 在显存利用和计算效率上均有显著提升,当前昇腾MindSpeed训练框架已适配FSDP2

本文将从初始化、优化器部署、训练循环等方面详细介绍 FSDP2 的使用方法,并提供可直接运行的完整脚本。

官方文档参考:https://docs.pytorch.org/tutorials/intermediate/FSDP_tutorial.html#how-fsdp2-works

快速开始: 将可执行的分布式脚本保存为 train.py 后,运行以下命令即可启动训练:

torchrun --nproc_per_node 2 train.py

其中 2 为可用的设备数量(如 NPU、GPU),可根据实际情况调整。如需直接获取可测试脚本,请跳转至本文第3节


1. 初始化

1.1 模型导入与参数初始化

在开始分布式训练前,需要准备好模型定义文件。本文以官方提供的 Transformer 实现为例,请先从 GitHub 下载或克隆 model.py 文件。为避免模块名与变量名冲突,建议将文件重命名为 model_Tr.py。此外,还需导入 ModelArgs 类来配置模型参数。

导入与初始化的方法如下:

import torch.distributed as dist
from model_Tr import Transformer, ModelArgs

local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])

# 初始化通信域
backend = "hccl" if torch_npu.npu.is_available() else "gloo"
dist.init_process_group(backend=backend, rank=rank, world_size=world_size)

# 绑定设备
device = torch.device(f"npu:{local_rank}")

# 参数初始化
args = ModelArgs()
vocab_size = args.vocab_size
seq_len = args.max_seq_len
batch_size = 4
epochs = 5

# 模型初始化
model = Transformer(args)

1.2 FSDP 分片初始化

与 DDP 不同,FSDP2 不仅需要在根模型上调用 fully_shard,还需要对子模块(如每个 Transformer 层)分别应用 fully_shard。这样,在计算某一层时,其余层的参数处于分片状态,从而大幅降低显存占用。

fully_shard(model) 会将未独立分片的参数归组,并为它们安排高效的 all-gather / reduce-scatter 操作。调用 fully_shard 后,分片后的模型会被自动放置到训练设备上。

from torch.distributed.fsdp import fully_shard, FSDPModule

# 按照官方文档顺序应用 FSDP2
for layer in model.layers:
    fully_shard(layer)
fully_shard(model)

1.3 完整初始化代码

结合前文的参数初始化与模型导入,完整的初始化代码如下:

from model_Tr import Transformer, ModelArgs
import torch
import torch_npu
import torch.distributed as dist
from torch.distributed.fsdp import fully_shard, FSDPModule

local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])

# 初始化通信域
backend = "hccl" if torch_npu.npu.is_available() else "gloo"
dist.init_process_group(backend=backend, rank=rank, world_size=world_size)

# 绑定设备
device = torch.device(f"npu:{local_rank}")

# 参数初始化
args = ModelArgs()
vocab_size = args.vocab_size
seq_len = args.max_seq_len
batch_size = 4
epochs = 5

# 模型实例化
model = Transformer(args)

# 应用 FSDP2
for layer in model.layers:
    fully_shard(layer)
fully_shard(model)

# 验证分片是否成功
assert isinstance(model, Transformer)
assert isinstance(model, FSDPModule)
print(model)

预期输出示例:

FSDPTransformer(
  (tok_embeddings): Embedding(...)
  ...
  (layers): 3 x FSDPTransformerBlock(...)
  (output): Linear(...)
)


2. 优化器部署与训练测试

2.1 优化器部署

应用 fully_shard 后,模型参数会变为 DTensor,其分片策略为 Shard(0)(沿第 0 维切分,即按设备数等分)。创建优化器前,可通过以下循环验证参数类型与分片方式:

from torch.distributed.tensor import DTensor, Shard

for param in model.parameters():
    assert isinstance(param, DTensor)
    assert param.placements == (Shard(0),)
    # 可通过 param.to_local() 查看分片后的参数

optim = torch.optim.Adam(model.parameters(), lr=1e-2)

注意: 优化器必须在 fully_shard 之后构造。此时模型和优化器的状态字典都将以 DTensor 表示,torch.optim.Adam 和 torch.nn.utils.clip_grad_norm_ 等 API 可直接用于 DTensor 参数,无需额外适配,使得单设备与分布式训练的代码保持一致。

2.2 前向/反向传播与预取(Prefetching)

fully_shard 会注册前向/反向钩子,在计算前自动执行 All-Gather 以获取完整参数,并在计算后立即重新分片(reshard)。All-Gather 是一种集合通信操作:它将各个 rank 上按 dim=0 切分的本地参数聚合为完整张量并广播至所有进程,实现“按需物化”。计算结束后,通过 reshard 丢弃冗余副本,仅保留本地分片。

这一机制将显存复杂度由 O(model_size) 降为 O(model_size / N),本质是用通信带宽换取显存容量。同时,All-Gather 可在独立 CUDA 流中异步执行,借助预取机制与计算重叠,通信延迟能够被有效掩盖。

隐式预取(默认)

CPU 线程会在计算第 i 层之前自动发起该层的 All-Gather,并放入非默认流中,使其与第 i 层的计算并行。反向传播时,All-Gather 按逆序发起。隐式预取无需额外配置,训练循环与单卡训练完全一致:

for _ in range(epochs):
    x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
    loss = model(x).sum()
    loss.backward()
    optim.step()
    optim.zero_grad()

建议先从隐式预取开始,以了解开箱即用的性能表现。

显式预取

通过 set_modules_to_forward_prefetch 和 set_modules_to_backward_prefetch 可手动控制预取顺序。在下例中,计算第 i 层时会提前发起第 i+1i+2 层的 All-Gather。显式预取在以下场景中效果显著:

  • CPU 密集型负载:隐式预取下,CPU 可能来不及在核函数执行前发起下一层通信,显式预取可提前调度。
  • 预取多层:一次预取多个层可进一步提升重叠效果,但会占用更多显存。
  • 提前发起首次 All-Gather:调用 model.unshard() 可让通信更早开始,避免在 model(x) 调用时才暴露延迟。
num_to_forward_prefetch = 2
for i, layer in enumerate(model.layers):
    if i >= len(model.layers) - num_to_forward_prefetch:
        break
    layers_to_prefetch = [
        model.layers[i + j] for j in range(1, num_to_forward_prefetch + 1)
    ]
    layer.set_modules_to_forward_prefetch(layers_to_prefetch)

num_to_backward_prefetch = 2
for i, layer in enumerate(model.layers):
    if i < num_to_backward_prefetch:
        continue
    layers_to_prefetch = [
        model.layers[i - j] for j in range(1, num_to_backward_prefetch + 1)
    ]
    layer.set_modules_to_backward_prefetch(layers_to_prefetch)

for _ in range(epochs):
    # 提前触发首次 all-gather,与 model(x) 之前的计算重叠
    model.unshard()
    x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
    loss = model(x).sum()
    loss.backward()
    optim.step()
    optim.zero_grad()


3. 脚本总结

请确保已将 model.py 下载并重命名为 model_Tr.py,置于当前工作目录下。

完整脚本(隐式预取)

适合快速验证分布式训练,无需手动配置预取,训练循环与单卡一致。

from model_Tr import Transformer, ModelArgs

import torch
import torch_npu
from torch.distributed.fsdp import fully_shard, FSDPModule
from torch.distributed.tensor import DTensor, Shard
def main():

	local_rank = int(os.environ["LOCAL_RANK"])
	rank = int(os.environ["RANK"])
	world_size = int(os.environ["WORLD_SIZE"])
    
	# 初始化通信域
	backend = "hccl" if torch_npu.npu.is_available() else "gloo"
	dist.init_process_group(backend=backend, rank=rank, world_size=world_size)
    
	# 绑定设备
	device = torch.device(f"npu:{local_rank}")

    #参数初始化
    args = ModelArgs()
    vocab_size = args.vocab_size
    seq_len = args.max_seq_len
    batch_size = 4
    epochs = 5
    local_rank = int(os.environ["LOCAL_RANK"])
	device = torch.device(f"npu:{local_rank}")


    # 创建模型
    model = Transformer(args)

    # 按照官方文档顺序应用 FSDP2
    for layer in model.layers:
        fully_shard(layer)
    fully_shard(model)

    assert isinstance(model, Transformer)
    assert isinstance(model, FSDPModule)
    print(model)
    #用于检测模型是否分片成功
    #——————————————输出样例————————————————
    #  FSDPTransformer(
    #    (tok_embeddings): Embedding(...)
    #    ...
    #    (layers): 3 x FSDPTransformerBlock(...)
    #    (output): Linear(...)
    #  )
    # 按照官方文档顺序应用 FSDP2
    #—————————————————————————————————————

    for param in model.parameters():
        assert isinstance(param, DTensor)
        assert param.placements == (Shard(0),)
        # 用于检测参数是否分片成功

    optim = torch.optim.Adam(model.parameters(), lr=1e-2)

    for _ in range(epochs):
        x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
        loss = model(x).sum()
        loss.backward()
        optim.step()
        optim.zero_grad()
        print(f"loss: {loss:.4f}")


if __name__ == "__main__":
    main()

完整脚本(显式预取)

通过多级预取和提前 unshard 进一步重叠通信与计算,以额外显存换取更高的吞吐。

from model_Tr import Transformer, ModelArgs

import torch
import torch_npu
from torch.distributed.fsdp import fully_shard, FSDPModule
from torch.distributed.tensor import DTensor, Shard

def main():

	local_rank = int(os.environ["LOCAL_RANK"])
	rank = int(os.environ["RANK"])
	world_size = int(os.environ["WORLD_SIZE"])
 
	# 初始化通信域
	backend = "hccl" if torch_npu.npu.is_available() else "gloo"
	dist.init_process_group(backend=backend, rank=rank, world_size=world_size)
 
	# 绑定设备
	device = torch.device(f"npu:{local_rank}")

    #参数初始化
    args = ModelArgs()
    vocab_size = args.vocab_size
    seq_len = args.max_seq_len
    batch_size = 4
    epochs = 5
    local_rank = int(os.environ["LOCAL_RANK"])
	device = torch.device(f"npu:{local_rank}")


    # 创建模型
    model = Transformer(args)

    for layer in model.layers:
        fully_shard(layer)
    fully_shard(model)

    assert isinstance(model, Transformer)
    assert isinstance(model, FSDPModule)
    print(model)
    #——————————————输出样例————————————————
    #  FSDPTransformer(
    #    (tok_embeddings): Embedding(...)
    #    ...
    #    (layers): 3 x FSDPTransformerBlock(...)
    #    (output): Linear(...)
    #  )
    # 按照官方文档顺序应用 FSDP2
    #—————————————————————————————————————

    for param in model.parameters():
        assert isinstance(param, DTensor)
        assert param.placements == (Shard(0),)
        # inspect sharded parameters with param.to_local()

    optim = torch.optim.Adam(model.parameters(), lr=1e-2)

    num_to_forward_prefetch = 2
    for i, layer in enumerate(model.layers):
        if i >= len(model.layers) - num_to_forward_prefetch:
            break
        layers_to_prefetch = [
            model.layers[i + j] for j in range(1, num_to_forward_prefetch + 1)
        ]
        layer.set_modules_to_forward_prefetch(layers_to_prefetch)

    num_to_backward_prefetch = 2
    for i, layer in enumerate(model.layers):
        if i < num_to_backward_prefetch:
            continue
        layers_to_prefetch = [
            model.layers[i - j] for j in range(1, num_to_backward_prefetch + 1)
        ]
        layer.set_modules_to_backward_prefetch(layers_to_prefetch)

    for _ in range(epochs):
        # trigger 1st all-gather earlier
        # this overlaps all-gather with any computation before model(x)
        model.unshard()
        x = torch.randint(0, vocab_size, (batch_size, seq_len), device=device)
        loss = model(x).sum()
        loss.backward()
        optim.step()
        optim.zero_grad()


if __name__ == "__main__":
    main()

将上述任一脚本保存为 train.py 后,运行以下命令即可开始训练:

torchrun --nproc_per_node 2 train.py
Logo

作为“人工智能6S店”的官方数字引擎,为AI开发者与企业提供一个覆盖软硬件全栈、一站式门户。

更多推荐