ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

modded-nanogpt 的 FlexAttention 记录解析:用 64K 上下文块级掩码把 NanoGPT 提速到 5 分钟

modded-nanogpt 的 FlexAttention 记录解析:用 64K 上下文块级掩码把 NanoGPT 提速到 5 分钟 人工智能大模型预训练分布式训练模型优化深度学习【免费下载链接】modded-nanogptNanoGPT (124M) in 90 seconds项目地址https://gitcode.com/GitHub_Trending/mo/modded-nanogpt点击查看免费下载本文以 modded-nanogpt 仓库 2024-11-19 FlexAttention 记录 及其完整运行日志8384493d-…txt为主体结合仓库源码还原这次“第 12 号世界纪录”的完整技术细节它如何用 PyTorch 2.6 的torch.nn.attention.flex_attention把注意力上下文从 1024 扩展到 64K、并叠加文档因果掩码与 1024 滑动窗口在 8 张 NVIDIA H100 上把 FineWeb 验证损失压到 3.28 以下。读完本文你将掌握 FlexAttention 的create_block_mask掩码函数定义与编译接入方式、64K 序列下的超参数batch size、序列长度、迭代数与 warmdown设计以及本记录方差波动约 0.005 stddev的来源分析可作为复现与进一步优化的直接参考。记录背景从“1024-ctx 稠密因果注意力”到“64K-ctx FlexAttention”modded-nanogpt 是一个公开的 NanoGPT speedrun 项目目标是“用 8 张 NVIDIA H100 训练出在 FineWeb 验证集上达到 3.28 交叉熵损失的模型”详见仓库根目录的 README.md。2024-11-19 的记录是速度榜单上的第 12 号记录由 KoszarskyB 贡献其官方描述为1024-ctx dense causal attention → 64K-ctx FlexAttention在它之前2024-11-10 的 U-Net 记录模型只能以 1024 token 的稠密因果注意力上下文训练本记录将注意力上下文一举扩展到 64K tokensequence_length 64*1024把上一记录的 7.2 分钟纪录推进到5.03 分钟这是把“更长的上下文”与“更稀疏的注意力模式”结合起来的代表性尝试。与记录配套的 README.md 明确记录了两个要点本文后续会逐一展开本训练的 run-to-run 方差约为 0.005 stddev因此并非每次运行都能降到 3.28 以下但均值约在 3.279该方差很可能源于上一记录U-Net 双倍学习率叠加本次显著缩短训练时长后的效应。核心依赖PyTorch 2.6 nightly 与flex_attention的编译接入本次运行日志的环境信息位于记录日志文件开头部分显示实验跑在Running pytorch 2.6.0.dev20241119cu124 compiled for CUDA 12.4、8 张 NVIDIA H100 80GB驱动 555.42.06 / CUDA 12.5之上。FlexAttention 是 PyTorch 2.5 引入的原生可编程注意力 API本记录直接复用了它的两个入口from torch.nn.attention.flex_attention import flex_attention, create_block_mask flex_attention torch.compile(flex_attention, dynamicFalse) create_block_mask torch.compile(create_block_mask, dynamicFalse)这段代码位于 记录日志文件 的开头部分。关键点有二对flex_attention本身做torch.compileflex_attention本身是一个逐元素注意力模板函数只有经过torch.compile编译才会针对给定掩码结构生成专门的 Triton 内核这是它能以 64K 序列长度跑出实用速度的前提dynamicFalse表示按静态形状编译避免动态 shape 带来的重编译开销。create_block_mask同样被编译掩码本身也从 Python 级掩码函数降级为可执行内核后续会看到它接受了_compileTrue参数二者共同把“掩码构造”这一步的开销从宿主端搬到了编译后的 GPU 内核上。需要说明的是这条记录处于 PyTorch 2.6.0 的开发版阶段所用 API 的签名与稳定版可能存在差异在更早的 PyTorch 2.5.x 中enable_cudnn_sdp(True)也已被本记录用于选择 cuDNN 注意力后端详见下文“注意力后端选择”小节。当前仓库主线已演进到 track_1_short/model/attention.py 的 FlashAttention-3 varlen 实现但在本记录当时FlexAttention 正是把长上下文落地的关键组件。掩码设计document causal mask 1024 滑动窗口掩码函数三段条件的与本记录的核心创新在于把1024-ctx 稠密因果注意力替换为一个更灵活的组合掩码。GPT.forward中定义了如下掩码函数完整上下文见 记录日志文件docs (idx 50256).cumsum(0) def document_causal_mask(b, h, q_idx, kv_idx): causal_mask q_idx kv_idx document_mask docs[q_idx] docs[kv_idx] window_mask q_idx - kv_idx 1024 return causal_mask document_mask window_mask S len(idx) block_mask create_block_mask(document_causal_mask, None, None, S, S, devicecuda, _compileTrue)三个条件的语义causal_mask q_idx kv_idx标准因果性query 只能看自己及之前的位置document_mask docs[q_idx] docs[kv_idx]文档级隔离。docs (idx 50256).cumsum(0)用 GPT-2 的 EOS token id50256对序列做累计计数把每个位置映射到它所属的文档编号掩码要求 query 与 key 属于同一文档从而禁止跨文档注意力——在 64K 长序列中FineWeb 天然包含多条文档这条掩码避免了模型在文档边界“偷看”后续文档window_mask q_idx - kv_idx 10241024 token 的滑动窗口即每个位置最多回溯 1024 个 token等价于把上一记录的稠密 1024 上下文保留下来同时让序列可以远比 1024 更长。最终掩码是三者取与于是 64K 上下文的训练在“稀疏度”上等价于窗口化的稠密注意力但序列内可以容纳完整的多文档结构。create_block_mask与_compileTruecreate_block_mask(document_causal_mask, None, None, S, S, devicecuda, _compileTrue)的第二个、第三个参数分别是 batch / num_heads 维度的掩码函数这里传None即掩码对所有 batch 与所有 head 相同第四、第五个参数是 query 与 key 的序列长度均为S 64*1024。_compileTrue让掩码构造也进入编译路径降低每个 step 重建块掩码的开销。前向中的注意力调用flex_attention(q, k, v, block_mask)掩码在CausalSelfAttention.forward中通过block_mask参数接入注意力计算记录日志文件y flex_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), block_maskblock_mask)q/k/v从 (B, T, n_head, head_dim) 转置为 (B, n_head, T, head_dim) 后送入flex_attentionFlexAttention 会基于block_mask的块稀疏结构跳过整块被屏蔽的区域从而把注意力计算量从O(T²)的稠密开销降到块级稀疏的实际用量这正是 64K 序列长度可行性的根本原因。值得一提的是本记录中的模型还带有一系列当代架构技巧Rotary 位置编码、QK-Norm、ReLU² 激活、zero-init 投影、logit softcap 30、value embedding 混合其中Rotary、apply_rotary_emb、CastedLinear、CausalSelfAttention、MLP、Block与GPT等类的完整实现均可直接在该日志文件中阅读。这些结构细节与本记录主题注意力稀疏化属于不同层面的优化本文不再展开。注意力后端选择cuDNN SDP而非默认 Flash日志中有一段值得注意的注意力后端配置from torch.backends.cuda import enable_cudnn_sdp, enable_flash_sdp, enable_math_sdp, enable_mem_efficient_sdp enable_cudnn_sdp(True) enable_flash_sdp(False) enable_mem_efficient_sdp(False) enable_math_sdp(False)日志注释给出的理由是“CUDNN attention 比 Flash 快约 4ms但在 PyTorch 2.5.1 中默认不会被选中”。也就是说在当时的 PyTorch 版本下torch.compile生成的 flex_attention 内核之外模型其余部分的 SDP 后端被显式锁定为 cuDNN这是记录在案的性能调优动作。64K 序列下的训练配置与超参数记录日志中的Hyperparametersdataclass 给出了本次运行的全部默认超参数记录日志文件参数默认值说明input_bindata/fineweb10B/fineweb_train_*.bin训练数据 shard 通配符input_val_bindata/fineweb10B/fineweb_val_*.bin验证数据 shard 通配符batch_size8全局跨设备batch size单位序列数device_batch_size1每张卡上的序列数sequence_length64*1024单序列长度即本次记录的核心变化点num_iterations1875训练迭代总数warmup_iters0无线性 warmupwarmdown_iters562线性 warmdown 步数约占总迭代的 30%weight_decay0未使用权重衰减val_loss_every125每 125 步评估一次验证损失val_tokens10485760验证 token 数固定为 10,485,760保证跨记录可比save_every0仅在结束时保存 checkpoint在此基础上训练循环使用梯度累积train_accumulation_steps batch_size // (B * ddp_world_size)本记录在 8 卡、每卡 B1 时即为 1。学习率调度为“无 warmup → 常数 → 线性 warmdown”的三段式get_lr(it)见日志文件warmdown 段从num_iterations - warmdown_iters 1313步开始线性降到 0。验证协议验证 token 数固定为val_tokens 10485760即 10M token并断言val_tokens % (B * T * ddp_world_size) 0每 125 步进行一次验证val_steps val_tokens // (B * T * ddp_world_size)。这与仓库 README.md 中“对 FineWeb val set 的前 10,485,760 个 token 赋予至少exp(-3.28 * 10485760)的概率”的目标度量一致。优化器分工本次运行沿用 modded-nanogpt 经典的“AdamW 处理 embedding/lm_head/scalarMuon 处理矩阵参数”分工记录日志文件optimizer1 Adam(wte.weight, lr0.6, betas(0.9, 0.95), fusedTrue)optimizer2 Adam(lm_head.weight, lr0.008, betas(0.9, 0.95), fusedTrue)optimizer3 Muon(matrix_params, lr0.04, momentum0.95)2D 参数Newton-Schulz 5 步正交化optimizer4 Adam(scalar_params, lr0.04, betas(0.9, 0.95), fusedTrue)训练循环还对 Muon 做了 momentum warmupfrac min(step/500, 1)使 momentum 从 0.85 线性过渡到 0.95。Muon 实现zeropower_via_newtonschulz5与分布式更新逻辑同样内嵌在日志文件中可对照阅读。运行结果5.03 分钟达 3.2783且逐 step 计时稳定日志记录了完整的 1875 步训练过程总计 2527 行含环境信息与验证日志。关键结果摘录如下最终验证step:1875/1875 val_loss:3.2783 train_time:301825ms step_avg:161.84ms——总用时约 302 秒即5.03 分钟验证损失3.2783 3.28达成 speedrun 目标计时方式前 10 步约 45 秒用于内核预热从第 11 步开始计时step_avg稳定在约 160-162 ms/步说明 FlexAttention 长序列训练在预热后吞吐非常平稳验证损失轨迹每 125 步10.8258step 0→ 4.4503125→ 4.0085250→ 3.8399375→ 3.7415500→ 3.6656625→ 3.6068750→ 3.5596875→ 3.51951000→ 3.49051125→ 3.46251250→ 3.43211375→ 3.38231500→ 3.33831625→ 3.29971750→ 3.27831875单调收敛到 3.28 以下训练过程中出现了数次训练损失异常尖峰如 step 918 的 4.8807、step 946 的 5.0975、step 1759 的 4.2259属于长序列训练的正常噪声不影响最终收敛。方差分析0.005 stddev 与 3.279 均值记录 README 明确指出README.mdThis training has significant variance of around 0.005 stddev between runs. So not all runs go beneath 3.28, though the mean is around 3.279.这句话的工程含义是单次运行是否低于 3.28 具有随机性统计上多次运行的平均验证损失约 3.279stddev 约 0.005。这意味着以 3.28 为门槛做验证时需要多次运行并取统计证据仓库 README.md 中的规则 2 也要求提交足够多的运行日志达到 p0.01 的显著性水平。README 同时给出了方差来源的推断This variance is probably caused by the previous record doubling the learning rate, plus this record significantly shortening the duration.即上一记录U-Net skip connections double lr2024-11-10引入了双倍学习率而本记录又把训练时长大幅缩短7.2 分钟 → 5.03 分钟两者叠加放大了 run-to-run 的随机性。这对后续实验的启示是当训练时长缩短、学习率偏高时单次运行结果并不能代表方法的真实水平评估时应以多次运行均值为准。与后续记录的关系FlexAttention → Window Warmup作为“长上下文 稀疏注意力”路线的起点本记录直接启发了 2024-11-24 的 Window Warmup 记录4.66 分钟后者在 FlexAttention 的窗口掩码基础上引入了“注意力窗口随训练逐步扩大”的 warmup 调度把窗口从短窗逐步增长到长窗进一步压低了训练时间。而注意力掩码/窗口体系在后续版本中持续演进从 2025-01-16 的 long-short 注意力Sub3Min 记录到 2025-09-03 正式切换到 FlashAttention-3 varlenFA3 记录再到当前主线 track_1_short/model/attention.py 中基于flash_attn_varlen_func与 YaRN 窗口扩展的实现。可以看出本记录验证的“程序化掩码定义长上下文注意力”思路构成了后续所有窗口化注意力优化的事实起点。复现要点与适用前提若希望在本仓库环境中复现该记录需要注意该记录运行于PyTorch 2.6.0.dev20241119cu124CUDA 12.4 编译并依赖当时 nightly 版本中的torch.nn.attention.flex_attentionAPI在更新版本上运行需要核对 API 兼容性。当前仓库主线的运行方式pip install -r requirements.txt后执行 run.sh已经切换到 FlashAttention-3 路线不再直接包含 FlexAttention 代码硬件前提为8 张 NVIDIA H100 80GB每卡一个 batch 序列序列长度 64K全局 batch size 8单次运行存在约 0.005 stddev 的随机性判断是否“达标 3.28”应基于多次运行的均值而非单次结果减少 GPU 数量可修改 run.sh 的--nproc_per_node但 FlexAttention 的 64K 序列长度受显存约束GPU 较少时需要按显存预算调低sequence_length会改变训练行为这一点在仓库 README.md 中亦有说明。综上所述2024-11-19 的 FlexAttention 记录是 modded-nanogpt 里程碑式的转折点它以 PyTorch 原生可编程注意力 API 为工具用一段十几行的掩码函数把 NanoGPT 的上下文从 1024 扩展到 64K在 8×H100 上 5.03 分钟即达 3.28 目标并由此开启了后续一系列窗口化注意力调度的探索。理解这份日志就能理解长上下文稀疏注意力在该 speedrun 项目中的完整演进起点。赞分享人工智能大模型预训练分布式训练模型优化深度学习【免费下载链接】modded-nanogptNanoGPT (124M) in 90 seconds项目地址https://gitcode.com/GitHub_Trending/mo/modded-nanogpt点击查看免费下载相关推荐AnythingSlider社区贡献指南如何参与这款jQuery轮播插件开发与维护AnythingSlider社区贡献指南如何参与这款jQuery轮播插件开发与维护 AnythingSlider 是一款功能强大的jQuery轮播插件自20Apache Arrow 贡献指南如何写出高质量的 Bug 报告与功能请求Apache Arrow 贡献指南如何写出高质量的 Bug 报告与功能请求 导读 Apache Arrow 是一个横跨 C、Python、R、Java、R人工智能大模型预训练分布式训练模型优化深度学习Modded-NanoGPT版本控制记录每次性能突破的代码变更Modded NanoGPT版本控制记录每次性能突破的代码变更 在人工智能模型训练领域版本控制不仅仅是代码管理的工具更是性能优化的历史档案。Modded人工智能大模型预训练分布式训练模型优化深度学习上一篇Apache Maka 计算机使用执行器加固语义绑定、失败关闭与物理输入隔离的演进实践下一篇NocoBase RunJS ctx.resource 完整指南用 FlowResource 在脚本中访问与操作数据创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表