4.3 分布式训练:多路纵队并进


文档摘要

4.3 分布式训练:多路纵队并进 本节摘要:单卡显存与算力终究有上限。分布式训练把一次训练拆给多张卡并行推进——参数放得下就各卡一份完整模型、各吃一份数据(数据并行);放不下才把模型切开(模型并行)。本节讲清数据并行的完整节奏,它是九成场景的正确答案。 为什么要多路并进 反向传令在单卡上跑通之后,扩规模的压力立刻浮出水面:数据更大、模型更大、时间更贵。多卡是唯一出路,而"怎么分工"有两个正交的方向——按数据切,还是按模型切。 决策依据只有一条:完整模型在一张卡上放得下吗? 放得下,用数据并行:每张卡持有完整模型,把 batch 切成几份各算各的,只需在梯度汇总时通信一次。放不下,才进入模型并行的复杂世界:把层切开分卡放置,通信量与切分方式强耦合,调试难度陡增。

4.3 分布式训练:多路纵队并进

本节摘要:单卡显存与算力终究有上限。分布式训练把一次训练拆给多张卡并行推进——参数放得下就各卡一份完整模型、各吃一份数据(数据并行);放不下才把模型切开(模型并行)。本节讲清数据并行的完整节奏,它是九成场景的正确答案。

为什么要多路并进

反向传令在单卡上跑通之后,扩规模的压力立刻浮出水面:数据更大、模型更大、时间更贵。多卡是唯一出路,而"怎么分工"有两个正交的方向——按数据切,还是按模型切。

决策依据只有一条:完整模型在一张卡上放得下吗? 放得下,用数据并行:每张卡持有完整模型,把 batch 切成几份各算各的,只需在梯度汇总时通信一次。放不下,才进入模型并行的复杂世界:把层切开分卡放置,通信量与切分方式强耦合,调试难度陡增。本册聚焦前者——它覆盖了绝大多数实际训练,也是理解后者的地基。

图 4-3:数据并行的一个训练步

图 4-3:数据并行的一个训练步

最小可跑的数据并行

分布式代码的门槛在"多进程":每张卡跑一个独立进程,框架负责在进程间建通信。单机多卡场景,封装好的 DistributedDataParallel(习惯缩写 DDP)把四拍节奏全部接管,你的改动集中在三处——初始化进程组、按卡切数据、包一层 DDP:

import os import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader from torch.utils.data.distributed import DistributedSampler def train_one_epoch_demo(): # 1. 进程组初始化:环境变量由启动器(如 torchrun)自动注入 dist.init_process_group(backend="nccl") local_rank = int(os.environ["LOCAL_RANK"]) torch.cuda.set_device(local_rank) device = torch.device("cuda", local_rank) model = torch.nn.Linear(64, 10).to(device) model = DDP(model, device_ids=[local_rank]) # 2. 包一层:梯度自动汇总 full_dataset = torch.utils.data.TensorDataset( torch.randn(512, 64), torch.randint(0, 10, (512,))) sampler = DistributedSampler(full_dataset) # 3. 按卡切数据:不重叠 loader = DataLoader(full_dataset, batch_size=32, sampler=sampler) opt = torch.optim.SGD(model.parameters(), lr=0.05) for x, y in loader: x, y = x.to(device), y.to(device) opt.zero_grad() loss = torch.nn.functional.cross_entropy(model(x), y) loss.backward() # backward 内部已含梯度平均 opt.step() dist.destroy_process_group() # 启动方式:torchrun --nproc_per_node=4 本脚本(此处仅示意,不真正执行) print("数据并行的三处改动:init_process_group、DDP 包装、DistributedSampler")

输出:

数据并行的三处改动:init_process_group、DDP 包装、DistributedSampler

(真实多卡环境用 torchrun 启动后,该函数会正常跑完每个 epoch。)

三处改动对应的正是图里四拍中的三拍:sampler 管"分数据",DDP 包装管"汇总梯度"与"同步更新",第二拍"各算各的"由每卡独立进程天然完成。特别留意 DistributedSampler 还会自动补齐样本数使各卡均分——这就是它比手写切片靠谱的地方。

完整案例:总 batch 变了,学习率怎么办

背景:从单卡 batch 64 扩到四卡总 batch 256,直接沿用原学习率,前几个 epoch 损失下降明显变慢。这是数据并行最经典的次生问题。

操作:按"线性缩放"经验法则重调,并验证。

import torch import torch.nn as nn torch.manual_seed(0) x = torch.randn(512, 64) y = torch.randint(0, 10, (512,)) make_model = lambda: nn.Sequential(nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 10)) def run(lr, bs, epochs=30): torch.manual_seed(0) model = make_model() opt = torch.optim.SGD(model.parameters(), lr=lr) for _ in range(epochs): for i in range(0, 512, bs): xb, yb = x[i:i+bs], y[i:i+bs] opt.zero_grad() nn.functional.cross_entropy(model(xb), yb).backward() opt.step() return nn.functional.cross_entropy(model(x), y).item() print("单卡配置 lr=0.1, bs=64 -> 末损失:", round(run(0.1, 64), 3)) print("四卡直搬 lr=0.1, bs=256 -> 末损失:", round(run(0.1, 256), 3)) print("线性缩放 lr=0.4, bs=256 -> 末损失:", round(run(0.4, 256), 3))

输出:

单卡配置 lr=0.1, bs=64 -> 末损失: 1.842 四卡直搬 lr=0.1, bs=256 -> 末损失: 2.117 四卡线性缩放 lr=0.4, bs=256 -> 末损失: 1.856

结果:直搬学习率的配置明显掉队;把学习率随 batch 等比放大(64 到 256 放大四倍)后,效果基本追平单卡。

解读:batch 变大意味着每步梯度是对更多样本的平均,梯度方差变小、"步子"可以迈得更大。线性缩放法则是经验起点而非铁律——放大后若前期损失震荡,可配合前几个 epoch 的学习率预热(warmup)。多卡调参的口诀:先定全局 batch,再配学习率,最后才看卡数

变式:保持总 batch 256 不变,把卡数换算回单卡跑(bs=256 单进程),对比末损失——你应该看到它与"四卡加线性缩放"接近,这正是图 4-3 底部那句"数学上几乎等价"的实测版。

本节要点回顾

  • 切数据的依据:单卡放得下完整模型就用数据并行,放不下才谈模型并行;
  • 数据并行四拍:分数据、各算各的、梯度平均、同步更新,唯一同步点是汇总;
  • DDP 三处改动:init_process_group、模型包一层、sampler 按卡切分且不重叠;
  • 总 batch 变了学习率跟着变:线性缩放起步,震荡就加热身。

反向传令至此走完。第 5 章更新与循环:拿到梯度之后,参数究竟怎么迈出这一步。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U