共计 2662 个字符,预计需要花费 7 分钟才能阅读完成。
背景痛点
许多开发者通过 3blue1brown 的《Transformer 视觉解说》视频学习 Transformer 架构时,虽然对自注意力机制(Self-Attention)的几何直观有了理解,但在实际编码中仍面临三大障碍:

- 维度变换抽象 :视频中高维空间的投影动画难以对应到代码中的张量操作(如
(batch_size, seq_len, d_model)的矩阵变换) - 矩阵计算细节缺失 :QKV(Query-Key-Value)矩阵的拆分、缩放点积(Scaled Dot-Product)的具体实现未展开
- 模块组合断层 :位置编码(Positional Encoding)、层归一化(LayerNorm)等模块如何与自注意力层协同工作缺乏示例
技术对比:论文 vs 视频
| 维度 | 原始论文《Attention is All You Need》 | 3blue1brown 视频解说 |
|---|---|---|
| 核心视角 | 数学公式与算法描述 | 高维空间几何可视化 |
| 注意力机制 | 矩阵乘法的分步推导 | 向量投影与 ” 信息检索 ” 类比 |
| 位置编码 | 直接给出三角函数公式 | 用螺旋线运动解释位置信息的注入 |
| 学习曲线 | 需线性代数基础 | 依赖空间想象力 |
核心实现
1. 位置编码(Positional Encoding)
视频中将位置编码比作 ” 给词向量添加时间维度 ” 的螺旋运动,代码需实现以下关键点:
import torch
import math
def positional_encoding(max_len: int, d_model: int) -> torch.Tensor:
"""
生成位置编码矩阵(可视化参考:https://projector.tensorflow.org/)Args:
max_len: 序列最大长度
d_model: 词向量维度
Returns:
pe: (max_len, d_model) 的位置编码矩阵
"""
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term) # 偶数维用 sin
pe[:, 1::2] = torch.cos(position * div_term) # 奇数维用 cos
return pe
关键细节 :
- 公式中的
10000.0控制波长范围,值越大则位置编码变化越平缓 - 交替使用 sin/cos 保证位置编码可被线性组合(参考视频中 ” 相对位置 ” 的解释)
2. 自注意力层实现
对应视频中 ” 通过 QKV 矩阵计算关联度 ” 的部分:
class SelfAttention(torch.nn.Module):
def __init__(self, d_model: int, head_size: int):
super().__init__()
self.query = torch.nn.Linear(d_model, head_size)
self.key = torch.nn.Linear(d_model, head_size)
self.value = torch.nn.Linear(d_model, head_size)
self.scale = head_size ** -0.5 # 缩放因子对应视频中的 "稳定梯度" 解释
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x 形状: (batch_size, seq_len, d_model)
Q = self.query(x) # (batch_size, seq_len, head_size)
K = self.key(x) # 形状同 Q
V = self.value(x) # 形状同 Q
attn_scores = torch.matmul(Q, K.transpose(-2, -1)) * self.scale
attn_weights = torch.softmax(attn_scores, dim=-1) # 视频中的 "权重分配"
out = torch.matmul(attn_weights, V)
assert out.shape == (x.shape[0], x.shape[1], self.head_size)
return out
3. 残差连接与 LayerNorm
视频未明确提及但至关重要的稳定训练技巧:
class TransformerBlock(torch.nn.Module):
def __init__(self, d_model: int, head_size: int):
super().__init__()
self.attention = SelfAttention(d_model, head_size)
self.norm1 = torch.nn.LayerNorm(d_model)
self.mlp = torch.nn.Sequential(torch.nn.Linear(d_model, 4 * d_model), # 扩展维度
torch.nn.GELU(),
torch.nn.Linear(4 * d_model, d_model) # 压缩回原维度
)
self.norm2 = torch.nn.LayerNorm(d_model)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# 残差连接 1(视频中未展示但实际必需)x = x + self.attention(self.norm1(x))
# 残差连接 2
x = x + self.mlp(self.norm2(x))
return x
避坑指南
- 注意力权重未做 softmax:
- 错误表现:直接使用
Q @ K.T的结果作为权重 -
修正方案:必须用
torch.softmax(..., dim=-1)归一化 -
位置编码错误相加 :
- 错误代码:
word_embedding + positional_encoding(可能导致信息覆盖) -
推荐方案:先通过线性层扩展维度再拼接
-
梯度爆炸问题 :
- 现象:训练中出现 NaN 值
- 解决方案:添加梯度裁剪
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
延伸思考
如何修改当前模型使其支持图像 patch 输入? 提示线索:
1. 将图像切分为 16×16 的 patch,每个 patch 视为一个 ” 词 ”
2. 用 CNN 或线性层将 patch 像素展平为向量
3. 位置编码需改为 2D 版本(参考 Vision Transformer 论文)
正文完
发表至: 未分类
近一天内
