使用Accelerator从检查点恢复训练导致损失值上升
从检查点恢复Stable Diffusion微调训练后损失突变、生成全噪声图像的排查问题
项目背景与问题
- 正在开展Stable Diffusion微调项目,引入自定义布局条件;冻结Hugging Face Stable Diffusion管线中所有组件,仅Unet和自定义的LayoutEmbeddeder保持可训练状态
- 训练过程中代码崩溃,从实现的检查点恢复训练后出现以下异常:
- 损失值大幅跳升(日志可见明显突变)
- 生成的验证图像完全是噪声,效果远差于训练期间的记录结果,甚至不如微调前的基础模型
相关信息
- 检查点通过Accelerator钩子实现,代码位于项目的
main.py中 - 检查点目录结构:
- 根目录包含多个以
checkpoint-xxx命名的子文件夹 - 单个checkpoint文件夹内包含
pytorch_model.bin、config.json等文件 - 文件夹内还存在Accelerate相关的状态文件
- 根目录包含多个以
- 恢复训练后的验证图像:全噪声输出,无有效内容
排查建议
- 检查检查点的参数范围:确认Accelerator是否仅保存了未冻结的Unet和LayoutEmbeddeder参数,而非整个模型。若错误保存并加载了冻结层(如Text Encoder、VAE)的参数,会覆盖预训练权重,直接导致生成失效
- 验证参数加载的正确性:恢复训练后,打印未冻结层的参数哈希值,对比训练中断前的参数,确认参数是否正确加载;同时检查冻结层的参数是否与预训练模型一致,避免被意外修改
- 核对优化器状态的保存与加载:Accelerator的检查点默认包含优化器状态,需确认恢复时是否同时加载了优化器状态。若仅加载模型参数而丢弃优化器状态,训练会失去之前的梯度更新状态,导致损失突变
- 确认随机种子一致性:恢复训练时必须设置与初始训练完全相同的随机种子,包括数据加载、噪声采样、模型初始化等环节的种子,否则会导致训练分布偏移
- 检查检查点文件完整性:验证
pytorch_model.bin的文件大小是否与训练中断时的预期一致,可尝试单独加载检查点进行推理,看是否能复现训练期间的生成效果,判断文件是否损坏 - 核对分布式训练配置:若使用多GPU训练,恢复时的GPU数量、进程数、Accelerate的分布式配置必须与初始训练一致,否则会出现参数分片不匹配的问题
- 添加检查点验证逻辑:在保存检查点后,立即加载并做一次小批量推理,验证生成结果是否正常,提前发现保存逻辑的问题
内容的提问来源于stack exchange,提问作者Suemayah Eldursi
相关产品推荐
相关产品推荐

