从零手撕Transformer:自注意力机制与PyTorch实现详解

📅 2026/7/28 19:47:09
从零手撕Transformer:自注意力机制与PyTorch实现详解
理解 Transformer 的核心原理是进入现代大模型世界的钥匙。无论是想深入理解 ChatGPT、GPT-4 等大语言模型,还是想学习 Vision Transformer、Swin Transformer 等视觉模型,亦或是想自己动手实现一个简易的 Transformer,都绕不开对其基础架构和核心思想的透彻掌握。本文将从零开始,系统性地拆解 Transformer 的每一个组件,结合数学公式、代码示例和直观解释,让你不仅能“看懂”,更能“手撕”一个 Transformer。我们将从 RNN 的瓶颈讲起,逐步深入到自注意力机制、多头注意力、位置编码、前馈网络、编码器-解码器架构,最后通过一个完整的 PyTorch 实现来巩固理解。无论你是刚入门深度学习的新手,还是希望夯实基础的中高级开发者,这篇文章都将为你提供清晰的路径。1. 背景与核心概念:从 RNN 到 “Attention is All You Need”在 Transformer 出现之前,处理序列数据(如文本、语音、时间序列)的主流模型是循环神经网络(RNN)及其变体 LSTM 和 GRU。这些模型按顺序处理输入,将之前步骤的信息保存在一个“隐藏状态”中传递给下一步。1.1 RNN 的瓶颈尽管 RSTM 在一定程度上缓解了长序列训练中的梯度消失/爆炸问题,但 RNN 家族存在两个根本性限制:顺序计算,难以并行:必须等待t-1时刻的计算完成,才能计算t时刻,这严重限制了计算效率,无法充分利用现代 GPU 的大规模并行计算能力。长程依赖问题:即使使用 LSTM,当序列非常长时,早期信息在传递过程中仍然会逐渐衰减或丢失。一个经典的例子是:在句子“The animal didn’t cross the street because it was too tired.”中,模型需要将 “it” 与远处的 “animal” 关联起来,这对 RNN 来说是挑战。1.2 注意力机制的引入为了解决信息瓶颈问题,注意力机制被引入。其核心思想是:在生成每个输出时,模型可以“回顾”输入序列的所有部分,并动态地决定关注哪些部分。这就像人在翻译句子时,会来回查看原文的不同部分。最初的注意力机制被用在基于 RNN 的编码器-解码器架构中,但它仍然建立在 RNN 的顺序计算之上。1.3 Transformer 的诞生2017 年,Google 的 Vaswani 等人在论文《Attention is All You Need》中提出了Transformer模型。它的革命性在于:完全摒弃了循环结构,仅依赖注意力机制来建模序列中元素之间的关系。实现了完全并行化,所有序列位置同时被处理,极大提升了训练速度。引入了“自注意力”,让序列中的每个元素都能直接与序列中所有其他元素交互,无论距离多远,从而更好地捕捉长程依赖。Transformer 迅速成为自然语言处理(NLP)的基石,并催生了 BERT、GPT、T5 等一系列划时代的模型,进而引发了当前的大模型(LLM)与生成式 AI 浪潮。其思想也被成功迁移到计算机视觉(Vision Transformer)、语音识别(Conformer)等多模态领域。2. 环境准备与版本说明为了后续的代码实践部分,我们需要搭建一个 Python 深度学习环境。本文将使用 PyTorch 框架,因为它动态图的特点更适合教学和理解。核心环境要求:Python: 3.8 或更高版本。深度学习框架: PyTorch 1.9+。辅助库: NumPy, Matplotlib (用于可视化)。你可以使用以下命令快速创建环境并安装依赖:# 使用 conda 创建环境(可选) conda create -n transformer-tutorial python=3.9 conda activate transformer-tutorial # 安装 PyTorch (请根据你的CUDA版本访问官网获取对应命令) # 例如,对于没有GPU或使用CPU的情况: pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu # 安装其他依赖 pip install numpy matplotlib版本说明:本文的代码示例基于 PyTorch 2.0+ 的 API 编写,核心逻辑与早期版本兼容。重点在于理解原理,实际项目中请根据你的硬件和项目需求选择合适的 PyTorch 版本。3. Transformer 核心组件原理拆解一个标准的 Transformer 模型(以原始论文中的编码器-解码器架构为例)主要由以下核心组件构成:输入嵌入(Input Embedding)与位置编码(Positional Encoding)多头注意力机制(Multi-Head Attention)前馈神经网络(Feed-Forward Network)残差连接(Residual Connection)与层归一化(Layer Normalization)编码器(Encoder)与解码器(Decoder)堆叠下面我们逐一深入。3.1 输入表示:从文本到向量Transformer 无法直接处理文本,需要先将文本转换为数字向量。1. 词元化(Tokenization)将输入文本分割成更小的单元,称为词元(Tokens)。常见方法有:词级:将每个单词作为一个词元。词汇表大,无法处理未登录词。子词级:如 Byte Pair Encoding (BPE), WordPiece。这是现代 Transformer(如 BERT, GPT)的主流方法,能有效平衡词汇表大小与未登录词问题。 例如,句子 “I love transformers.” 可能被分词为[“I”, “love”, “transform”, “ers”, “.”]。2. 词嵌入(Embedding)每个词元通过一个可学习的查找表(Embedding Matrix)被映射为一个高维向量(例如 512 维)。这个向量旨在捕获该词元的语义信息。import torch import torch.nn as nn # 假设词汇表大小为 10000,嵌入维度为 512 vocab_size = 10000 d_model = 512 embedding_layer = nn.Embedding(vocab_size, d_model) # 输入是一个包含词元ID的序列,形状为 (batch_size, seq_len) # 例如 batch_size=2, seq_len=5 input_ids = torch.LongTensor([[101, 2054, 2003, 1037, 102], [101, 2023, 2003, 1037, 102]]) # 输出形状为 (batch_size, seq_len, d_model) embedded_output = embedding_layer(input_ids) print(embedded_output.shape) # torch.Size([2, 5, 512])3.2 位置编码(Positional Encoding)由于 Transformer 没有循环或卷积结构,它本身无法感知序列中词元的顺序信息。位置编码就是为了注入序列的顺序信息。原始 Transformer 使用正弦和余弦函数来生成位置编码: 对于位置pos和维度i(i为偶数或奇数),编码如下:$PE_{(pos, 2i)} = sin(pos / 10000^{2i/d_{model}})$ $PE_{(pos, 2i+1)} = cos(pos / 10000^{2i/d_{model}})$其中d_model是嵌入维度。这种编码方式的好处是,对于固定的偏移量k,PE(pos+k)可以表示为PE(pos)的线性函数,这使得模型能够轻松学习到相对位置信息。import math def get_positional_encoding(seq_len, d_model): """生成位置编码矩阵""" pe = torch.zeros(seq_len, d_model) position = torch.arange(0, seq_len, dtype=torch.float).unsqueeze(1) # (seq_len, 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 # (seq_len, d_model) # 示例:生成长度为10,维度为512的位置编码 pe = get_positional_encoding(seq_len=10, d_model=512) print(pe.shape) # torch.Size([10, 512])生成的位置编码矩阵会直接加到词嵌入向量上:input = embedding_output + positional_encoding。3.3 缩放点积注意力(