对应课程: course_A W3 (注意力机制)
难度: ★★★ (L3 进阶)
知识基础: Build LLM from Scratch Ch3 — 编码注意力机制
> 理解模型"在看什么" — 审计模型可解释性的核心技能
每个 token 通过点积计算与其他所有 token 的相关性。
import torch
import torch.nn as nn
# 模拟输入: 3个token, 每个6维
inputs = torch.tensor([
[0.43, 0.15, 0.89, 0.34, 0.66, 0.12], # Token A
[0.55, 0.87, 0.66, 0.55, 0.12, 0.44], # Token B
[0.78, 0.23, 0.11, 0.94, 0.33, 0.22], # Token C
])
# 计算注意力分数: 每个token与所有token的点积
attn_scores = inputs @ inputs.T # (3,3)
print("注意力分数:")
print(attn_scores)
# softmax 归一化为权重 (和为1)
attn_weights = torch.softmax(attn_scores, dim=-1)
print("\n注意力权重:")
print(attn_weights)
# 加权求和: 每个token的新表示 = 所有token的加权平均
context_vec = attn_weights @ inputs
print(f"\n上下文向量形状: {context_vec.shape}")
print("Token A 的上下文向量:", context_vec[0])核心直觉: Token A 的最终表示 = 它自己 + 与它相关的其他token的信息。
权重越高 = 相关性越强。
对应书籍 3.4节: 引入 $W_q$, $W_k$, $W_v$ 三个权重矩阵。
Query 问: "我该关注谁?"
Key 答: "我有这些信息"
Value 给: "把我的信息给你"
class SelfAttentionV1(nn.Module):
"""带可训练权重的自注意力 (原始版)"""
def __init__(self, d_in, d_out):
super().__init__()
self.W_q = nn.Parameter(torch.randn(d_in, d_out))
self.W_k = nn.Parameter(torch.randn(d_in, d_out))
self.W_v = nn.Parameter(torch.randn(d_in, d_out))
def forward(self, x):
# x: (batch, num_tokens, d_in)
queries = x @ self.W_q
keys = x @ self.W_k
values = x @ self.W_v
attn_scores = queries @ keys.T # (num_tokens, num_tokens)
attn_weights = torch.softmax(attn_scores / keys.shape[-1]**0.5, dim=-1)
return attn_weights @ values
class SelfAttentionV2(nn.Module):
"""使用 nn.Linear 的更规范版本"""
def __init__(self, d_in, d_out, qkv_bias=False):
super().__init__()
self.W_q = nn.Linear(d_in, d_out, bias=qkv_bias)
self.W_k = nn.Linear(d_in, d_out, bias=qkv_bias)
self.W_v = nn.Linear(d_in, d_out, bias=qkv_bias)
def forward(self, x):
queries = self.W_q(x)
keys = self.W_k(x)
values = self.W_v(x)
attn_scores = queries @ keys.transpose(-2, -1)
attn_weights = torch.softmax(attn_scores / keys.shape[-1]**0.5, dim=-1)
return attn_weights @ values
# 测试
torch.manual_seed(42)
sa_v2 = SelfAttentionV2(d_in=6, d_out=6)
output = sa_v2(inputs.unsqueeze(0)) # (1, 3, 6)
print(f"输入: {inputs.shape} → 输出: {output.shape}")
print("\n缩放点积注意力 (Scaled Dot-Product Attention):")
print("attn = softmax(Q @ K^T / sqrt(d_k)) @ V")为什么除以 $\sqrt{d_k}$?
当维度 d_k 很大时,点积的值会变得很大,softmax 的梯度会消失。
除以 $\sqrt{d_k}$ 把方差控制回 1,保持梯度稳定。
对应书籍 3.4节末尾: GPT 生成时只能看前面的token,不能看后面的。
方法: 加一个上三角掩码矩阵,把未来位置的分数设为 $-\infty$
class CausalAttention(nn.Module):
"""因果注意力: 每个token只能关注自己及之前的token"""
def __init__(self, d_in, d_out, context_length, dropout=0.0):
super().__init__()
self.W_q = nn.Linear(d_in, d_out, bias=False)
self.W_k = nn.Linear(d_in, d_out, bias=False)
self.W_v = nn.Linear(d_in, d_out, bias=False)
self.dropout = nn.Dropout(dropout)
# 因果掩码: 上三角矩阵 (对角线以上为 -inf)
mask = torch.triu(torch.ones(context_length, context_length), diagonal=1)
self.register_buffer('mask', mask.bool())
def forward(self, x):
b, num_tokens, d_in = x.shape
queries = self.W_q(x) # (b, n, d_out)
keys = self.W_k(x) # (b, n, d_out)
values = self.W_v(x) # (b, n, d_out)
attn_scores = queries @ keys.transpose(1, 2) # (b, n, n)
# 应用因果掩码
attn_scores.masked_fill_(
self.mask.bool()[:num_tokens, :num_tokens], -torch.inf
)
attn_weights = torch.softmax(attn_scores / keys.shape[-1]**0.5, dim=-1)
attn_weights = self.dropout(attn_weights)
return attn_weights @ values
# 演示因果掩码
context_len = 5
mask = torch.triu(torch.ones(context_len, context_len), diagonal=1)
print("因果掩码 (1=被屏蔽的位置):")
print(mask)
print("\nToken 0 能看到: [0]");
print("Token 1 能看到: [0,1]");
print("Token 2 能看到: [0,1,2]");
print("...")# 测试因果注意力
torch.manual_seed(123)
ca = CausalAttention(d_in=6, d_out=6, context_length=3, dropout=0.0)
with torch.no_grad():
out = ca(inputs.unsqueeze(0))
print(f"因果注意力输出形状: {out.shape}")
print(f"Token 0 输出: {out[0, 0]}")对应书籍 3.5节: 并行多个注意力头,每个头关注不同的关系。
比如: 一个头关注语法关系,另一个头关注语义关系。
class MultiHeadAttention(nn.Module):
"""多头注意力: n_heads 个并行的因果注意力"""
def __init__(self, d_in, d_out, context_length, dropout=0.0, n_heads=2):
super().__init__()
assert d_out % n_heads == 0, "d_out 必须能被 n_heads 整除"
self.d_out = d_out
self.n_heads = n_heads
self.head_dim = d_out // n_heads # 每个头的维度
# 单一大矩阵比 n_heads 个小矩阵更高效
self.W_q = nn.Linear(d_in, d_out, bias=False)
self.W_k = nn.Linear(d_in, d_out, bias=False)
self.W_v = nn.Linear(d_in, d_out, bias=False)
self.out_proj = nn.Linear(d_out, d_out) # 输出投影
self.dropout = nn.Dropout(dropout)
mask = torch.triu(torch.ones(context_length, context_length), diagonal=1)
self.register_buffer('mask', mask.bool())
def forward(self, x):
b, num_tokens, d_in = x.shape
# 1. 线性投影 + 拆分为多头
queries = self.W_q(x).view(b, num_tokens, self.n_heads, self.head_dim)
keys = self.W_k(x).view(b, num_tokens, self.n_heads, self.head_dim)
values = self.W_v(x).view(b, num_tokens, self.n_heads, self.head_dim)
# 转置: (b, n_heads, num_tokens, head_dim)
queries = queries.transpose(1, 2)
keys = keys.transpose(1, 2)
values = values.transpose(1, 2)
# 2. 每个头独立计算注意力
attn_scores = queries @ keys.transpose(-2, -1)
# 应用因果掩码 (对每个头)
mask = self.mask.bool()[:num_tokens, :num_tokens].unsqueeze(0).unsqueeze(0)
attn_scores.masked_fill_(mask, -torch.inf)
attn_weights = torch.softmax(attn_scores / self.head_dim**0.5, dim=-1)
attn_weights = self.dropout(attn_weights)
# 3. 加权求和 + 合并多头
context_vec = (attn_weights @ values).transpose(1, 2)
context_vec = context_vec.contiguous().view(b, num_tokens, self.d_out)
context_vec = self.out_proj(context_vec)
return context_vec
# 测试
mha = MultiHeadAttention(d_in=6, d_out=6, context_length=3, n_heads=2)
out = mha(inputs.unsqueeze(0))
print(f"多头注意力输出形状: {out.shape}")
print(f"(batch=1, tokens=3, d_out=6)")多头注意力 vs 单头:
让模型生成一段文本,然后分析它"关注"了什么。
# %%ai openai-chat-custom:Qwen3.5-9B-Q4_K_M.gguf --format markdown
# 用浅显的类比解释: 为什么Transformer需要"缩放点积注意力"而不是直接用点积?
# 为什么除以 sqrt(d_k) 对训练稳定这么重要?# 对比实验: 缩放 vs 不缩放
import math
dims = [4, 16, 64, 256, 1024]
for d in dims:
q = torch.randn(1, d)
k = torch.randn(1, d)
score = (q @ k.T).item()
scaled = (q @ k.T / math.sqrt(d)).item()
print(f"d={d:5d} | 原始点积: {score:.2f} | 缩放后: {scaled:.2f}")07_GPT_Architecture.ipynb — 把注意力装进完整的 GPT 模型> L4 映射: 注意力机制帮助你理解模型审计的"可解释性"维度—你能回答"模型为什么关注这里"。
注意力机制本身是"公平的" — 它根据数据学习关注模式。但如果训练数据中某些群体出现频率远高于其他群体,模型会学会"不关注"少数群体。
审计要点: 检查注意力权重分布 — 如果某个群体(token/特征)在所有样本中都被忽略,说明训练数据有偏差
中国算法备案要求提供"模型可解释性说明"。注意力权重是满足这一要求的关键证据:
因果注意力掩码只关注上文 — 这看起来是"安全设计",但如果提示词注入攻击者在输入中藏了恶意指令,注意力机制会忠实地关注它。
检查方法: 用对抗性输入测试 — 注入"忽略之前指令"看模型是否会关注
生产环境中,建议记录每次推理的注意力分布摘要: token级注意力方差、最大注意力值。这些数据在出现争议时可用于追溯。