PyTorch提供分布式训练方案。
import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP # 初始化 dist.init_process_group(backend='nccl') # 包装模型 model = DDP(model)
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP model = FSDP(model, sharding_strategy=ShardingStrategy.FULL_SHARD)
DDP: 内存占用高,速度快
FSDP: 内存占用低,速度略慢
PyTorch分布式训练加速大模型训练。