模型精选

NanoGPT 训练纪录刷新至 39.9 秒,Deven Pzak 引入 flop 级稀疏优化

精选理由

Deven Pzak 把 NanoGPT 训练干到 39.9 秒,思路不是堆算力而是按 flop 精打细算,嵌入参数撒到 65B 还只用 124M 活跃参数,搞训练的都该看看。

Deven Pzak 将 NanoGPT 在 8xH100 上的训练纪录从 67.6 秒降到 39.9 秒。核心思路是 flop 级优化:跳过低价值计算,包括采样 softmax(约省 8 秒)、稀疏优化器更新、n-gram 表分片通信等手段。配合 EMA(约省 4 秒)和新优化器 Anvil2(约省 1 秒)进一步提升。稀疏嵌入参数规模从 640M 扩到 65B,占总收益的 25%。

原文 · Thomas Wolf

impressive Larry Dial @classiclarryd New historic NanoGPT record at 39.9s (-27.7s) from @DevenPzak , obliterating the prior record of 67.6s! This record introduces a new paradigm of thinking to NanoGPT: instead of optimizing matmuls or adding more expressive operations, optimize at the individual flop level with incredibly clever engineering and ML judgement. If a flop is low value on a particular step, skip it. Specifically: -(~8s) Sampled softmax. If a token doesn’t appear in a batch, skip its lm_head fwd/bwd some fraction of the time. -Sparse values. Only run an optimizer step for ngram embeddings that occurred in the batch. Set beta1 to zero to enable this. Beta2 is applied retroactively when the row is later used. -Sparse updates. Only update ngram and value embeddings once every 4 steps instead of once every 2. -Sparse communication. Shard the n-gram table across GPUs, and only pass the rows receiving updates on each step. -Sparse optimizer states. For the n-gram table, reduce from 2 floats in Adam optimizer per param, to 1 float per 768 params. -Hand-rolled flash attention for 64 dim heads. There are several additions that add accuracy too: -(~4s) EMA during last 300 steps, combined with lifting final_lr to 0.3 instead of 0.15. -(~1s) A new optimizer, Anvil2, which expands muon via a second tracked momentum buffer, improves the ortho coefficients, and modifies the cautious weight decay application. -A couple additional dynamic skip connections in the network. The most striking consequence of the ‘flop aware paradigm’ is you can grow parameters arbitrarily large, only limited by the available memory, since you can selectively choose how to expend flops on those parameters on each step. NanoGPT has kept active parameters below 124M, but total is unbounded, and has grown to 640M through embedding sparsity over the last year. This PR takes that to its logical conclusion on the 8xH100, scaling up to 65B sparse embedding parameters, which accounts for 25% of the PR’s gains. At frontier scale, where one is not bounded by an 8xH100, one could imagine where this paradigm could lead. github.com/KellerJordan/m… As this was a very notable PR, I spoke with Deven for an hour to learn how he did it. Here’s his story on the changes: hyperstition.cc/training-nanog… 🔗 View Quoted Tweet 💬 1 🔄 0 ❤️ 2 👀 1103 📊 1 ⚡