更新于 2026年7月23日

在上一篇文章「LLaMA 张量并行 PyTorch 实现:基于 FairScale 的 ColumnParallel/RowParallel 教程」中,我们详细介绍了张量并行中的 CollumnParallelLinear 和 RowParallelLinear 的从零实现过程,同时,我们也知道在多头注意力中最核心就是4个线性变换的组合,而这也是我们前面花了大量篇幅来介绍 CollumnParallelLinear 和 RowParallelLinear 原理与实现的原因。

在接下来的这篇文章中,将详细给大家介绍如何基于 CollumnParallelLinear 和 RowParallelLinear 来实现多头注意力机制。

在正式介绍之前,我们先来简单通过图示回顾一下整个多头注意力机制的计算过程。同时,我们还需要明白的是,对于每个多头注意力机制的计算过程,在基于 CollumnParallelLinear 和 RowParallelLinear 的实现情况下,我们只能以头为单位将其分配到不同设备上并行计算,而不能把一个头放到多个设备上进行并行。

两个 GPU 设备时并行计算示意图

例如,一共两个设备8个头,那么可以在每个设备上分配4个头来完成整个并行计算过程。

关键字:多头注意力机制、CollumnParallelLinear 、RowParallelLinear、张量并行、FairScale

1. 多头注意力计算过程#

假定现在原始序列经过 Embedding 后的结果,即多头注意力的输入分别为 $Q$、$K$ 和 $V$ ,那么可以通过式(1) 所示的过程来完成整个多头注意力的计算过程

$$ \begin{aligned} \text{MultiHead}(Q,K,V)=\text{Concat}(\text{head}_1,...,\text{head}_h)W^O\\\text{where}\;\;\text{head}_i=\text{Attention}(QW_i^Q,KW_i^K,VW_i^V) \end{aligned} \tag{1} $$

其中$Q\in\mathbb{R}^{n\times d_{q}}$,$K\in\mathbb{R}^{n\times d_{k}}$,$V\in\mathbb{R}^{n\times d_{v}}$,三者的每一行均可以看做是一个字符的向量表示。

对于式(1)中的 $\text{Attention()}$,其计算过程为

$$ \text{Attention}(Q,K,V)=\text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V\in\mathbb{R}^{n\times d_v} \tag{2} $$

下面,我们再通过一个两头的计算过程来从感官上回顾一下上面的计算过程,同时也有利于理解后续的并行计算。

图 1. 多头注意力线性变换计算图(图中整数序号仅表示位置记录,下同)
图 1. 多头注意力线性变换计算图(图中整数序号仅表示位置记录,下同)

如图1所示,输入 $X$ 会先经过3个线性变换分别得到 $Q$、$K$ 和 $V$ ,然后根据式(2)来计算注意力权重,如图2所示。

图 2. 多头注意力权重计算图
图 2. 多头注意力权重计算图

如图2所示,因为是两个头,所以会计算得到两个注意力权重矩阵。

进一步,注意力权重矩阵分别同 $V$ 进行计算得到 $z_i$ ,如图3所示。

图 3. 多头注意力权重作用图
图 3. 多头注意力权重作用图

在根据多个权重矩阵计算得到 $z_1$ 和 $z_2$ 以后,再将其拼接起来得到 $Z$,并再进行一次线性变换得到最终的表示 $Y$,如图4所示。

图 4. 编码结果计算图
图 4. 编码结果计算图

如图4所示,经过一次线性变换后便得到了序列最终的编码表示。

更多关于 Transformer 内容的介绍,可参见专题文章「Transformer 位置编码教程:图解正弦 Positional Encoding 与编解码过程」 中的内容。

2. 图解多GPU并行计算#

在回顾完多头注意力机制的原理后,我们再来看如何通过 CollumnParallelLinear 和 RowParallelLinear 来完成整个并行计算过程。

从上面的图示计算过程我们可以知道,对于每个头的计算过程来说,其一个包含有4个线性层的计算,即:图1中的3个和图4中的1个。所以,最后我们在完成并行计算的时候一定是 CollumnParallelLinear 和 RowParallelLinear 加起来一共有4个。

那到底是怎么样的呢?

根据在前一篇中对于 CollumnParallelLinear 和 RowParallelLinear 各自特性的介绍可知,通常 CollumnParallelLinear 会在 RowParallelLinear 的前面执行,所以在整个并行计算过程中会先通过3个 CollumnParallelLinear 来完成图1中的3个线性变换过程,然后在通过1个 RowParallelLinear 来完成图4中的线性变换过程。

整个计算过程在两个 GPU 设备上的计算过程如图5所示。

图 5. 两个 GPU 设备时并行计算示意图
图 5. 两个 GPU 设备时并行计算示意图

如图5所示便是在两个 GPU 上两个头的并行计算示意图。从图中可以看出整体上分为三大块:① 3个 CollumnParallelLinear 分别计算3个线性变换;② 并行计算各个头注意力权重矩阵,并进一步计算得到 $z_i$;③ 1个 RowParallelLinear 计算得到最终的想来表示。

下面来逐一进行介绍。

2.1 CollumnParallelLinear 计算过程#

如图6所示,3个 CollumnParallelLinear 都将同时在每个 GPU 上进行并行计算。

图 6.  CollumnParallelLinear 计算过程
图 6. CollumnParallelLinear 计算过程

也就是图1中的每个线性变换都将通过 CollumnParallelLinear 被拆分到了两个 GPU 设备上进行并行计算了,且拆分的逻辑一定是以一个头为整体。例如假定此时有3个 GPU 设备,那么你并不能将整个权重 $W$ 拆成3份(即每个 GPU 上占2列)放到3个 GPU 上,因为此时每个 GPU 计算的并不是一个完整的头。

2.2 头注意力权重矩阵计算过程#

在并行完成 CollumnParallelLinear 的计算过程以后,我们将得到原始输入 $X$ 分别经过3个线性变换后的结果 $Q$、$K$ 和 $V$ ,并且 $Q[:,:3]$、$K[:,:3]$ 和 $V[:,:3]$ 与 $Q[:,3:]$、$K[:,3:]$ 和 $V[:,3:]$ 分别表示两个注意力头,分布在两个 GPU 设备上,如图7所示。

图 7. 注意力权重矩阵计算图
图 7. 注意力权重矩阵计算图

进一步,对于每个 GPU 设备来说,各自完成式(2)中的计算过程,此时也是并行在完成计算。这部分完成计算后将会得到每个头对应的输出,即 $Z[:,:3]$ 和 $Z[:,3:]$,且分别分布在两个 GPU 设备上,这也刚好符合 RowParallelLinear 的前置计算条件。

2.3 RowParallelLinear 计算过程#

在分别计算得到 $Z[:,:3]$ 和 $Z[:,3:]$ 以后,则可以利用 RowParallelLinear 来完成最后一个线性变换,即图4所示的过程。

图 8. RowParallelLinear 计算过程
图 8. RowParallelLinear 计算过程

如图8所示, $Z[:,:3]$ 和 $Z[:,3:]$ 分别位于两个 GPU 设备上,此时应在 $W^O$ 的输入维度上对其进行拆分成两部分 $W^O[:3,:]$ 和 $W^O[3:,:]$ ,进一步并行完成计算得到 $Y_0$ 和 $Y_1$。

最后,只需要进行 Reduce-Sum 操作便可以得到原始序列最后的编码输出 $Y$。

以上就是多头注意力机制通过 CollumnParallelLinear 和 RowParallelLinear 在多卡上进行并行计算的过程。

不过聪明的你此时可能会提出一个疑问:如果你有100 张卡,但是只有 64 个头,难道还不能充分利用这些硬件资源来同时并行计算?

由于张量并行是将 n_heads * head_dim 这个维度在多张 GPU 上进行 列并行(ColumnParallel)切分,而注意力机制中,每个头其实是一个 head_dim 维度的子向量,所以一个完整的注意力头不会被拆到多张卡上,但每张卡上可以有若干个完整的注意力头。

之所以不能将一个注意力头拆分到多个 GPU 上是因为:

① 计算耦合度太高:一个头的 $QK^T$ 运算要依赖全量向量,如果拆开,会有大量通信。

② 无法高效并行:拆头的粒度太细,带宽、同步成本远高于收益。

③ 工程复杂度暴涨:把一个头拆成分片 Wq/Wk/Wv,每个都要同步还原,维护代价极高。

所以,张量并行这种设计本质上是一种 head-level 粒度的并行,所以我们可以同时采用其它并行策略来充分利用硬件资源,后续遇到时再给大家介绍。

3 从零实现多头注意力并行计算#

下面,我们再来看看如何基于 CollumnParallelLinear 和 RowParallelLinear 来实现多头注意力,完整示例代码可参见 Code/C04_Parallel/multi_head.py 文件。

3.1 初始化方法定义#

首先,我们需要定义整个类的初始化方法,即 3 个 CollumnParallelLinear 和 1 个 RowParallelLinear,代码如下:

 1 class Attention(nn.Module):
 2     def __init__(self, args: ModelArgs):
 3         super().__init__()
 4         world_size = get_model_parallel_world_size()
 5         self.n_local_heads = args.n_heads // world_size
 6         self.head_dim = args.dim // args.n_heads
 7         assert self.head_dim * args.n_heads == args.dim
 8         assert self.n_local_heads * world_size == args.n_heads
 9 
10         self.wq = ColumnParallelLinear(args.dim, args.n_heads * self.head_dim,
11                     bias=False, gather_output=False, init_method=lambda x: x)
12         self.wk = ColumnParallelLinear(args.dim, args.n_heads * self.head_dim,
13                     bias=False, gather_output=False, init_method=lambda x: x)
14         self.wv = ColumnParallelLinear(args.dim, args.n_heads * self.head_dim,
15                     bias=False, gather_output=False, init_method=lambda x: x)
16         self.wo = RowParallelLinear(args.n_heads * self.head_dim, args.dim,
17                     bias=False, input_is_parallel=True, init_method=lambda x: x)

在上述代码中,第4~8行是定义相关的变量及超参数,其中第5行和第8行便是用于判断多头的个数是否能刚好分配到 world_size 个设备上,即 n_local_heads 表示每个设备上多头的数量。第10~15行是分别实例化3个 ColumnParallelLinear ,会在输出维度上对权重进行划分,用于完成对输入的3个线性变换,其中 init_method=lambda x: x 表示 ColumnParallelLinear 内部 _initialize_affine_weight() 函数初始化时保持不变。第16~17行是实例化 RowParallelLinear ,会在输入维度上对权重进行划分。

更多关于 CollumnParallelLinear 和 RowParallelLinear 的介绍可以参见文章 「这张图让你弄懂 LLaMA 中张量并行从零实现过程!如何在 CPU 上模拟多卡并行计算?

3.2 前向传播实现#

对于多头注意力的并行计算前向传播实现过程,整体思路上同之前「第3讲 Multi-Head Attention 代码实现教程:基于 PyTorch 手写多头注意力与掩码」介绍的一致,唯一需要区别的地方在于并行计算的处理上,示例代码如下:

 1     def forward(self, x: torch.Tensor, mask: Optional[torch.Tensor] = None):
 2         bsz, seq_len, _ = x.shape
 3         xq, xk, xv = self.wq(x), self.wk(x), self.wv(x)
 4         xq = xq.view(bsz, seq_len, self.n_local_heads, self.head_dim)
 5         keys = xk.view(bsz, seq_len, self.n_local_heads, self.head_dim)
 6         values = xv.view(bsz, seq_len, self.n_local_heads, self.head_dim)
 7 
 8         xq = xq.transpose(1, 2)  # [bsz, n_local_heads, seq_len , head_dim]
 9         keys = keys.transpose(1, 2)  # [bsz, n_local_heads, seq_len , head_dim]
10         values = values.transpose(1, 2)  # [bsz, n_local_heads, seq_len , head_dim]
11         scores = torch.matmul(xq, keys.transpose(2, 3)) / math.sqrt(self.head_dim)
12         if mask is not None:
13             scores = scores + mask
14         scores = F.softmax(scores.float(), dim=-1).type_as(xq)
15         output = torch.matmul(scores, values)
16         output = output.transpose(1, 2).contiguous().view(bsz, seq_len, -1)
17         return self.wo(output)

在上述代码中,第2行取输入 x 的形状 [bsz, seq_len, dim]。第3行是分别进行3个线性变换,执行结束后的形状均为 [bsz, seq_len, n_local_heads * head_dim] 。这里需要注意的一点是,虽然上面第3.1节中3个线性变换中权重矩阵的形状均为 [dim, n_heads * head_dim],但是每个设备上并行计算时只会划分一部分权重,形状为 [dim, n_heads * head_dim//world_size],即 [dim, head_dim * n_local_heads]。所以线性变换计算时的维度变换为 [bsz, seq_len, dim] @ [dim, head_dim * n_local_heads] = [bsz, seq_len, n_local_heads * head_dim]

第4~6行是将最后一个维度拆分开,变成 [bsz, n_local_heads, seq_len , head_dim],以便下一步进行注意力权重矩阵的计算。第11行是计算注意力权重矩阵,即式(2)中的过程,score 的形状为 [bsz, n_local_heads, seq_len, seq_len]。第12~13行是施加掩码注意力,即仅允许模型看到过去的信息不能看到未来的信息。

第14~15行分别是归一化注意力权重矩阵及计算得到 $z_i$,即图7中所对应的过程,最后 output 的形状为 [bsz, n_local_heads, seq_len, head_dim]。第16行是先交换维度然后放到连续存储空间并合并最后一个维度,即先变为 [bsz, seq_len ,n_local_heads, head_dim],再变为 [bsz, seq_len, n_local_heads * head_dim]

注:在 PyTorch 中,一些操作(如 .transpose().permute())虽然不会改变张量的数据,但会改变张量在内存中的存储顺序。这会导致张量变为“非连续”的(non-contiguous)。但是某些底层操作(如 .view())要求输入必须是连续的张量,否则就会报错。这时候就需要使用 .contiguous()

第17行则是进行最后一个基于 RowParallelLinear 的线性变换,这里需要注意的是 RowParallelLinear 内部的前向传播计算完成后会对所有设备上的结果进行 Reduce-sum 操作,所以返回的是完整的结果,也即每个设备上的结果都是一样的,形状为 [bsz, seq_len, dim]

3.3 运行结果#

完成上述实现过程以后,我们便可以通过如下方式在进行运行,示例代码如下:

 1 def run(rank, world_size):
 2     print(f"当前进程 PID: {os.getpid()},父进程 PPID: {os.getppid()}")
 3     print(f"Rank: {rank} initializing model parallel with size {world_size}", flush=True)
 4     # 设置分布式环境变量
 5     os.environ['MASTER_ADDR'] = 'localhost'
 6     os.environ['MASTER_PORT'] = '12355'
 7     os.environ['RANK'] = str(rank)
 8     os.environ['WORLD_SIZE'] = str(world_size)
 9     # 初始化分布式进程组(CPU 上使用 gloo)
10     dist.init_process_group(backend="gloo", rank=rank, world_size=world_size)
11 
12     # dist.init_process_group(backend='nccl', rank=rank, world_size=world_size)
13     # torch.cuda.set_device(rank) # 把当前进程绑定到 rank 对应的 GPU 上
14 
15     initialize_model_parallel(world_size)
16     config = ModelArgs()
17     attn = Attention(config).to(config.device)
18     x = torch.randn(2, 128, config.dim).to(config.device)
19     out = attn(x)
20 
21     print(f"Rank: {rank} output shape: {out.shape}")  # [bsz, seq_len ,dim]
22     print(f"Rank: {rank} wq shape: {attn.wq.weight.shape}")  # [dim//world_size, dim]
23     print(f"Rank: {rank} wo shape: {attn.wo.weight.shape}")  # [dim, dim//world_size]
24     dist.destroy_process_group()
25 
26 def main():
27     world_size = 2
28     mp.spawn(run, args=(world_size,), nprocs=world_size, join=True)

在上述代码中,第3~9行是设置相关换家变量,之前的文章我们已经介绍过这里就不在赘述。第10行和第12~13行分别是针对 CPU 和 GPU 的设置,可以根据自己的实际情况进行选择。第21~23行便是相关变量的输出结果,如下所示

当前进程 PID: 25108父进程 PPID: 25102
Rank: 0 initializing model parallel with size 2
当前进程 PID: 25109父进程 PPID: 25102
Rank: 1 initializing model parallel with size 2
> initializing model parallel with size 2
> initializing ddp with size 1
> initializing pipeline with size 1
Rank: 1 output shape: torch.Size([2, 128, 512])
Rank: 0 output shape: torch.Size([2, 128, 512])
Rank: 0 wq shape: torch.Size([256, 512])
Rank: 1 wq shape: torch.Size([256, 512])
Rank: 0 wo shape: torch.Size([512, 256])
Rank: 1 wo shape: torch.Size([512, 256])

3.4 速度对比#

在介绍完基于并行机制下的多头注意力实现以后,我们还可以来对比一下并行和非并行在同样的超参数设置下的运行时间。下面,我们对比在4卡并行和单卡情况下,执行若干次前向传播所耗费的时间。完整示例代码可参见 Code/C04_Parallel/time_comp.py 文件。

迭代次数 4卡并行 单卡非并行
50 11.3 秒 3.4 秒
100 13.2 秒 4.3 秒
500 26.5 秒 25 秒
1000 44.1 秒 51.6 秒
5000 3分02秒 4分16秒

如上表所示便是对比实验后的结果。从结果可以看出,在迭代次数晓雨500次以前,单卡的速度是快于4卡并行的;当迭代次数超过500以后,4卡并行的速度就超过了单卡情况。不过根据这个实验来看,4卡并行的速速并没有想象中的那么快。

询问 ChatGPT 以后得知,因为上面的实验设置仅仅只是针对前向传播过程,纯推理不涉及反向传播所以不能体现并行对内存及算力的减负优势。所以,大家如果有兴趣可以自行再设计一个涉及反向传播的实验进行对比。

到此,基于 ColumnParallelLinear 和 RowParallelLinear 的多头注意力实现过程就介绍完了,感谢您的阅读!

引用#

[1] https://github.com/facebookresearch/fairscale

阅读 --

第1节 多头注意力机制原理

本文是《Attention Is All You Need》精读系列第1讲,从 RNN 串行瓶颈出发讲清 Transformer 为什么必须引入自注意力,再用图解与公式拆解 Scaled Dot-Product …

第2节 位置编码与编码解码过程

本文是 Transformer 精读系列第2讲,用图解与对比示例讲清 Transformer 为什么必须加入位置编码,并拆解正弦 Positional Encoding 公式与高低频维度直觉。随后通过可视化方式演示 …