最近在拿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遇显存爆炸,梯度检查点也救不回来,求支招
全部回复
共 75 条我之前跑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值,比调参靠谱。
我之前跑别的模型也撞过这种前100步正常然后突然nan的鬼情况,最后查出来是数据里有一条超长文本把位置编码撑爆了。你可以先写个脚本扫一下序列长度分布,看看是不是有极端值混进去了,顺便把tokenizer的truncation策略改成max_length硬截断。另外qlora的scale参数确实值得怀疑,试下降到4或者8,有时候默认值在低精度下会放大梯度异常。还有个小技巧,把loss改成fp32累加,能帮定位是不是混合精度下的溢出。
我遇到过类似的,不过当时是bf16的问题,换回fp16就好了,A100上bf16的精度有时候反而更敏感。你可以先加个log看nan出现前那几步的梯度范数,如果暴增就锁定是某个样本。另外20w条数据其实不算多,试试把qlora的r值降到16或者8,减少低秩矩阵的累积误差。检查下数据预处理是不是有脏字符,比如不可见unicode或者特别长的重复串,这些容易让embedding算崩。
我之前跑7b也遇到过一模一样的,loss炸之前那几步loss曲线其实会先小幅抖动,你往前翻翻tensorboard,大概率能定位到是某个特定batch的数据搞的鬼。可以先把你数据里包含超长重复片段或者特殊符号的样本筛出来过一遍,尤其是那种几百个数字连在一起的,qlora的4bit下很容易出inf。另外scale参数我建议你试试固定成8或者16,别用默认的,然后检查一下是不是某个embedding维度更新特别大。要是还不行,就把序列长度砍到256跑几百步看看还炸不炸,能快速排除是不是长度问题。
我之前也踩过类似的坑,不过是在llama-2上,loss突然变nan大概率不是lr的锅,更像是数据里混进了极端离群值,尤其是那种超长重复片段或者全是特殊字符的样本,在attention计算时容易把中间值推到inf。你可以写个脚本把loss异常步附近的样本捞出来看看,或者直接按token长度分布截断一下尾部,20w条里哪怕只有几条这种垃圾数据就够炸了。另外qlora的scale参数确实值得怀疑,4bit下如果某个低秩矩阵的奇异值分布太宽,加上double quant的量化误差,可能在特定激活下放大溢出,试试把lora的r从16降到8,或者把alpha调成跟r一样,别用默认的2倍关系。还有个小技巧,用bf16代替fp16混合精度,能显著缓解溢出,A100对bf16支持很好,loss曲线会更稳。如果还不行,就先在干净的小数据集上跑通验证一下流程,别一上来就全量。梯度检查点救不回显存的话,可以试试把注意力改成flash-attention-2,能省不少显存,而且数值上更稳。你要是方便,可以把出问题那几步的loss和梯度norm打出来,对比一下是突然跳变还是渐进发散,这个信息对定位很有用。
先扫一遍数据里有没有超长重复片段,我之前就是被某条脏数据搞炸的,过滤掉就稳了。
之前跑别的模型碰到过类似情况,最后查出来是数据里几条超长样本的某些token在4bit下量化误差被放大,导致activation突然爆炸。你可以先试着把序列长度砍到256跑几十步看看还崩不崩,或者用fp16跑一小段对比下,能快速定位是不是量化的问题。另外qlora的scale参数确实值得查一下,特别是如果用了新版本transformers,默认值可能有变化,手动设成16或者32试试。
查一下是不是有样本label和input重叠太狠,之前我遇到过脏数据导致loss炸的情况。
换个思路,先用小批量跑一遍数据清洗,nan前那几步的batch单独拎出来看,八成是长尾token的embedding爆了。
我之前跑别的模型也撞到过这种前一半正常然后突然nan的情况,查了半天发现是数据里混了几条超长重复片段,把tokenizer的max_length撑爆后embedding直接溢出。你可以先写个脚本扫一遍数据,看有没有长度接近512或者包含异常字符的样本,单独拎出来试试。另外qlora的scale可以试着降到1倍以下,有时候默认参数在低精度下确实容易炸,但主要还得先排除数据问题。
先查下数据里有没有超长或乱码样本,我之前也是200步左右炸,筛完就好了。