在大模型训练中,序列并行(Sequence Parallelism)通过将输入序列切分到多个GPU上,能够有效降低单张卡上的激活显存占用。但在实际训练中,常常会观察到不同设备之间的显存占用存在明显差异,这种不均衡不仅降低了显存利用率,还可能导致某些卡先触发OOM,从而限制整体batch size。解决这一问题的关键思路之一,就是在序列并行的基础上进一步对注意力头(Head)维度进行切分,让每张卡上缓存的激活张量形状完全一致。这种做法在保持训练正确性的前提下,能显著平滑显存使用曲线。

序列并行显存不均的根因
序列并行的典型做法是将输入序列沿长度维度均匀切分,每个设备只处理自己负责的那一段序列。在Transformer的注意力模块中,Q、K、V投影的输出形状通常为 [batch, seq_len, num_heads, head_dim]。序列并行会把 seq_len 变为原来的 1/N,但 num_heads 和 head_dim 保持不变。也就是说,每个设备仍然需要保存完整的注意力头信息,这导致显存占用与头数线性相关,而头数在所有设备上是相同的,理论上不应产生不均。然而,真正的问题出现在注意力计算和后续网络中。
注意力权重矩阵的形状为 [batch, num_heads, seq_len, seq_len],在序列并行下,每个设备计算局部的 query 与全局 key 之间的相似度,因此需要缓存完整的 key 和 value。如果某些层的注意力头数量较少,而另一些层头数较多(例如某些模型采用混合头数设计),那么不同层产生的激活张量大小就会不同。当这些层被分配到不同设备时,显存占用自然会出现差异。此外,dropout mask、残差连接缓存以及前馈网络的中间激活等也会因为序列切分方式不同而加剧这种不均衡。另一个隐性来源是,一些框架在实现序列并行时只切分输入序列,但保留了完整的张量形状用于通信,导致某些设备上同时存在两份不同长度的激活,形成显存尖峰。
要消除这种不均衡,单纯调整序列切分策略是不够的,因为问题根源在于注意力头维度被默认完整保留。因此,需要从 Head 维度入手,将头数也纳入并行切分的范围,使每个设备上的激活形状不仅序列长度一致,头数也一致。
Head维度切分的实现原理
Head 维度切分的核心思想是把多头注意力中的 num_heads 按照并行度均匀拆分,每个设备只计算属于自己的那部分头。假设总头数为 32,并行度为 4,那么每个设备负责 8 个头。这样 Q、K、V 投影在每个设备上的输出形状就变为 [batch, seq_len, 8, head_dim],而不是完整的 [batch, seq_len, 32, head_dim]。由于每个设备负责的头数相同,显存占用也就变得均匀。更进一步,投影矩阵的参数也可以按 head 维度切分,每个设备只保存自己那部分头的权重,从而减少参数显存。
在注意力计算阶段,每个设备独立完成局部头的 attention 操作,得到局部输出。之后需要将各设备的结果合并,以恢复完整的 head 维度。合并方式有两种:如果后续输出投影是完整的,则需要先执行 all-gather 将所有设备的头输出拼接起来;如果输出投影也按 head 切分,则每个设备可以只对自己负责的 head 做投影,最后通过 all-reduce 求和得到最终结果。后一种方式通信量更小,且能更好地与张量并行融合。
下面是使用 PyTorch 实现 Head 维度切分的一个简化示例,展示了如何按 head 切分并完成局部注意力计算:
import torch
import torch.nn as nn
import torch.distributed as dist
class HeadParallelAttention(nn.Module):
def __init__(self, hidden_size, num_heads, head_dim, world_size):
super().__init__()
self.num_heads = num_heads
self.head_dim = head_dim
self.hidden_size = hidden_size
self.world_size = world_size
self.num_heads_per_rank = num_heads // world_size
# 参数可以按head切分,这里为了清晰保留完整投影
self.q_proj = nn.Linear(hidden_size, num_heads * head_dim, bias=False)
self.k_proj = nn.Linear(hidden_size, num_heads * head_dim, bias=False)
self.v_proj = nn.Linear(hidden_size, num_heads * head_dim, bias=False)
self.o_proj = nn.Linear(num_heads * head_dim, hidden_size, bias=False)
def forward(self, hidden_states, seq_parallel_group):
batch, seq_len, _ = hidden_states.shape
q = self.q_proj(hidden_states)
k = self.k_proj(hidden_states)
v = self.v_proj(hidden_states)
rank = dist.get_rank(seq_parallel_group)
start = rank * self.num_heads_per_rank
end = start + self.num_heads_per_rank
q = q.view(batch, seq_len, self.num_heads, self.head_dim)[:, :, start:end, :]
k = k.view(batch, seq_len, self.num_heads, self.head_dim)[:, :, start:end, :]
v = v.view(batch, seq_len, self.num_heads, self.head_dim)[:, :, start:end, :]
q = q.permute(0, 2, 1, 3)
k = k.permute(0, 2, 1, 3)
v = v.permute(0, 2, 1, 3)
scale = self.head_dim ** 0.5
attn_weights = torch.matmul(q, k.transpose(-2, -1)) / scale
attn_weights = torch.softmax(attn_weights, dim=-1)
attn_output = torch.matmul(attn_weights, v)
attn_output = attn_output.permute(0, 2, 1, 3).contiguous()
attn_output = attn_output.view(batch, seq_len, self.num_heads_per_rank * self.head_dim)
# 收集所有rank的head输出,恢复完整hidden_size
full_attn_output = torch.empty(
batch, seq_len, self.num_heads * self.head_dim,
device=attn_output.device, dtype=attn_output.dtype
)
dist.all_gather_into_tensor(full_attn_output, attn_output, group=seq_parallel_group)
output = self.o_proj(full_attn_output)
return output
上述代码中,每个 rank 只计算自己负责的 head,注意力权重和输出形状都缩减为原来的 1/world_size。all_gather_into_tensor 负责将所有 rank 的输出拼接回完整的 head 维度,再送入输出投影。实际工程中为了减少通信,可以将 o_proj 也按 head 切分,并用 all-reduce 替代 all-gather,这里为了可读性保留了 all-gather 形式。
工程实践与性能权衡
Head 维度切分虽然能带来显存均衡,但引入了额外的通信开销。all-gather 或 all-reduce 的数据量为 batch * seq_len * hidden_size,与序列并行本身所需的通信量相当,因此整体通信量大约翻倍。为了降低这部分开销,可以采取通信与计算重叠的策略,例如在计算局部 attention 的同时开始准备通信缓冲区,或者利用 NVIDIA 的 NCCL 库进行异步聚合。如果模型同时使用了张量并行,Head 切分可以与其合并,因为张量并行本身就包含对 head 或 hidden 维度的切分,两者可以共用通信组,避免重复同步。
另一个需要注意的问题是头数无法整除并行度的情况。例如总头数为 30,并行度为 8,就无法均匀分配。常见的处理方式是限制并行度必须能整除头数,或者允许部分 rank 多计算一个头,但这样又会引入新的显存不均。更稳妥的方案是在模型设计阶段就将头数设置为并行度的倍数,或者使用动态 padding 补齐头数。此外,现代 GPU 上高度优化的 FlashAttention 通常期望完整的 head 维度,直接做 Head 切分可能需要额外的适配,比如在局部头上调用 FlashAttention 的内核,或者将切分后的 head 组合成连续块再调用。
从实际效果来看,Head 维度切分在长序列训练中收益尤其明显。以 32 头、序列长度 8192 的模型为例,序列并行度为 4 时,未做 Head 切分前每张卡需缓存大约 8192/4 * 32 * 64 的 QKV 激活,而 Head 切分后每张卡再除以 4,激活显存下降为原来的四分之一。更重要的是,所有设备上的激活形状完全一致,不会再出现某张卡显存先打满的情况。在 A100 或 H100 等硬件上配合 bf16 精度和通信融合,整体训练吞吐通常能提升 10% 到 25%,具体取决于序列长度和并行规模。
落地建议与总结
在实现 Head 维度切分时,建议先梳理模型中的注意力层,确认所有注意力头的数量以及各层之间的差异。可以先从推理阶段验证正确性,再逐步引入训练。对于使用 PyTorch 的开发者,可以基于张量并行的通信组来管理 head 切分,利用 torch.distributed 提供的 all_gather 和 all_reduce 原语。对于大模型训练框架如 Megatron-LM 或 DeepSpeed,它们已经在张量并行中实现了 head 维度的切分,实际使用时需要将序列并行与张量并行的配置正确组合,确保通信组一致。
总结来说,序列并行显存不均的主要原因是注意力头维度被完整保留,导致不同层或不同设备间的激活形状存在差异。通过将 Head 维度纳入并行切分,可以显著平滑显存使用,提升 batch size 和训练效率。代价是增加一次合并通信,但可以通过算子融合和异步通信抵消大部分开销。对于长序列大模型训练,这是一个值得优先采用的优化方向。