首页 / 资讯中心 / 文章详情

配对多头注意力(Paired Head Attention):modded-nanogpt 提速 3.5 秒的新记录实战解析

配对多头注意力(Paired Head Attention):modded-nanogpt 提速 3.5 秒的新记录实战解析 ★ FEATURED ARTICLE
人工智能大模型预训练分布式训练模型优化深度学习【免费下载链接】modded-nanogptNanoGPT (124M) in 90 seconds项目地址https://gitcode.com/GitHub_Trending/mo/modded-nanogpt点击查看免费下载本文基于 modded-nanogpt 仓库 track 1short赛道 2026-01-07_PairedHeadAttention 记录文档完整拆解配对多头注意力这一架构改动如何让一个 query 同时关注同一位置的多份表征、如何通过 K/Q/V 交错把注意力序列拉长一倍、为什么它能净省约 3.5 秒训练时间-65 steps并结合当前仓库源码model/gpt.py、model/attention.py、perf/kernels/qkv_rope.py还原其从模型配置、RoPE 内核到 varlen FlashAttention-3 调用的完整实现链路。读完本文你可以直接理解 paired-head 的设计动机、参数取舍与源码落点并掌握用 scipy 做同类 A/B 计时验证的方法。记录背景一次零新增参数的架构提速这条记录是 modded-nanogpt 训练竞速124M GPT-2目标 90 秒量级过程中的一次里程碑更新PR 标题为New Record: Paired Head Attention (-3.5s, -65 steps)包含两项改动配对多头注意力Paired Head Attention——核心架构改动也是本文主题更快的 RoPE 实现——将 rotary 计算收敛为单行表达式额外节省约 0.5 秒。文档给出的最终节奏是新方案 5 次运行的均值耗时约 108.80 秒、均值 loss 3.2788与此前记录的重新计时 112.262 秒相比净提升约 3.5 秒、步数 -65 steps。为与 3.5s 的收益口径保持一致记录按 109.2s 的标准口径合并。需要强调的是paired head attention 是在既有模型结构上重排 Q/K/V而非引入新的可学习参数——这也是它能同时带来训练加速的原因之一。配对多头注意力的核心思想让 query 关注同一位置的多种表征标准多头注意力中每个 head 独立投影出各自的 K/Q/Vquery 只能关注到自己那份 key 序列。paired head attention 的出发点是让 queries 关注每一个位置的多份表征并让同一个位置产生多个 logits 进入同一个 softmax。具体实现方式是将相邻两个 head 的 k、q、v 交错排布构成两倍长的序列。文档给出了直观示例原始排布为[k1_h1, k2_h1, k3_h1], [k1_h2, k2_h2, k3_h2]交错之后变成[k1_h1, k1_h2, k2_h1, k2_h2, k3_h1, k3_h2]即不再按先整段 head1、再整段 head2排列而是按位置维度交替位置 1 的 head1/head2、位置 2 的 head1/head2……对 q 和 v 做同样的重排。经过这种交错注意力函数内部的序列长度翻倍head 数量减半两个 head 合并为一对而每个 query 在 softmax 中同时面对两个 head 对同一位置的 key 表征从而让模型在同样的 softmax 窗口内获得更丰富的同位置多视角信息。因果掩码下的交错语义因为数据按顺序存储且注意力使用 causal mask交错直接改写了能看多远的语义。文档以位置 4 为例head 1 的 query 4只能关注到 head 2 的keys 1-3head 2 的 key 4 在时间顺序上更靠后被掩码遮挡head 2 的 query 4可以关注到 head 1 的keys 1-4。换句话说两个配对 head 在回看的视野上天然错开了一步这种不对称恰好构成了配对 head 之间信息互补的基础。实现细节零拷贝 RoPE 与延迟 reshape单行 rotary()PyTorch 可以不复制数据地算 x_flip文档指出 PyTorch 有能力在不进行任何数据拷贝的情况下计算 x_flip因此把 rotary 收敛成了单行实现def rotary(self, x_BTHD): assert self.factor1.size(0) x_BTHD.size(-3) factor1, factor2 ( self.factor1[None, : x_BTHD.size(-3), None, :], self.factor2[None, : x_BTHD.size(-3), None, :], ) x_flip x_BTHD.view(*x_BTHD.shape[:-1], x_BTHD.shape[-1] // 2, 2).flip(-1).view(x_BTHD.shape) return factor1 * x_BTHD factor2 * x_flip关键技巧在于view - flip(-1) - viewview只是改变张量的元数据解释flip(-1)在 PyTorch 中也是惰性的返回带负步幅的视图因此x_flip从未真正物化一份副本最终只产生一次factor1 * x factor2 * x_flip的融合计算。这一改动带来约 0.5s 的提速。为 PairedHeadAttention 定制 rotary把 reshape 推迟到旋转之后标准流程中Q/K 投影后立即 reshape 出头维度reshape 一旦发生后续 rotary 的输出若要再 reshape 就可能触发一次数据拷贝。文档给出的 paired 专属路径刻意把 q/k 的 reshape 推迟到 rotary 之后# delay q,k reshape until rotary makes data contiguous, to enable view (non-copy) q q.view(B, T, self.num_heads // 2, self.head_dim * 2) k k.view(B, T, self.num_heads // 2, self.head_dim * 2) v v.reshape(B, T*2, self.num_heads//2, self.head_dim) q, k yarn.rotary(q), yarn.rotary(k) q q.view(B, T*2, self.num_heads//2, self.head_dim) k k.view(B, T*2, self.num_heads//2, self.head_dim)设计意图是先让 rotary 输出变为连续内存再用view非拷贝完成最终形状变换。注意 v 因为要直接进入交错布局用的是reshape允许拷贝提前排好(B, T*2, heads//2, head_dim)。yarn.rotary(q)在这里作用于配对宽度head_dim * 2的 Q/K旋转后再切成两倍长、一半头的序列——这正是交错发生的地方。交错旋转staggered rotary文档还提到配对用的 rotary 被设计成交错偏移head 1 的旋转相位相对 head 2 有一个偏移offset在有限的测试中表现更好。由于配对的 rotary 表把位置 2t 与 2t1 并排打包奇数/偶数 head 自然读取同一行的不同半段从而获得互相错开的旋转角度当前源码中即factor_offset ((logical_head % 2) * qk_dim)的取半逻辑。序列长度与窗口的处理交错后每条序列长度翻倍因此文档明确要求对输入做配对修正# paired head correction seqlens 2 * seqlens max_len 2 * max_len同时作者有意不改动 window size这意味着原本 window128 的滑动窗口实际上只回看64 个位置head1 回看 64、head2 回看 64合计 128 个 key 槽位。这是一处刻意的权衡——用每 head 视野减半换取同位置双表征进 softmax。层选择策略短窗口层优先逐步加层作者最初不确定 paired head 与长窗口层上的 partial key offset 如何相互作用因此只先应用到部分短窗口层随后逐层追加、持续变好最终在 4 层处停止文档注明或许更多层更好但为控制 PR 范围而收手。这一策略在当前仓库中固化为了明确的层拓扑常量见 track_1_short/model/gpt.pyATTN_LAYERS tuple(i for i in range(NUM_LAYERS) if i not in NO_ATTN_LAYERS) # (0, 1, 2, 3, 5, 8, 10) LONG_WINDOW_LAYERS (3, 10) PAIRED_HEAD_LAYERS (0, 2, 5)即 11 层模型中注意力层为 (0,1,2,3,5,8,10)其中 (3,10) 是带 partial key offset 的长窗口层paired head 只落在短窗口层 (0, 2, 5)与文档应用到部分 short window 层、4 层左右的描述一致。当前仓库中的源码级落点1. 模型配置model/gpt.py每个注意力层按自己的是否配对标志构建模块并按 (query/key 宽度, 是否配对) 维护三张独立的 RoPE 表self.attn nn.ModuleDict({ str(layer): CausalSelfAttention( num_heads, head_dim, qk_dimself.attn_qk_dim(layer), v_dimself.attn_v_dim(layer), val_max_seq_lenmax_seq_len, pairedlayer in PAIRED_HEAD_LAYERS, ) for layer in ATTN_LAYERS }) self.yarn Yarn(NARROW_HEAD_DIM, max_seq_len, attn_scaleNARROW_ATTN_SCALE, deviceself.device) self.yarn_paired_head Yarn( NARROW_HEAD_DIM, max_seq_len, pairedTrue, attn_scaleNARROW_ATTN_SCALE, deviceself.device, ) self.yarn_wide Yarn(head_dim, max_seq_len, attn_scaleWIDE_ATTN_SCALE, deviceself.device) assert not set(PAIRED_HEAD_LAYERS) set(WIDE_QK_LAYERS), no paired rotary table at full width窄 head64 维的 paired 层使用专门的yarn_paired_headattn_scale与普通窄层一致0.13长窗口宽 head 层与 paired 层互斥保证不存在全宽度下的配对 rotary 表。2. 配对 rotary 表的构建model/attention.py配对的 Yarn 表把位置 2t 与 2t1 并排打包进一行所以表仍只有max_seq_len行但每行宽度是2 * head_dimt_even 2 * t t_odd t_even 1 theta1 torch.outer(t_even, self.angular_freq) theta2 torch.outer(t_odd, self.angular_freq) self.factor1[lo:hi].copy_(torch.cat((theta1.cos(), theta2.cos()), dim-1)) self.factor2[lo:hi].copy_(torch.cat((theta1.sin(), theta2.sin()), dim-1))偶数位置与奇数位置的 cos/sin 分列左右两半——头部奇偶偏移stagger正是从这里读出来的。由于每行只依赖位置与当前频率YaRN 窗口变化时的部分重建rebuild_rows与 validation 前的ensure_full()补齐机制对配对表同样适用见 Yarn 类注释与 GPT.complete_yarn_tables。3. 前向的 paired 分支model/attention.pyif not self.paired: if aux_v is not None: v v aux_v else: # Paired heads: adjacent heads queries attend to each others keys. Two copies of the # input stream are interleaved (q, k already are, by the norm/rotary kernel), which # doubles each sequences length and halves the effective window. v v.reshape(B, T * 2, H // 2, self.v_dim) if aux_v is not None: v v aux_v.reshape(v.shape) seqlens 2 * seqlens max_len 2 * max_len y flash_attn_interface.flash_attn_varlen_func(q[0], k[0], v[0], cu_seqlens_qseqlens, cu_seqlens_kseqlens, max_seqlen_qmax_len, max_seqlen_kmax_len, causalTrue, softmax_scaleyarn.attn_scale, window_size(bm_size, 0))可以看到源码注释与文档完全对应q/k 的交错已由 norm/rotary 内核完成见下v 在此处显式 reshape 成(B, T*2, H//2, v_dim)随后seqlens 2 * seqlens、max_len 2 * max_len正是文档中的paired head correction最终交给 varlen FlashAttention-3 调用。4. Triton 内核中的交错索引perf/kernels/qkv_rope.pypaired 交错的底层由融合 QK-norm RoPE 的 Triton 内核实现把逻辑 head 映射到交错后的 (token, head) 坐标——if PAIRED: # Paired heads: head h of token t lands at token 2t h // (num_heads/2), so adjacent heads # attend to each others keys; odd heads read the second half of the paired rotary row. output_token 2 * token logical_head // (num_heads // 2) output_head logical_head % (num_heads // 2) factor_offset ((logical_head % 2) * qk_dim)[:, None]token 交错output_token 2*token h // (heads/2)让相邻两个 head 落到2t与2t1两个 token 槽与文档的 interleaving 图示一一对应head 减半output_head h % (heads/2)rotary 偏移factor_offset使奇数 head 读取配对表行的后半段实现 staggered rotary。该内核同时承担了 RoPE 的寄存器内翻转x_flip在寄存器中形成避免二次读 HBM以及 QK RMS-norm其 backwardqkv_rope.py也用相同的PAIRED分支保持索引一致性。5. 与其它机制的协同与约束与 partial key offset 互斥paired 层的 qk_dim64 等于ROTARY_DIM本身不具备 key offset 所需的静止维度内核与PackedFP8QKVFunction中均有assert not (paired and key_offset)qkv_rope.py与文档不确定与长窗口 partial key offset 如何相互作用的考量一致。XSA 不应用于 paired 层gated XSA 需要按 head 对齐 v 的逐位置形状paired 层 v 形状不同因此仅在非 paired 层生效attention.py。FP8 打包路径的适配训练走 FP8 时paired 层把 packed [Q;K;V] 拆成两次列切片 GEMMQK 一次、V 一次让 V 从诞生起就是稠密布局避免合并 token 轴与 head 轴的物化拷贝qkv_rope.py。计时与验证如何严谨地证明更快且不更差竞速记录最关心两件事耗时是否真的下降、loss 是否没有变差。文档给出了完整的验证脚本scipy.stats 单侧 t 检验 torch 均值/标准差import scipy.stats import torch losses [3.2793, 3.2796, 3.2783, 3.2782, 3.2784] times [108.695, 108.844, 108.775, 108.795, 108.87] print(p%.4f % scipy.stats.ttest_1samp(losses, 3.28, alternativeless).pvalue) # p0.0063 print(losses:, torch.std_mean(torch.tensor(losses))) # losses: (tensor(0.0006), tensor(3.2788)) print(time:, torch.std_mean(torch.tensor(times))) # time: (tensor(0.0679), tensor(108.7958))结果解读loss5 次运行均值 3.2788、标准差仅 0.0006相对 3.28 目标值做单侧 t 检验得到p0.0063说明新方案 loss 显著低于 3.28 基线在统计上是可信的即提速没有以质量下降为代价time均值 108.7958s、标准差 0.0679s运行间抖动极小0.1s 量级计时稳定性很好。作为对照此前记录的重新计时为 112.262s三次运行 112.187 / 112.278 / 112.321因此 paired head attention 的净收益约3.5 秒。文档最后注明若无其他变化将按 109.2s 的口径合并以与 3.5s 的收益声明保持一致。附带的独立实验Adam mantissa文档另外记录了一项独立于本 PR的实验为 Adam 增加 mantissa尾数追踪——bf16 参数只有 7 位尾数通过维护额外的 uint16 mantissa 与参数拼成 fp32 精度做更新这一思路在后续记录中演化为 NorMuon 的 cautious weight decay mantissa 更新 等优化器实现见cautious_wd_and_update_inplace内核。该实验似乎有轻微提升但为保持本 PR 范围聚焦而暂时搁置没有合入 paired head attention 记录。总结Paired Head Attention 是一次典型的重排数据而非增加参数的架构优化通过把相邻 head 的 K/Q/V 按位置交错让每个 query 在同一 softmax 中面对同一位置的多份表征注意力序列长度翻倍、head 数量减半配套的单行 rotary 与延迟 reshape 技巧把内存拷贝降到最低paired 专属的 staggered rotary 提供位置相位的互补偏移序列长度修正seqlens 2 * seqlens、max_len 2 * max_len与窗口不变即每 head 视野减半的取舍共同定义了对滑动窗口语义的改写策略上先只用于部分短窗口层并逐步加层当前仓库固化于PAIRED_HEAD_LAYERS (0, 2, 5)且与 partial key offset 互斥、与 XSA 不叠加通过 scipy 单侧 t 检验p0.0063与毫秒级稳定的计时108.80 ± 0.07s严谨地证明了相对旧记录 112.262s 的约 3.5 秒净收益。对希望复现或继续探索的读者建议从 track_1_short/model/gpt.py 的层拓扑常量入手依次阅读 model/attention.py 的 paired 前向分支与 perf/kernels/qkv_rope.py 的 Triton 交错索引即可完整掌握该机制从配置到内核的每一环。赞分享人工智能大模型预训练分布式训练模型优化深度学习【免费下载链接】modded-nanogptNanoGPT (124M) in 90 seconds项目地址https://gitcode.com/GitHub_Trending/mo/modded-nanogpt点击查看免费下载相关推荐modded-nanogpt 稀疏注意力门控Sparse Attention Gate解析替代 Attention Sink 的上下文感知机制与 3.28 验证记录modded nanogpt 稀疏注意力门控Sparse Attention Gate解析替代 Attention Sink 的上下文感知机制与 3.28人工智能大模型预训练分布式训练模型优化深度学习TypeSpec Java 客户端生成器修复 Javadoc 中 */* 内容类型导致的注释提前终止问题TypeSpec Java 客户端生成器修复 Javadoc 中 / 内容类型导致的注释提前终止问题 导读 本文基于 TypeSpec 仓库中的变更日志 .c人工智能大模型预训练分布式训练模型优化深度学习fairseq 注意力头选择Attention Head Selection实战指南多语言与多领域序列建模fairseq 注意力头选择Attention Head Selection实战指南多语言与多领域序列建模 本指南基于 fairseq 仓库中 examp人工智能深度学习预训练NLP语音上一篇深入解析 go-retryablehttp 变更历史buildkit 中自动重试 HTTP 客户端的演进与源码实现下一篇change_detection.pytorch3 步跑通遥感图像变化检测完整模型创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
阅读完成 · 觉得有帮助?
咨询建站