尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

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

梯度检查点+梯度累积:train-llm-from-scratch 如何省出 LLM 训练显存——全技巧完整指南

梯度检查点+梯度累积:train-llm-from-scratch 如何省出 LLM 训练显存——全技巧完整指南 梯度检查点梯度累积train-llm-from-scratch 如何省出 LLM 训练显存——全技巧完整指南【免费下载链接】train-llm-from-scratchA straightforward method for training your LLM, from downloading data to generating text.项目地址: https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratchtrain-llm-from-scratch 是一个从数据下载到生成文本的完整 LLM 训练项目本文带你一次搞懂它的两大显存节省核心技巧——梯度检查点Gradient Checkpointing与梯度累积Gradient Accumulation并附上混合精度等可叠加技巧让中小显存显卡也能跑起数亿参数的模型 为什么训练 LLM 会爆显存训练时显存主要被三块吃掉显存去向大致规模能否靠本文技巧压缩参数 优化器状态约 16 字节/参数fp32 权重梯度Adam 两个动量❌ 固定开销激活值前向中间结果随批大小、序列长度暴涨✅ 梯度检查点大批次的瞬时占用随 batch 线性增长✅ 梯度累积其中最有杀伤力的是激活值教学版注意力会把每层(B, n_head, T, T)的注意力矩阵物化出来序列长度 T 一翻倍这部分显存直接涨好几倍所以思路很清晰省激活值用梯度检查点省批大小用梯度累积。技巧一梯度检查点——用计算换显存 梯度检查点的原理一句话前向时不保存每个 Transformer Block 的中间激活反向传播到该 Block 时重新算一遍。显存占用从保存所有层激活降到只保存少量检查点代价是多约 1/3 的前向计算。项目里它被做成了一个开关核心实现只有两行见 src/models/transformer.py 中的标志位前向循环中self.gradient_checkpointing为真时对每个 block 调用torch.utils.checkpointsrc/models/transformer.py并用use_reentrantFalse保证兼容性。启用方式默认关闭行为与数值完全不变配置文件config/config.py 中把USE_GRADIENT_CHECKPOINTING True命令行--grad-checkpointing定义于 scripts/train_transformer.py 适用场景显存差一截、又愿意多花点训练时间时。这是无损技巧——收敛结果不变只是慢一些。技巧二梯度累积——不涨显存也能等效大 batch 梯度累积的思路把一个等效大 batch 拆成 N 个 micro-batch逐个前向反向把梯度累加起来全部算完才执行一次optimizer.step()。显存只需容纳一个 micro-batch训练效果却等价于大 batch。训练循环在 scripts/train_transformer.py先optimizer.zero_grad(set_to_noneTrue)对每个 micro-batch 计算loss / grad_accum关键除以 N 后累加出的梯度正好等于整个大 batch 的平均梯度N 个 micro-batch 全部累加完后统一做梯度裁剪clip_grad_norm_ max_norm1.0再optimizer.step()等效批大小的公式非常好记等效 batch batch_size × 序列长度 × grad_accum× GPU 数真实案例来自 docs/02_pretraining.md在 2×H100 上训练约 4 亿参数模型时因为注意力矩阵占显存大头作者选择--batch_size 8 --grad_accum 12的组合——单卡轻松塞下 batch 8等效 batch 再用累积补齐两卡吞吐约 32k tokens/s。启用方式配置文件config/config.py 中GRAD_ACCUM_STEPS命令行--grad-accum N定义于 scripts/train_transformer.py⚠️ 注意调大grad_accum会增大等效 batch可能需要同步调整学习率单步耗时也会近似变为 N 倍。技巧三叠加这 3 招显存再省一截 ️train-llm-from-scratch 还内置了配套工具和上面两招叠加效果最佳混合精度AMP--amp启用 bf16/fp16 自动混合精度默认 bf16CUDA 专用前向计算显存近乎减半bf16 动态范围足够无需 GradScaler只有 fp16 才启用scripts/train_transformer.py显存预算预览--report-memory在开训前打印参数优化器状态的显存估算帮你预判会不会 OOMscripts/train_transformer.py输出里还会贴心提示reduce with --grad-checkpointing / --grad-accum内存碎片治理多卡启动时加PYTORCH_CUDA_ALLOC_CONFexpandable_segments:True减少显存碎片导致的额外占用docs/howto/train.md快速上手命令速查若需克隆仓库git clone https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratch# 单卡预训练开显存优化全家桶 python scripts/train_transformer.py --grad-checkpointing --grad-accum 8 --amp --report-memory # 双卡预训练小 batch 大累积等效 batch 8 × 12 × 2 PYTORCH_CUDA_ALLOC_CONFexpandable_segments:True \ torchrun --standalone --nproc_per_node2 scripts/pretrain_base.py \ --batch_size 8 --grad_accum 12 --train_steps 50000参数优先级CLI 参数 配置文件全部开关默认关闭不影响原始教学行为scripts/train_transformer.py。想深入原理建议阅读项目自带的基础教程尤其是优化与训练系统章节 docs/foundations/optimization.md怎么确认省显存后训练没变味✅看峰值显存每次评估步都会打印Peak VRAM allocated / reservedscripts/train_transformer.py对比开/关优化即可量化收益看 loss 曲线只要 loss 稳定下降、dev loss 无异常说明优化开关没有破坏收敛开启 AMP/累积后训练步数变慢属正常现象一句话总结 技巧省什么代价开关梯度检查点激活值显存约 1/3 额外计算--grad-checkpointing梯度累积单批瞬时显存单步耗时 ×N--grad-accum N混合精度计算显存几乎无--amp显存不够先开--grad-accum缩 micro-batch再叠--grad-checkpointing最后补--amp——三连招组合拳消费级显卡也能把 LLM 跑起来 【免费下载链接】train-llm-from-scratchA straightforward method for training your LLM, from downloading data to generating text.项目地址: https://gitcode.com/GitHub_Trending/tr/train-llm-from-scratch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表