标准多头注意力机制维度变换解析
深入理解标准多头注意力机制张量维度变换
import torch
import torch.nn as nn
# 设置参数
batch_size = 2
num_tokens = 3
d_in = 4
d_out = 8
num_heads = 2
head_dim = d_out // num_heads # 4
print("="*70)
print("多头注意力机制 - 逐步演示")
print("="*70)
# 创建输入
torch.manual_seed(42)
x = torch.randn(batch_size, num_tokens, d_in)
print(f"\n1. 输入 x 形状: {x.shape}") # (2, 3, 4)
# 第1步:线性投影
W_key = nn.Linear(d_in, d_out, bias=False)
W_query = nn.Linear(d_in, d_out, bias=False)
W_value = nn.Linear(d_in, d_out, bias=False)
keys = W_key(x)
queries = W_query(x)
values = W_value(x)
print(f"\n2. 线性投影后 keys 形状: {keys.shape}") # (2, 3, 8)
# 第2步:reshape 为多头
keys_reshaped = keys.view(batch_size, num_tokens, num_heads, head_dim)
queries_reshaped = queries.view(batch_size, num_tokens, num_heads, head_dim)
values_reshaped = values.view(batch_size, num_tokens, num_heads, head_dim)
print(f"\n3. Reshape 为多头 keys 形状: {keys_reshaped.shape}") # (2, 3, 2, 4)
print(" 含义: (batch, tokens, heads, head_dim)")
# 验证数据没有变
print(f"\n 验证:reshape 前后数据是否相同?")
print(f" keys[0, 0, :] = {keys[0, 0, :]}")
print(f" keys_reshaped[0, 0, :, :] 展平 = {keys_reshaped[0, 0, :, :].flatten()}")
print(f" 是否相同: {torch.allclose(keys[0, 0, :], keys_reshaped[0, 0, :, :].flatten())}")
# True ✅
# 第3步:转置
keys_transposed = keys_reshaped.transpose(1, 2)
queries_transposed = queries_reshaped.transpose(1, 2)
values_transposed = values_reshaped.transpose(1, 2)
print(f"\n4. 转置后 keys 形状: {keys_transposed.shape}") # (2, 2, 3, 4)
print(" 含义: (batch, heads, tokens, head_dim)")
# 第4步:计算注意力分数
attn_scores = queries_transposed @ keys_transposed.transpose(2, 3)
print(f"\n5. 注意力分数形状: {attn_scores.shape}") # (2, 2, 3, 3)
print(" 含义: (batch, heads, query_tokens, key_tokens)")
# 查看具体的注意力分数
print(f"\n Batch 0, Head 0 的注意力分数:")
print(attn_scores[0, 0])
# tensor([[ 0.8503, -0.3421, 1.2341],
# [-0.2341, 0.6789, -0.1234],
# [ 1.1234, -0.4567, 0.7890]])
# 第5步:应用因果掩码
mask = torch.triu(torch.ones(num_tokens, num_tokens), diagonal=1).bool()
print(f"\n6. 因果掩码:")
print(mask.int())
# tensor([[0, 1, 1],
# [0, 0, 1],
# [0, 0, 0]])
attn_scores_masked = attn_scores.clone()
attn_scores_masked = attn_scores_masked.masked_fill(mask, float('-inf'))
print(f"\n 应用掩码后 Batch 0, Head 0:")
print(attn_scores_masked[0, 0])
# tensor([[ 0.8503, -inf, -inf],
# [-0.2341, 0.6789, -inf],
# [ 1.1234, -0.4567, 0.7890]])
# 第6步:Softmax
scale = head_dim ** 0.5 # √4 = 2
attn_weights = torch.softmax(attn_scores_masked / scale, dim=-1)
print(f"\n7. Softmax 后注意力权重 Batch 0, Head 0:")
print(attn_weights[0, 0])
# tensor([[1.0000, 0.0000, 0.0000], ← 只关注自己
# [0.3527, 0.6473, 0.0000], ← 关注 token 0 和 1
# [0.3245, 0.2134, 0.4621]]) ← 关注所有 token
# 验证每行和为 1
print(f"\n 每行的和: {attn_weights[0, 0].sum(dim=-1)}")
# tensor([1.0000, 1.0000, 1.0000]) ✅
# 第7步:加权求和
context_vec = attn_weights @ values_transposed
print(f"\n8. 加权求和后 context_vec 形状: {context_vec.shape}") # (2, 2, 3, 4)
# 第8步:转置回原顺序并合并头
context_vec = context_vec.transpose(1, 2) # (2, 3, 2, 4)
print(f" 转置后形状: {context_vec.shape}")
context_vec = context_vec.reshape(batch_size, num_tokens, d_out) # (2, 3, 8)
print(f" 合并头后形状: {context_vec.shape}")
print("\n" + "="*70)
print("完成!多头注意力机制的所有步骤演示完毕")
print("="*70)
第1步:线性投影生成 Q、K、V
输入
x: (2, 3, 4) # batch=2, tokens=3, d_in=4
经过线性层
keys = W_key(x) # (2, 3, 8) ← d_out=8
queries = W_query(x) # (2, 3, 8)
values = W_value(x) # (2, 3, 8)
图示(以 keys 为例,只看 batch 0):
keys[0] 形状: (3, 8)
┌ ┐
│ k0_1 k0_2 k0_3 k0_4 k0_5 k0_6 k0_7 k0_8 │ ← token 0
│ k1_1 k1_2 k1_3 k1_4 k1_5 k1_6 k1_7 k1_8 │ ← token 1
│ k2_1 k2_2 k2_3 k2_4 k2_5 k2_6 k2_7 k2_8 │ ← token 2
└ ┘
3 tokens × 8 维
这 8 维包含了 2 个头的信息(每个头 4 维)
第2步:隐式分割为多头(view/reshape)
keys = keys.view(b, num_tokens, num_heads, head_dim)
```python
# (2, 3, 8) -> (2, 3, 2, 4)
原始形状: (2, 3, 8)
↓ view(2, 3, 2, 4)
新形状: (2, 3, 2, 4)
↑ ↑ ↑ ↑
b t h d
含义:
```python
- b=2: 2 个 batch
- t=3: 3 个 token
- h=2: 2 个 attention head
- d=4: 每个 head 的维度是 4
原来: keys[0] 形状 (3, 8)
┌ ┐
│ [k0_1 k0_2 k0_3 k0_4] [k0_5 k0_6 k0_7 k0_8] │ ← token 0
│ [k1_1 k1_2 k1_3 k1_4] [k1_5 k1_6 k1_7 k1_8] │ ← token 1
│ [k2_1 k2_2 k2_3 k2_4] [k2_5 k2_6 k2_7 k2_8] │ ← token 2
└ ┘
←-- head 0 ---→ ←-- head 1 ---→
重塑后: keys[0] 形状 (3, 2, 4)
┌ ┐
│ token 0: │
│ head 0: [k0_1 k0_2 k0_3 k0_4]
│ head 1: [k0_5 k0_6 k0_7 k0_8]
│ │
│ token 1: │
│ head 0: [k1_1 k1_2 k1_3 k1_4]
│ head 1: [k1_5 k1_6 k1_7 k1_8]
│ │
│ token 2: │
│ head 0: [k2_1 k2_2 k2_3 k2_4]
│ head 1: [k2_5 k2_6 k2_7 k2_8]
└ ┘
关键:数据没有变,只是重新解释了维度!
内存布局完全相同,只是"看待"数据的方式变了
为什么叫"隐式分割"?
我们没有实际切割数据成两个独立的张量
而是通过改变视图,让 PyTorch "认为"数据已经分成了多个头
就像一本书:
- 原来:看作 3 章,每章 8 页
- 现在:看作 3 章,每章 2 节,每节 4 页
书的内容没变,只是组织方式变了!
第3步:转置以适配矩阵乘法
keys = keys.transpose(1, 2)
# (2, 3, 2, 4) -> (2, 2, 3, 4)
# ↑ ↑ ↑ ↑
# dim1 dim2 交换这两个维度
为什么要转置?
转置前: (batch, tokens, heads, head_dim)
转置后: (batch, heads, tokens, head_dim)
目的:让 heads 维度靠近 batch,方便批量矩阵乘法
转置前 keys 形状: (2, 3, 2, 4)
(b, t, h, d)
Batch 0:
┌ ┐
│ Token 0: │
│ Head 0: [k00_1 ... k00_4]
│ Head 1: [k00_1 ... k00_4]
│ Token 1: │
│ Head 0: [k01_1 ... k01_4]
│ Head 1: [k01_1 ... k01_4]
│ Token 2: │
│ Head 0: [k02_1 ... k02_4]
│ Head 1: [k02_1 ... k02_4]
└ ┘
转置后 keys 形状: (2, 2, 3, 4)
(b, h, t, d)
Batch 0:
┌ ┐
│ Head 0: │ ← 现在按头组织
│ Token 0: [k00_1 ... k00_4]
│ Token 1: [k01_1 ... k01_4]
│ Token 2: [k02_1 ... k02_4]
│ │
│ Head 1: │
│ Token 0: [k00_1 ... k00_4]
│ Token 1: [k01_1 ... k01_4]
│ Token 2: [k02_1 ... k02_4]
└ ┘
优势:
- 每个 head 的数据连续存储
- 可以对 (batch*heads) 做批量矩阵乘法
- 形状可以看作 (4, 3, 4),即 4 个独立的 3×4 矩阵
第4步:计算缩放点积注意力
attn_scores = queries @ keys.transpose(2, 3)
queries 形状: (2, 2, 3, 4)
(b, h, t, d)
keys.transpose(2,3): (2, 2, 4, 3)
(b, h, d, t)
↑ ↑ 交换最后两维
矩阵乘法:
(2, 2, 3, 4) @ (2, 2, 4, 3)
↑ ↑
都是 4,匹配!✅
结果 attn_scores: (2, 2, 3, 3)
(b, h, t_q, t_k)
详细计算过程(简化:只看 batch 0, head 0)
queries[0, 0] 形状: (3, 4) ← batch 0, head 0
┌ ┐
│ q00_1 q00_2 q00_3 q00_4 │ ← token 0 的 query
│ q01_1 q01_2 q01_3 q01_4 │ ← token 1 的 query
│ q02_1 q02_2 q02_3 q02_4 │ ← token 2 的 query
└ ┘
keys[0, 0] 形状: (3, 4) ← batch 0, head 0
┌ ┐
│ k00_1 k00_2 k00_3 k00_4 │ ← token 0 的 key
│ k01_1 k01_2 k01_3 k01_4 │ ← token 1 的 key
│ k02_1 k02_2 k02_3 k02_4 │ ← token 2 的 key
└ ┘
keys[0, 0].transpose(0, 1) 形状: (4, 3)
┌ ┐
│ k00_1 k01_1 k02_1 │
│ k00_2 k01_2 k02_2 │
│ k00_3 k01_3 k02_3 │
│ k00_4 k01_4 k02_4 │
└ ┘
计算 attn_scores[0, 0] = queries[0, 0] @ keys[0, 0].T
形状: (3, 4) @ (4, 3) = (3, 3)
┌ ┐ ┌ ┐ ┌ ┐
│ q00_1 q00_2 q00_3 q00_4 │ │ k00_1 k01_1..│ │ s00 s01 s02 │
│ q01_1 q01_2 q01_3 q01_4 │ @ │ k00_2 k01_2..│ = │ s10 s11 s12 │
│ q02_1 q02_2 q02_3 q02_4 │ │ k00_3 k01_3..│ │ s20 s21 s22 │
└ ┘ │ k00_4 k01_4..│ └ ┘
└ ┘
其中:
s00 = q00·k00 = q00_1*k00_1 + q00_2*k00_2 + q00_3*k00_3 + q00_4*k00_4
s01 = q00·k01 = q00_1*k01_1 + q00_2*k01_2 + q00_3*k01_3 + q00_4*k01_4
s02 = q00·k02 = q00_1*k02_1 + q00_2*k02_2 + q00_3*k02_3 + q00_4*k02_4
含义:
s00 = token 0 对 token 0 的注意力分数
s01 = token 0 对 token 1 的注意力分数
s02 = token 0 对 token 2 的注意力分数
完整的注意力分数矩阵
attn_scores 形状: (2, 2, 3, 3)
(batch, head, query_token, key_token)
Batch 0, Head 0:
┌ ┐
│ s00 s01 s02 │ ← token 0 对所有 token 的分数
│ s10 s11 s12 │ ← token 1 对所有 token 的分数
│ s20 s21 s22 │ ← token 2 对所有 token 的分数
└ ┘
Batch 0, Head 1:
┌ ┐
│ s'00 s'01 s'02 │ ← head 1 的分数(不同的子空间)
│ s'10 s'11 s'12 │
│ s'20 s'21 s'22 │
└ ┘
Batch 1, Head 0:
┌ ┐
│ s''00 ... │
│ ... │
└ ┘
Batch 1, Head 1:
┌ ┐
│ s'''00 ... │
│ ... │
└ ┘
总共:2 batches × 2 heads = 4 个独立的 3×3 注意力矩阵
第5步:应用因果掩码
mask_bool = self.mask.bool()[:num_tokens, :num_tokens]
attn_scores.masked_fill_(mask_bool, -torch.inf)
因果掩码的作用
假设 context_length = 3
mask = torch.triu(torch.ones(3, 3), diagonal=1)
mask:
┌ ┐
│ 0 1 1 │ ← diagonal=1,上三角(不含对角线)
│ 0 0 1 │
│ 0 0 0 │
└ ┘
转换为布尔:
┌ ┐
│ F T T │ ← True 表示需要遮蔽的位置
│ F F T │
│ F F F │
└ ┘
物理意义:
- token 0 只能看到 token 0(自己)
- token 1 可以看到 token 0, 1
- token 2 可以看到 token 0, 1, 2
这就是"因果":只能看过去和现在,不能看未来!
应用掩码的过程
原始 attn_scores[0, 0]:
┌ ┐
│ 2.1 1.5 0.8 │
│ 0.9 1.8 1.2 │
│ 0.5 0.7 1.5 │
└ ┘
应用 mask(将 True 位置填为 -inf):
┌ ┐
│ 2.1 -inf -inf │ ← token 0 只能关注自己
│ 0.9 1.8 -inf │ ← token 1 可以关注 0,1
│ 0.5 0.7 1.5 │ ← token 2 可以关注 0,1,2
└ ┘
然后 softmax:
softmax([2.1, -inf, -inf]) = [1.0, 0.0, 0.0] ✅
softmax([0.9, 1.8, -inf]) = [0.29, 0.71, 0.0] ✅
softmax([0.5, 0.7, 1.5]) = [0.18, 0.22, 0.60] ✅
每行的和都是 1.0,且未来位置的权重为 0
更多推荐



所有评论(0)