从3blue1brown《Transformer视觉解说》入门:手把手实现最简Transformer模型

1次阅读
没有评论

共计 2662 个字符,预计需要花费 7 分钟才能阅读完成。

image.webp

背景痛点

许多开发者通过 3blue1brown 的《Transformer 视觉解说》视频学习 Transformer 架构时,虽然对自注意力机制(Self-Attention)的几何直观有了理解,但在实际编码中仍面临三大障碍:

从 3blue1brown《Transformer 视觉解说》入门:手把手实现最简 Transformer 模型

  1. 维度变换抽象 :视频中高维空间的投影动画难以对应到代码中的张量操作(如 (batch_size, seq_len, d_model) 的矩阵变换)
  2. 矩阵计算细节缺失 :QKV(Query-Key-Value)矩阵的拆分、缩放点积(Scaled Dot-Product)的具体实现未展开
  3. 模块组合断层 :位置编码(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

避坑指南

  1. 注意力权重未做 softmax
  2. 错误表现:直接使用 Q @ K.T 的结果作为权重
  3. 修正方案:必须用 torch.softmax(..., dim=-1) 归一化

  4. 位置编码错误相加

  5. 错误代码:word_embedding + positional_encoding(可能导致信息覆盖)
  6. 推荐方案:先通过线性层扩展维度再拼接

  7. 梯度爆炸问题

  8. 现象:训练中出现 NaN 值
  9. 解决方案:添加梯度裁剪
    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 论文)

正文完
 0
评论(没有评论)