在这里插入图片描述

PyTorch Scala 高校计算机硕士研一课程

章节 13: 分布式训练与并行

现代深度学习模型经常超出单个GPU的内存容量,并且在大量数据集上训练可能需要不切实际的时间。本章侧重于PyTorch中的分布式训练和并行技术,以应对这些挑战。

我们将研究跨多个GPU和节点扩展训练的方法。主要内容包括:

  • 与训练相关的分布式计算基本思想。
  • 数据并行,使用DistributedDataParallel (DDP)。
  • 处理极大型模型的方法,例如张量模型并行和流水线并行。
  • 内存高效的完全分片数据并行 (FSDP)。
  • 使用不同后端配置分布式环境。
  • 直接使用PyTorch的底层通信原语 (torch.distributed)。

在本章结束时,您将明白如何应用各种并行处理策略,以更高效地训练更大的模型。

分布式计算基本原理

训练当前的深度学习模型面临着重大的计算障碍。拥有数十亿参数的模型可能无法在单个加速器的内存中容纳,而在一台机器上处理数TB数据可能使训练时间从数天延长到数周甚至数月。分布式计算提供了一条出路,它能够汇集多台机器(或单台机器)上的多个设备(如GPU)的资源,以解决这些大规模问题。在实施诸如 DistributedDataParallel (DDP) 或完全分片数据并行 (FSDP) 等特定的PyTorch策略之前,掌握这些分布式配置中使用的基本术语和通信模式非常必要。

分布式训练中的核心术语

在讨论分布式训练时,一些术语经常出现。理解它们的准确含义对于配置和调试分布式任务很重要。

  • 节点 (Node): 指您的配置中的一台独立计算机器。这可以是机架中的物理服务器,也可以是云中的虚拟机实例。一个节点通常包含一个或多个处理单元(CPU、GPU)。
  • 进程 / 工作器 (Process / Worker): 运行在节点上的Python训练脚本的一个独立实例。在典型的基于GPU的训练中,通常为每个GPU启动一个进程,以最大限度地发挥硬件效能。这些进程并行执行,需要协调。
  • 秩 (Rank): 分配给参与分布式计算的每个进程的唯一整数标识符。秩通常从0到 N−1N−1 变化,其中 NN 是所涉及的进程总数。按照惯例,秩为0的进程通常承担特殊职责,例如日志记录或保存检查点,但这并非严格要求。
  • 大小 (Size): 在分布式训练任务中配合工作的进程总数 NN。如果您在4个节点上进行训练,每个节点有8个GPU,并且每个GPU运行一个进程,则总大小为 4×8=324×8=32。
  • 进程组 (Process Group): 所有进程的一个已定义子集(即组)。默认情况下,所有进程都属于一个组。但是,PyTorch允许创建子组,这对于更复杂的并行方案(如混合数据并行和模型并行)很有用,因为在这些方案中,不同类型的工作器集合之间可能会发生不同类型的通信。
  • 后端 (Backend): 促进进程间消息传递的底层通信库。PyTorch的 torch.distributed 包支持多种后端:
    • NCCL (NVIDIA Collective Communications Library): 针对NVIDIA硬件上基于GPU训练的首选后端。它对GPU间通信进行了高度优化,无论是在节点内部(使用NVLink)还是跨节点(使用InfiniBand或以太网等网络接口)。
    • Gloo: 一个平台无关的后端,适用于基于CPU的通信,以及在NCCL可能不佳或不可用的不同节点类型或网络设置中GPU之间的通信。它也支持GPU,但对于GPU集合通信通常比NCCL慢。
    • MPI (消息传递接口): 高性能计算通信的标准。如果您的集群环境已配置为MPI,则可以使用它,但NCCL或Gloo在PyTorch生态系统中更常见。

集体通信操作

分布式训练高度依赖集体通信操作,其中多个进程同时同步并交换数据。这些是构建DDP等高级策略的底层原语。torch.distributed 提供了这些操作的函数:

  • 广播 (torch.distributed.broadcast):将一个张量从一个指定的进程(src 秩)发送到组中的所有其他进程。这通常在训练开始时使用,以确保所有工作器都以完全相同的初始模型参数开始。
  • 归约 (torch.distributed.reduce):从组中所有进程收集张量,应用指定的归约操作(如 SUMAVGMAXMIN),并将结果存储在单个目标进程(dst 秩)上。
  • All-Reduce (torch.distributed.all_reduce):与归约类似,但归约操作的最终结果会分发回组中的所有进程。这是DDP的核心,每个工作器独立计算的梯度会在所有工作器之间求平均,确保模型在各处得到一致的更新。
  • 分散 (torch.distributed.scatter):从单个源进程(src)获取一个张量列表,并将列表中的一个张量分发给组中的每个进程(包括自身)。列表中的第 ii 个张量会发送给秩为 ii 的进程。
  • 收集 (torch.distributed.gather):分散的逆操作。每个进程将其张量发送到指定的目的进程(dst),后者将它们收集成一个按秩排序的张量列表。
  • All-Gather (torch.distributed.all_gather):与收集类似,但组中的每个进程都会收到所有其他进程的连接张量列表。当每个工作器都需要从所有其他工作器那里获得完整的计算结果时很有用,例如并行计算的嵌入。
  • 归约分散 (torch.distributed.reduce_scatter):对所有进程上的一组输入张量执行元素级归约(类似于All-Reduce),然后分散归约后的结果,使每个进程接收到最终归约张量的一部分。在某些情况下,这可能比单独的归约和分散操作更有效。

将这些操作可视化有助于理解。考虑一个在4进程设置中对梯度求和的简单All-Reduce操作:

All-Reduce (求和)进程 0(秩 0)梯度: G0集体通信(例如,环形All-Reduce)发送 G0进程 1(秩 1)梯度: G1发送 G1进程 2(秩 2)梯度: G2发送 G2进程 3(秩 3)梯度: G3发送 G3接收 和(G0…G3)接收 和(G0…G3)接收 和(G0…G3)接收 和(G0…G3)

每个进程计算其局部梯度 (GiG**i)。在All-Reduce步骤中,这些梯度在所有进程之间进行通信并求和。最终求和的梯度随后在每个进程上可用,为优化器步骤做好准备。

将原理与训练策略联系起来

这些基本原理直接对应本章后面讨论的分布式训练技术:

  • 数据并行 (DDP): 大量依赖 All-Reduce 来平均不同数据批次在工作器上计算的梯度。广播 最初用于同步模型权重。
  • 模型/流水线并行: 涉及更有针对性的点对点通信(使用 send/recv 原语,此处未详细说明但属于 torch.distributed 的一部分)或在持有模型不同部分或处理不同微批次的特定秩之间进行特定的集体操作,如 分散收集
  • 完全分片数据并行 (FSDP): 使用 All-Gather 的组合来重建层内前向/后向传播的完整参数,以及 归约分散 来平均梯度并有效地将它们分片回工作器。

理解这些构成要素——节点、进程、秩、后端和集体通信模式——为在PyTorch中有效实施和排除分布式训练工作流提供了坚实的根本。在建立了这些术语后,我们可以继续检查PyTorch如何为不同的并行化策略安排这些元素。

使用 DistributedDataParallel (DDP) 进行数据并行

训练大型模型或使用大量数据集时,在单个 GPU 上顺序处理所有数据会很快成为性能瓶颈。数据并行是一种策略,将相同的模型复制到多个处理单元(通常是 GPU)上,每个单元处理输入数据批次的不同子集。尽管 PyTorch 提供了直接的 torch.nn.DataParallel (DP) 模块,但由于 Python 全局解释器锁 (GIL) 的限制以及其梯度聚合的集中式方法,它在性能上常常不足。

为了高效、可扩展的数据并行,特别是在多 GPU 和多节点环境中,torch.nn.parallel.DistributedDataParallel (DDP) 是推荐的方案。DDP 使用多进程,为每个 GPU 分配一个独立的 Python 进程。这绕过了 GIL,实现了真正的并行执行。此外,它采用高效的集合通信操作(如 all-reduce),由 NCCL(适用于 NVIDIA GPU)或 Gloo(适用于 CPU 或 NCCL 不可用时)等后端管理,以直接在 GPU 之间同步梯度,在反向传播期间将通信与计算重叠,以提高性能。

理解 DDP 工作流程

DDP 的核心思想巧妙而强大:

  1. 初始化: 使用 torch.distributed.init_process_group 设置分布式环境。每个参与的进程被分配一个唯一的 rank(从 0 到 world_size - 1),它们协调通信,通常通过指定的后端(如 NCCL)。world_size 指的是训练中涉及的进程总数。
  2. 模型复制: 模型架构在每个进程(GPU)上完全相同地复制。
  3. 数据划分: 输入数据批次被分成更小的分片。每个进程接收一个分片。这通常通过 torch.utils.data.distributed.DistributedSampler 来管理,它确保每个进程在每个 epoch 中看到数据集的独特且不重叠的部分。
  4. 前向传播: 每个模型副本独立地对其本地数据分片执行前向传播。
  5. 梯度计算: 在反向传播 (loss.backward()) 期间,梯度在每个副本上进行本地计算。
  6. 梯度平均: 这是 DDP 的亮点所在。在梯度计算时,DDP 自动介入自动求导引擎。它在后台启动一个 all-reduce 集合操作。此操作汇总所有副本中每个参数的梯度,然后除以 world_size,从而有效地进行平均。结果会分发回所有副本。重要的是,DDP 将通信与梯度计算重叠,从而隐藏了通信延迟。
  7. 参数更新: 每个进程上的优化器 (optimizer.step()) 使用相同的平均梯度更新其本地模型副本的参数。因为所有副本都以相同的权重开始并接收相同的平均梯度,它们的参数在整个训练过程中保持同步,更新后无需显式参数广播。

初始化进程 0 (GPU 0)进程 1 (GPU 1)进程 N (GPU N)init_process_group(rank, world_size, backend)模型副本模型副本模型副本本地梯度 ∇L₀反向数据分片 0前向优化器梯度 All-ReduceΣ(∇Lᵢ) / world_size本地梯度 ∇L₁反向数据分片 1前向优化器本地梯度 ∇Lɴ反向数据分片 N前向优化器平均梯度平均梯度平均梯度完整数据集批次DistributedSampler拆分拆分拆分

工作流程说明了通过 DistributedSampler 进行数据分片、模型副本上的独立前向/反向传播,以及在每个进程的优化器步骤之前,用于梯度平均的 all-reduce 核心操作。

在实际中实现 DDP

将 DDP 集成到标准 PyTorch 训练脚本中需要进行一些修改:

  1. 环境配置: 您需要一种方式来启动多个 Python 进程,每个 GPU 一个。像 torchrun(推荐)或旧的 torch.distributed.launch 这样的标准工具可以处理此任务。它们负责设置 init_process_group 所需的环境变量,例如 MASTER_ADDRMASTER_PORTRANKWORLD_SIZE。您还需要确定 local_rank,它通常对应于当前进程应使用的 GPU 索引。

  2. 初始化进程组: 在脚本早期,初始化分布式后端:

    import torch
    import torch.distributed as dist
    import os
    
    // 假设环境变量 RANK, WORLD_SIZE, LOCAL_RANK 已由启动器设置
    val rank = os.environ['RANK'].toInt
    val world_size = os.environ['WORLD_SIZE'].toInt
    val local_rank = os.environ['LOCAL_RANK'].toInt
    
    // 初始化进程组
    dist.init_process_group(backend='nccl', // 'nccl' 用于 GPU, 'gloo' 用于 CPU
                            rank=rank,
                            world_size=world_size)
    
    // 为当前进程设置设备
    torch.cuda.set_device(local_rank)
    val device = torch.device(f"cuda:{local_rank}")
    

    强烈建议在 NVIDIA GPU 训练中使用 nccl,因为它性能优越。

  3. 准备分布式数据加载器: 修改您的数据加载以使用 DistributedSampler。这个采样器确保每个进程获得数据的一个不同部分,且不重叠。

    import torch.utils.data.{DataLoader, Dataset}
    import torch.utils.data.distributed.DistributedSampler
    // 假设 'train_dataset' 是您的 torch.utils.data.Dataset 实例
    val train_sampler = DistributedSampler(train_dataset, num_replicas=world_size, rank=rank)
    
    // 假设 'train_dataset' 是您的 torch.utils.data.Dataset 实例
    val train_sampler = DistributedSampler(train_dataset, num_replicas=world_size, rank=rank)
    
    // 重要:shuffle=False 因为 DistributedSampler 处理了洗牌
    // 重要:pin_memory=True 可以加快主机到设备的传输速度
    val train_loader = DataLoader(train_dataset,
                              batch_size=per_device_batch_size,
                              sampler=train_sampler,
                              num_workers=num_workers_per_process,
                              pin_memory=true,
                              shuffle=false) // 采样器处理洗牌
    

    请注意,DataLoader 中的 batch_size 现在指的是每个进程的批次大小。所有 GPU 上的总有效批次大小是 per_device_batch_size * world_size。请记住在 DataLoader 中设置 shuffle=False,因为 DistributedSampler 会负责在每个 epoch 中适当地打乱数据。

  4. 包装模型: 实例化您的模型并将其移动到当前进程的指定设备上,然后再用 DDP 包装。

    import torch.nn as nn
    import torch.nn.parallel.DistributedDataParallel as DDP
    
    // 实例化您的模型
    val model = YourModel().to(device) // 首先将模型移动到正确的 GPU
    
    // 使用 DDP 包装模型
    // device_ids 应包含此进程的单个 GPU ID
    // output_device 应与 device_ids[0] 相同
    val model = DDP(model, device_ids=List(local_rank), output_device=local_rank)
    

    device_ids 告知 DDP 此进程管理哪个或哪些 GPU(通常只有一个,即 local_rank),output_device 指定模型输出应放置的位置(通常是同一设备)。

  5. 训练循环调整: 核心训练循环基本保持不变。主要区别在于 loss.backward() 现在隐式地触发所有进程间的梯度同步。

    // 优化器
    val optimizer = YourOptimizer(model.parameters(), lr=learning_rate)
    
    for epoch in range(num_epochs):
        // 为采样器设置 epoch,以确保每个 epoch 之间数据正确洗牌
        train_sampler.set_epoch(epoch)
    
        model.train()
        for batch_idx, (data, target) in enumerate(train_loader):
            // 将数据移动到进程的 GPU
            val data = data.to(device)
            val target = target.to(device)
    
            optimizer.zero_grad()
            // 前向传播
            val output = model(data)
            val loss = loss_fn(output, target)
    
            loss.backward() // 触发梯度同步
    
            optimizer.step() # 使用平均梯度更新本地副本
    
            if rank == 0 and batch_idx % log_interval == 0: // 只在 rank 0 上记录日志
                println(f"Epoch: {epoch} | Batch: {batch_idx} | Loss: {loss.item()}")
    
        // 验证循环(通常只在 rank 0 上执行或使用分布式采样器)
        // ...
    
    // 清理进程组
    dist.destroy_process_group()
    

    通常的做法是,日志记录、保存检查点或验证等操作主要在一个进程(通常是 rank == 0)上执行,以避免冗余操作和混乱的输出。请记住在每个 epoch 开始时调用 train_sampler.set_epoch(epoch),以确保在使用 DistributedSampler 时每个 epoch 的数据洗牌不同。最后,在训练结束时调用 dist.destroy_process_group() 来清理资源。

DDP 的注意事项

  • 保存与加载: 保存 DDP 模型时,通常只需要一个进程(例如 rank == 0)来保存状态字典。DDP 包装了原始模型,因此保存的状态字典的键会带有 module. 前缀。将状态字典加载回非 DDP 模型时,您需要处理这个前缀,或者通过 model.module.state_dict() 访问底层模型。

    // 保存 (仅在 rank 0 上)
    if rank == 0 then
        torch.save(model.module.state_dict(), "model_checkpoint.pt")
    
    // 加载 (在所有 rank 上)
    val map_location = Map("cuda:0" -> s"cuda:$local_rank") // 映射到当前设备
    val checkpoint = torch.load("model_checkpoint.pt", map_location=map_location)
    // 首先创建模型实例
    val model_instance = YourModel().to(device)
    // 将状态字典加载到原始模型中
    model_instance.load_state_dict(checkpoint)
    // 然后用 DDP 包装
    val ddp_model = DDP(model_instance, device_ids=List(local_rank), output_device=local_rank)
    
  • 批量归一化: DDP 默认正确处理了跨进程的批量归一化统计信息同步。与 DataParallel 不同,您通常不需要 torch.nn.SyncBatchNorm,尽管在需要时也可以使用。

  • 混合精度: DDP 与 torch.cuda.amp(自动混合精度)兼容。在实例化 GradScaler 之后用 DDP 包装模型,但要在训练循环中遵循标准的 AMP 模式。

  • find_unused_parameters 如果您的模型存在在反向传播期间未接收到梯度的参数(例如,由于 forward 方法中的条件逻辑),DDP 的反向传播同步可能会停滞,等待永远不会到达的梯度。在 DDP 构造函数中设置 find_unused_parameters=True 可以解决此问题,但这会增加一些开销。通常最好确保所有需要梯度的参数都参与损失计算(如果可能)。

DistributedDataParallel 提供了一种高性能机制,用于在多个 GPU 和节点上扩展训练。通过了解其多进程架构、梯度平均对集合通信的依赖以及对数据加载和模型包装的必要调整,您可以有效地训练更大、更复杂的模型,比以往任何时候都快。

张量模型并行

数据并行在不同设备间复制模型并处理不同的数据批次时,当模型本身过大,无法载入单个加速器(GPU)的内存时,这种方法就无能为力了。针对此类情况,我们需要模型并行,它将一个模型划分到多个设备上。张量模型并行(TMP)是一种模型并行方式,它将单个或层内的特定操作拆分到多个设备上。

这种方法对于那些在特定层中参数量巨大的模型尤其适用,例如现代Transformer架构中常见的大型嵌入表或前馈网络(FFN)层。

基本原理:张量拆分

张量模型并行的基本原理是有序地将层的权重张量(有时也包括激活张量)拆分到多个GPU上。然后,计算在每个GPU上部分地执行,接着通过通信步骤来同步或合并结果,最终得到与原始未拆分层相同的输出。

我们以一个标准线性层为例,其操作定义为 Y=XA+bY=X**A+b,其中 XX 是输入激活,AA 是权重矩阵,bb 是偏置,YY 是输出激活。张量模型并行提供不同的策略来并行化此操作。

列并行

在列并行中,权重矩阵 AA 按列拆分到 NN 个GPU上。设 A=[A1,A2,…,AN]A=[A1,A2,…,A**N],其中每个 AiA**i 位于不同的GPU上。

  1. 输入广播:输入 XX 广播(或已可用)到模型并行组中的所有GPU。
  2. 并行计算:每个GPU ii 计算一个部分结果 Yi=XAiY**i=XAi
  3. 收集/拼接:部分结果 YiY**i 沿列维度收集并拼接,形成最终输出 Y=[Y1,Y2,…,YN]Y=[Y1,Y2,…,Y**N]。偏置 bb 也可以按列拆分,b=[b1,b2,…,bN]b=[b1,b2,…,b**N],并在收集步骤前在每个GPU上独立添加。

GPU (N个设备)输入 XX * A1 + b1广播X * A2 + b2广播…广播输出 Y = [Y1, Y2, …, Yn]拼接拼接拼接

列并行应用于线性层。输入 XX 是共享的,权重 AA 和偏置 bb 按列拆分,部分输出被拼接。

这种方法在并行计算后需要通信(收集或全收集操作),用于组装完整的输出张量 YY

行并行

或者,我们可以按行拆分权重矩阵 AA:A=[A1A2⋮AN]A=A1A2⋮A**N

  1. 输入拆分/分散:输入 XX 沿其最后一个维度(特征维度)在GPU间拆分:X=[X1,X2,…,XN]X=[X1,X2,…,X**N]。注意,这与列并行不同,在列并行中,每个GPU都需要完整的 XX
  2. 并行计算:每个GPU ii 使用其部分的输入和权重计算一个部分结果:XiAiXiA**i
  3. 全归约:部分结果使用全归约操作在所有GPU上求和,生成最终输出 Y=∑iXiAiY=∑iXiAi。偏置 bb 通常只在一个GPU上(例如,rank 0)在全归约求和之后添加。

GPU (N个设备)输入 XX1 * A1拆分X2 * A2拆分…拆分输出 Y = Sum(Xi * Ai) + b全归约全归约全归约

行并行应用于线性层。输入 XX 被拆分,权重 AA 按行拆分,计算部分结果,然后通过全归约求和。偏置 bb 在归约后添加。

行并行在矩阵乘法之后需要通信(一个全归约操作)。一个重要好处是,输出激活 YY 在所有参与的GPU上都得到复制,这可能是后续层(例如层归一化或另一个行并行线性层)所需的输入格式。

Transformer中列并行与行并行的结合

Transformer模型通常结合使用这些技术。例如,在一个由两个线性层组成的标准前馈网络(FFN)块中:

  1. 第一个线性层可能使用列并行。这将较大的中间维度拆分到多个GPU上。输出激活是分布的(分片的)。
  2. 第二个线性层可能使用行并行。它将分片激活作为输入(如果拆分维度匹配则无需通信),并生成通过全归约求和的输出,使FFN块的最终输出在GPU之间复制,为Transformer的下一部分(如残差连接或层归一化)做好准备。

这种策略性组合有助于最大限度地减少通信开销,通过在FFN块内的两个线性层之间保持激活分片。

嵌入并行化

对于词汇量非常大的模型,嵌入表可能成为一个显著的内存瓶颈。在这里,张量模型并行可以通过将嵌入表按行(沿词汇维度)在GPU之间拆分来应用。当查找输入token ID的嵌入时:

  1. 每个GPU接收完整的输入ID序列。
  2. 每个GPU仅查找其所持有的词汇部分对应的嵌入。如果某个ID对应于另一个GPU持有的行,它会产生一个零向量。
  3. 全归约操作会汇总所有GPU上的部分嵌入向量。由于对于任何给定的token ID,只有一个GPU生成了实际的嵌入向量(其他GPU生成了零),因此总和有效地收集了正确的嵌入。

实现细节与通信

手动实现张量模型并行需要仔细处理张量分片、计算和同步,使用 torch.distributed 包中的原语:

  • torch.distributed.broadcast:将张量从一个进程发送到所有其他进程。
  • torch.distributed.scatter:将张量块分散到各个进程。
  • torch.distributed.gather:将张量从所有进程收集到一个进程。
  • torch.distributed.all_gather:将张量从所有进程收集到所有进程。
  • torch.distributed.reduce_scatter:在进程间执行操作(如求和)并分散结果。
  • torch.distributed.all_reduce:在进程间执行操作(如求和)并使结果在所有进程上可用。

通常会创建专用函数或包装器来封装这些操作,以适用于特定层类型(例如 ColumnParallelLinearRowParallelLinear)。像NVIDIA的Megatron-LM这样的库率先使用了许多这些技术,并且其中一部分功能正在通过诸如 torch.distributed.tensor.parallel 等模块集成到PyTorch核心中,这些模块提供了更高级的API来简化这些实现。

考虑一个使用辅助函数的列并行线性层示例:

import torch
import torch.nn as nn
import torch.distributed as dist

// 假设这些辅助函数用于管理并行组和通信
from .parallel_utils import (
    get_tensor_model_parallel_group,
    get_tensor_model_parallel_rank,
    get_tensor_model_parallel_world_size,
    copy_to_tensor_model_parallel_region, # 处理输入广播/拆分
    gather_from_tensor_model_parallel_region # 处理输出收集/归约
)

class ColumnParallelLinear extends nn.Module:
    def __init__(self, input_size: Int, output_size: Int, bias: Boolean = true, **kwargs):
        super().__init__()
        world_size = get_tensor_model_parallel_world_size()
        //确保 output_size 可以被 world_size 整除
        assert output_size % world_size == 0
        val output_size_per_partition = output_size // world_size
        val input_size = input_size

        // 权重矩阵沿输出维度(列)拆分
        val weight = nn.Parameter(torch.empty(
            output_size_per_partition, input_size, **kwargs
        ))
        // 初始化权重...(例如,使用 init.kaiming_uniform_)

        if bias:
            // 偏置也沿输出维度拆分
            val bias = nn.Parameter(torch.empty(
                output_size_per_partition, **kwargs
            ))
            // 初始化偏置...(例如,使用 init.zeros_)
        else:
            register_parameter('bias', None)

    def forward(self, input_: Tensor):
        // 如果前一层输出已复制(例如 LayerNorm),
        // 输入可能需要广播或已可用。
        // 此函数处理必要的通信。
        val parallel_input = copy_to_tensor_model_parallel_region(input_)

        // 执行局部矩阵乘法
        val output_parallel = nn.functional.linear(parallel_input, weight, bias)

        // 从张量并行组中的所有GPU收集结果
        // 沿列维度拼接。
        val output_ = gather_from_tensor_model_parallel_region(output_parallel)

        return output_

列并行线性层的实现草图。请注意权重/偏置的显式分片以及通信包装器的使用。

权衡与考虑

  • 复杂性:手动实现张量模型并行需要对模型架构代码进行大量修改,这与通常作为包装器运行的DDP不同。调试并行逻辑可能具有挑战性。
  • 通信开销:张量模型并行在单个层的前向和反向传播中引入了通信。全收集或全归约等操作可能成为瓶颈,特别是当GPU之间的互连带宽(例如NVLink与PCIe)受限时。
  • 粒度:对于参数量大或激活值相对于执行的计算量大的层,它最有效。
  • 组合性:张量模型并行通常与数据并行(DDP)和流水线并行结合使用,以训练真正大型的模型,形成混合并行策略。

总而言之,张量模型并行是一种不可或缺的技术,当单个模型层超出单个设备的内存限制时。通过在多个设备上拆分层内的权重和计算,它使得训练比以往可能更大的模型成为可能,尽管代价是增加了实现复杂性和通信开销。它是当前扩展大型神经网络的基础组成部分。

流水线并行实现

收藏

当神经网络变得过大以至于单个设备无法容纳其任何一层,或者需要以不同方式重叠计算和通信时,流水线并行提供了一种替代的扩展策略。不同于复制整个模型或拆分单个层的方法,流水线并行将模型本身按顺序分配到多个设备上。每个设备或设备组都成为流水线中的一个“阶段”,负责运行模型层的一个子集。

流水线并行的机制

设想一个由多个连续层或模块组成的模型。在流水线并行中,你将连续的模块分配给不同的设备。例如,一个在四个GPU上运行的四层模型:

  1. GPU 0 (阶段 0): 运行第 1 层。
  2. GPU 1 (阶段 1): 运行第 2 层。
  3. GPU 2 (阶段 2): 运行第 3 层。
  4. GPU 3 (阶段 3): 运行第 4 层并计算损失。

输入数据进入第一阶段(GPU 0)。处理后,输出激活被发送到第二阶段(GPU 1)。这会一直持续,直到最后阶段计算出输出和损失。随后,梯度以相反的顺序反向流经流水线。GPU 3 计算第 4 层的梯度,并将第 3 层输出的梯度发回给 GPU 2,然后 GPU 2 计算第 3 层的梯度并将其发回给 GPU 1,依此类推,直到梯度到达第一阶段。

流水线气泡问题

此过程的简单实现效率低下。思考时间线:当阶段 1 处理第一个数据批次时,阶段 0 处于空闲状态,等待下一个批次。类似地,当阶段 2 处理时,阶段 0 和 1 处于空闲状态(假设只有一个批次流过)。在反向传播期间,各阶段会再次处于空闲状态,因为它们在等待来自后续阶段的梯度。这种空闲时间,被称为“流水线气泡”,大幅降低了硬件利用率。

设备时间步长GPU 0(阶段 0)GPU 1(阶段 1)GPU 2(阶段 2)GPU 3(阶段 3)T1T2正向 0T3T4T5T6T7正向 1正向 2正向 3反向 3反向 2反向 1反向 0T8

单批次简单流水线运行示意图,显示了在时间步长(T1-T8)内的正向(Fwd)和反向(Bwd)传播过程中,GPU上存在大量空闲时间(气泡)。

通过微批次处理减少气泡

解决流水线气泡问题的标准方法是微批次处理。我们不是一次性将整个小批次数据送入流水线,而是将其分成更小的块,称为微批次。流水线并发处理这些微批次。

一旦阶段 0 完成处理第一个微批次并将其激活发送到阶段 1,阶段 0 就可以立即开始处理第二个微批次。这使得多个微批次可以在流水线中同时“飞行”,重叠各阶段的计算,并大幅减少空闲时间。微批次的数量(mm)是一个超参数;更大的 mm 通常会带来更好的利用率,但会增加通信开销,并可能因存储每个微批次的中间激活和梯度而增加内存使用量。

设备时间步长(微批次 m0-m3)GPU 0(阶段 0)GPU 1(阶段 1)GPU 2(阶段 2)GPU 3(阶段 3)T1T2F0(m0)T3F0(m1)T4F0(m2)T5F0(m3)T6T7T8T9T10T11T12B0(m0)T13B0(m1)B0(m2)F1(m0)F1(m1)F1(m2)F2(m0)F2(m1)F1(m3)F2(m2)F3(m0)F3(m1)F2(m3)F3(m2)F3(m3)B3(m0)B3(m1)B3(m2)B3(m3)B2(m0)B2(m1)B2(m2)B2(m3)B1(m0)B1(m1)B1(m2)B1(m3)B0(m3)T14

使用微批次处理(m0-m3)的流水线运行示意图。不同微批次的正向(F)和反向(B)传播在各阶段(GPU)之间重叠,与简单方法相比,减少了空闲时间。初始填充和最终排空阶段仍存在一些气泡。

在PyTorch中实现流水线并行

PyTorch提供了实现流水线并行的工具,尽管它通常比DDP需要更多手动设置。其主要构成部分包括:

  1. 模型划分: 将你的 nn.Module 手动拆分为连续的 nn.Sequential 模块,每个阶段一个。将每个模块放置在其指定设备上。
  2. 通信: 使用 torch.distributed.sendtorch.distributed.recv 操作,在相邻阶段之间正向传输激活和反向传输梯度。请记住这些操作是阻塞的。
  3. 微批次调度: 实现一个循环,遍历微批次,管理每个微批次在各阶段的正向和反向传播。这需要仔细同步。
  4. 梯度同步: 每个微批次在反向传播期间计算的梯度需要累积。优化器步骤通常仅在小批次中的所有微批次都完成其正向和反向传播后才执行。

我们来概述一个带有微批次处理的简单两阶段流水线(GPU 0 和 GPU 1)的流程:

import torch
import torch.nn as nn
import torch.distributed as dist

// 假设分布式环境已初始化(rank 0 在 GPU 0,rank 1 在 GPU 1)
// 假设模型已拆分为 stage0 和 stage1,并放置在各自的设备上

def run_pipeline_step(stage0, stage1, micro_batches_data, micro_batches_labels, loss_fn, optimizer):

    val num_micro_batches = len(micro_batches_data)
    val activations_storage = Array.fill(num_micro_batches)(None) // 存储用于反向传播的激活
    val gradients_storage = Array.fill(num_micro_batches)(None) // 存储用于反向传播的梯度

    val current_rank = dist.get_rank()
    val world_size = dist.get_world_size() // 在此示例中假设 world_size = 2

    // --- 正向传播 ---
    for i in range(num_micro_batches):
        val micro_batch = micro_batches_data(i)

        if current_rank == 0: // 第一阶段
            // 计算阶段 0 的激活
            val activations = stage0(micro_batch.to(current_rank))

            // 将激活发送到下一阶段(rank 1)
            dist.send(activations.cpu(), dst=1, tag=i) // 发送 CPU 张量以避免 GPU 同步问题
            activations_storage(i) = activations // 存储用于反向传播

        elif current_rank == 1: // 最后阶段
            // 从上一阶段(rank 0)接收激活
            val received_activations = torch.empty_like(some_prototype_tensor_shape, device='cpu') // 需要形状信息
            dist.recv(received_activations, src=0, tag=i) // 接收 CPU 张量以避免 GPU 同步问题
            received_activations = received_activations.to(current_rank)
            received_activations.requires_grad_() # 重要:为接收到的张量启用梯度

            // 计算阶段 1 的激活(最终输出)
            val outputs = stage1(received_activations)

            // 计算损失
            val labels = micro_batches_labels(i).to(current_rank)
            val loss = loss_fn(outputs, labels)

            // 存储反向传播所需信息
            activations_storage(i) = received_activations // 此阶段的输入
            gradients_storage(i) = loss // 存储损失以供稍后启动反向传播

    // --- 反向传播 ---
    // 为确保正确性,对微批次进行反向迭代(GPipe 调度)
    for i in range(num_micro_batches - 1, -1, -1):
        if current_rank == 1: # 最后阶段
            val loss = gradients_storage(i)
            val input_activation = activations_storage(i)

            // 为此微批次的损失启动反向传播
            // 局部计算阶段 1 参数的梯度
            // 如果不是此阶段输入的最后一个微批次,则需要保留计算图
            val retain_graph_flag = (i != 0) 
            loss.backward(retain_graph=retain_graph_flag) 

            // 将输入激活的梯度发回上一阶段(rank 0)
            val grad_to_send = input_activation.grad.cpu() 
            dist.send(grad_to_send, dst=0, tag=i)

        elif current_rank == 0: // 第一阶段
            // 从下一阶段(rank 1)接收梯度
            val grad_received = torch.empty_like(some_prototype_grad_shape, device='cpu') // 需要形状信息
            dist.recv(grad_received, src=1, tag=i)
            grad_received = grad_received.to(current_rank)

            // 使用接收到的梯度继续反向传播
            val output_activation = activations_storage(i)
            // 局部计算阶段 0 参数的梯度
            output_activation.backward(gradient=grad_received)

    // --- 优化器步骤 ---
    // 处理完所有微批次后,更新权重
    optimizer.step()
    optimizer.zero_grad()

// --- 注意 ---
// 1. 这是一个简化示例(GPipe 风格的调度)。
// 2. 需要机制来确定接收缓冲区中的张量形状。
// 3. 错误处理、正确的设备放置和同步很关键。
// 4. 存在更高级的调度(例如,交错式)。
// 5. 像 `torch.distributed.pipeline`(实验性)这样的库旨在简化此过程。

这种手动实现突出了其复杂性:显式通信调用,管理每个微批次的中间激活及其梯度,以及阶段之间细致的同步。

PyTorch 还有一个实验性的 torch.distributed.pipeline.sync.Pipe 模块,旨在抽象化部分这种复杂性,根据拆分为阶段的模型定义,自动处理微批次处理、通信和梯度传播。然而,使用基本操作理解手动过程,能帮助更好了解其内部机制。

考量与权衡

  • 负载均衡: 每个阶段的计算成本应大致相等,以避免阶段成为瓶颈。这通常需要仔细的模型划分。
  • 通信开销: 在设备之间传输激活和梯度会引入开销。如果张量尺寸较大或互连带宽较低,则这种开销会更加明显。微批次处理由于更频繁、更小的传输而增加了这种开销。
  • 内存: 虽然每个设备的峰值内存有所减少(因为每个设备只保存模型的一部分),但存储多个正在处理的微批次的激活的需求,与简单的流水线相比,可能增加整体内存需求。
  • 复杂性: 正确实现流水线并行,尤其是在微批次处理和优化调度方面,比使用 DDP 复杂得多。

流水线并行对于超大型模型最有用,即使张量并行也不足以应对,或者需要对设备运行和内存进行精细控制时。它通常与数据并行结合使用(例如,在每个流水线阶段内运行 DDP),以实现进一步的扩展。

Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐