目录
- Self-Attention单头注意力
-
- [Single Head Attention代码](#Single Head Attention代码)
- [Single Head Attention流程图](#Single Head Attention流程图)
- 掩码矩阵的形式
- [多头注意力MHA(Multi-Head Attention)](#多头注意力MHA(Multi-Head Attention))
-
- [Multi-Head Attention代码](#Multi-Head Attention代码)
- MHA流程图
- [线性注意力(Linear Attention)](#线性注意力(Linear Attention))
-
- [Linear Attention与普通 Self-Attention 的对比](#Linear Attention与普通 Self-Attention 的对比)
- [MLA(Multi-head Latent Attention)](#MLA(Multi-head Latent Attention))
-
- 背景知识
-
- [1. 低秩分解](#1. 低秩分解)
- [2. ColumnParallelLinear和RowParallelLinear函数](#2. ColumnParallelLinear和RowParallelLinear函数)
-
- [2.1 ColumnParallelLinear](#2.1 ColumnParallelLinear)
- [2.2 RowParallelLinear](#2.2 RowParallelLinear)
- [2.3 通过一个例子来解释](#2.3 通过一个例子来解释)
-
- [2.3.1 先看"不并行"时,正常是怎么算的](#2.3.1 先看“不并行”时,正常是怎么算的)
- [2.3.2 现在切开(RowParallelLinear 登场)](#2.3.2 现在切开(RowParallelLinear 登场))
- [2.3.3 为什么叫"部分和(Partial Sum)"?](#2.3.3 为什么叫“部分和(Partial Sum)”?)
- [2.3.4 救星登场:All-Reduce 通信](#2.3.4 救星登场:All-Reduce 通信)
- [2.3.5 最终结果:](#2.3.5 最终结果:)
- [2.4 为什么ColumnParallelLinear和RowParallelLinear是"天生一对"?](#2.4 为什么ColumnParallelLinear和RowParallelLinear是“天生一对”?)
- [Multi-head Latent Attention代码](#Multi-head Latent Attention代码)
Self-Attention单头注意力
Single Head Attention代码
python
import torch
import torch.nn as nn
import torch.nn.functional as F
class SelfAttention(nn.Module):
def __init__(self, embed_size):
super(SelfAttention, self).__init__()
# embed_size:输入嵌入向量的维度,即每个 token 的特征向量长度。
self.embed_size = embed_size
# 三个线性变换层
# 创建Value投影矩阵W^V。输入输出维度都是 embed_size,没有偏置 bias=False。
# 将输入 x 映射为 "值" 向量,代表每个token实际携带的信息内容。
self.values = nn.Linear(embed_size, embed_size, bias=False)
# 创建Key投影矩阵W^K。将输入 x 映射为 "键" 向量,代表每个 token "是什么",用于被查询匹配。
self.keys = nn.Linear(embed_size, embed_size, bias=False)
# 创建Query投影矩阵W^Q。将输入 x 映射为 "查询" 向量,代表每个token"在找什么",用于去匹配其他token的key。
self.queries = nn.Linear(embed_size, embed_size, bias=False)
# 直观理解:
# 想象你在图书馆找书 ------ Query 是你的需求("我想看科幻"),Key 是每本书的标签("科幻/文学/历史"),
# Value 是书的实际内容。Query 和 Key 匹配度越高,对应的 Value 就越被关注。
# 前向传播函数。参数:
# x:输入张量,形状 (N, seq_length, embed_size)
# mask:可选的掩码张量,用于屏蔽某些位置(如 padding 或因果掩码)
def forward(self, x, mask):
# 解构输入 x 的形状:
# N:batch size(一个批次有多少个样本)
# seq_length:序列长度(每个样本有多少个token)
# _: 丢弃 embed_size(不需要显式使用)
N, seq_length, _ = x.size()
# Q、K、V的投影计算
# x的shape原本为(N, seq_length, embed_size)
# 经过后计算得到的queries的shape仍然是(N, seq_length, embed_size)
queries = self.queries(x)
# keys和values的计算与queries类似
keys = self.keys(x) # shape仍然是(N, seq_length, embed_size)
values = self.values(x) # shape仍然是(N, seq_length, embed_size)
# 计算注意力分数
# 计算每个 query 与所有 key 的相似度
# keys.transpose(1, 2):将keys从(N, seq_length, embed_size)转置为(N, embed_size, seq_length),实现行列互换。
# torch.bmm(batch matrix multiply):批量矩阵乘法。对于每个batch的矩阵进行运算:
# queries形状(seq_length, embed_size)点乘keys^T形状(embed_size, seq_length)
# 结果形状(seq_length, seq_length) ---- 即每个 token 对每个 token 的注意力原始分数
# 最终attention_scores 形状为 (N, seq_length, seq_length)
attention_scores = torch.bmm(queries, keys.transpose(1, 2))
# 掩码处理
# mask == 0:找出 mask 中值为 0 的位置(表示这些位置需要被屏蔽),
# 将这些位置的注意力分数设为负无穷大(用 -1e20 近似)
# 为什么这样做?
# 因为后续的 softmax 会将 e^(-∞) ≈ 0,使得被屏蔽位置的注意力权重趋近于 0,从而实现"忽略这些位置"的效果。
if mask is not None:
attention_scores = attention_scores.masked_fill(mask == 0, float('-1e20'))
# (self.embed_size ** 0.5)也就是√d_k
# 当d_k很大时,点积结果方差变大,softmax会趋向于极端的one-hot分布,梯度很小。除以√d_k可以缓解这个问题。
# F.softmax(..., dim=-1)沿最后一个维度(即每个query对所有key的那一行)做softmax
# 将缩放后的分数转化为概率分布(每个token对其他token的关注权重,和为1)
# 结果attention形状为 (N, seq_length, seq_length),即注意力权重矩阵。
attention = F.softmax(attention_scores / (self.embed_size ** 0.5), dim=-1)
# 加权求和(输出)
# 用注意力权重对 values 做加权求和:
# attention 形状 (N, seq_length, seq_length)----权重
# values 形状 (N, seq_length, embed_size)----被加权的内容
# torch.bmm 结果形状 (N, seq_length, embed_size)
output = torch.bmm(attention, values)
# 在每个batch里,attention与values的点乘,也就是:
# token i = 所有token的value向量按注意力权重加起来的和。
return output
Single Head Attention流程图

掩码矩阵的形式
因果掩码是一个上三角矩阵,对角线及以下为1(可见),对角线以上为0(屏蔽):

多头注意力MHA(Multi-Head Attention)
Multi-Head Attention代码
python
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, embed_size, num_heads):
super(MultiHeadAttention, self).__init__()
# 嵌入维度大小
self.embed_size = embed_size
# 注意力头的数量
self.num_heads = num_heads
# 每个注意力头的维度。例如 embed_size=512, num_heads=8 → head_dim=64。
# 核心思想:把一个大维度拆成多个小维度,每个头在小空间里独立做注意力。
self.head_dim = embed_size // num_heads
# 确保嵌入尺寸可以被头数整除
assert (
self.head_dim * num_heads == embed_size
), "嵌入尺寸需要能够被头数整除"
# 为多头注意力机制定义线性层
self.values = nn.Linear(embed_size, embed_size, bias=False)
self.keys = nn.Linear(embed_size, embed_size, bias=False)
self.queries = nn.Linear(embed_size, embed_size, bias=False)
# 定义一个线性层用于合并多个头的输出,这个有bias(默认bias=True)
self.fc_out = nn.Linear(embed_size, embed_size)
def forward(self, x, mask):
N = x.shape[0] # 批量大小
seq_length = x.shape[1] # 序列长度
# x的shape是(N, seq_length, embed_size)
# 将输入x通过线性变换后分割成num_heads份,每份是head_dim维。
# x的shape (N, seq_length, embed_size)最后一个维度embed_size被拆分成了 (num_heads, head_dim)
# 变换后的shape是(N, seq_length, num_heads, head_dim)
values = self.values(x).view(N, seq_length, self.num_heads, self.head_dim)
keys = self.keys(x).view(N, seq_length, self.num_heads, self.head_dim)
queries = self.queries(x).view(N, seq_length, self.num_heads, self.head_dim)
# 把num_heads提前到第二维,这样每个头就是独立的 (seq_len, head_dim) 矩阵
# 后续可以并行计算所有头的注意力,互不干扰。
# 从(N, seq_length, num_heads, head_dim) 变为 (N, num_heads, seq_length, head_dim)
values = values.permute(0, 2, 1, 3)
keys = keys.permute(0, 2, 1, 3)
queries = queries.permute(0, 2, 1, 3)
# 计算缩放点积注意力得分(爱因斯坦求和)
# n: batch维度N q:query length k:key length h:head d:head_dim
# 在 d 上做内积 → 结果 nhqk:(N, heads, Q_seq, K_seq) 相对于前面的单头的注意力,多了一个heads维度
attention_scores = torch.einsum("nhqd,nhkd->nhqk", [queries, keys])
if mask is not None:
# 如果有掩码,则应用掩码到注意力得分上
attention_scores = attention_scores.masked_fill(mask.unsqueeze(1) == 0, float('-1e20'))
# 应用softmax函数并根据head_dim开平方根进行缩放
attention = F.softmax(attention_scores / (self.head_dim ** 0.5), dim=-1)
# 将注意力权重与值向量相乘
# attention的shape (N, num_heads, Q_seq_length, K_seq_length)
# values的shape (N, num_heads, K_seq_length, head_dim)
# 在K维度上求和 → nqhd:(N, Q, heads, D)
out = torch.einsum("nhqk,nhkd->nqhd", [attention, values]).reshape(
N, seq_length, self.embed_size
)
# 通过线性层组合来自所有头的信息
out = self.fc_out(out)
return out
MHA流程图

线性注意力(Linear Attention)

Linear Attention与普通 Self-Attention 的对比

MLA(Multi-head Latent Attention)
近些年,Attention的模型出现了很多改进版本,但主要思想是,减少KV的数量,达到减少KV Cache显存的使用量,MLA的思想是对KV的维度降维,实现减少对显存的占用率。
背景知识
1. 低秩分解
在讲解MLA之前,先了解一下什么是低秩分解 。
举一个例子:
直接用一个矩阵做投影:
输入 (5120维) × W_big (5120×24576)──▶ 输出 (24576维)
↑
W_big有1.26亿个参数!
这太重了。
低秩分解的想法:这个大矩阵可以近似为两个小矩阵的乘积:
W_big ≈ W2 × W1
(5120×24576) ≈ (5120×1536) × (1536×24576)
由标准模式的1步转变为低秩模式的2步,参数量由1.26亿降为4560万。
标准模式(1 步)
x (5120) ────────── wq ──────────▶ Q (24576)
5120 × 24576
参数量 ≈ 1.26 亿
低秩模式(2 步)
x (5120) ── wq_a ──▶ q_compressed (1536) ── RMSNorm ── wq_b ──▶ Q (24576)
5120×1536 1536×24576
≈ 787 万 ≈ 3775 万
└────── 合计 ≈ 4560 万 ──────┘
为什么叫"低秩"?
任何一个大矩阵 W (5120×24576),如果它的秩不高(即信息有大量冗余),就可以用两个小矩阵的乘积来逼近。
中间维度 r 就叫秩(rank),这里 r = q_lora_rank = 1536。秩越小,压缩越狠,参数越少,但逼近精度越差。1536 是 DeepSeek 找到的平衡点。
直觉理解:
- wq_a 把信息"收拢"到一个低维瓶颈 1536
- wq_b 再把信息"展开"到各头 24576
- 瓶颈层迫使模型只保留最关键的信息,丢弃冗余
- 类似于自动编码器的 Encoder-Decoder 结构
2. ColumnParallelLinear和RowParallelLinear函数
ColumnParallelLinear 和 RowParallelLinear 是分布式大模型训练和推理中(特别是基于 Megatron-LM 提出的张量并行 Tensor Parallelism, TP)极其重要的两个核心算子。
它们的作用是将一个巨大的线性层(全连接层)的参数矩阵切分到多张 GPU 上,使得每张 GPU 只需要计算和存储一部分权重,从而突破单张显存的瓶颈。这两个函数通常是成对使用的。
2.1 ColumnParallelLinear
将权重矩阵按 列(Column) 进行切分,分配给多张 GPU。
假设我们有一个线性运算 Y = W ⋅ X Y = W \cdot X Y=W⋅X。
如果 W W W 非常大,我们把 W W W 沿着列方向切分为多块(假设有 2 张 GPU): W = W 1 , W 2 W = W_1, W_2 W=W1,W2。
- GPU 1 负责计算 Y 1 = W 1 ⋅ X Y_1 = W_1 \cdot X Y1=W1⋅X
- GPU 2 负责计算 Y 2 = W 2 ⋅ X Y_2 = W_2 \cdot X Y2=W2⋅X
输入与输出特征
- 输入 X X X:在所有 GPU 上都是完全相同的拷贝(前向传播时,每张卡都会拿到完整输入)。
- 输出 Y Y Y:每张 GPU 算出的结果是最终输出的"一部分"。GPU 1 拿到了特征向量的左半部分,GPU 2 拿到了右半部分。
通信代价
Tip:
零通信! 在前向传播(Forward)阶段,不需要任何 GPU 间的通信。每张显卡独立计算自己分得的列。
代码中的体现 :
self.wq_b = ColumnParallelLinear(self.q_lora_rank, self.n_heads * self.qk_head_dim)
self.wkv_b = ColumnParallelLinear(self.kv_lora_rank, self.n_heads * (self.qk_nope_head_dim + self.v_head_dim))在 MLA(多头隐向量注意力)中,生成 Q、K、V 的投影矩阵(如
self.wq,self.wkv_b)用的就是列并行。在注意力机制中,切分列本质上就是把"不同的注意力头(Attention Heads)"分配给了不同的 GPU。
2.2 RowParallelLinear
将权重矩阵按 行(Row) 进行切分,分配给多张 GPU。
假设我们要进行下一步计算 Z = Y ⋅ W o u t Z = Y \cdot W_{out} Z=Y⋅Wout,其中 Y Y Y 是刚才列并行的输出(已经切分为 Y 1 Y_1 Y1 和 Y 2 Y_2 Y2 散落在两张卡上)。
Y = Y 1 , Y 2 Y = Y_1, Y_2 Y=Y1,Y2
此时,我们将 W o u t W_{out} Wout 沿着行方向切分:
W o u t = W o u t 1 W o u t 2 W_{out} = \begin{bmatrix} W_{out1} \\ W_{out2} \end{bmatrix} Wout=Wout1Wout2
计算结果则为: Z = Y 1 ⋅ W o u t 1 + Y 2 ⋅ W o u t 2 Z = Y_1 \cdot W_{out1} + Y_2 \cdot W_{out2} Z=Y1⋅Wout1+Y2⋅Wout2。
输入与输出特征
- 输入 :GPU 1 使用它刚才计算出的 Y 1 Y_1 Y1 乘以 W o u t 1 W_{out1} Wout1;GPU 2 使用 Y 2 Y_2 Y2 乘以 W o u t 2 W_{out2} Wout2。
- 输出 :每张显卡各自算出来的只是一部分 部分和(Partial Sum)。
通信代价
IMPORTANT
需要 All-Reduce! 为了得到最终完整的 Z Z Z,必须在这个算子计算完成后,执行一次跨 GPU 的 All-Reduce 操作,把所有 GPU 上算出来的部分和累加起来。这样所有 GPU 才能拿到完整且一致的 Z Z Z 结果。
代码中的体现 :
self.wo = RowParallelLinear(self.n_heads * self.v_head_dim, self.dim)注意力机制最后一步的输出投影矩阵
self.wo使用的就是行并行。前面各张 GPU 分别算完了属于自己的注意力头的输出,通过按行切分的wo相乘后,执行 All-Reduce 将结果累加,输出最终完整的隐状态向量。
2.3 通过一个例子来解释
2.3.1 先看"不并行"时,正常是怎么算的
假设我们现在要做一个极其简单的矩阵乘法(向量乘以矩阵): Z = X ⋅ W Z = X \cdot W Z=X⋅W。
为了方便理解,假设 Y Y Y 是一个有 4 个数字的向量, W W W 是一列有 4 个数字的权重:
- X = x 1 , x 2 , x 3 , x 4 X = x_1, x_2, x_3, x_4 X=x1,x2,x3,x4
- W = w 1 w 2 w 3 w 4 W = \begin{bmatrix} w_1 \\ w_2 \\ w_3 \\ w_4 \end{bmatrix} W= w1w2w3w4
按照正常的矩阵乘法(点乘)规则,最终结果 Z Z Z 是把对应位置相乘,然后全部加起来 :
Z = x 1 w 1 + x 2 w 2 + x 3 w 3 + x 4 w 4 Z = x_1 w_1 + x_2 w_2 + x_3 w_3 + x_4 w_4 Z=x1w1+x2w2+x3w3+x4w4
这就是最终完整、正确的结果
2.3.2 现在切开(RowParallelLinear 登场)
在上一层(列并行)结束时,输入特征 Y Y Y 已经被切成了两半散落在两张显卡上了(这是列并行的特性造成的):
- GPU 1 手里只拿着 Y Y Y 的左半边: Y l e f t = y 1 , y 2 Y_{left} = y_1, y_2 Yleft=y1,y2
- GPU 2 手里只拿着 Y Y Y 的右半边: Y r i g h t = y 3 , y 4 Y_{right} = y_3, y_4 Yright=y3,y4
这时候遇到了我们要解释的 RowParallelLinear(行并行) 。它的规则是把权重 W W W 横向(按行)切断: - GPU 1 分配到 W W W 的上半部: W t o p = w 1 w 2 W_{top} = \begin{bmatrix} w_1 \\ w_2 \end{bmatrix} Wtop=w1w2
- GPU 2 分配到 W W W 的下半部: W b o t = w 3 w 4 W_{bot} = \begin{bmatrix} w_3 \\ w_4 \end{bmatrix} Wbot=w3w4
2.3.3 为什么叫"部分和(Partial Sum)"?
现在,两张显卡开始自己算自己的(不互相通信)。
按照矩阵乘法规则,它们各自手里的向量和权重正好能乘起来:
- GPU 1 自己算的输出(设为 Z 1 Z_1 Z1) :
Z 1 = Y l e f t ⋅ W t o p = y 1 , y 2 ⋅ w 1 w 2 = y 1 w 1 + y 2 w 2 Z_1 = Y_{left} \cdot W_{top} = y_1, y_2 \cdot \begin{bmatrix} w_1 \\ w_2 \end{bmatrix} = \mathbf{y_1 w_1 + y_2 w_2} Z1=Yleft⋅Wtop=y1,y2⋅w1w2=y1w1+y2w2 - GPU 2 自己算的输出(设为 Z 2 Z_2 Z2) :
Z 2 = Y r i g h t ⋅ W b o t = y 3 , y 4 ⋅ w 3 w 4 = y 3 w 3 + y 4 w 4 Z_2 = Y_{right} \cdot W_{bot} = y_3, y_4 \cdot \begin{bmatrix} w_3 \\ w_4 \end{bmatrix} = \mathbf{y_3 w_3 + y_4 w_4} Z2=Yright⋅Wbot=y3,y4⋅w3w4=y3w3+y4w4
你对比一下正常的最终结果 Z Z Z 和 GPU 各自算出来的 Z 1 , Z 2 Z_1, Z_2 Z1,Z2:
正常的 Z = ( y 1 w 1 + y 2 w 2 ) + ( y 3 w 3 + y 4 w 4 ) Z = (y_1 w_1 + y_2 w_2) + (y_3 w_3 + y_4 w_4) Z=(y1w1+y2w2)+(y3w3+y4w4)
所以:Z = Z 1 + Z 2 Z = Z_1 + Z_2 Z=Z1+Z2
看出来了吗?
GPU 1 算出来的 Z 1 Z_1 Z1 不是最终结果,它只是最终结果"前半部分的加和"。
GPU 2 算出来的 Z 2 Z_2 Z2 也不是最终结果,它只是最终结果"后半部分的加和"。
这就是为什么我们把它们叫做 部分和(Partial Sum)。任何一张单卡拿着自己的 Partial Sum 是没法往下做后续模型计算的,因为数据不完整。
2.3.4 救星登场:All-Reduce 通信
为了让模型继续往下算,两张卡必须把完整的 Z Z Z 拼凑出来。这时候就需要调用底层分布式通信协议(比如 NCCL)里的 All-Reduce 操作。
!NOTE
All-Reduce(全规约) 的动作可以用大白话翻译为:
"兄弟们,把你们手里的结果都拿出来,我们统一加在一起,然后把最终的总和,抄送给每一个人。"
- 步骤 A (Reduce) :GPU 1 拿出 Z 1 Z_1 Z1,GPU 2 拿出 Z 2 Z_2 Z2,通过显卡间的网线(如 NVLink)传数据,在某处相加得到完整的 Z = Z 1 + Z 2 Z = Z_1 + Z_2 Z=Z1+Z2。
- 步骤 B (All) :把相加后完整的 Z Z Z,广播发回给 GPU 1 和 GPU 2。
2.3.5 最终结果:
执行完 All-Reduce 之后,GPU 1 手里的内存从毫无意义的 Z 1 Z_1 Z1 变成了完整的 Z Z Z;GPU 2 手里的内存也从 Z 2 Z_2 Z2 变成了完整的 Z Z Z。
两张卡重新拥有了完全相同且正确的数据,又可以愉快地进行下一层网络(如新一轮的列并行)的计算了!
通过这种数学上的"拆解分配 -> 独立相乘 -> 汇聚相加",模型成功绕过了单卡算力/显存的瓶颈,这正是张量并行(Tensor Parallelism)的魔法所在。
2.4 为什么ColumnParallelLinear和RowParallelLinear是"天生一对"?
在 Transformer 的注意力层(Attention)或前馈网络层(FFN/MLP)中,这两个算子总是采用 列并行 → \rightarrow → 行并行 的组合策略。
这么设计的精妙之处在于将分布式训练的通信开销降到了最低。其数据流转过程如下:
| 步骤 | 阶段 | 使用算子 / 操作 | 是否需要 GPU 通信 |
|---|---|---|---|
| 1 | 投影阶段 | 经过 ColumnParallelLinear 生成 Q/K/V | ❌ 零通信 |
| 2 | 计算阶段 | 执行 Attention 打分或激活函数计算 | ❌ 零通信 |
| 3 | 聚合阶段 | 经过 RowParallelLinear 计算部分和 | ❌ 零通信 |
| 4 | 同步阶段 | 模块出口处执行一次 All-Reduce 汇聚结果 | ✅ 一次通信 |
如果随意切分矩阵,可能导致每做一步矩阵乘法都需要在 GPU 之间传输大量数据,挤爆网络带宽。列-行组合保证了一个庞大的计算模块(如整个 Attention Layer)内部是完全独立的"零通信"状态,只有在离开该模块时才进行一次高效的数据汇聚。
Multi-head Latent Attention代码
以下代码摘自DeepSeek V3:
传统的 Transformer(如 LLaMA 使用的 MHA/GQA)在推理时需要缓存每个 Token 的完整 K 和 V 矩阵(即 KV Cache),这会占用巨大的显存。MLA 的核心思想是通过"低秩投影"将 K 和 V 压缩到一个共享的低维潜在空间(Latent Space)中,从而将 KV Cache 的显存占用降低 10 倍以上,同时由于独特的推导,它在数学上完全等价于标准注意力,不损失性能。
python
class MLA(nn.Module):
"""
Multi-Headed Attention Layer (MLA).
Attributes:
dim (int): Dimensionality of the input features. 模型隐藏维度
n_heads (int): Number of attention heads.
n_local_heads (int): Number of local attention heads for distributed systems.
q_lora_rank (int): Rank for low-rank query projection.
kv_lora_rank (int): Rank for low-rank key/value projection.
qk_nope_head_dim (int): Dimensionality of non-positional query/key projections.
qk_rope_head_dim (int): Dimensionality of rotary-positional query/key projections.
qk_head_dim (int): Total dimensionality of query/key projections.
v_head_dim (int): Dimensionality of value projections.
softmax_scale (float): Scaling factor for softmax in attention computation.
"""
# 核心思想:标准多头注意力每个头存完整的 K、V,总 KV Cache = 2 × n_heads × head_dim。
# MLA 把 KV 压缩到共享的低秩潜在空间(kv_lora_rank = 512),再按需展开到各头,KV Cache 只存压缩后的向量和位置编码,显存减少 10 倍以上。
def __init__(self, args: ModelArgs):
super().__init__()
self.dim = args.dim
self.n_heads = args.n_heads
# 张量并行(Tensor Parallelism)设置:n_local_heads代表当前GPU负责计算的注意力头数量
# world_size:分布式训练的GPU数量,每个 GPU 只负责部分头
self.n_local_heads = args.n_heads // world_size
# Q压缩后的潜在维度
self.q_lora_rank = args.q_lora_rank
# KV压缩后的潜在维度
self.kv_lora_rank = args.kv_lora_rank
# 无位置编码的QK维度
self.qk_nope_head_dim = args.qk_nope_head_dim
# 旋转位置编码的维度
self.qk_rope_head_dim = args.qk_rope_head_dim
# qk_head_dim:每个头完整 Q/K 的维度 = 无位置部分 + RoPE 部分。例如 128 + 64 = 192
# Q 和 K 被拆成两部分处理:不施加位置编码的部分(nope)+ 施加 RoPE 的部分(rope)
self.qk_head_dim = args.qk_nope_head_dim + args.qk_rope_head_dim
self.v_head_dim = args.v_head_dim
if self.q_lora_rank == 0:
# 如果 q_lora_rank 为 0,退化为标准的多头注意力Q投影
self.wq = ColumnParallelLinear(self.dim, self.n_heads * self.qk_head_dim)
else:
# 否则,对 Q 也进行低秩压缩:先压缩到 q_lora_rank,做归一化,再展开到各个头
# 类似于前面提到的 LoRA 思想:
# x (5120) ── wq_a ──▶ q_compressed (1536) ── RMSNorm ── wq_b ──▶ Q (24576)
self.wq_a = Linear(self.dim, self.q_lora_rank)
self.q_norm = RMSNorm(self.q_lora_rank)
self.wq_b = ColumnParallelLinear(self.q_lora_rank, self.n_heads * self.qk_head_dim)
# wkv_a 负责将输入 x 压缩成:Latent KV (维度 kv_lora_rank) + 用于 RoPE 的 K (维度 qk_rope_head_dim)
# 注意:这里只有一个矩阵,所有注意力头共享这个压缩后的 Latent 向量!
self.wkv_a = Linear(self.dim, self.kv_lora_rank + self.qk_rope_head_dim)
self.kv_norm = RMSNorm(self.kv_lora_rank)
# wkv_b 负责将压缩后的 Latent KV 展开(解压)成各个头独享的 K_nope 和 V
self.wkv_b = ColumnParallelLinear(self.kv_lora_rank, self.n_heads * (self.qk_nope_head_dim + self.v_head_dim))
# 最终的输出投影矩阵
self.wo = RowParallelLinear(self.n_heads * self.v_head_dim, self.dim)
# Attention 里的 softmax 缩放因子 (1 / sqrt(d))
self.softmax_scale = self.qk_head_dim ** -0.5
# 针对长文本推断(Context Extension)的缩放调整(类似 Yarn 算法的乘子)
if args.max_seq_len > args.original_seq_len:
mscale = 0.1 * args.mscale * math.log(args.rope_factor) + 1.0
self.softmax_scale = self.softmax_scale * mscale * mscale
# kv cache显存分配
if attn_impl == "naive":
# 朴素实现:开辟标准的 K 和 V Cache。维度极大 (seq_len, heads, head_dim)
self.register_buffer("k_cache",
torch.zeros(args.max_batch_size, args.max_seq_len, self.n_local_heads, self.qk_head_dim), persistent=False)
self.register_buffer("v_cache",
torch.zeros(args.max_batch_size, args.max_seq_len, self.n_local_heads, self.v_head_dim), persistent=False)
else:
# 优化实现(实际部署用):KV Cache只需要存压缩后的 kv_lora_rank 向量,以及一小段共享的pe(位置编码)向量。
# 这就是 MLA 能省十几倍显存的根本原因!
self.register_buffer("kv_cache",
torch.zeros(args.max_batch_size, args.max_seq_len, self.kv_lora_rank), persistent=False)
self.register_buffer("pe_cache",
torch.zeros(args.max_batch_size, args.max_seq_len, self.qk_rope_head_dim), persistent=False)
def forward(self, x: torch.Tensor, start_pos: int, freqs_cis: torch.Tensor, mask: Optional[torch.Tensor]):
"""
Forward pass for the Multi-Headed Attention Layer (MLA).
Args:
x (torch.Tensor): Input tensor of shape (batch_size, seq_len, dim).
start_pos (int): Starting position in the sequence for caching.
freqs_cis (torch.Tensor): Precomputed complex exponential values for rotary embeddings.
mask (Optional[torch.Tensor]): Mask tensor to exclude certain positions from attention.
Returns:
torch.Tensor: Output tensor with the same shape as the input.
"""
bsz, seqlen, _ = x.size() # (batch_size, seq_length, embedding_size)
end_pos = start_pos + seqlen
# ================= 1. 计算 Query =================
if self.q_lora_rank == 0: # 退化为标准的多头注意力
q = self.wq(x)
else: # 低秩压缩
q = self.wq_b(self.q_norm(self.wq_a(x)))
# q的shape (batch_size, seq_length, embed_size)最后一个维度embed_size被拆分成了 (num_heads, head_dim)
# 变换后的shape是(N, seq_length, num_heads, head_dim) 这一步与前面的MHA的forward中的shape转换是一致的
q = q.view(bsz, seqlen, self.n_local_heads, self.qk_head_dim)
# 将 Q 拆分为两半:不带位置编码的部分 (q_nope) 和 需要应用位置编码的部分 (q_pe)
q_nope, q_pe = torch.split(q, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)
# 仅对 q_pe 施加旋转位置编码 (RoPE)
q_pe = apply_rotary_emb(q_pe, freqs_cis)
# ================= 2. 计算并压缩 KV =================
kv = self.wkv_a(x) # [batch, seq, kv_lora_rank + qk_rope_head_dim]
# 拆分为:Latent KV (内容向量) 和 k_pe (需要位置编码的 K)
kv, k_pe = torch.split(kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1)
# 仅对 k_pe 施加旋转位置编码 (注意这里 k_pe 在各个头上是共享的,所以 unsqueeze 扩展了一维)
k_pe = apply_rotary_emb(k_pe.unsqueeze(2), freqs_cis)
if attn_impl == "naive":
q = torch.cat([q_nope, q_pe], dim=-1)
kv = self.wkv_b(self.kv_norm(kv))
kv = kv.view(bsz, seqlen, self.n_local_heads, self.qk_nope_head_dim + self.v_head_dim)
k_nope, v = torch.split(kv, [self.qk_nope_head_dim, self.v_head_dim], dim=-1)
k = torch.cat([k_nope, k_pe.expand(-1, -1, self.n_local_heads, -1)], dim=-1)
self.k_cache[:bsz, start_pos:end_pos] = k
self.v_cache[:bsz, start_pos:end_pos] = v
scores = torch.einsum("bshd,bthd->bsht", q, self.k_cache[:bsz, :end_pos]) * self.softmax_scale
else:
wkv_b = self.wkv_b.weight if self.wkv_b.scale is None
else weight_dequant(self.wkv_b.weight, self.wkv_b.scale, block_size)
wkv_b = wkv_b.view(self.n_local_heads, -1, self.kv_lora_rank)
q_nope = torch.einsum("bshd,hdc->bshc", q_nope, wkv_b[:, :self.qk_nope_head_dim])
self.kv_cache[:bsz, start_pos:end_pos] = self.kv_norm(kv)
self.pe_cache[:bsz, start_pos:end_pos] = k_pe.squeeze(2)
scores = (torch.einsum("bshc,btc->bsht", q_nope, self.kv_cache[:bsz, :end_pos]) +
torch.einsum("bshr,btr->bsht", q_pe, self.pe_cache[:bsz, :end_pos])) * self.softmax_scale
if mask is not None:
scores += mask.unsqueeze(1)
scores = scores.softmax(dim=-1, dtype=torch.float32).type_as(x)
if attn_impl == "naive":
x = torch.einsum("bsht,bthd->bshd", scores, self.v_cache[:bsz, :end_pos])
else:
x = torch.einsum("bsht,btc->bshc", scores, self.kv_cache[:bsz, :end_pos])
x = torch.einsum("bshc,hdc->bshd", x, wkv_b[:, -self.v_head_dim:])
x = self.wo(x.flatten(2))
return x