使用 PyTorch 从零构建类 ChatGPT 的 Transformer
外来客 • 2026-08-21 23:51:10
声明:本文为对公开内容的摘要整理,
未经本站独立核实,可能与原内容存在出入,不代表本站立场、观点或建议;
观点与版权归原作者及原平台所有。
如涉及版权问题,请联系我们,核实后立即删除。
[ 免责声明 ]
(原标题: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 论文做法,不添加偏置项。
- 前向传播计算流程:
- 生成 Q, K, V:将编码后的 Token 分别通过三个线性层得到 Q、K、V 矩阵。
- 计算相似度:使用 `torch.matmul` 计算 Q 与 K 的转置的乘积,得到相似度矩阵 `sims`。
- 缩放:将相似度除以 `d_model` 的平方根,防止梯度消失,这是 2017 年原始论文的标准做法。
- 应用掩码:
- 掩码矩阵中 `True` 对应需要忽略的位置(即未来 Token)。
- 使用 `masked_fill` 将 `True` 位置填充为负无穷大(近似值 -1e9),`False` 位置填充为 0。
- 将此掩码加到缩放后的相似度上,确保早期 Token 无法“偷看”后续 Token。
- Softmax:对掩码处理后的相似度应用 `softmax`,得到注意力百分比 `attention_percents`。
- 加权求和:将注意力百分比与 V 矩阵相乘,得到最终的注意力得分 `attention_scores`。
🏗️ 仅解码器 Transformer 架构
- 类定义:创建 `DecoderOnlyTransformer` 类,继承自 `LightningModule` 以利用 Lightning 的训练功能。
- 初始化组件:
- 词嵌入:`nn.Embedding`,维度由词汇表大小和 `d_model` 决定。
- 位置编码:实例化前述的 `PositionEncoding` 类。
- 注意力层:实例化前述的 `Attention` 类。
- 全连接层:`nn.Linear`,输入输出维度均为 `d_model`。
- 损失函数:使用交叉熵损失(Cross Entropy Loss),该函数内部自动执行 Softmax。
- 前向传播逻辑:
- 将输入 Token ID 转换为词嵌入向量。
- 添加位置编码。
- 生成掩码:
- 使用 `torch.ones` 创建全 1 矩阵。
- 使用 `torch.tril`(下三角)保留下三角的 1,上三角变为 0。
- 将 0 转换为 `True`,1 转换为 `False`,形成用于注意力计算的布尔掩码。
- 计算注意力:将位置编码后的向量同时作为 Q、K、V 的输入(自注意力),并传入掩码。
- 残差连接:将注意力输出与输入相加。
- 全连接输出:通过全连接层得到最终输出,直接返回(Softmax 由损失函数处理)。
🚀 训练与推理流程
- 优化器配置:
- 使用 `Adam` 优化器,学习率设为 0.1(针对此简单模型加速训练,常规默认值为 0.001)。
- 传入模型所有可训练参数。
- 训练步骤:
- 定义 `training_step` 方法,接收批次数据和索引。
- 分离输入和标签。
- 调用 `forward` 方法计算输出。
- 计算输出与标签之间的交叉熵损失。
- 返回损失值供 Lightning 进行反向传播。
- 推理生成逻辑:
- 初始预测:输入提示词(如 “What is StatQuest EOS”),模型生成每个位置的预测。
- 提取下一 Token:取最后一个输入 Token(EOS)对应的输出向量,使用 `argmax` 找到概率最大的 Token ID。
- 循环生成:
- 将新生成的 Token 追加到输入序列中。
- 重新运行模型,基于完整上下文(原输入 + 已生成输出)预测下一个 Token。
- 重复此过程,直到生成 `EOS` 或达到最大长度限制。
- 结果转换:将生成的 Token ID 映射回文本。
- 训练前后对比:
- 训练前:输入 “What is StatQuest EOS”,模型直接输出 “EOS”,未生成预期答案。
- 训练后:使用 Lightning Trainer 训练 30 个 Epoch 后,输入相同提示词,模型正确输出 “Awesome EOS”。
- 验证:输入 “StatQuest is what EOS”,模型同样正确输出 “Awesome EOS”,证明模型成功学习了双向提示词到固定回答的映射。
👤 同一博主
来自父亲的另外三条人生经验
单纯形算法的数学细节
线性规划优化与单纯形算法
StatQuest:随机森林第二部分:缺失值与聚类
假发现率(FDR)详解
线性回归的本质
用超简单的方式解释AI的工作原理
StatQuest:来自科技行业领袖的职业建议
人类反馈强化学习(RLHF)详解
基于神经网络的强化学习:数学细节
🧭 类似博主
-
最佳顺畅稳定机场 1000多Mbps,支持Win 安卓 IOS Mac 全平台,GPT/Gemini
-
人类正式进入埃米时代?!台积电1.6nm进入量产倒计时,深度剖析A16的技术密码
-
无影无踪,核周报9.5
-
取代 OpenAI 模型的方法:我換了這套繁中字幕流程
-
苹果9月9日iPhone发布会:我们期待的一切!
-
Cassie Coppersmith(卡茜·科珀史密斯)的伪考古学谬论
-
一个视频搞懂DeepSeek Harness!
-
约翰·特纳斯(John Ternus)出任CEO、法律纠纷及iPhone Fold登上AppleIn
-
为什么 DDR5 内存的性能表现往往不如预期?
-
明天就要開會,我卻什麼都沒準備!Genspark Super Agent 能救我嗎?
0 条评论
发表评论
请先 登录 后参与讨论。