深入理解标准多头注意力机制张量维度变换

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,10.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

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐