Transformer 核心机制详解(QKV 与注意力机制)
Transformer 核心机制详解(QKV 与注意力机制)
以下是关于 Transformer 中 QKV、注意力机制、维度变化、不对称注意力以及 KV Cache 的完整笔记,内容通俗易懂,同时保留精确的数学描述和例子。
- QKV 的来源与维度变化(标准多头自注意力)
基本参数示例
batch_size = 2
seq_len = 10(序列长度)
embed_dim = 512(d_model)
num_heads = 8
head_dim = 512 // 8 = 64
输入 X 形状:[2, 10, 512]
生成 Q、K、V
三个独立的线性投影:
textQ = X @ W^Q + b^Q → [2, 10, 512]
K = X @ W^K + b^K → [2, 10, 512]
V = X @ W^V + b^V → [2, 10, 512]
分成多头
PythonQ = Q.view(batch_size, seq_len, num_heads, head_dim).transpose(1, 2)
结果形状:[2, 8, 10, 64]
K、V 同理。
核心注意力公式(Scaled Dot-Product Attention)
每个头独立计算:
textscores = Q @ K.transpose(-2, -1) → [2, 8, 10, 10]
scaled_scores = scores / √head_dim
attn_weights = softmax(scaled_scores, dim=-1)
output_per_head = attn_weights @ V → [2, 8, 10, 64]
多头拼接 + 最终投影
拼接所有头:[2, 10, 512]
通过 W^O 线性层 → 最终输出仍为 [2, 10, 512]
- 维度变化完整表格
步骤张量形状说明输入X[2, 10, 512]原始 token 嵌入线性投影Q, K, V[2, 10, 512]三个独立投影分头Q, K, V[2, 8, 10, 64]准备多头计算Q @ K^Tscores[2, 8, 10, 10]注意力原始分数缩放 + softmaxattn_weights[2, 8, 10, 10]归一化权重attn_weights @ Vhead_output[2, 8, 10, 64]每个头输出转置 + 拼接concat[2, 10, 512]合并所有头最终线性投影 W^Ooutput[2, 10, 512]Multi-Head Attention 最终输出
3. 最灵活的核心:Q 和 KV 长度可以不一样(不对称注意力)
关键结论:
输出维度永远跟 Q 的序列长度和特征维度 一致
Q 的长度 和 K/V 的长度 可以完全不同
K 和 V 的序列长度必须相同(因为它们是配对的键-值)
示例(完全正确)
Q:100 × 2048
K^T:2048 × 500
V:500 × 2048
计算过程:
textscores = Q @ K^T → 100 × 500
attn_weights = softmax(scores / √d) → 100 × 500
output = attn_weights @ V → 100 × 2048
结果:100 个查询,每个查询从 500 个键值对中加权汇总信息
两种极端用法
用法Q 长度KV 长度典型场景更多 Q 查询更少 KV多少编码器自注意力(对称)更少 Q 查询更多 KV少多检索、记忆、池化(如 Pons Adapter)自回归生成(GPT 式)1越来越长每步只算新 token 的 Q,用历史所有 KV
4. 如何在代码中提取 KV(Hugging Face 示例)
Pythonoutputs = model(**inputs, use_cache=True)
past_key_values = outputs.past_key_values # 就是 KV cache
结构:tuple of length = num_layers
每层:(key, value),形状 [batch, num_heads, seq_len, head_dim]
key_layer0, value_layer0 = past_key_values[0]
继续生成时直接传回:
Pythonoutputs2 = model(next_token, past_key_values=past_key_values, use_cache=True)
5. 通俗比喻总结
Q:你想问的问题(查询向量)
K:别人能回答什么问题的“标签”(键)
V:别人真实提供的信息内容(值)
注意力 = 每个查询去找最匹配的键,然后取出对应的值加权平均
多头 = 从多个不同角度(主题)同时看匹配关系
不对称长度 = 你可以问1个问题查1000条记忆,也可以同时问100个问题只查10条记忆
这就是 Transformer 注意力机制的全部数学与实现精髓!
更多推荐
所有评论(0)