arXivDaily arXiv每日学术速递 周一至周五更新
arXiv周末暂无论文更新,休息一下吧,周末愉快~~
arXiv 2609.27193cs.DC

LayerCheck:面向大语言模型后训练的逐层自适应检查点技术

LayerCheck: Adaptive Layer-wise Checkpointing for Large Language Model Post-training

  • University of Delaware(特拉华大学)
  • RIKEN Center for Computational Science(理化学研究所计算科学中心)
  • Pacific Northwest National Laboratory(太平洋西北国家实验室)

机构由 AI 辅助整理,请以论文原文为准。

Minqiu Sun, Xin Huang, Luanzheng Guo, Nathan R. Tallent, Kento Sato, Dong Dai

AI总结:

针对LLM训练中检查点开销大且各层更新不均的问题,提出LayerCheck逐层自适应检查点框架,选择性持久化更新超阈值的层,实现检查点大小减少22.6倍、训练时间减少1.31倍,且恢复后损失偏差仅0.54%。

AI中文摘要:

随着训练大语言模型(LLM)的计算和资金成本不断上升,检查点技术(周期性地存储模型状态以便恢复)对于容错变得至关重要。传统检查点方法在检查点频率(I/O开销)和计算恢复(恢复时间)之间存在严重权衡。最先进的方法通过流水线化检查点I/O、差分检查点或内存持久化来缓解这一成本,但均未利用LLM训练动态的独特特征,即模型权重更新在Transformer各层之间呈非均匀分布。这一观察表明,每次保存所有权重可能并不高效。受此启发,我们提出了LayerCheck,一个逐层自适应检查点框架,它选择性地持久化更新超过阈值的层。该设计通过随时间分散逐层检查点写入,避免了周期性的I/O突发,从而产生更平滑、更均衡的I/O剖面。恢复时,LayerCheck通过聚合每层最近持久化的版本及其匹配的优化器状态,重建一个混合时间戳的复合模型状态。在有界逐层陈旧性保护下,这引入了受控扰动:在标准Adam假设下,它增加了一个有界陈旧性项,且经验上,重启后的损失与无故障轨迹的偏差最多为0.54%。在多个开源LLM及不同数据集上的实证结果进一步表明,恢复后的模型保持了原始收敛行为和准确性,同时大幅降低了检查点开销。具体而言,与最先进系统相比,LayerCheck实现了总检查点大小最多减少22.6倍,端到端训练时间减少1.31倍,显著降低了检查点成本。

英文摘要:

With the rising computational and monetary costs of training large language models (LLMs), checkpointing---periodically storing model states for recovery---becomes essential for fault tolerance. Conventional checkpointing entails a severe trade-off between checkpoint frequency (I/O overhead) and computational recovery (recovery time). State-of-the-art approaches mitigate this cost through pipelining checkpoint I/Os, differential checkpointing, or in-memory persistence, yet none leverage the distinct characteristics of LLM training dynamics, where model weight updates are non-uniformly distributed across transformer layers. This observation implies that saving all weights each time might not be efficient. Inspired by this observation, we present LayerCheck, a layer-wise adaptive checkpointing framework that selectively persists layers whose updates exceed a threshold. This design avoids periodic I/O bursts by distributing layer-wise checkpoint writes over time, resulting in smoother and more balanced I/O profiles. Upon recovery, LayerCheck reconstructs a mixed-timestamp composite model state by aggregating the most recently persisted versions of each layer together with their matching optimizer states. Under a bounded per-layer staleness guard, this introduces a controlled perturbation: under standard Adam assumptions it adds a bounded staleness term, and empirically the post-restart loss deviates from the failure-free trajectory by at most 0.54%. Empirical results on multiple open-source LLMs with different datasets further demonstrate that recovered models preserve the original convergence behavior and accuracy while substantially reducing checkpoint overheads. Specifically, LayerCheck achieves up to 22.6x reduction in total checkpoint size and 1.31x reduction in end-to-end training time compared to state-of-the-art systems, significantly lowering the cost of checkpointing.

补充信息

↑