我们是由枫哥组建的IT技术团队,成立于2017年,致力于帮助IT从业者提供实力,成功入职理想企业,我们提供一对一学习辅导,由知名大厂导师指导,分享Java技术、参与项目实战等服务,并为学员定制职业规划,全面提升竞争力,过去8年,我们已成功帮助数千名求职者拿到满意的Offer:IT枫斗者IT枫斗者-Java面试突击


一、开篇前言

2017 年,Google 研究团队发表里程碑式论文 《Attention Is All You Need》,正式提出 Transformer 架构。

该架构彻底摒弃了传统 RNN/LSTM 的递归时序结构,完全依托自注意力机制(Self-Attention)捕获文本长距离语义依赖,一举解决了递归模型长序列梯度消失、无法并行计算的痛点。

时至今日,所有主流大模型均基于 Transformer 衍生而来:

  • BERT:仅使用 Transformer Encoder
  • GPT、LLaMA、Qwen:仅使用 Transformer Decoder
  • T5、BART:完整 Encoder-Decoder 架构

本文将通过350行纯PyTorch极简可运行代码,从零逐层拆解 Transformer 所有核心组件,公式+代码+原理三重解析,无冗余封装,适合新手入门、原理调试、二次开发学习。

二、Transformer 核心组件总览

Transformer 所有能力均由七大基础组件构成,各组件功能、时间复杂度一目了然:

核心组件 核心功能 时间复杂度
Scaled Dot-Product Attention Q/K/V 相似度计算与加权聚合 O(n·dₖ)
Multi-Head Attention 多表示空间并行注意力,捕捉多维语义 O(n·d_model)
Position-wise FFN 单位置非线性特征变换,提升模型容量 O(d_model·d_ff)
Positional Encoding 为序列注入位置时序信息 O(max_len·d_model)
Layer Norm 特征维度规范化,稳定训练 O(d_model)
Encoder Layer 自注意力编码全局语义,双向理解 O(n²·d_model)
Decoder Layer 掩码自注意力+交叉注意力,序列生成 O(n²·d_model)

三、核心组件从零实现与深度解析

1. Scaled Dot-Product Attention(缩放点积注意力)

1.1 核心公式与原理

缩放点积注意力是 Transformer 的最小核心计算单元,核心逻辑:通过 Query(查询)与 Key(键)的相似度,对 Value(值)进行加权聚合

核心公式:

A t t e n t i o n ( Q , K , V ) = s o f t m a x ( Q K ⊤ d k ) V Attention(Q,K,V)=softmax(\frac{QK^\top}{\sqrt{d_k}})V Attention(Q,K,V)=softmax(dk QK)V

1.2 为什么需要除以 d k \sqrt{d_k} dk

当特征维度 d k d_k dk 较大时,Q、K 点积结果的方差会随维度增大而急剧升高,数值会变得极大。此时 softmax 函数会落入梯度饱和区,梯度趋近于0,导致训练停滞。

除以 d k \sqrt{d_k} dk 可将点积结果方差稳定在1左右,保证 softmax 梯度有效,让训练过程更稳定。

1.3 完整代码实现
import math
import torch
import torch.nn as nn
import torch.nn.functional as F

class ScaledDotProductAttention(nn.Module):
    """
    缩放点积注意力
    Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V
    """
    def __init__(self, dropout: float = 0.1):
        super().__init__()
        self.dropout = nn.Dropout(dropout)

    def forward(self, Q, K, V, mask=None):
        # 获取特征维度
        d_k = Q.size(-1)
        # 计算注意力分数并缩放
        scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)

        # 掩码处理:padding位置置为-∞,softmax后权重为0
        if mask is not None:
            scores = scores.masked_fill(mask == 0, float('-inf'))

        # 计算注意力权重并dropout正则化
        attn_weights = F.softmax(scores, dim=-1)
        attn_weights = self.dropout(attn_weights)
        # 加权聚合Value特征
        output = torch.matmul(attn_weights, V)
        
        return output, attn_weights
1.4 核心关键点总结
  • masked_fill:将padding无效位置填充为负无穷,经过softmax后权重趋近于0,彻底屏蔽无效字符干扰
  • 注意力后dropout:经典正则化方案,有效防止注意力过拟合
  • 返回attn_weights:用于训练可视化、注意力权重调试分析
  • 性能瓶颈 Q K ⊤ QK^\top QK 矩阵乘法时间复杂度为 O ( n 2 ⋅ d k ) O(n^2 \cdot d_k) O(n2dk),序列长度n越大,计算量爆炸,是Transformer核心性能瓶颈

2. Multi-Head Attention(多头注意力)

2.1 核心原理

单头注意力仅能在单一特征空间计算语义关联,表征能力有限。多头注意力将 Q/K/V 投影至多个独立特征子空间,并行计算多组注意力,最后拼接融合所有子空间特征,捕捉丰富的多维语义关系。

核心公式:

M u l t i H e a d ( Q , K , V ) = C o n c a t ( h e a d 1 , h e a d 2 , . . . , h e a d h ) W O MultiHead(Q,K,V)=Concat(head_1,head_2,...,head_h)W_O MultiHead(Q,K,V)=Concat(head1,head2,...,headh)WO

其中单个头: h e a d i = A t t e n t i o n ( Q W Q i , K W K i , V W V i ) head_i = Attention(QW_{Q_i},KW_{K_i},VW_{V_i}) headi=Attention(QWQi,KWKi,VWVi)

关键设计优化:不单独定义3h个线性层,而是先整体投影再拆头,仅用4个线性层即可实现等价效果,大幅减少参数量与计算量。

2.2 完整代码实现
class MultiHeadAttention(nn.Module):
    """
    多头注意力机制
    MultiHead(Q, K, V) = Concat(head_1, ..., head_h) @ W_O
    """
    def __init__(self, d_model: int, n_heads: int, dropout: float = 0.1):
        super().__init__()
        # 特征维度必须整除头数,保证每个头维度均等
        assert d_model % n_heads == 0, "d_model 必须能被 n_heads 整除"

        self.d_model = d_model
        self.n_heads = n_heads
        self.d_k = d_model // n_heads

        # 4个核心投影矩阵
        self.W_Q = nn.Linear(d_model, d_model, bias=False)
        self.W_K = nn.Linear(d_model, d_model, bias=False)
        self.W_V = nn.Linear(d_model, d_model, bias=False)
        self.W_O = nn.Linear(d_model, d_model, bias=False)

        # 复用缩放点积注意力
        self.attention = ScaledDotProductAttention(dropout)

    def forward(self, Q, K, V, mask=None):
        batch_size = Q.size(0)

        # 1. 整体线性投影
        Q = self.W_Q(Q)
        K = self.W_K(K)
        V = self.W_V(V)

        # 2. 拆分多头:(batch, seq_len, d_model) -> (batch, n_heads, seq_len, d_k)
        Q = Q.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
        K = K.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)
        V = V.view(batch_size, -1, self.n_heads, self.d_k).transpose(1, 2)

        # 3. 并行计算多头注意力
        attn_output, attn_weights = self.attention(Q, K, V, mask)

        # 4. 拼接所有头特征,恢复维度
        attn_output = attn_output.transpose(1, 2).contiguous().view(
            batch_size, -1, self.d_model
        )

        # 5. 最终融合投影
        output = self.W_O(attn_output)
        return output
2.3 核心关键点总结
  • 维度约束:d_model % n_heads == 0,保证每个注意力头分配均等维度
  • view+transpose:拆分特征维度,将多头维度前置,实现并行计算
  • contiguous():transpose仅修改视图、不改变内存布局,该方法重置内存,避免后续维度变换报错
  • 广播掩码:mask格式为(batch, 1, 1, seq_len),可直接与注意力分数矩阵广播运算,无需额外维度扩展

3. Position-wise Feed-Forward Network(位置前馈网络)

3.1 核心原理

注意力层负责捕捉序列全局关联,而 FFN 负责对每个位置特征独立做非线性变换,等价于 kernel_size=1 的一维卷积,大幅提升模型特征表达容量。

核心公式:

F F N ( x ) = R e L U ( x W 1 + b 1 ) W 2 + b 2 FFN(x)=ReLU(xW_1+b_1)W_2+b_2 FFN(x)=ReLU(xW1+b1)W2+b2

3.2 完整代码实现
class PositionWiseFeedForward(nn.Module):
    """
    位置式前馈网络
    维度变换:d_model → d_ff(升维)→ d_model(降维)
    """
    def __init__(self, d_model: int, d_ff: int, dropout: float = 0.1):
        super().__init__()
        self.linear1 = nn.Linear(d_model, d_ff)
        self.linear2 = nn.Linear(d_ff, d_model)
        self.dropout = nn.Dropout(dropout)

    def forward(self, x):
        # 升维+非线性激活+dropout+降维
        return self.linear2(self.dropout(F.relu(self.linear1(x))))
3.3 核心关键点总结
  • 维度设计:d_ff 远大于 d_model(论文512→2048),通过升维提升非线性拟合能力
  • dropout位置:置于ReLU激活之后,是工业界最优实践
  • 激活函数迭代:原始论文用ReLU,后续GPT、LLaMA等模型升级为GELU,拟合效果更优
  • 两层最优解:单层FFN表达能力不足,三层及以上收益极低、算力消耗剧增,两层是性能与算力的最优平衡

4. Positional Encoding(位置编码)

4.1 核心原理

Self-Attention 是排列不变性结构,无法感知序列顺序。为模型注入时序位置信息,论文采用正余弦位置编码,无需训练、可泛化超长序列。

核心公式:

P E ( p o s , 2 i ) = s i n ( p o s 10000 2 i / d m o d e l ) PE(pos,2i)=sin(\frac{pos}{10000^{2i/d_{model}}}) PE(pos,2i)=sin(100002i/dmodelpos)

P E ( p o s , 2 i + 1 ) = c o s ( p o s 10000 2 i / d m o d e l ) PE(pos,2i+1)=cos(\frac{pos}{10000^{2i/d_{model}}}) PE(pos,2i+1)=cos(100002i/dmodelpos)

4.2 正余弦编码优势
  • 可外推性:支持推理时出现训练集未见过的超长序列
  • 无参数:无需训练参数,节省内存与算力
  • 相对位置感知:利用三角函数和差公式,可线性表达相对位置偏移
4.3 完整代码实现
class PositionalEncoding(nn.Module):
    """
    正余弦位置编码
    PE(pos, 2i)   = sin(pos / 10000^(2i/d_model))
    PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
    """
    def __init__(self, d_model: int, max_len: int = 5000, dropout: float = 0.1):
        super().__init__()
        self.dropout = nn.Dropout(dropout)

        # 初始化位置编码矩阵
        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)
        )

        # 偶数维度sin,奇数维度cos
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)

        # 注册为缓冲区:不参与梯度更新,随模型迁移设备
        pe = pe.unsqueeze(0)
        self.register_buffer('pe', pe)

    def forward(self, x):
        # 广播相加,注入位置信息
        x = x + self.pe[:, :x.size(1), :]
        return self.dropout(x)
4.4 核心关键点总结
  • 数值稳定性:采用指数运算替代直接幂运算,避免数值溢出/下溢
  • register_buffer:位置编码是固定常量,不参与参数更新,仅随模型迁移CPU/GPU
  • 相加融合:特征嵌入与位置编码直接广播相加,轻量化融合位置信息

5. Layer Normalization(层归一化)

5.1 核心原理

层归一化对单个样本的所有特征维度做归一化,解决训练过程中特征分布偏移问题。与BatchNorm不同,LN不依赖batch大小,对变长NLP序列极其友好,是Transformer的标配归一化方案。

核心公式:

L a y e r N o r m ( x ) = γ ⋅ x − μ σ 2 + ϵ + β LayerNorm(x)=\gamma \cdot \frac{x-\mu}{\sqrt{\sigma^2+\epsilon}}+\beta LayerNorm(x)=γσ2+ϵ xμ+β

5.2 完整代码实现
class LayerNorm(nn.Module):
    """
    手写层归一化,便于底层原理理解
    生产环境可直接使用 nn.LayerNorm
    """
    def __init__(self, d_model: int, eps: float = 1e-6):
        super().__init__()
        # 可学习缩放、偏移参数
        self.gamma = nn.Parameter(torch.ones(d_model))
        self.beta = nn.Parameter(torch.zeros(d_model))
        self.eps = eps

    def forward(self, x):
        # 对最后一维(特征维)计算均值、方差
        mean = x.mean(dim=-1, keepdim=True)
        std = x.std(dim=-1, keepdim=True, unbiased=False)
        # 归一化+缩放偏移
        return self.gamma * (x - mean) / (std + self.eps) + self.beta
5.3 核心关键点总结
  • unbiased=False:使用有偏样本标准差,与原始Transformer论文实现完全对齐
  • eps极小值:防止分母为0,避免数值报错
  • 可学习参数:gamma、beta让归一化后特征具备自适应拟合能力

6. Encoder Layer(编码器层)

6.1 核心结构

编码器用于双向语义理解,每层由两个子层组成,搭配残差连接+层归一化,堆叠形成深层特征提取网络:

输入 → 多头自注意力 → 残差连接+LN → FFN → 残差连接+LN → 输出

6.2 完整代码实现
class EncoderLayer(nn.Module):
    """Transformer 编码器单层结构"""
    def __init__(self, d_model: int, n_heads: int, d_ff: int, dropout: float = 0.1):
        super().__init__()
        self.self_attn = MultiHeadAttention(d_model, n_heads, dropout)
        self.ffn = PositionWiseFeedForward(d_model, d_ff, dropout)
        self.norm1 = LayerNorm(d_model)
        self.norm2 = LayerNorm(d_model)
        self.dropout1 = nn.Dropout(dropout)
        self.dropout2 = nn.Dropout(dropout)

    def forward(self, x, mask=None):
        # 1. 自注意力子层 + 残差归一化
        attn_output = self.self_attn(x, x, x, mask)
        x = x + self.dropout1(attn_output)
        x = self.norm1(x)

        # 2. 前馈网络子层 + 残差归一化
        ffn_output = self.ffn(x)
        x = x + self.dropout2(ffn_output)
        x = self.norm2(x)

        return x
6.3 核心关键点总结
  • 残差连接x + sublayer(x),打通梯度直通通道,彻底解决深层网络梯度消失问题
  • Post-LN结构:先残差、后归一化,严格对齐原始论文实现
  • 双向自注意力:Q/K/V均来自输入自身,可感知上下文双向语义

7. Decoder Layer(解码器层)

7.1 核心结构

解码器用于自回归序列生成,在编码器基础上新增交叉注意力,共三个子层,严格防止未来信息泄露:

输入 → 掩码自注意力 → 残差LN → 交叉注意力 → 残差LN → FFN → 残差LN → 输出

7.2 完整代码实现
class DecoderLayer(nn.Module):
    """Transformer 解码器单层结构"""
    def __init__(self, d_model: int, n_heads: int, d_ff: int, dropout: float = 0.1):
        super().__init__()
        self.self_attn = MultiHeadAttention(d_model, n_heads, dropout)
        self.cross_attn = MultiHeadAttention(d_model, n_heads, dropout)
        self.ffn = PositionWiseFeedForward(d_model, d_ff, dropout)

        self.norm1 = LayerNorm(d_model)
        self.norm2 = LayerNorm(d_model)
        self.norm3 = LayerNorm(d_model)

        self.dropout1 = nn.Dropout(dropout)
        self.dropout2 = nn.Dropout(dropout)
        self.dropout3 = nn.Dropout(dropout)

    def forward(self, x, enc_output, src_mask=None, tgt_mask=None):
        # 1. 掩码自注意力:屏蔽未来位置信息
        attn_output = self.self_attn(x, x, x, tgt_mask)
        x = x + self.dropout1(attn_output)
        x = self.norm1(x)

        # 2. 交叉注意力:Decoder查Encoder全局特征
        attn_output = self.cross_attn(x, enc_output, enc_output, src_mask)
        x = x + self.dropout2(attn_output)
        x = self.norm2(x)

        # 3. 前馈网络
        ffn_output = self.ffn(x)
        x = x + self.dropout3(ffn_output)
        x = self.norm3(x)

        return x
7.3 核心关键点总结
  • Target掩码:下三角掩码,屏蔽当前token之后的所有位置,杜绝未来信息泄露,保障自回归生成逻辑
  • 交叉注意力机制:Q来自解码器当前输出,K/V来自编码器全局编码结果,实现输入序列对生成序列的语义引导
  • 三层残差归一化:相比编码器多一组子层结构,适配生成任务复杂度

8. 完整 Transformer 模型拼装

8.1 整体网络流程

编码器流程:源序列 → 嵌入层 → 位置编码 → N层编码器堆叠 → 全局语义特征

解码器流程:目标序列 → 嵌入层 → 位置编码 → N层解码器堆叠 → 全连接输出词表概率

8.2 完整模型代码
class Transformer(nn.Module):
    """
    完整 Transformer 模型
    适配机器翻译、序列生成等经典任务
    """
    def __init__(self, src_vocab, tgt_vocab, d_model=512,
                 n_heads=8, d_ff=2048, n_layers=6,
                 dropout=0.1, max_len=5000):
        super().__init__()
        # 嵌入层
        self.encoder_embed = nn.Embedding(src_vocab, d_model)
        self.decoder_embed = nn.Embedding(tgt_vocab, d_model)
        # 位置编码
        self.pos_encoding = PositionalEncoding(d_model, max_len, dropout)

        # 堆叠多层编码器、解码器
        self.encoder_layers = nn.ModuleList([
            EncoderLayer(d_model, n_heads, d_ff, dropout)
            for _ in range(n_layers)
        ])
        self.decoder_layers = nn.ModuleList([
            DecoderLayer(d_model, n_heads, d_ff, dropout)
            for _ in range(n_layers)
        ])

        # 最终输出层:映射到词表维度
        self.fc_out = nn.Linear(d_model, tgt_vocab)

    def forward(self, src, tgt, src_mask=None, tgt_mask=None):
        # 编码器前向传播
        src_emb = self.pos_encoding(self.encoder_embed(src))
        for layer in self.encoder_layers:
            src_emb = layer(src_emb, src_mask)

        # 解码器前向传播
        tgt_emb = self.pos_encoding(self.decoder_embed(tgt))
        for layer in self.decoder_layers:
            tgt_emb = layer(tgt_emb, src_emb, src_mask, tgt_mask)

        # 输出每个位置的词表概率分布
        return self.fc_out(tgt_emb)
8.3 核心关键点总结
  • ModuleList:正确注册多层网络参数,保证模型可训练、参数可更新
  • 双嵌入层设计:编码器、解码器独立词嵌入,适配源域、目标域不同语义分布
  • 输出映射:将512维特征映射至词表维度,用于token预测与序列生成

四、全文核心总结

  1. 核心本质:Transformer 彻底抛弃递归结构,完全依靠注意力机制实现全局语义建模,支持全并行训练
  2. 缩放注意力:除以 d k \sqrt{d_k} dk 是训练稳定的关键,解决高维点积梯度消失问题
  3. 多头机制:多空间并行建模,丰富语义表征,是模型强大拟合能力的核心
  4. 位置编码:正余弦编码实现无参数、可外推的序列位置注入
  5. 编解码分工:Encoder双向理解、Decoder自回归生成,适配各类NLP任务
  6. 残差+LN:深层网络训练的基础保障,解决梯度消失、特征分布偏移问题

⭐️推荐:

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐