最近在拿llama-3-8b做领域微调,数据量就20w条,用的qlora(4bit+double quant),batch size已经调到1了,序列长度512。但跑了200步loss突然飙到nan,然后显存直接OOM(A100 40G)。诡异的是前100步loss正常下降,我怀疑是某个样本触发了数值溢出。已经排除学习率问题(试过1e-4和2e-4),也试过梯度裁剪。有人遇到过类似情况吗?是不是需要检查数据集里的异常长尾token?或者qlora的scale参数要跟着调?现在卡在这两天了,求有经验的大佬指点一下排查思路。
微调LLaMA-3遇显存爆炸,梯度检查点也救不回来,求支招
全部回复
共 67 条我之前跑bloom-7b也撞过一模一样的鬼门关,nan出现在200步附近大概率是某个batch里混进了极端长token或重复文本,建议先写个脚本扫一下数据集的token长度分布,把超过512截断后还剩超长残段的样本单独拎出来看。qlora的scale我倒没动过,但你可以试试把4bit的量化改成8bit看还会不会炸,能缩小排查范围。另外A100 40G跑8b居然会OOM,你确认下是不是activation checkpointing和qlora的显存释放有冲突,有时候这俩叠加反而会爆。
之前跑别的模型遇到过类似情况,最后发现是数据里混了几条超长文本,embedding之后某些位置的值特别大,直接就把loss炸了。你可以先扫一遍tokenizer后的长度分布,再单独把那几条特别长的样本拿出来跑一下看看。qlora的scale我一般不动,但如果你用了double quant,可以试试把4bit的normalize改成false,有时候是量化误差累积的问题。另外nan出现后显存OOM大概率是优化器状态污染了,建议把checkpoint回滚到正常步数再调数据重跑,别在原基础上硬续。
这现象我遇到过,大概率不是lr的问题,你试试定位到出nan那一步的具体样本,用dataloader的seed固定住然后逐步排查。另外qlora的scale可以调小到16或者8试试,有时候4bit下那个常数确实会放大异常值。还有个小坑,llama3的rope对长尾token很敏感,建议先看看数据里有没有超长重复片段或者非法编码。
同款问题遇到过,不过我是7b模型,最后定位到是数据里几条特别长的样本,embedding层的某些token id在4bit量化下激活值异常大,直接把loss顶爆了。你试试把序列长度砍到256跑一遍,如果稳定了基本就是长尾样本的锅,或者干脆写个脚本把token长度超过480的样本过滤掉。
另外qlora的scale参数确实值得看一眼,默认值在低rank下有时候会放大异常梯度,我习惯把lora_alpha从16降到8再配个更小的学习率,虽然收敛慢点但稳很多。你也可以开amp的grad_scaler看看是不是fp16下溢出的问题,有时候bf16反而更安全。
排查顺序建议先做数据清洗,再用最小数据集跑过拟合测试,最后才动模型配置。你那个20w条数据里如果混着乱码或者特殊符号,比调参更容易出这种诡异现象。
我前两天刚踩过类似的坑,最后发现是数据里有一条超长重复片段,tokenizer把那个位置编码撑爆了。你可以先扫一遍数据,看看有没有长度异常或者特殊字符堆叠的样本,单独拎出来试试能不能复现。另外qlora的scale参数确实影响数值稳定性,可以试试调低到0.1或者0.05,同时把4bit的double quant关掉看看。梯度检查点救不了OOM的话,先别急着扩显存,检查下是不是loss spike之后模型权重已经崩了,这时候直接加载checkpoint重新跑反而更有效。
之前跑bert也这样,后来发现是数据里混了超长文本,清洗一遍就好了。
跑200步才炸八成是数据里有极端长尾,扫一下token长度分布和embedding的norm值,比调参靠谱。