代码实现 GQA Attention

从零理解 GQA Attention:结合 PyTorch 源码详解

首先先看下MiniMind的实现

def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1):
    def rotate_half(x): return torch.cat((-x[..., x.shape[-1] // 2:], x[..., : x.shape[-1] // 2]), dim=-1)
    q_embed = ((q * cos.unsqueeze(unsqueeze_dim)) + (rotate_half(q) * sin.unsqueeze(unsqueeze_dim))).to(q.dtype)
    k_embed = ((k * cos.unsqueeze(unsqueeze_dim)) + (rotate_half(k) * sin.unsqueeze(unsqueeze_dim))).to(k.dtype)
    return q_embed, k_embed

def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:
    bs, slen, num_key_value_heads, head_dim = x.shape
    if n_rep == 1: return x
    return (x[:, :, :, None, :].expand(bs, slen, num_key_value_heads, n_rep, head_dim).reshape(bs, slen, num_key_value_heads * n_rep, head_dim))

class Attention(nn.Module):
    def __init__(self, config: MiniMindConfig):
        super().__init__()
        self.num_key_value_heads = config.num_attention_heads if config.num_key_value_heads is None else config.num_key_value_heads
        self.n_local_heads = config.num_attention_heads
        self.n_local_kv_heads = self.num_key_value_heads
        self.n_rep = self.n_local_heads // self.n_local_kv_heads
        self.head_dim = config.head_dim
        self.is_causal = True
        self.q_proj = nn.Linear(config.hidden_size, config.num_attention_heads * self.head_dim, bias=False)
        self.k_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)
        self.v_proj = nn.Linear(config.hidden_size, self.num_key_value_heads * self.head_dim, bias=False)
        self.o_proj = nn.Linear(config.num_attention_heads * self.head_dim, config.hidden_size, bias=False)
        self.q_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)
        self.k_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)
        self.attn_dropout = nn.Dropout(config.dropout)
        self.resid_dropout = nn.Dropout(config.dropout)
        self.dropout = config.dropout
        self.flash = hasattr(torch.nn.functional, 'scaled_dot_product_attention') and config.flash_attn

    def forward(self, x, position_embeddings, past_key_value=None, use_cache=False, attention_mask=None):
        bsz, seq_len, _ = x.shape
        xq, xk, xv = self.q_proj(x), self.k_proj(x), self.v_proj(x)
        xq = xq.view(bsz, seq_len, self.n_local_heads, self.head_dim)
        xk = xk.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim)
        xv = xv.view(bsz, seq_len, self.n_local_kv_heads, self.head_dim)
        xq, xk = self.q_norm(xq), self.k_norm(xk)
        cos, sin = position_embeddings
        xq, xk = apply_rotary_pos_emb(xq, xk, cos, sin)
        if past_key_value is not None:
            xk = torch.cat([past_key_value[0], xk], dim=1)
            xv = torch.cat([past_key_value[1], xv], dim=1)
        past_kv = (xk, xv) if use_cache else None
        xq, xk, xv = (xq.transpose(1, 2), repeat_kv(xk, self.n_rep).transpose(1, 2), repeat_kv(xv, self.n_rep).transpose(1, 2))
        if self.flash and (seq_len > 1) and (not self.is_causal or past_key_value is None) and (attention_mask is None or torch.all(attention_mask == 1)):
            output = F.scaled_dot_product_attention(xq, xk, xv, dropout_p=self.dropout if self.training else 0.0, is_causal=self.is_causal)
        else:
            scores = (xq @ xk.transpose(-2, -1)) / math.sqrt(self.head_dim)
            if self.is_causal: scores[:, :, :, -seq_len:] += torch.full((seq_len, seq_len), float("-inf"), device=scores.device).triu(1)
            if attention_mask is not None: scores += (1.0 - attention_mask.unsqueeze(1).unsqueeze(2)) * -1e9
            output = self.attn_dropout(F.softmax(scores.float(), dim=-1).type_as(xq)) @ xv
        output = output.transpose(1, 2).reshape(bsz, seq_len, -1)
        output = self.resid_dropout(self.o_proj(output))
        return output, past_kv

这段代码实际上就是一个 Llama/Qwen 风格的 GQA Attention 实现

如果只看代码,很多人会迷失在:

view
transpose
repeat_kv
flash attention
cache

里面。

实际上整个流程只有 8 步:

输入 x
 ↓
Q投影
K投影
V投影
 ↓
分Head
 ↓
RoPE
 ↓
拼接KV Cache
 ↓
GQA扩展KV
 ↓
Attention
 ↓
输出投影

先从整体架构理解,再逐行拆解。


一、Attention初始化

配置参数

假设:

hidden_size = 512
num_attention_heads = 8
num_key_value_heads = 2
head_dim = 64

那么:

Q Head数 = 8
KV Head数 = 2

这就是 GQA。


代码:

self.n_local_heads = config.num_attention_heads

即:

8

self.n_local_kv_heads = self.num_key_value_heads

即:

2

计算:

self.n_rep = self.n_local_heads // self.n_local_kv_heads

得到:

8 // 2 = 4

表示:

每个KV Head
被4个Q Head共享

即:

Q0 Q1 Q2 Q3 -> KV0

Q4 Q5 Q6 Q7 -> KV1

这是 GQA 最核心的参数。


二、QKV Projection

输入:

x.shape = [B, L, hidden_size]

例如:

[2,128,512]

Query

xq = self.q_proj(x)

线性层:

512
↓
8 × 64
↓
512

输出:

[2,128,512]

Key

xk = self.k_proj(x)

注意:

num_key_value_heads=2

所以:

512
↓
2 × 64
↓
128

输出:

[2,128,128]

Value

同理:

xv = [2,128,128]

这就是 GQA 的第一处省显存。

传统 MHA:

Q=512
K=512
V=512

GQA:

Q=512
K=128
V=128

直接减少:

75%

三、拆成多个Head

Query

xq = xq.view(
    bsz,
    seq_len,
    self.n_local_heads,
    self.head_dim
)

变成:

[2,128,8,64]

即:

8个Q Head

Key

xk.view(...)

变成:

[2,128,2,64]

即:

2个KV Head

Value

同理:

[2,128,2,64]

此时:

Q : 8 heads
K : 2 heads
V : 2 heads

已经是标准 GQA 结构。


四、QK Norm

代码:

xq = self.q_norm(xq)
xk = self.k_norm(xk)

对应:

QK-Norm

这是 Llama3 引入的重要改进。


对每个 Head 做 RMSNorm:

RMS(x)=1dixi2

然后:

x=x/RMS(x)

作用:

避免QK数值爆炸
训练更稳定

五、RoPE

代码:

xq,xk=apply_rotary_pos_emb(...)

进入:

q_embed = q * cos + rotate_half(q) * sin

即:

q=qcosθ+R(q)sinθ

其中:

R(q)=[q2,q1]

本质:

把位置编码旋转进Q和K

最终:

Q包含位置信息
K包含位置信息
V不包含

六、KV Cache

推理时:

假设已经生成:

Hello world

缓存:

past_key_value =
(
 old_k,
 old_v
)

新token:

!

对应:

new_k
new_v

代码:

xk = torch.cat(
    [past_key_value[0], xk],
    dim=1
)

变成:

old_k + new_k

同理:

xv

拼接。

最终:

所有历史KV

保存在缓存中。


七、GQA核心

终于来到最关键部分。


当前:

xq = [B,L,8,64]

xk = [B,L,2,64]

xv = [B,L,2,64]

Head数对不上:

Q=8

K=2

Attention没法算。


于是:

repeat_kv(xk, n_rep=4)

进入:

x[:, :, :, None, :]

原来:

[2,128,2,64]

变:

[2,128,2,1,64]

expand:

expand(
    bs,
    slen,
    2,
    4,
    64
)

变:

[2,128,2,4,64]

逻辑:

KV0
↓
复制4份

KV1
↓
复制4份

reshape:

[2,128,8,64]

变成:

KV0 KV0 KV0 KV0
KV1 KV1 KV1 KV1

对应:

Q0 → KV0
Q1 → KV0
Q2 → KV0
Q3 → KV0

Q4 → KV1
Q5 → KV1
Q6 → KV1
Q7 → KV1

这就是 GQA。


为什么这样不会增加KV Cache?

很多人第一次会疑惑:

你不是复制了吗?

实际上:

复制发生在:

当前Forward

阶段。

缓存中保存的仍然是:

[2,128,2,64]

而不是:

[2,128,8,64]

所以:

KV Cache大小仍然是2个Head

没有增加。


八、Transpose

代码:

xq.transpose(1,2)

从:

[B,L,H,D]

变:

[B,H,L,D]

得到:

Q = [2,8,128,64]

Attention实现都喜欢:

Batch
Head
Seq
Dim

这种布局。


九、Attention计算

Flash Attention路径

F.scaled_dot_product_attention(...)

直接调用:

PyTorch FlashAttention

内部融合:

QKᵀ
Softmax
AV

速度最快。


十、普通Attention路径

Step1

scores = (xq @ xk.transpose(-2,-1)) / sqrt(head_dim)

得到:

[B,H,L,L]

即:

Attention Score

数学上:

S=QKTd

Step2

因果Mask

scores += upper_triangle(-inf)

变成:

未来token不可见

例如:

I love

不能偷看:

you

Step3

Softmax

A = softmax(scores)

对应:

A=softmax(S)

Step4

乘V

output = A @ xv

对应:

O=softmax(QKT)V

输出:

[B,H,L,D]

十一、合并Head

代码:

output.transpose(1,2)

回到:

[B,L,H,D]

然后:

reshape(
    bsz,
    seq_len,
    -1
)

变:

[B,L,8×64]

即:

[2,128,512]

十二、输出投影

output = self.o_proj(output)

对应:

OWO

即 Transformer 最后一层:

Concat Heads
 ↓
Linear
 ↓
Hidden Size

输出:

[2,128,512]

整个GQA流程图

Pasted image 20260625143931.png

从源码实现角度看,GQA 真正新增的代码其实只有 repeat_kv() 这一处。其它流程(QKV 投影、RoPE、Mask、Softmax、输出投影)与普通 MHA 基本一致。

因此可以把 GQA 理解为:

训练和推理时只存储少量 KV Head,在计算 Attention 的瞬间,通过 repeat_kv 将 KV 逻辑上扩展到与 Query Head 数量一致,从而兼顾 KV Cache 节省和多头表达能力。

这也是为什么 Llama 2/3、Qwen、DeepSeek 等模型都采用 num_heads > num_kv_heads 的设计。真正节省显存的不是 Attention 计算本身,而是长期保存的 KV Cache