Context Parallelism
了解上下文并行如何将长序列分配到各个 GPU 上,以减少内存使用、扩展 Transformer 训练,并支持长文档和视频 AI 工作负载。
上下文并行是一种分布式计算技术,它将长的输入序列拆分到多个加速器上。每个 GPU 在注意力计算期间与其他 GPU 协同工作,仅处理序列的一部分。这降低了每个设备的激活内存,使 transformer 能够在超过单 GPU 内存限制的输入(例如超长文档、长视频或大量图像切片集合)上进行训练。
与单纯扩大模型的上下文窗口不同,上下文并行不会改变架构理论上能接收的信息量。相反,它通过分发序列维度,使处理该上下文在计算上变得切实可行。
上下文并行是如何工作的#
假设一个序列包含 32,000 个 token,并且上下文并行使用了四个 GPU。每个设备最初接收大约 8,000 个 token 并存储对应的中间激活值。
诸如归一化和前馈层之类的大多数操作可以独立处理这些本地序列块。难点在于注意力机制:本地查询可能需要关注由其他所有设备持有的键和值。
因此,实现方案会在 GPU 之间交换键值(即 KV)块。在环形注意力中,每个设备使用其本地数据计算部分注意力,将一个 KV 块传递给下一个设备,并重复此过程直到处理完完整的序列。PyTorch 上下文并行教程通过分布式缩放点积注意力展示了这一行为,而NVIDIA 上下文并行包则描述了全聚合(all-gather)、归约散布(reduce-scatter)和基于环的通信选项。
最终结果在数学上等同于全序列注意力(受正常数值差异影响),但没有单个 GPU 需要保留每个序列激活值。
为什么上下文并行很重要#
长序列会产生两个主要的扩展问题。首先,随着序列长度增长,保存的激活值会消耗更多内存。其次,标准的自注意力机制会跨序列比较 token,产生大量的计算和临时数据。
上下文并行通过在设备之间分配激活值来解决内存问题。它还可以分配注意力工作,尽管通信会引入新的成本。高效的系统会使用NCCL 集合操作指南中记录的操作,将 KV 传输与计算重叠进行。
当引起内存溢出(OOM)错误的是序列长度而不是模型权重或批次大小时,该技术最具价值。它属于分布式训练这一更广阔的领域,通常与其他策略结合使用以同时扩展多个维度。
上下文并行与相关技术的对比#
- **张量并行**在单个层内划分操作或权重矩阵。而上下文并行则沿序列维度划分 token。
- **流水线并行**将不同的模型层组分配给不同的设备。它划分的是模型深度而不是序列长度。
- 数据并行复制模型并为每个副本提供不同的训练样本。PyTorch 分布式概览建议在完整模型和每个样本都能放入一个 GPU 时使用它。
- 序列并行通常为与张量并行相关的选定操作分片激活值。而上下文并行则更广泛地将序列分区应用于网络输入和激活。
这些方法是相辅相成的。NVIDIA 并行策略指南展示了上下文、张量、流水线和数据并行如何共同构成多维设备布局。
实际应用#
-
**长文档 AI:**法律或医疗语言模型可能需要处理整个案卷、病史或技术手册。上下文并行将数千个文档 token 分布到各个加速器上,在保持远距离部分之间注意力的同时降低激活内存压力。
-
**长视频与多模态理解:**视频 transformer 和大型视觉模型可能会将帧、图像切片、音频片段和文本表示为一个长 token 序列。分发该序列有助于模型分析扩展的录像,而无需过度降低帧率或空间细节。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 提供了更简单的云端和本地工作流,用于数据集标注、训练、部署和监控,而无需进行手动上下文并行配置。






