上周做一个大模型微调项目,训练到一半模型突然崩了,损失飙升到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的时候,我的第一反应是:完了,模型崩了。

我立刻做了几件事:

  1. 停止训练,防止覆盖更多的checkpoint
  2. 查看训练日志,找到损失开始异常的时间点
  3. 检查保存的checkpoint,看有没有能用的
  4. 查看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。