上周做一个大模型微调项目,训练到一半模型突然崩了,损失飙升到NaN,两天的工作差点白费。
那是一个周四的晚上,我训练了一个GPT-2的文本生成模型,用的是公司的业务数据。训练了6个小时,损失一直在下降,看起来很顺利。我想着再训练两个小时就能收工了,就去吃了个饭。回来一看,监控面板上的损失曲线突然变成了一条直线——NaN。
当时我心里一沉。NaN意味着模型参数已经变成了非数字,训练彻底废了。更糟糕的是,我只保存了最新的checkpoint,而那个checkpoint已经是NaN了。两天的工作,差点全部白费。
经过十几个小时的排查,终于找到了原因,也挽救了大部分工作。本文完整复盘这次故障,从故障发生、排查过程、根因分析、解决方案,到经验总结,帮你避免同样的坑。
一、故障发生
先说说故障是怎么发生的。
1. 项目背景
这个项目是用GPT-2做业务文本的生成微调。数据量大概10万条,每条文本平均200字。用的是Hugging Face的Transformers库,模型是GPT-2 Medium(355M参数),GPU是RTX 3090(24G显存)。
训练配置:
- batch size: 8(梯度累积4步,等效32)
- 学习率: 5e-5
- 训练轮数: 3个epoch
- 优化器: AdamW
- 混合精度: fp16
2. 故障现象
训练了6个小时,大概1.5个epoch,损失从3.2降到了1.8,一切正常。
然后,损失突然开始飙升:
- step 5000: loss 1.8
- step 5010: loss 3.5
- step 5020: loss 15.2
- step 5030: loss NaN
从损失开始异常到变成NaN,只用了30步,大概5分钟。
同时,梯度范数(gradient norm)也飙升到了1000以上(正常应该在1-5之间)。
3. 第一反应
看到NaN的时候,我的第一反应是:完了,模型崩了。
我立刻做了几件事:
- 停止训练,防止覆盖更多的checkpoint
- 查看训练日志,找到损失开始异常的时间点
- 检查保存的checkpoint,看有没有能用的
- 查看GPU状态,确认不是硬件问题
幸运的是,我设置了每1000步保存一个checkpoint,所以step 5000的checkpoint还是好的。虽然损失了500步的训练,但至少不用从头开始。
二、排查过程
接下来是漫长的排查过程。
1. 排查方向一:数据问题
首先怀疑的是数据问题。是不是某条数据有问题,导致模型计算异常?
我做了以下检查:
- 查看step 5000左右的训练数据,有没有异常文本
- 检查数据中有没有特殊字符、乱码、超长文本
- 检查数据的tokenize结果,有没有异常的input_ids
检查结果:数据看起来正常,没有发现明显的异常。
但我还是不放心,写了一个脚本,对所有数据做了全面的检查:
- 过滤掉长度超过1024的文本(GPT-2的最大长度是1024)
- 过滤掉空文本和纯标点的文本
- 检查input_ids的范围,确保在vocab范围内
检查发现,有几条数据的长度超过了1024,被截断了。但截断应该不会导致NaN。
结论:数据不是直接原因,但超长文本可能是诱因之一。
2. 排查方向二:学习率问题
然后怀疑是学习率的问题。学习率太大会导致梯度爆炸,损失飙升。
我检查了学习率调度器的配置:
from transformers import get_linear_schedule_with_warmup
scheduler = get_linear_schedule_with_warmup(
optimizer,
num_warmup_steps=500,
num_training_steps=total_steps
)学习率是5e-5,warmup 500步,然后线性衰减。这个配置看起来没问题。
但我突然想到一个问题:我用了梯度累积,等效batch size是32。但学习率是按batch size 8设置的。等效batch size变大了,学习率是不是也应该相应调整?
一般来说,batch size变大,学习率可以适当变大(线性缩放规则)。但我没有调,学习率还是5e-5。这应该不会导致梯度爆炸,反而可能偏小。
结论:学习率配置基本合理,不是直接原因。
3. 排查方向三:混合精度问题
然后怀疑是混合精度(fp16)的问题。fp16的数值范围比fp32小,容易溢出。
我检查了fp16的配置:
from transformers import Trainer, TrainingArguments
training_args = TrainingArguments(
fp16=True,
fp16_backend="auto",
...
)用的是Transformers的Trainer,fp16用的是默认配置。
我查了一下,fp16训练中,损失飙升到NaN是一个常见问题。原因可能是:
- 梯度溢出(fp16的最大值是65504,超过就变成inf)
- 损失缩放(loss scaling)参数不合适
- 某些层的计算不适合fp16
Transformers的Trainer默认会做动态损失缩放(dynamic loss scaling),应该能处理大部分情况。但如果梯度突然变得很大,损失缩放可能来不及调整。
结论:fp16可能是原因之一,需要进一步验证。
4. 排查方向四:梯度爆炸
然后我重点排查了梯度爆炸的问题。
从日志看,梯度范数在损失飙升之前,已经开始缓慢上升:
- step 4900: grad_norm 2.3
- step 4950: grad_norm 3.1
- step 4980: grad_norm 5.8
- step 5000: grad_norm 12.5
- step 5010: grad_norm 156.2
- step 5020: grad_norm 1024.0
梯度范数在损失飙升之前,已经开始异常上升了。这说明,确实发生了梯度爆炸。
但为什么梯度会突然爆炸呢?
我做了一个实验:用step 5000的checkpoint,重新训练,看看能不能复现。
结果,重新训练后,在step 5030左右,梯度又开始上升。虽然没有立刻变成NaN,但趋势很明显。
这说明,梯度爆炸不是随机的,而是有确定的原因。
5. 排查方向五:数据顺序
重新训练的时候,我注意到一个细节:每次梯度爆炸,都发生在同一个数据批次附近。
这让我怀疑,是不是某一批数据有问题?
我把那个批次的数据找出来,仔细检查。终于发现了问题:
那个批次里,有一条数据,是一个超长的文本(大概2000字),被截断成了1024个token。而且,这条文本的内容很特殊——它是一个代码文件,里面有大量的重复字符和特殊符号。
我单独用这条数据做了一个测试:只训练这一条数据,看看梯度的变化。
结果,只用这一条数据训练,梯度范数直接飙升到了500多!
找到了!就是这条数据导致的梯度爆炸!
三、根因分析
找到了问题数据,接下来分析为什么这条数据会导致梯度爆炸。
1. 数据的特殊性
这条数据有几个特点:
- 超长:2000字,被截断成1024 token
- 重复内容:大量重复的字符和模式
- 特殊符号:代码中的特殊字符,tokenize后产生了很多罕见的token
- 数值内容:代码中有很多数字,tokenize后变成了很多单独的数字token
2. 为什么会导致梯度爆炸
我分析了一下,原因可能是:
原因一:罕见token的embedding梯度大
GPT-2的词表中,很多罕见token的embedding训练不充分。当输入大量罕见token时,这些embedding的梯度会很大,导致梯度爆炸。
原因二:重复内容导致梯度累积
重复的内容,会让模型在多个位置产生相似的梯度。这些梯度叠加起来,就会变大。
原因三:fp16溢出
fp16的数值范围小。当梯度变大时,fp16无法表示,就会溢出变成inf,然后传播成NaN。
原因四:损失缩放来不及调整
虽然用了动态损失缩放,但梯度是突然变大的,损失缩放器来不及调整,导致梯度溢出。
3. 根本原因
根本原因是:数据预处理不严格,没有过滤掉异常数据,加上fp16训练的数值稳定性问题,导致梯度爆炸。
如果数据预处理时过滤掉了超长文本和异常内容,或者用了梯度裁剪(gradient clipping),这次故障就不会发生。
四、解决方案
找到原因后,我采取了以下解决方案。
1. 数据清洗
首先,对数据做了更严格的清洗:
def clean_text(text):
# 过滤超长文本
if len(text) > 500: # 降低长度阈值
return None
# 过滤空文本
if not text.strip():
return None
# 过滤重复内容过多的文本
if len(set(text)) < len(text) * 0.1:
return None
# 过滤特殊字符过多的文本
special_chars = sum(1 for c in text if not c.isalnum() and not c.isspace())
if special_chars > len(text) * 0.5:
return None
return text清洗后,10万条数据剩下了9.5万条,过滤掉了5000条异常数据。
2. 梯度裁剪
加了梯度裁剪,限制梯度的最大范数:
training_args = TrainingArguments(
max_grad_norm=1.0, # 梯度裁剪,默认是1.0
...
)其实Transformers的Trainer默认就有梯度裁剪(maxgradnorm=1.0),但我之前为了"训练更快",把它改成了0(关闭了)。这是一个致命的错误。
梯度裁剪是防止梯度爆炸的最后一道防线,一定要开启!
3. fp16优化
对fp16做了优化:
training_args = TrainingArguments(
fp16=True,
fp16_backend="auto",
fp16_opt_level="O1", # 混合精度优化级别
...
)另外,我还加了一个参数:
training_args = TrainingArguments(
dataloader_drop_last=True, # 丢弃最后一个不完整的batch
...
)防止最后一个batch数据量太少,导致梯度异常。
4. 更频繁的checkpoint
把checkpoint的保存频率从每1000步改成每500步:
training_args = TrainingArguments(
save_steps=500,
save_total_limit=5, # 只保留最近5个checkpoint
...
)这样即使训练崩了,也最多损失500步。
5. 监控告警
加了监控告警:
- 梯度范数超过10时告警
- 损失连续5步上升时告警
- 损失变成NaN时自动停止训练
用Transformers的TrainerCallback实现:
from transformers import TrainerCallback
class NaNCallback(TrainerCallback):
def on_log(self, args, state, control, logs=None, **kwargs):
if logs and 'loss' in logs:
if logs['loss'] != logs['loss']: # NaN检查
print("Loss is NaN, stopping training!")
control.should_training_stop = True五、重新训练
修复了问题之后,从step 5000的checkpoint重新训练。
这次训练很顺利:
- 梯度范数稳定在1-3之间
- 损失稳步下降,没有再出现飙升
- 训练了3个epoch,最终损失降到了1.2
- 生成的文本质量也不错
总共花了8个小时,完成了训练。虽然比原计划多花了时间,但至少结果是好的。
六、经验总结
这次故障,让我总结了很多经验。
1. 数据预处理是重中之重
大模型微调,数据是基础。数据质量不好,后面全白搭。
- 一定要过滤异常数据:超长、空文本、重复内容、特殊字符过多
- 数据清洗要严格,不要心存侥幸
- 训练前要对数据做全面的统计分析:长度分布、token分布、异常值
2. 梯度裁剪一定要开启
梯度裁剪是防止梯度爆炸的最后一道防线,一定要开启。
不要为了所谓的"训练更快"而关闭梯度裁剪。梯度裁剪对正常训练几乎没有影响,但能在关键时刻救你一命。
3. fp16训练要小心
fp16能加快训练、节省显存,但数值稳定性不如fp32。
- 重要的训练任务,可以先用fp32验证,再用fp16
- 开启动态损失缩放
- 监控梯度范数和损失,发现异常及时停止
- 如果fp16经常出问题,可以换成bf16(如果GPU支持)或fp32
4. 频繁保存checkpoint
不要只保存最新的checkpoint,要保存多个历史checkpoint。
- 每500-1000步保存一次
- 保留最近3-5个checkpoint
- 最好在每个epoch结束时也保存
这样即使训练崩了,也能从最近的好checkpoint恢复,损失最小。
5. 监控和告警
训练过程中,要监控关键指标:
- 损失(train loss、eval loss)
- 梯度范数
- 学习率
- GPU利用率和显存
设置告警,发现异常及时处理。不要等训练崩了才发现问题。
6. 不要在训练时离开太久
这次故障,我就是因为去吃饭了,没有及时发现损失异常。等我回来的时候,已经变成NaN了。
虽然有自动保存,但如果能及时发现,可能只需要回退几十步,而不是500步。
重要的训练任务,训练过程中要时不时看一下监控。
7. 复现问题
遇到问题,要想办法复现。
- 找到问题发生的时间点
- 用那个时间点的checkpoint重新训练
- 看能不能复现问题
- 如果能复现,就可以逐步排查
能复现的问题,就一定能解决。
七、写在最后
这次大模型微调故障,是一次惊心动魄的经历。
从看到NaN时的心一沉,到十几个小时的排查,到找到原因时的如释重负,到重新训练成功时的开心,整个过程像坐过山车一样。
但这次故障也让我学到了很多。以前我总觉得,大模型微调就是"调包侠"的工作,加载模型、喂数据、训练,完事。但真正做了才知道,里面的坑太多了。数据、模型、训练、优化,每一步都可能出问题。
2022年了,大模型越来越火,越来越多的人在做大模型微调。但很多人只是跟着教程走,遇到问题就不知道怎么办了。希望这次故障复盘,能帮你避免同样的坑。
最后,用一句话总结:"大模型微调,90%的问题都出在数据和细节上。把数据处理好,把细节做到位,再加上完善的监控和容错,才能训出好模型。"
愿大家的大模型训练,都能一帆风顺,不遇NaN。
评论(0)
暂无评论,快来抢沙发~
评论功能仅对会员开放,请先登录
登录