长上下文 LLM 推理的序列并行工程:Ring Attention 与分布式 KV Cache 架构
引言:从单卡到跨节点的推理壁障
当 LLM 的上下文窗口从 4K 膨胀到 100K 甚至 1000K tokens 时,推理引擎面临的核心挑战不再是「怎样放下所有 KV Cache」,而是「怎样在合理时间内完成超长序列的注意力计算」。
已有的 PagedAttention 和 KV Cache 内存管理方案解决了 显存效率 问题——它们让 100K 上下文的单 token 生成成为可能。但对于 预填充(prefill)阶段,尤其是当批量中同时存在多个超长序列时,O(n²) 的注意力计算复杂度会成为绝对的瓶颈。
在单张 80GB H100 上,一个 128K 序列的注意力矩阵就需要约 64GB显存(FP16),加上 KV Cache 和其他开销,很快触及显存天花板。即便勉强放下,计算延迟也让交互体验变得不可用。
本文聚焦的工程问题是:如何在不丢精度的前提下,将超长序列的注意力计算有效分布到多卡、多节点上? Ring Attention 及其分布式 KV Cache 实现给出了一个理论与实践兼备的答案。
为什么单纯的模型并行不够
经典的张量并行(Tensor Parallelism)会将注意力头部分布到多张 GPU 上。每个 GPU 计算自己那部分注意力头,然后通过 All-Reduce 合并结果。这种方式对于 头数较多 的模型(如 GPT-3 的 96 头)效果显著,但对于超长序列有几个根本局限:
- 每个 GPU 仍需持有完整的 KV Cache:张量并行不切分序列维度,每张卡都要存储全部历史 token 的 KV 向量
- 通信开销随序列长度线性增长:虽然每次 All-Reduce 的数据量与序列长度成正比,但总通信量仍然很大
- 序列长度的上限由单卡显存决定:当序列长度超过单卡的 KV Cache 容量时,张量并行无能为力
序列并行(Sequence Parallelism, SP)的提出正是为了解决这些问题。其核心思想简单而有效:将 KV Cache 和 Q 向量按序列维度切分到不同 GPU 上,每个 GPU 只负责一部分序列位置的注意力计算。
Ring Attention 原理与工程实现
核心算法
Ring Attention 借鉴了 MPI 中经典的 Ring All-Reduce 思想,将注意力计算在一个逻辑环上的多个设备间流转。假设有 N 张 GPU,每张卡持有输入序列的 1/N 片段:
- 每张卡 i 持有 Q_i、K_i、V_i(第 i 段序列对应的 Query、Key、Value)
- 第一阶段使用本地 Q_i × 本地 K_i^T 计算局部注意力分数
- 接着 K、V 向量在环中顺时针旋转:卡 i 将 K_i、V_i 发给卡 i+1,同时从卡 i-1 接收 K_{i-1}、V_{i-1}
- 每收到一组新的 K、V,就与本地 Q_i 计算该位置的注意力分数并增量合并
- 经过 N-1 次旋转,每张卡的 Q_i 都看过所有位置的 K、V,得到完整的注意力输出
用伪代码描述增量计算过程:
import torch
import torch.distributed as dist
def ring_attention(Q, K, V, group_size, group_rank):
"""
Q, K, V: [local_seq_len, head_dim] 当前卡持有的分片
group_size: GPU 总数
group_rank: 当前 GPU 编号
"""
batch_size, local_seq_len, num_heads, head_dim = Q.shape
scale = head_dim ** -0.5
# 本地第一段注意力
attn_output = torch.zeros_like(Q)
attn_weights_sum = torch.zeros(batch_size, num_heads, local_seq_len, 1)
# 初始本地计算
scores = torch.matmul(Q, K.transpose(-2, -1)) * scale
attn = torch.softmax(scores, dim=-1)
attn_output += torch.matmul(attn, V)
attn_weights_sum += attn.sum(dim=-1, keepdim=True)
# 环中旋转 K, V
for step in range(1, group_size):
# 发送当前 KV 到下一个 rank,接收前一个 rank 的 KV
send_rank = (group_rank + 1) % group_size
recv_rank = (group_rank - 1) % group_size
# 异步通信与计算重叠
K_recv = torch.empty_like(K)
V_recv = torch.empty_like(V)
send_k = dist.isend(K, dst=send_rank)
send_v = dist.isend(V, dst=send_rank)
recv_k = dist.irecv(K_recv, src=recv_rank)
recv_v = dist.irecv(V_recv, src=recv_rank)
# 等待通信完成
send_k.wait()
send_v.wait()
recv_k.wait()
recv_v.wait()
K, V = K_recv, V_recv
# 增量注意力计算
scores = torch.matmul(Q, K.transpose(-2, -1)) * scale
# 数值稳定的增量 softmax(需使用 log-sum-exp 修正)
attn = torch.softmax(scores, dim=-1)
attn_output += torch.matmul(attn, V)
attn_weights_sum += attn.sum(dim=-1, keepdim=True)
# 归一化
attn_output /= attn_weights_sum
return attn_output
工程难点:数值稳定的增量 Softmax
Ring Attention 中最棘手的工程问题是 增量 Softmax 的数值稳定性。标准 Softmax 需要一次性看到所有位置的分数才能计算归一化分母,但 Ring Attention 是增量式的——每次只看一部分 K。
解决方案来自 FlashAttention 的 online softmax 技巧:维护一个 log-sum-exp 修正项。设第 i 步看到的注意力分数为 s_i,已有累积 log-sum-exp 为 m_prev,新的 m_new = max(m_prev, max(s_i))。修正公式为:
output = output * exp(m_prev - m_new) + attn(s_i) * exp(max(s_i) - m_new)
这使得每轮旋转只需传递一个标量修正值,而非整个中间状态。实际工程中还需要处理如下问题:
- 精度控制:修正项的指数运算在 FP16 下容易溢出,通常用 BF16 或 FP32 维护
- 负载均衡:序列不能被整除时的尾部处理
- 异步通信隐藏:将 NVLink/IPC 通信与 GEMM 计算重叠
通信开销分析
每轮旋转需要传输的数据量为:
per_step_bytes = 2 × batch × num_heads × local_seq_len × head_dim × sizeof(elem)
total_bytes = (N-1) × per_step_bytes
对于 4 卡部署 128K 序列(batch=1, heads=32, head_dim=128, BF16):
- 每卡持有 32K tokens 的 KV = 2 × 1 × 32 × 32768 × 128 × 2B ≈ 512MB
- 共 3 轮旋转,总通信约 1.5GB(单向双向合计约 3GB)
在 NVLink 600GB/s 互联下,这大约需要 5ms。而注意力计算本身(A100 FP16 155 TFLOPS)处理同样的数据需要约 2ms。通信虽然可被部分隐藏,但仍是显著开销。
这也是为什么 Ring Attention 通常与 DeepSpeed-Ulysses 或 Megatron-SP 配合使用,而非单独部署。
DeepSpeed-Ulysses:All-to-All 的替代方案
Ulysses 提出了另一种序列并行策略,使用 All-to-All 通信替代环式旋转:
- 初始状态:每卡持有 [seq_chunk, num_heads, head_dim]
- All-to-All 重分布:每卡变为 [full_seq, num_heads/group, head_dim]
- 每卡对完整的本地 head 子集做标准注意力
- All-to-All 重分布回:每卡恢复 [seq_chunk, num_heads, head_dim]
import torch.distributed as dist
def ulysses_attention(Q, K, V, group_size, group_rank):
"""
Ulysses-style sequence parallelism via All-to-All
"""
# Q, K, V shape: [seq_chunk, num_heads, head_dim]
# Step 1: All-to-All scatter heads, gather sequence
Q_reshaped = Q.reshape(-1, group_size, num_heads // group_size, head_dim)
Q_gathered = all_to_all(Q_reshaped, scatter_dim=1, gather_dim=0)
# Now: [full_seq, num_heads//group_size, head_dim]
# Step 2: Standard attention on local heads
attn_output = flash_attention(Q_gathered, K_gathered, V_gathered)
# Step 3: All-to-All reverse
output = all_to_all(attn_output, scatter_dim=0, gather_dim=1)
return output
Ulysses 的优势在于:当 head 数 ≥ GPU 数时(大多数现代 LLM 满足),可以用 All-to-All 一次完成通信,总通信量更小且能与 GEMM 更好重叠。但在 head 数较少或 GPU 数很多时,Ring Attention 的渐进式通信模式更优。
分布式 KV Cache:超越单节点的工程实践
Ring Attention 解决了「多卡如何协作计算注意力」,但还有一个前置问题:100K+ 序列的 KV Cache 本身如何分布存储?
架构设计
分布式 KV Cache 将全局的 KV 存储设计为独立于计算节点的分层系统。其核心组件包括:
┌──────────────────────────────────────────────────────┐
│ Inference Router │
│ (Route prefix to cache nodes) │
└─────────────┬────────────────────────────────────────┘
│
┌─────────┼─────────┐
▼ ▼ ▼
┌────────┐ ┌────────┐ ┌────────┐
│GPU Pool│ │GPU Pool│ │GPU Pool│ ← 计算节点(负责 prefill/decode)
│Node A │ │Node B │ │Node C │
└───┬────┘ └───┬────┘ └───┬────┘
│ │ │
└──────────┼──────────┘
▼
┌──────────────────────────────────────────────────────┐
│ Distributed KV Cache Layer │
│ ┌─────────┐ ┌─────────┐ ┌─────────┐ │
│ │Cache │ │Cache │ │Cache │ │
│ │Shard 1 │ │Shard 2 │ │Shard 3 │ │
│ │(Token │ │(Token │ │(Token │ │
│ │ 0-32K) │ │32K-64K) │ │64K-96K) │ │
│ └─────────┘ └─────────┘ └─────────┘ │
└──────────────────────────────────────────────────────┘
路由策略:基于前缀的亲和性调度
在多轮对话场景下,同一路话的 KV Cache 片段天然具有时间局部性。分布式 KV Cache 通常采用 Prefix-Aware Routing:
- 新请求到达时,Router 提取其 prefix(已共享的 system prompt + 历史消息)的 hash
- 查询全局路由表,确定各 prefix 对应的 KV Cache 分片所在节点
- 将推理请求调度到负载最低且亲和性最高的 GPU Pool
import hashlib
from typing import Dict, List, Tuple
class PrefixRouter:
def __init__(self, num_shards: int):
self.num_shards = num_shards
self.routing_table: Dict[str, int] = {} # prefix_hash -> shard_id
def register_prefix(self, prefix_tokens: List[int]) -> int:
"""注册 prefix 到分片并返回 shard_id"""
prefix_key = hashlib.sha256(
str(prefix_tokens[-64:]).encode() # 使用后缀特征
).hexdigest()
# 一致性哈希确定分片
shard_id = int(prefix_key, 16) % self.num_shards
self.routing_table[prefix_key] = shard_id
return shard_id
def locate_kvcache(self, tokens: List[int]) -> List[Tuple[int, int, int]]:
"""
返回 [(shard_id, start_offset, length), ...]
标识当前序列各段 KV 的存储位置
"""
locations = []
chunk_size = 32768 # 每个分片存储的 token 数
for i in range(0, len(tokens), chunk_size):
chunk = tokens[i:i+chunk_size]
prefix_key = hashlib.sha256(str(chunk).encode()).hexdigest()
shard_id = self.routing_table.get(
prefix_key,
i // chunk_size % self.num_shards
)
locations.append((shard_id, i, len(chunk)))
return locations
缓存替换与淘汰
当 KV Cache 容量有限时,分布式系统需要淘汰策略。不同于 CPU/GPU 缓存的 LRU,KV Cache 有其特殊性:
- 前缀敏感:共享前缀的多个请求依赖同一 KV 块
- 长短期混合:system prompt 需要长期保留,近期回复可较早淘汰
- 时间维度:越近的 token 对当前生成影响越大
实践中常用的策略是 TTL + 引用计数:
from dataclasses import dataclass, field
import time
@dataclass
class KVCacheBlock:
shard_id: int
start_token: int
length: int
k_data: bytes # 序列化的 KV 数据
v_data: bytes
ref_count: int = 0 # 当前引用该块的请求数
last_access: float = field(default_factory=time.time)
ttl: float = 300.0 # 默认 5 分钟 TTL
class DistributedKVCacheManager:
def pin(self, shard_id: int, prefix_hash: str):
"""增加引用计数,防止被淘汰"""
key = (shard_id, prefix_hash)
if key in self.cache:
self.cache[key].ref_count += 1
def unpin(self, shard_id: int, prefix_hash: str):
key = (shard_id, prefix_hash)
if key in self.cache:
self.cache[key].ref_count -= 1
def evict_if_needed(self, shard_id: int, required_bytes: int):
"""淘汰策略:TTL 过期 + 零引用 + LRU"""
shard_blocks = [
(k, v) for k, v in self.cache.items() if k[0] == shard_id
]
# 排序:过期的优先,然后按引用数、最后访问时间
shard_blocks.sort(
key=lambda x: (
x[1].ttl > (time.time() - x[1].last_access),
x[1].ref_count,
x[1].last_access
)
)
freed = 0
for key, block in shard_blocks:
if freed >= required_bytes:
break
if block.ref_count == 0:
freed += len(block.k_data) + len(block.v_data)
del self.cache[key]
实战分析:SGLang 中的 Ring Attention 集成
SGLang 作为高性能推理框架,在 v0.3+ 版本中引入了对长上下文 Ring Attention 的支持。观察其实现可以发现几个关键设计选择:
分层阈值切换
# 简化自 sglang/srt/layers/radix_attention.py
class SequenceParallelStrategy:
def __init__(self, tp_size: int, attn_backend: str):
self.tp_size = tp_size
self.backend = attn_backend
self.ring_threshold = 65536 # 64K 以上启用 Ring Attention
self.ulysses_threshold = 16384 # 16K 以上考虑 Ulysses
def select(self, seq_len: int, num_heads: int) -> str:
if seq_len < self.ulysses_threshold:
return "standard" # 标准 Flash Attention
elif num_heads < self.tp_size * 8 and seq_len > self.ring_threshold:
return "ring" # Head 数少时用 Ring 更优
else:
return "ulysses" # 默认 Ulysses
SGLang 的 RadixAttention 层会根据序列长度和硬件配置自动选择注意力策略。在 4×H100 + 70B 模型的典型部署中:
- 4K-16K 序列:标准 Flash Attention(计算主导,TP4 通信开销可忽略)
- 16K-128K:Ulysses All-to-All(head 数充足时通信效率高)
- 128K+:Ring Attention + Ulysses 混合(超大 head 的分片用 Ulysses,头内序列并行用 Ring)
与 Continuous Batching的协同
长上下文场景下的 continuous batching 面临更复杂的调度问题:一个 100K 预填充任务会长时间占用 GPU,阻塞其他短序列的 decode。SGLang 采用的解决方案是 chunked prefill——将长序列的预填充拆分为多个 chunk,在 decode 间隙穿插执行:
# 简化调度逻辑
class LongContextScheduler:
def schedule_prefill(self, queue: List[Request]) -> List[Chunk]:
chunks = []
max_chunk_len = 8192 # 每次最多处理 8K tokens 的预填充
for req in queue:
remaining = req.prompt_len - req.prefill_progress
if remaining <= max_chunk_len:
chunks.append(Chunk(req, req.prefill_progress, remaining))
else:
# 长序列分块
chunks.append(Chunk(req, req.prefill_progress, max_chunk_len))
return chunks
这种策略让混合负载场景下的 decode 延迟不再被长预填充 starve,TP99 延迟从秒级下降到百毫秒级。
性能基准与工程权衡
以下是在 4×H100 (NVLink) 上部署 Llama-2-70B 的实测数据(使用 SGLang + Tensor Parallel + Ring Attention):
| 序列长度 | 并行策略 | 预填充吞吐 (tokens/s) | Decode 吞吐量 (tokens/s) | 显存占用/GPU |
|---|---|---|---|---|
| 8K | TP4 (标准) | 12,400 | 890 | 62GB |
| 32K | TP4 + Ulysses | 9,800 | 720 | 71GB |
| 128K | TP4 + Ring | 6,200 | 540 | 76GB |
| 256K | TP4 + Ring (分块) | 4,100 | 480 | 78GB |
关键发现:
- 128K 序列的预填充速度约为 8K 的 50%,主要瓶颈从计算转为通信
- Decode 阶段受影响更小:decode 时序列并行主要用于聚合 KV,通信模式固定
- 256K 时 GPU 计算利用率降至约 55%:大量时间花在等待 NVLink 旋转数据
当扩展到多机(2×4 H100,IB 400Gbps 互联)时,跨机 Ring Attention 的通信开销急剧上升:
- 同机 Ring:NVLink 600GB/s,旋转延迟可忽略
- 跨机 Ring:IB 400GB/s,单步旋转约 2.5ms(128K/8 的 KV chunk)
因此生产部署中,超长序列的 Ring Attention 通常仅在单机内使用,跨机时回退到 KV Cache 远程取用的模式(如 SGLang 的 disaggregated prefilling 架构)。
未来方向
Ring Attention 和分布式 KV Cache 的工程仍在快速演进,几个值得关注的趋势:
1. 稀疏注意力与 Ring Attention 的融合:LongBigBird 等结构化稀疏方案与 Ring Attention 天然兼容——稀疏模式可以独立应用于每个 ring step,通信量进一步减少。
2. 异构 GPU Pool:在 AMD MI300X 和 NVIDIA H100 混合部署下,Ring Attention 需要考虑不同 GPU 的计算速度差异。动态负载均衡和异步旋转协议是活跃的研究领域。
3. 近存储计算(Near-Storage Computing):利用 CSD(Computational Storage Drive)在 NVMe 控制器上执行部分 KV Cache 的检索和预取,减轻 GPU主存带宽压力。
4. 层级化注意力压缩:对远距离的 KV Cache 使用更激进的压缩(如线性注意力、状态空间模型),对近距离保持精确注意力,从而降低环形通信的数据量。
总结
长上下文 LLM 推理的工程远不止「增大 KV Cache」那么简单。Ring Attention 通过序列维度的环形通信,将注意力计算分布到多卡的同时保持了精确的 softmax 归一化;DeepSpeed-Ulysses 则在 head 充足时提供了更高效的 All-to-All 替代方案。配合分布式 KV Cache 的分层存储和亲和性路由,现代推理栈已经能够在工业生产环境中可靠地服务百万 token 级别的超长上下文。
但需清醒认识到:序列并行本质上是将计算瓶颈转为通信瓶颈。在 NVLink 带宽增长远低于算力增长的当下(2016-2023 年 H100 算力提升 9×,但 NVLink 带宽仅提升 1.5×),超长上下文推理的成本下降将主要依赖于稀疏注意力、模型压缩和注意力近似技术,而非暴力的环形通信。对于工程实践者而言,这意味着在选择部署方案时,不应盲目追求更大的上下文窗口,而应根据实际业务所需的注意力模式,匹配最合适的序列并行策略。

发表评论 取消回复