ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

重新绑定 LM Head 与 Embed 权重:modded-nanogpt 的 FP8 尺度重调与步数缩减实战

重新绑定 LM Head 与 Embed 权重:modded-nanogpt 的 FP8 尺度重调与步数缩减实战 人工智能大模型预训练分布式训练模型优化深度学习【免费下载链接】modded-nanogptNanoGPT (124M) in 90 seconds项目地址https://gitcode.com/GitHub_Trending/mo/modded-nanogpt点击查看免费下载导读本文基于 modded-nanogpt 仓库中 2025-12-19_RetieLMHead 这一训练记录深入解析两项关键优化一是把此前被解绑的 LM Head语言模型头与第一层 Embed 权重重新绑定weight tying二是针对共享权重重新调整 FP8 量化尺度。通过这两项改动该记录在步时几乎不变的情况下减少了 55 个训练步并在同机对比中将单次完整训练时间缩短约 3.2 秒。读完本文你将掌握权重绑定在分布式 AdamDistAdam通信优化中的动机与实现方式理解 FP8 尺度scale应该如何依据权重幅值分布p999进行实验性选取以及如何用 t 检验与多次复跑来严谨验证优化结果。背景从解绑Record 8回到绑定在了解本次改动前需要先回顾历史。modded-nanogpt 的早期记录Record 8即 2024-11-03 的 UntieEmbed曾做出两项与 Embed 相关的改动解绑 LM Head 与 Embed 层同时将 LM Head 初始化为零并伴随在 Embedding 之后增加 RMS Norm。也就是说彼时模型把输入词嵌入矩阵与输出 logits 投影矩阵拆成两个独立的参数矩阵。本记录2025-12-19_RetieLMHead则反其道而行之恢复 LM Head 与 Embed 的权重绑定并把 LM Head 的初始化从零初始化改回小方差正态初始化。作者的动机非常直接——在 8 卡分布式训练中Adam 每步都需要在卡间通信所有被它管理的参数把两个大矩阵合并成一个共享矩阵可以减少需要通信的参数总量从而缩短 Adam 步的耗时。从结果看这个 tradeoff 出人意料地划算作者原本预计绑定会以步数增加为代价实际上重绑 Embed 与 LM Head 显著减少了步数而单步耗时几乎不受影响。README 中对这一现象给出了一个推测LM Head 的权重在反向传播中既是第一个梯度元素又是最后一个梯度元素这种首尾同时出现的特性会破坏 DistAdam 中部分异步逻辑的并行性。时序与验证Timing and Validation本记录用严谨的统计手段验证了优化收益在同一台机器上对最终验证损失和单次训练总耗时各进行多次复跑并用 scipy 的单样本 t 检验确认损失改善的显著性。import scipy.stats import torch losses [3.2777, 3.2790, 3.2776, 3.2792, 3.2760, 3.2792, 3.2767, 3.2763, 3.2770] times [123.853, 123.969, 123.929, 123.933, 123.914, 123.906, 123.970, 123.914, 123.964] print(p%.4f % scipy.stats.ttest_1samp(losses, 3.28, alternativeless).pvalue) # p0.0002 print(losses:, torch.std_mean(torch.tensor(losses))) # losses: (std0.0013, mean3.2776) print(time:, torch.std_mean(torch.tensor(times))) # time: (std0.0375, mean123.9280)上一记录同机计时的参考数据为import scipy.stats import torch times [127.051, 127.139, 127.049, 127.147, 127.163, 127.161] print(time:, torch.std_mean(torch.tensor(times))) # time: (std0.0537, mean127.1183)两组数据的解读最终验证损失9 次复跑均值为3.2776、标准差仅0.0013单样本 t 检验原假设为不小于 3.28得到p0.0002说明新记录在统计意义上显著优于 3.28 的损失线且复现性极好。单次完整训练耗时新记录均值123.928 秒上一记录均值127.118 秒改善约3.2 秒同时记录提到本记录比上一记录少了 55 个训练步步时几乎相同见日志step_avg保持在约 33~60ms 的渐进区间最终步2035/2035val_loss 恰好收敛到3.2777与上表复跑均值一致。换言之收益来自更少步数而非更快单步——这正是重绑权重后 Adam 通信量下降带来的结构性红利。该记录的实际运行日志见 0828d309-ecfe-4442-9ee9-68fed3a4b599.txt其中完整记录了从 warmup、每 250 步验证val_loss:10.8369 → 4.2657 → 4.0017 → 3.8428 → 3.6846 → 3.5677 → 3.4545 → 3.3569 → 3.2847 → 3.2777到峰值显存allocated 29958 MiB / reserved 38816 MiB的全部过程。源码剖析Tied Embed and LM Head 的实现在训练脚本中权重绑定最直观的体现是模型forward中不再单独使用nn.Embedding而是直接用lm_head.weight做查表# weight-tied: use lm_head.weight for embedding lookup x F.embedding(input_seq, self.lm_head.weight)对应地GPT.__init__中lm_head的声明见 0828d309 记录脚本为use_fp8 not os.environ.get(DISABLE_FP8, False) self.lm_head CastedLinear(model_dim, vocab_size, use_fp8use_fp8, x_s100/448, w_s1.5/448, grad_s0.75/448) nn.init.normal_(self.lm_head.weight, mean0, std0.005) self.lm_head.weight.label lm_head几个值得注意的实现细节词表对齐vocab_size next_multiple_of_n(vocab_size, n128)把 GPT-2 的 50257 个 token 向上取整到 128 的倍数即 50304以保证 FP8 矩阵乘法的 shape 对齐效率。初始化nn.init.normal_(..., std0.005)取代了 Record 8 的零初始化。这是本记录相对 Record 8反向的第二个改动——LM Head 不再从全零出发。标签驱动优化器分组self.lm_head.weight.label lm_head是关键。DistAdam的构造函数按label_order [lm_head, value_embed, scalars]显式分组见 记录脚本并在每个参数上注册post_accumulate_grad_hook通过reduce_scatter_tensor把梯度按行切分到各卡、异步完成all_gather。训练循环中head_params [model.lm_head.weight]被单独收集并与其他 Embed、标量参数一起交给DistAdamoptimizer1 DistAdam(embed_params scalar_params head_params, lr0.008, betas(0.65, 0.95), eps1e-8, weight_decay0.005)。由于lm_head.weight同时承担 Embedding 查表和 logits 投影两个角色它在反向传播里会同时收到来自词嵌入梯度与输出投影梯度两端的贡献——这正是 README 中推测破坏 DistAdam 异步逻辑的结构性原因该权重既是反传的第一个参数也是最后一个参数其 reduce-scatter 与 all-gather 的时间窗口无法像普通参数那样被前向/反向计算完全遮蔽。为什么绑定能减少步数README 明确指出重绑带来的另一个收益是步数下降且作者坦承并不知道确切原因。这属于实验观察而非结论性事实。从源码结构可以推断一个合理的解释路径绑定后 LM Head 梯度被强制共享同一套 Adam 状态与同一学习率等价于引入了输出空间与输入空间共享几何的正则化约束客观上限制了词嵌入空间漂移使训练更稳定、更快收敛。但这只是推断本文如实标注为推测不做事实断言。下调 FP8 Scales实验驱动的量化尺度选取第二个核心改动是重新调整 FP8 量化尺度。在 8 卡训练中CastedLinear走的是自定义算子torch.ops.nanogpt.mmFP8 矩阵乘其前向把激活和权重分别除以x_s、w_s后转成float8_e4m3fn反传则把梯度除以grad_s后转成float8_e5m2x_f8 x.div(x_s).to(torch.float8_e4m3fn) w_f8 w.div(w_s).to(torch.float8_e4m3fn) out torch._scaled_mm(x_f8, w_f8.T, out_dtypetorch.bfloat16, scale_ax.new_tensor(x_s, dtypetorch.float32), scale_bx.new_tensor(w_s, dtypetorch.float32), ...)本记录为该共享的 LM Head/Embed 权重选取的相对尺度为激活尺度x_s 100/448权重尺度w_s 1.5/448梯度尺度grad_s 0.75/448选取依据是一组针对共享 LM Head/Embed 权重做的实验图见下图。该图包含三个子图横轴均为训练步Mean p999_x_in vs stepLM Head 输入特征 99 百分位p999的均值在约 1000 步前快速上升之后趋于平稳Mean p999_weight vs stepLM Head 权重的 p999 均值在 1500 步范围内持续上升、尚未完全收敛说明权重幅值在整个训练期都在增长Mean p999_grad_out vs stepLM Head 输出梯度的 p999 均值约 1500 步后进入小幅波动稳定段。这三个子图传递的核心信息是共享权重不同分量输入、权重、梯度的动态范围随训练阶段变化且收敛节奏各不相同。因此 FP8 尺度不能用静态值一刀切而应依据对应分量在关键训练阶段的幅值上限p999来设置——既不能过小导致溢出也不能过大浪费 FP8 的精度。README 同时指出这些曲线还暗示简单的线性调度linear schedule可能是 FP8 尺度的最优方案但作者用简单尝试未能使其奏效因此最终仍采用固定尺度这一探索性结论原样保留供读者参考。值得注意的是本记录同时吸收了上游 PR#172 的成果而 FP8 尺度重调与 PR#172 是互相受益的作者先独立发现了绑定收益随后发现绑定与 PR#172 都能从 FP8 尺度重调中获益。为了隔离变量作者把本记录的尺度单独应用到 PR#172不做 LM Head 绑定上复跑 5 次import scipy.stats import torch losses [3.2754, 3.2756, 3.2779, 3.2746, 3.2740] print(losses:, torch.std_mean(torch.tensor(losses))) # losses: (std0.0015, mean3.2755)对比可见仅靠 FP8 权重尺度重调不绑定损失均值即可达到3.2755优于本记录绑定方案自身的3.2776同时方差仍然偏大std0.0015。README 对此的解读是FP8 权重尺度重调本身提升了均值并缓解了由 Adam 参数上的 cautious weight decay谨慎权重衰减带来的高方差。经验总结与可复制要点权重绑定是分布式训练中的通信税优化在 8 卡 DistAdam 场景下参数通信量与参数个数强相关将两个大矩阵合二为一是降低通信量的直接手段。但绑定收益步数减少超出通信量本身的解释说明它还改变了优化轨迹值得在类似小模型超速训练任务中复现。FP8 尺度必须看数据选取以共享权重的 p999 幅值为参照系为x_s、w_s、grad_s分别选取合适数值不同分量激活/权重/梯度幅值收敛节奏不同不能互相套用。验证要可统计、可复现同机多次复跑 scipy.stats.ttest_1samp显著性检验 torch.std_mean报告均值与方差是判断一个改动是否真正有效的规范做法。改动之间会互相耦合本记录中绑定 FP8 尺度重调与PR#172 FP8 尺度重调均优于基线但两者叠加并非简单加和——做消融ablation时务必单变量隔离。如需进一步研读可对照本仓库当前实现 train_gpt.py其中lm_head_f8_col、sampled_softmax.gather等机制反映了后续对 LM Head FP8 拷贝的持续演进以及本记录所在的 track_1_short 记录体系了解完整的速度竞赛脉络。赞分享人工智能大模型预训练分布式训练模型优化深度学习【免费下载链接】modded-nanogptNanoGPT (124M) in 90 seconds项目地址https://gitcode.com/GitHub_Trending/mo/modded-nanogpt点击查看免费下载相关推荐把流程图写进文档Mermaid 文本绘图上手与团队实践把流程图写进文档Mermaid 文本绘图上手与团队实践 文档里的图总是落后于代码流程改了图片还停在旧版本想改一个节点得找原图作者重新导出一版。Merm图表库前端数据可视化modded-nanogpt 的 Logit Rescale 记录一次 −2.9 秒的 Softcap 参数重调与统计验证modded nanogpt 的 Logit Rescale 记录一次 −2.9 秒的 Softcap 参数重调与统计验证 本篇文章以 modded nano人工智能大模型预训练分布式训练模型优化深度学习Cautious Weight Decay 在 Adam 上的落地modded-nanogpt 将谨慎权重衰减扩展到标量参数之外Cautious Weight Decay 在 Adam 上的落地modded nanogpt 将谨慎权重衰减扩展到标量参数之外 导读本文围绕 modd人工智能大模型预训练分布式训练模型优化深度学习上一篇omp Collab 实时会话共享端到端加密、跨终端与浏览器联动的 Coding Agent 实战指南下一篇一次扫码免费导出QQ空间历史说说GetQzonehistory 本地存档完整指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表