博主头像

使用 PyTorch 从零构建类 ChatGPT 的 Transformer

外来客 • 2026-08-21 23:51:10

分享
𝕏 f
声明:本文为对公开内容的摘要整理, 未经本站独立核实,可能与原内容存在出入,不代表本站立场、观点或建议; 观点与版权归原作者及原平台所有。 如涉及版权问题,请联系我们,核实后立即删除。 [ 免责声明 ]

(原标题:Coding a ChatGPT Like Transformer From Scratch in PyTorch)

📚 核心目标与前置准备

  • 核心目标:使用 PyTorch 从零构建一个仅解码器(Decoder-only)Transformer 模型,该架构是 ChatGPT 的基础。
  • 环境依赖:导入 `torch` 用于创建张量和辅助函数;导入 `torch.nn` 获取 `Linear` 和 `Embedding` 类;导入 `torch.nn.functional` 访问 `softmax` 函数;导入 `Adam` 优化器用于反向传播训练;导入 `TensorDataset` 和 `DataLoader` 处理大规模数据;导入 `Lightning` 以简化代码编写并支持云端自动优化。
  • 数据准备:构建一个极简训练集,仅包含两个提示词(Prompt):“What is StatQuest?” 和 “StatQuest is what?”,期望两者的回答均为 “Awesome”。
  • 词汇表映射:定义词汇表包含 `What`、`is`、`StatQuest`、`Awesome` 和 `EOS`(结束符)。建立 `token_to_id` 和 `id_to_token` 字典,因为 PyTorch 的嵌入层仅接受数字输入。
  • 输入输出逻辑
  • 输入序列由提示词处理阶段和输出生成阶段的 Token 组成。
  • 例如提示词 “What is StatQuest?”,输入张量编码为 `What, is, StatQuest, EOS, Awesome`。
  • 标签(Label)序列为 `is, StatQuest, EOS, Awesome, EOS`,即每个输入 Token 预测下一个 Token。
  • 使用 `TensorDataset` 封装输入和标签,再通过 `DataLoader` 创建数据加载器。

📐 位置编码实现细节

  • 原理:使用交替的正弦和余弦函数计算每个 Token 的位置值,以保留序列顺序信息。
  • 公式参数:`pos` 代表 Token 在输入中的位置索引;`i` 代表嵌入值的索引;`d_model` 代表每个 Token 的嵌入维度。
  • 预计算策略:为避免每次前向传播时重复计算,预先计算位置编码矩阵并存储。
  • 代码实现步骤
  • 定义继承自 `nn.Module` 的 `PositionEncoding` 类。
  • 初始化参数:`d_model`(嵌入维度,示例中设为 2)和 `max_len`(最大 Token 数,示例中设为 6)。
  • 创建全零矩阵 `pe`,形状为 `(max_len, d_model)`。
  • 生成位置列向量 `position`(0 到 max_len-1)和嵌入索引行向量 `embedding_position`(步长为 2,即 0, 2, 4...)。
  • 计算除数项 `div_term`,用于调整正弦和余弦函数的频率。
  • 将正弦函数值填入 `pe` 矩阵的偶数列(0, 2...),余弦函数值填入奇数列(1, 3...)。
  • 使用 `register_buffer` 确保矩阵随模型移动到 GPU。
  • 在 `forward` 方法中,将预计算的位置编码值直接加到词嵌入值上。

🔄 掩码自注意力机制

  • 核心组件:计算查询(Query, Q)、键(Key, K)和值(Value, V)。
  • 权重矩阵
  • 使用 `nn.Linear` 创建三个线性层 `W_Q`、`W_K`、`W_V`。
  • 输入和输出特征维度均为 `d_model`。
  • 设置 `bias=False`,遵循原始 Transformer 论文做法,不添加偏置项。
  • 前向传播计算流程
  1. 生成 Q, K, V:将编码后的 Token 分别通过三个线性层得到 Q、K、V 矩阵。
  2. 计算相似度:使用 `torch.matmul` 计算 Q 与 K 的转置的乘积,得到相似度矩阵 `sims`。
  3. 缩放:将相似度除以 `d_model` 的平方根,防止梯度消失,这是 2017 年原始论文的标准做法。
  4. 应用掩码
  • 掩码矩阵中 `True` 对应需要忽略的位置(即未来 Token)。
  • 使用 `masked_fill` 将 `True` 位置填充为负无穷大(近似值 -1e9),`False` 位置填充为 0。
  • 将此掩码加到缩放后的相似度上,确保早期 Token 无法“偷看”后续 Token。
  1. Softmax:对掩码处理后的相似度应用 `softmax`,得到注意力百分比 `attention_percents`。
  2. 加权求和:将注意力百分比与 V 矩阵相乘,得到最终的注意力得分 `attention_scores`。

🏗️ 仅解码器 Transformer 架构

  • 类定义:创建 `DecoderOnlyTransformer` 类,继承自 `LightningModule` 以利用 Lightning 的训练功能。
  • 初始化组件
  • 词嵌入:`nn.Embedding`,维度由词汇表大小和 `d_model` 决定。
  • 位置编码:实例化前述的 `PositionEncoding` 类。
  • 注意力层:实例化前述的 `Attention` 类。
  • 全连接层:`nn.Linear`,输入输出维度均为 `d_model`。
  • 损失函数:使用交叉熵损失(Cross Entropy Loss),该函数内部自动执行 Softmax。
  • 前向传播逻辑
  1. 将输入 Token ID 转换为词嵌入向量。
  2. 添加位置编码。
  3. 生成掩码
  • 使用 `torch.ones` 创建全 1 矩阵。
  • 使用 `torch.tril`(下三角)保留下三角的 1,上三角变为 0。
  • 将 0 转换为 `True`,1 转换为 `False`,形成用于注意力计算的布尔掩码。
  1. 计算注意力:将位置编码后的向量同时作为 Q、K、V 的输入(自注意力),并传入掩码。
  2. 残差连接:将注意力输出与输入相加。
  3. 全连接输出:通过全连接层得到最终输出,直接返回(Softmax 由损失函数处理)。

🚀 训练与推理流程

  • 优化器配置
  • 使用 `Adam` 优化器,学习率设为 0.1(针对此简单模型加速训练,常规默认值为 0.001)。
  • 传入模型所有可训练参数。
  • 训练步骤
  • 定义 `training_step` 方法,接收批次数据和索引。
  • 分离输入和标签。
  • 调用 `forward` 方法计算输出。
  • 计算输出与标签之间的交叉熵损失。
  • 返回损失值供 Lightning 进行反向传播。
  • 推理生成逻辑
  1. 初始预测:输入提示词(如 “What is StatQuest EOS”),模型生成每个位置的预测。
  2. 提取下一 Token:取最后一个输入 Token(EOS)对应的输出向量,使用 `argmax` 找到概率最大的 Token ID。
  3. 循环生成
  • 将新生成的 Token 追加到输入序列中。
  • 重新运行模型,基于完整上下文(原输入 + 已生成输出)预测下一个 Token。
  • 重复此过程,直到生成 `EOS` 或达到最大长度限制。
  1. 结果转换:将生成的 Token ID 映射回文本。
  • 训练前后对比
  • 训练前:输入 “What is StatQuest EOS”,模型直接输出 “EOS”,未生成预期答案。
  • 训练后:使用 Lightning Trainer 训练 30 个 Epoch 后,输入相同提示词,模型正确输出 “Awesome EOS”。
  • 验证:输入 “StatQuest is what EOS”,模型同样正确输出 “Awesome EOS”,证明模型成功学习了双向提示词到固定回答的映射。

0 条评论

发表评论

请先 登录 后参与讨论。