多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

Partial Key Offset 让 NanoGPT 提速:modded-nanogpt 128.8s 世界纪录中的注意力机制改造实录

Partial Key Offset 让 NanoGPT 提速:modded-nanogpt 128.8s 世界纪录中的注意力机制改造实录 人工智能大模型预训练分布式训练模型优化深度学习【免费下载链接】modded-nanogptNanoGPT (124M) in 90 seconds项目地址https://gitcode.com/GitHub_Trending/mo/modded-nanogpt点击查看免费下载导读本篇文章围绕 modded-nanogpt 在 2025-12-14 创造的 128.8 秒8×H100、124M 参数、fineweb10B训练纪录展开核心剖析该纪录中包含的五项更新重点深挖其中的Partial Key Offset部分键偏移注意力机制改造它如何通过只平移 Key 的静止维让一层注意力即可完成归纳induction从而在不显著增加计算的前提下压低验证损失。文章同时结合仓库源码track_1_short/model/gpt.py、track_1_short/model/attention.py、track_1_short/perf/kernels/qkv_rope.py给出可复现的实现细节、统计验证方法与完整训练脚本解读。读完本文你将理解为何 key offset 只作用在长滑动窗口与静止头维、x0_lambda合并的代数化简、batch size 调度与窗口调度如何精确对齐以及这类纪录型训练脚本如何用严格的统计检验来验收每一次提速。纪录背景128.8s 是怎么来的原纪录 README 的第一行即点明主题New WR 128.8s: Partial Key Offset新的世界纪录 128.8 秒靠的是部分键偏移。该纪录由五项更新叠加而成合计带来约 2.4 秒的提速其分解如下Partial Key Offset为长滑动窗口实现部分键偏移详见下一节合并 layer 0 的残差缩放把第 0 层的x_lambda*x x0_lambda*x0化简为(x_lambdax0_lambda)*x并顺带清理代码使 11 层模型的结构表达得更清晰减少 50 步训练每步约 60ms直接节省约 3 秒墙钟时间这是更快与更低损失之间的权衡取舍对齐 batch size 调度与窗口调度让 batch size 的切换时刻与滑动窗口尺寸的更新时刻严格同步因为 0.33 ≠ 1/3原本两者存在错位初始 value embeddings 清零对最终指标影响很小但可能降低训练方差同时符合零初始化作为最低假设配置的原则。这些更新对应的完整可运行脚本保存在 records/track_1_short/2025-12-14_PartialKeyOffset/ 目录下同名训练脚本有 11 份为同一配置的不同运行副本README 位于 records/track_1_short/2025-12-14_PartialKeyOffset/README.md。注以下验证损失、训练时间等数字均直接引自该纪录 README 与同目录训练日志反映的是 2025-12-14 当日 8×H100 环境下的观测结果。Partial Key Offset一层注意力也能完成归纳从键偏移说起在仅含因果掩码的标准注意力里token 位置 t 只能看到 ≤ t 的键。若想让注意力抄写前一个 token例如序列abc...中预测b时要能直接关注到a的表示通常需要多层叠加第一层把相邻信息搬运进表示后续层再利用位置编码实现归纳。而key offset的做法更直接把 Key 张量在时间维上整体前移一个位置使位置 t 的 query 恰好能对齐位置 t-1 的 key对 stationary 维而言k[t] k[t-1]从而在单层注意力内就建立了相邻 token 键值对齐的捷径让模型仅凭一层即可完成 1-layer induction。为什么只偏移静止维本纪录的关键洞察是key offset 只应用于静止的 head 维stationary head dims而不是所有维度。在 head_dim 128 的配置下前 64 维承载 RoPE 旋转其中前 32 维为完整旋转频率、后 32 维为半截断旋转的零频率维而后 64 维是静止的。README 中明确说明The partial key offset is only applied to the stationary head dims (32-64 and 96-128). This was found to perform better than applying it to all dims. This approach gives the queries more freedom to attend to multiple positions at once through a single key.也就是说只偏移静止维后同一个 key 的旋转维仍携带自身位置的精确信息而静止维则携带上一个 token 的语义信息。这样一来query 可以同时通过旋转维关注当前位置、通过静止维关注上一位置一个 key 同时服务于两个相邻位置给 query 更大的注意力自由度若把所有维度都偏移则会破坏位置唯一性效果反而更差。实现代码记录脚本内嵌版原 README 给出了最核心的两行实现位于注意力 forward 中、RoPE 之后见记录脚本150d40bf...txt第 943-946 行if key_shift: # shift keys forward for the stationary head dims. Enables 1-layer induction. k[:, 1:, :, self.head_dim//4:self.head_dim//2] k[:, :-1, :, self.head_dim//4:self.head_dim//2] k[:, 1:, :, self.head_dim//4self.head_dim//2:] k[:, :-1, :, self.head_dim//4self.head_dim//2:]以 head_dim128 代入head_dim//4 32、head_dim//2 64所以这两行分别平移了[32:64)与[96:128)两个区间——恰好是两个静止的 32 维块合计 64 维。k[:, 1:, ...] k[:, :-1, ...]表示位置 t 取位置 t-1 的值首位置t0保持自身值不变。何时启用只给长滑动窗口关键设计决策是并非所有层都启用 key shift只作用于承担归纳任务的长窗口层。在记录脚本的GPT.forward中150d40bf...txt第 1079-1081 行bm_sizes [short_bm, short_bm, short_bm, long_bm, short_bm, short_bm, None, short_bm, short_bm, short_bm, long_bm] assert len(bm_sizes) self.num_layers key_shift [blong_bm for b in bm_sizes] # apply key shift to long windows即 11 层中第 3 层与第 10 层为长窗口层其余为短窗口第 6 层无注意力key shift 仅在这两个长窗口层开启。README 给出了背后的实验依据The lowest loss was achieved when applying it to every layer, but the cost:speed ratio seems best when only applying to the long windows, which are the ones primarily responsible for induction.也就是说全层应用能得到最低损失但会引入额外开销而长窗口正是归纳induction的主力把 key offset 限定在长窗口层能在每步约 60ms的极致速度预算下取得最好的性价比。README 还补充了消融结论只对部分 head 子集应用 key offset 效果更差。从纪录脚本到正式仓库Triton 内核化纪录时的实现是 PyTorch 的切片赋值由torch.compile融合编译而这一机制在后来的正式仓库中被下沉为 Triton kernel并演进出更细的规则。这为我们理解其工程演化提供了绝佳参照。正式模型中的开关与分层规则在 track_1_short/model/gpt.py 中模型现在用LONG_WINDOW_LAYERS显式声明长窗口层并把 key offset 与窗口一一绑定第 560-562 行# sliding-window sizes and key shift: the long windows get the partial key offset bm_sizes [ws_long if i in LONG_WINDOW_LAYERS else ws_short for i in range(self.num_layers)] key_offset [i in LONG_WINDOW_LAYERS for i in range(self.num_layers)]而 track_1_short/model/attention.py 第 37-38 行给出LONG_WINDOW_LAYERS (3, 10)——与纪录脚本中的两个长窗口层完全一致该文件中AttnArgs携带key_offset: bool字段CausalSelfAttention.forward中还有一条重要前置断言第 174 行# The partial key offset shifts a keys non-rotating dims, so only a head with some can carry it. assert not key_offset or self.qk_dim Yarn.ROTARY_DIM即只有 qk_dim 大于旋转维64的 head 才允许做 key offset——这与 README 中只偏移静止维的语义严格对应。Triton 内核中的逐元素实现正式实现位于 track_1_short/perf/kernels/qkv_rope.py 的_qk_norm_rope_forward_kernel第 72-82 行if KEY_OFFSET: # Partial key offset: a keys stationary (non-rotating) dims come from the previous token. shift_row (input_head num_heads) (token 0) previous_token tl.maximum(token - 1, 0) x_previous tl.load( qk previous_token[:, None] * stride_qkt input_head[:, None] * stride_qkh offs_d[None, :], maskmask shift_row[:, None], other0.0, ).to(tl.float32) previous_rstd tl.rsqrt(tl.sum(x_previous * x_previous, axis1) / qk_dim 1.1920928955078125e-7) shift shift_row[:, None] (offs_d[None, :] rotary_dim) y tl.where(shift, x_previous * previous_rstd[:, None], y)这里 QK 是打包成[tokens, 2*num_heads, qk_dim]的单个张量前半是 Q、后半是 K因此input_head num_heads即 K 行内核在归一化与 RoPE 的同一个 epilogue 里完成 key offsetshift_row限定只处理 K 行且 token 0从token-1行重新加载原始 QK 值并按 qk_dim 重新计算 rstd因为被平移的维同样经过 QK normshift掩码offs_d rotary_dim精确限定只有非旋转静止维才被替换为上一 token 的值最后把旋转维与静止维拼回同一个 y 输出。这一实现与 README 中仅静止维、时间前移一位、首位置保持不动的语义完全一致。此外qk_norm_rope_forward外层第 97-99 行还有两个工程细节assert not (paired and key_offset)paired head 不兼容 key offset以及block_m 4 if key_offset else 8偏移需要多读一行 token因此减小 tile 尺寸。向后兼容的PackedFP8QKVFunction也把key_offset作为参数贯穿 FP8 训练路径。其余四项更新逐条拆解1. 第 0 层残差缩放合并原实现中每一层都执行x resid_lambdas[i] * x x0_lambdas[i] * x0其中x0是归一化前的首层输入。由于第 0 层的x0_lambdas[0]恰好作用于x0自身可以把两项合并为一次乘法。纪录脚本第 1110-1113 行if i 0: x (resid_lambdas[0] x0_lambdas[0]) * x else: x resid_lambdas[i] * x x0_lambdas[i] * x0这里resid_lambdas与x0_lambdas均来自同一个可学习scalars参数初始化1.1 * ones(num_layers)与0 * ones(num_layers)见150d40bf...txt第 1035-1047 行。合并后第 0 层少一次乘加且由于二者是同一个张量的相邻切片(resid_lambdas[0] x0_lambdas[0])可以提前融合计算对 11 层模型的代码结构也更清晰。2. 减少 50 步训练训练步数由原来的 ~2160 步缩减为 2110 步num_scheduled_iterations 2070num_extension_iterations 40见150d40bf...txt第 1329-1331 行。按 README 所述每步 60ms计算50 步约合 3 秒考虑到 key offset 本身带来的额外计算最终总时长仍达到 128.8s。这体现了该类纪录脚本的典型权衡在不破坏最终验证损失的前提下通过裁剪尾部低收益步数换取墙钟时间验证部分见第四节。3. 对齐 batch size 与窗口调度调度代码的核心是get_ws与get_bs均按同一比例x step / num_scheduled_iterations分段150d40bf...txt第 1442-1458 行def get_ws(step: int): if step args.num_scheduled_iterations: return args.ws_final // 2, args.ws_final x step / args.num_scheduled_iterations ws_idx int(len(args.ws_schedule) * x) return args.ws_schedule[ws_idx] // 2, args.ws_schedule[ws_idx] def get_bs(step: int): if step args.num_scheduled_iterations: return args.train_bs_extension x step / args.num_scheduled_iterations bs_idx int(len(args.train_bs_schedule) * x) return args.train_bs_schedule[bs_idx]三项调度ws_schedule (3, 7, 11)、train_bs_schedule (8*2048*8, 16*2048*8, 24*2048*8)、LR 调度都使用同一个分段边界。README 特别指出更新前的 bug0.33 ! 1/3——旧代码用0.33之类的近似常量做分段比例导致 batch size 切换时刻与窗口尺寸更新时刻错位修正为按int(len(schedule) * x)统一计算后两个调度的切换点精确对齐避免在窗口/批大小不匹配的过渡期浪费步骤。4. 初始 value embeddings 清零模型包含 3 个 token value embeddingsself.value_embeds按012...012结构分布于各层纪录脚本在第 1022-1023 行对它们做零初始化for embed in self.value_embeds: nn.init.zeros_(embed.weight)README 的评价是影响非常轻微但可能降低训练方差且零初始化符合最低假设配置lowest assumption config原则——初始时不向 V 注入任何先验信息让网络自行决定何时、如何利用 value embedding。计时与验证用统计检验验收提速纪录型训练脚本的验收必须同时回答两个问题损失是否真的更低计时是否可信原 README 给出了可复制的验证脚本scipy.statstorchimport scipy.stats import torch losses [3.2788,3.2774,3.2786,3.2792,3.2762,3.2769,3.2781,3.2778,3.2761,3.2783,3.2809] times [128.892,128.907,128.912,128.844,128.822,128.869,128.818,128.882,128.886,128.95,128.946] print(p%.4f % scipy.stats.ttest_1samp(losses, 3.28, alternativeless).pvalue) # p0.0004 print(losses:, torch.std_mean(torch.tensor(losses))) # losses: (tensor(0.0014), tensor(3.2780)) print(time:, torch.std_mean(torch.tensor(times))) # time: (tensor(0.0443), tensor(128.8844))解读如下11 次独立运行的最终验证损失均值3.2780 ± 0.0014全部落在 3.28 下方对 H0真实均值 ≥ 3.28做单侧 t 检验p 0.0004远小于常规显著性阈值说明该配置显著低于 3.28在统计上是稳健的11 次运行墙钟时间均值128.88 ± 0.044 秒方差极小约 0.3%表明计时高度可复现。README 还补充了旧纪录的复测数据131.2: [131.270, 131.241, 131.213]并注明这次似乎拿到了一台略慢的机器appears I got a slightly slower machine this time——即 128.8s 相对 131.2s 的约 2.4s 提升中部分窗口是在不同批次硬件上测得的这也是所有秒级纪录对比都必须声明的环境前提。如何在当前仓库中查看与复现该纪录的完整训练脚本保存在 records/track_1_short/2025-12-14_PartialKeyOffset/ 目录含 README 与 11 份脚本运行副本脚本开头会自动把自身源码写入日志以便审计。复现要点与环境要求如下硬件8×NVIDIA H100 80GB记录日志中的nvidia-smi快照显示为 H100 80GB HBM3CUDA 12.6 驱动 560.35.03软件栈Python 3.10.12、PyTorch 2.10.0.devCUDA 12.6 编译、Triton 3.6.0数据为 fineweb10B 的.bin分片data/fineweb10B/fineweb_train_*.bin分布式启动脚本依赖torchrun设置RANK/WORLD_SIZE/LOCAL_RANK要求world_size为 8 的约数grad_accum_steps 8 // world_sizePYTORCH_ALLOC_CONFexpandable_segments:True已内置关键超参num_iterations 21102070 调度步 40 扩展步、cooldown_frac 0.55、block_size 128、ws_schedule (3, 7, 11)、ws_final 13、train_bs_schedule (8*2048*8, 16*2048*8, 24*2048*8)、val_tokens 10485760模型50257 词表向上取整到 128 的倍数、11 层、6 头、head_dim128、model_dim768优化器为DistAdamembed/scalars/lm_headlr0.008与NorMuonattn/mlp/gatelr0.023weight_decay1.2双优化器Muon 动量按步数做 warmup/cooldown且 LR 与动量调度全部与步数联动。若希望直接体验 key offset 在当前主干模型中的实现可阅读 track_1_short/model/gpt.py分层窗口与key_offset的绑定、track_1_short/model/attention.pyAttnArgs.key_offset与CausalSelfAttention的消费以及 track_1_short/perf/kernels/qkv_rope.pyTriton 内核中的逐元素实现与PackedFP8QKVFunction的 FP8 路径。小结Partial Key Offset 是 modded-nanogpt 速度纪录迭代中的一个典型样本它用只偏移 Key 的静止维、且只作用于长窗口层两个约束把单层归纳这一能力以近乎零额外成本的代价嵌入注意力换取更低的验证损失再配合第 0 层残差缩放合并、步数裁剪、调度对齐与 value embedding 零初始化四项小改进合计约 2.4s最终达成 128.8s 的 124M 训练纪录。更重要的是README 给出的统计验收脚本提醒我们秒级纪录的每一个数字都应该用多次运行的标准差与假设检验来背书。这份从实验到工程、再到统计验收的完整闭环正是 modded-nanogpt 持续刷新纪录的方法论所在。赞分享人工智能大模型预训练分布式训练模型优化深度学习【免费下载链接】modded-nanogptNanoGPT (124M) in 90 seconds项目地址https://gitcode.com/GitHub_Trending/mo/modded-nanogpt点击查看免费下载相关推荐Webpack Bundle Size Analyzer常见问题解决压缩后大小显示不准确的终极指南Webpack Bundle Size Analyzer常见问题解决压缩后大小显示不准确的终极指南 Webpack Bundle Size Analyzer是modded-nanogpt 的 FlexAttention 记录解析用 64K 上下文块级掩码把 NanoGPT 提速到 5 分钟modded nanogpt 的 FlexAttention 记录解析用 64K 上下文块级掩码把 NanoGPT 提速到 5 分钟 本文以 modded n人工智能大模型预训练分布式训练模型优化深度学习modded-nanogpt 稀疏注意力门控Sparse Attention Gate解析替代 Attention Sink 的上下文感知机制与 3.28 验证记录modded nanogpt 稀疏注意力门控Sparse Attention Gate解析替代 Attention Sink 的上下文感知机制与 3.28人工智能大模型预训练分布式训练模型优化深度学习上一篇老Mac焕新终极指南用OpenCore Legacy Patcher轻松升级最新macOS下一篇ONNX Runtime 推理排障五步定位问题完全指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表