深入解析Transformer架构:从数学原理到工程实现

1次阅读
没有评论

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

image.webp

序列建模的痛点与 Transformer 的诞生

在自然语言处理领域,传统 RNN 和 LSTM 面临着两个核心挑战:

深入解析 Transformer 架构:从数学原理到工程实现

  1. 长距离依赖问题 :随着序列长度增加,RNN 难以有效捕捉远距离单词间的关系,梯度消失 / 爆炸现象频发。LSTM 通过门控机制缓解了该问题,但实验表明其在超过 100 个 token 的序列上表现仍会显著下降
  2. 并行计算限制 :RNN 的时序依赖性导致必须按顺序计算,无法充分利用 GPU 的并行计算能力。即便 LSTM 的单个 cell 计算仅需 $O(1)$ 时间,整个序列仍需 $O(n)$ 时间步完成计算

Transformer 通过完全基于注意力机制的架构解决了上述问题:

  • 全局依赖性 :自注意力层使任意两个 token 都能直接建立联系,理论最大路径长度仅为 $O(1)$
  • 并行计算 :所有位置的 attention 计算可同时进行,训练速度比 LSTM 快 5 -10 倍

Self-Attention 机制数学解析

核心计算公式

$$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$$

其中:
– $Q \in \mathbb{R}^{n \times d_k}$ (Query)
– $K \in \mathbb{R}^{n \times d_k}$ (Key)
– $V \in \mathbb{R}^{n \times d_v}$ (Value)

计算流程分解

  1. 线性投影 :将输入 embedding $X$ 通过三个权重矩阵投影
    $$Q = XW_Q, \quad K = XW_K, \quad V = XW_V$$

  2. 相似度计算 :通过点积衡量 query 与 key 的关联程度
    $$S = QK^T \in \mathbb{R}^{n \times n}$$

  3. 缩放与归一化 :防止点积结果过大导致 softmax 梯度消失
    $$S_{scaled} = \frac{S}{\sqrt{d_k}}$$

  4. 注意力权重 :通过 softmax 获得归一化的注意力分布
    $$A = \text{softmax}(S_{scaled})$$

  5. 加权求和 :根据注意力权重聚合 value 信息
    $$\text{Output} = AV$$

PyTorch 完整实现

import torch
import torch.nn as nn
import math

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model=512, n_heads=8):
        super().__init__()
        assert d_model % n_heads == 0
        self.d_k = d_model // n_heads
        self.n_heads = n_heads

        # 线性投影层
        self.W_q = nn.Linear(d_model, d_model)
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)
        self.W_o = nn.Linear(d_model, d_model)

    def forward(self, x, mask=None):
        """
        Args:
            x: [batch_size, seq_len, d_model]
            mask: [batch_size, seq_len, seq_len]
        """
        batch_size = x.size(0)

        # 1. 线性投影并分头
        Q = self.W_q(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1,2)
        K = self.W_k(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1,2)
        V = self.W_v(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1,2)

        # 2. 计算缩放点积注意力
        scores = torch.matmul(Q, K.transpose(-2,-1)) / math.sqrt(self.d_k)

        # 3. 应用 mask(解码器使用)if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e9)

        # 4. softmax 归一化
        attn_weights = torch.softmax(scores, dim=-1)

        # 5. 加权求和
        context = torch.matmul(attn_weights, V)

        # 6. 合并多头输出
        context = context.transpose(1,2).contiguous()\
                   .view(batch_size, -1, self.n_heads * self.d_k)

        return self.W_o(context)

计算复杂度分析

训练阶段

  • 自注意力层
    $$4nd^2 + 2n^2d$$
    (其中 $n$ 是序列长度,$d$ 是 embedding 维度)

  • 前馈网络
    $$8nd^2$$

推理阶段

对于自回归生成任务(如 GPT),需要缓存之前的 key 和 value:

  • 第 $t$ 步计算量
    $$4d^2 + (2t+2)d$$

生产环境优化策略

显存优化

  1. 梯度检查点

    from torch.utils.checkpoint import checkpoint
    
    def forward(self, x):
        return checkpoint(self._forward, x)

  2. 激活值压缩 :使用混合精度训练

    scaler = torch.cuda.amp.GradScaler()
    
    with torch.cuda.amp.autocast():
        outputs = model(inputs)
        loss = criterion(outputs, targets)
    
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

序列处理

  • 动态 padding:按 batch 内最大长度 padding
  • Bucket 策略 :将相似长度的样本分组处理

开放性问题

  1. 如何优化 attention 的 $O(n^2)$ 计算复杂度?(参考:稀疏注意力、局部窗口)
  2. 位置编码能否完全替代传统的位置信息建模?(参考:相对位置编码、旋转位置编码)
  3. 在多模态场景下如何统一设计 attention 机制?(参考:Cross-modality Attention)
正文完
 0
评论(没有评论)