Context Parallelism
了解上下文并行如何将长序列分布到多个 GPU 上,以降低内存使用量、扩展 Transformer 训练,并支持长文档和视频 AI 工作负载。
上下文并行是一种分布式计算技术,可将长输入序列拆分到多个加速器上。每个 GPU 只处理序列的一部分,并在注意力计算期间与其他 GPU 协同工作。这会减少每个设备的激活值内存,使 Transformer 能够在超出单个 GPU 内存容量的输入上训练,例如超长文档、长视频或大型图像块集合。
与单纯增大模型的上下文窗口不同,上下文并行不会改变架构理论上能够接受的信息量。它通过分发序列维度,让处理这类上下文在计算上变得切实可行。
上下文并行的工作原理#
假设一个序列包含 32,000 个标记,并使用 4 个 GPU 进行上下文并行。每个设备最初会接收大约 8,000 个标记,并存储相应的中间激活值。
大多数运算(例如归一化和前馈层)都可以独立处理本地序列分块。难点在于注意力机制:本地查询可能需要关注其他所有设备上的键和值。
因此,实现通常会在 GPU 之间交换键值(KV)块。在环形注意力中,每个设备使用本地数据计算部分注意力,将一个 KV 块传递给下一个设备,然后重复此过程,直到处理完整个序列。PyTorch 上下文并行教程通过分布式缩放点积注意力演示了这一行为,而 NVIDIA 上下文并行包则介绍了全收集、规约散播和环形通信选项。
最终结果在数学上等同于完整序列注意力,但会有常见的数值差异;不过,没有任何一个 GPU需要保留所有序列激活值。
上下文并行为何重要#
长序列会带来两个主要的扩展问题。首先,随着序列长度增加,已保存的激活值会占用更多内存。其次,标准自注意力会比较序列中的标记,从而产生大量计算和临时数据。
上下文并行通过在设备间拆分激活值来解决内存问题。它也可以拆分注意力计算,但通信会带来新的开销。高效系统会使用 NCCL 集合通信操作指南中所述的操作,将 KV 传输与计算重叠执行。
当导致内存不足错误的原因是序列长度,而不是模型权重或批次大小时,这项技术最有价值。它属于更广泛的分布式训练领域,通常会与其他策略结合,同时扩展多个维度。
上下文并行与相关技术的比较#
- 张量并行会拆分单个层中的运算或权重矩阵。上下文并行则沿序列维度拆分标记。
- 流水线并行会将不同的模型层组分配给不同设备。它拆分的是模型深度,而不是序列长度。
- 数据并行会复制模型,并将不同的训练样本分配给每个副本。PyTorch 分布式概览建议在完整模型和每个样本都能装入单个 GPU 时使用数据并行。
- 序列并行通常会为与张量并行相关的特定运算分片激活值。上下文并行则在网络输入和激活值中更广泛地应用序列分区。
这些方法可以互相补充。NVIDIA 并行策略指南展示了上下文并行、张量并行、流水线并行和数据并行如何构成多维设备布局。
实际应用#
-
长文档 AI:法律或医疗语言模型可能需要处理完整的案件档案、患者病史或技术手册。上下文并行会将数千个文档标记分发到多个加速器上,在保留远距离章节间注意力的同时,降低激活值内存压力。
-
长视频和多模态理解:视频 Transformer 和大型视觉模型可能会将帧、图像块、音频片段和文本表示为一个长标记序列。分发该序列有助于模型分析较长的录制内容,而不必大幅减少帧数或空间细节。AWS Neuron 上下文并行概览展示了加速器组如何为这些长上下文工作负载交换 KV 分片。
对于 Ultralytics YOLO26 这类紧凑型计算机视觉架构,通常不需要上下文并行。通过 Ultralytics 模型训练工作流使用标准的多 GPU 数据并行,通常是加速训练的更合适方法。
实际应用与权衡#
下面这个最简示例使用 PyTorch 的实验性上下文并行 API 和缩放点积注意力控制。将示例保存为 cp_example.py,然后使用 torchrun --standalone --nproc-per-node=2 cp_example.py 在两块 GPU 上启动。
import os
import torch
import torch.distributed as dist
import torch.nn.functional as F
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor.experimental import context_parallel
from torch.nn.attention import SDPBackend, sdpa_kernel
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(rank)
torch.cuda.manual_seed(0)
dist.init_process_group("nccl")
mesh = init_device_mesh("cuda", (world_size,))
qkv = [torch.randn(1, 4, 4096, 64, device="cuda", dtype=torch.bfloat16, requires_grad=True) for _ in range(3)]
with sdpa_kernel(SDPBackend.FLASH_ATTENTION), context_parallel(mesh, buffers=tuple(qkv), buffer_seq_dims=(2, 2, 2)):
output = F.scaled_dot_product_attention(*qkv, is_causal=True)
output.float().square().mean().backward()
dist.destroy_process_group()这里的维度 2 是序列维度,因此每个进程都会接收一个序列分片,同时注意力计算会在设备网格中协调进行。
在实践中,工程师应确认节省的内存能够抵消通信开销,使用高速互连,并对有代表性的序列长度进行基准测试。对于标准视觉项目,Ultralytics Platform提供更简单的云端和本地工作流,用于数据集标注、训练、部署和监控,无需手动配置上下文并行。









