多头注意力(Multi-head Attention,MHA)
原始transformer架构的多头注意力是将QKV拆成多个头,对于(batch_size,seq_len,d_model)维度的QKV向量拆成:(batch_size,seq_len,head_num,d_k),相应的每个QKV都会有一个大小为(d_model,head_num*d_k)的参数矩阵。
多查询注意力(Multi-Query Attention,MQA)
动机
MQA的原论文的题目是:Fast Transformer Decoding: One Write-Head is All You Need,从题目可以看到MQA出现的原因是为了加快推理解码速度的。引言第一句:
As we will discuss, the speed of incremental Transformer inference on modern computing hardware is limited by the memory bandwidth necessary to reload the large "keys" and "values" tensors which encode the state of the attention layers.
Transformer推理的时候由于每次都得加载巨大的key和values向量,造成了推理速度的减慢,这里注意,是reload,而不是重新计算,所以肯定当时就有KV cache的技术的存在了。但是重新加载也很费时间。这里就可以明白MQA出现的动机就是为了省KV向量的内存,然后让推理更快
实现
MQA出自的实现很简单,就是所有的Q头共享一个K_head和V_head。
以下是原论文的伪代码:

有点抽象了,一堆爱因斯坦求和,而且还是tensorflow(谷歌的工作)这里我们可以跳过直接看下边的GQA(MQA就是GQA的特殊情况)
组查询注意力(Grouped-Query Attention,GQA)
动机
出自于2023年的EMNLP:GQA:Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoint
引言就把动机说的很清晰了

自回归解码器的推理是 Transformer 模型中的一个严重瓶颈,因为在每一个解码步骤中,都需要加载解码器的模型权重,以及注意力机制中所有历史的 Key 和 Value,这会带来巨大的内存带宽开销(memory bandwidth overhead)
MQA可以显著减少加载 Key 和 Value 所产生的内存带宽开销,但是:MQA会导致模型质量的下降和训练不稳定。
所以GQA就是综合了MHA和MQA的优缺点。论文中的图也很清晰:

所有头的Q共享一组K,V头的向量损失太严重,那么就给他分组,每几个Q共享一组K,V头。这样既没有MHA那样耗费内存,也没有MQA损失性能严重
实现
以下摘自modeling_llama.py的源码(删去了一些无关内容以及防御性代码)
首先要看一个函数:
def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
"""
This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from
(batch,num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
kv头的个数要复制到和q头的个数一样
"""
batch, num_key_value_heads, slen, head_dim = hidden_states.shape
if n_rep == 1:# 当n_rep表示复制几份,n_rep=q_head_num//k_v_head_num
return hidden_states
hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
#(Batch_size,num_kv_heads,seq_len,head_dim)--->(Batch_size,num_kv_heads,1,seq_len,head_dim)--->(Batch_size,num_kv_heads,n_rep,seq_len,head_dim)
return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
#(Batch_size,num_kv_heads,n_rep,seq_len,head_dim)-->(Batch_size,num_kv_heads*n_rep,seq_len,head_dim) 但是expand通过复制视图的方式扩展张量,而不会实际复制数据由于K和V的头数少于Q,在计算时需要将 KV "复制"到与 Q 匹配:但这里复制不能深拷贝,使用.expand进行浅拷贝
假设q_head_num=6,k_v_head_num=2,则需要给kv头向量复制3份,但又要他们指向同一地址,所以n_rep=6/2=3。
┌── K0
Q0 用 ──────┤
Q1 用 ──────┤── 同一块 K0 数据
Q2 用 ──────┘
┌── K1
Q3 用 ──────┤
Q4 用 ──────┤── 同一块 K1 数据
Q5 用 ──────┘下边是类的定义
class LlamaAttention(nn.Module):
"""Multi-headed attention from 'Attention Is All You Need' paper"""
def __init__(self, config: LlamaConfig, layer_idx: Optional[int] = None):
super().__init__()
self.config = config
self.attention_dropout = config.attention_dropout
self.hidden_size = config.hidden_size
self.num_heads = config.num_attention_heads
self.head_dim = self.hidden_size // self.num_heads
self.num_key_value_heads = config.num_key_value_heads
self.num_key_value_groups = self.num_heads // self.num_key_value_heads
self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=config.attention_bias)
self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=config.attention_bias)
self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=config.attention_bias)
self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=config.attention_bias)可以看到k、v的num_head和q的num_head是不一样的。
前向传播:
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
past_key_value: Optional[Cache] = None,
output_attentions: bool = False,
use_cache: bool = False,
**kwargs,
) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
bsz, q_len, _ = hidden_states.size()
query_states = self.q_proj(hidden_states)
key_states = self.k_proj(hidden_states)
value_states = self.v_proj(hidden_states)
query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)
key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)
value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2)
#初始化之后,q和kv的维度是不一样的
kv_seq_len = key_states.shape[-2]
# 这里通过repeat_kv复制函数
key_states = repeat_kv(key_states, self.num_key_value_groups)
value_states = repeat_kv(value_states, self.num_key_value_groups)
#计算注意力分数
attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim)
attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype)
attn_weights = nn.functional.dropout(attn_weights, p=self.attention_dropout, training=self.training)
attn_output = torch.matmul(attn_weights, value_states)
attn_output = attn_output.transpose(1, 2).contiguous()
attn_output = attn_output.reshape(bsz, q_len, self.hidden_size)
attn_output = self.o_proj(attn_output)
return attn_output了解了repeat_kv函数,后边的实现就和正常的attention计算差不多了。当num_key_value_heads=self.num_heads时,就是原始的MHA,当num_key_value_heads=1时,就是MQA。
一个有点蠢的问题
关于GQA和MQA 为啥都是k和v可以共享一份权重 为什么q不可以共享?这是我最开始直接看MQA和GQA时发出的疑问。首先。。从MQA和GQA的动机我们可以发现,他们是为了省下KV的内存开销,加快推理,所以q共享不共享的没啥区别。其次,如果真的q共享了,那就没什么意义了。多头注意力的意义最开始就是让q向量从不同角度去了解语意,但是现在q都共享了,那多头有什么意义呢,一个head复制N份,共享向量,那和一个头有啥区别。。。








