最近在拿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 条跑200步才炸大概率是数据问题,扫一遍loss为nan的样本看看是不是有超长重复片段。qlora的scale一般不用动,先试试把4bit换成8bit排除量化误差。
我也遇到过类似情况,最后发现是数据里混了几条超长文本,虽然截断到512但某些token的embedding值特别大,把中间层的激活值直接冲爆了。建议你先跑个脚本统计一下输入token的id分布,看看有没有异常高频的稀有token,或者直接按loss飙升那几步的样本索引反查数据。另外qlora的scale可以试试调低到0.1或0.05,我这边调低之后稳定了不少,但前提是先把数据清洗一遍。
查一下是不是有脏数据混进标签了,之前我也遇到过,清洗完就好了。
先查数据里有没有极端长度的样本,截断到512前可能混入超长token把数值炸了。
我之前跑bloom-7b也撞过一模一样的墙,loss突然变nan然后显存爆掉,查了半天最后定位到是数据里有一条超长重复片段,tokenizer把那个序列编码出了异常大的id,embedding算出来直接溢出。建议你先别急着调参,把训练集里所有样本的token长度分布和id最大值拉出来看看,特别关注那些极端值。另外qlora的scale参数确实会影响数值稳定性,我后来把lora_alpha从16降到8,同时把r从8提到16,反而更稳,你可以试试反着来。还有个偏方,把optimizer换成adamw的eps调大一点,比如1e-6改到1e-5,有时候能兜住很微小的梯度爆炸。如果还不行,就写个钩子,每步检查一下loss和梯度范数,定位到具体是哪个batch出的问题,直接扔掉那批数据再继续训练。你前100步正常说明模型本身没问题,大概率是数据里的脏样本在作怪,别急着怀疑qlora配置。
我上周刚踩过一模一样的坑,最后查出来是数据里几条超长样本的token ids在4bit反量化后数值特别大,直接干爆了loss。你试试按token长度排序,把top0.1%的样本单独拎出来看下,或者干脆截断到480。另外qlora的scale参数一般不用动,但你可以把alpha从16降到8看看,有时候是低秩矩阵的scale太大导致中间激活值溢出。
我之前跑bloom-7b也撞过一模一样的墙,loss突然nan然后显存跟着爆,排查了半天发现是数据里有一条超长token序列,embedding直接溢出。你那个怀疑方向我觉得靠谱,建议先写个脚本扫一遍数据,看有没有长度超过512截断后还残留异常值的样本,特别是那种URL堆叠或者base64编码的长串,很容易让layer norm炸掉。另外qlora的scale参数确实值得查,默认值在4bit下有时候会偏激进,尤其你用了double quant之后数值范围变化更大,可以试着把lora_alpha调低到8或者16看看。还有个骚操作是开bf16混合精度,A100对bf16支持很好,能缓解一部分数值问题,但注意loss scaling要重新设。实在不行就加个梯度累积,虽然batch size已经1了,但累积几步相当于变相增大batch,有时候反而能稳定训练。最后建议你load checkpoint回到nan前一步,把梯度打印出来看哪层先爆,这样定位最快,别从头重跑。
先查下loss爆炸那几步的输入数据,八成是长尾token或特殊字符导致的,清洗一下试试。
这情况我碰到过,多半不是lr或者grad clip的事。你试下把qlora的alpha调低到8或16,scale跟着ratio走,有时候默认32配4bit会放大异常梯度。另外建议用dataloader的collate_fn做下token长度分布统计,超480的样本单独拎出来看,长尾token特别容易在深层attention爆掉。我上次是这么定位到3条脏数据的,清洗后稳得一批。
我之前也踩过类似的坑,但不是LLaMA-3,是调一个更小的模型时遇到的。loss突然变nan然后OOM,我后来定位到是数据里有几条约等于重复的超长文本,某些位置上的token对应的embedding梯度异常大,直接冲垮了优化器状态。你可以先写个脚本统计一下数据里有没有极端长度或者重复度高的样本,单独抽出来看看是不是它们触发的。另外qlora的scale参数确实值得怀疑,默认值在4bit下对某些分布极端的激活值可能不够稳,你可以试试手动调小一点,比如从原来的值减半,同时把梯度裁剪阈值再收紧到0.5看看。还有个小技巧,把loss的打印频率改成每步都打,配合tensorboard盯住梯度范数,如果发现某一步梯度范数突然暴涨,基本就能锁定是数据问题了。实在不行先换个损失函数试试,比如标签平滑,有时候能压住尖峰。别急,这种问题多半不是模型本身,数据清洗一遍大概率能解决。
我之前跑7B也遇到过,查查是不是数据里有超长重复片段,换成字节级BPE分词试试。
检查下loss spike那几步的输入,八成是脏数据,过滤掉就稳了。
我之前用7B模型也撞到过一模一样的情况,loss突然跳nan然后显存爆掉,排查了很久发现是数据里混了几条超长文本,虽然截断到512但某几个样本的tokenizer出来的input_id里有特别极端的值,embedding层直接炸了。你可以先别急着调qlora参数,把跑挂的那几步的数据捞出来看看到底是哪条样本,单独喂进去看看能不能复现,大概率就是数据问题。另外qlora的scale我倒没调过,但如果你是用了新版transformers的话,检查一下rope scaling或者attention实现是不是有变动,有时候版本更新会引入一些数值边界问题。还有个笨办法,把batch size再降到1但gradient accumulation调大,同时把loss改成fp16的混合精度计算,有时候能暂时绕开溢出点。你试试把数据集里所有样本的token长度分布画出来,看看有没有特别离谱的尾巴,我那次就是发现有几千条样本的tokenizer居然产生了重复的异常id,清掉就正常了。如果数据没问题,那就得考虑是不是某个层的权重初始化导致特定输入下激活值爆炸,可以加个hook监控每层输出的max norm,定位到具体哪一层再针对性做梯度裁剪或者改量化配置。
我前两天刚踩过一模一样的坑,最后定位到是数据里混了几条超长重复片段,tokenize之后把某个位置激活值顶到爆了。你试试把训练集按loss值排序,把前几轮loss特别高的样本单独拎出来看下。另外qlora的scale如果跟着秩走,4bit下默认值有时候确实会偏大,可以试着砍半或者调到0.25看看。
我之前跑别的模型也遇到过这种前100步正常然后突然爆loss的情况,最后查出来是数据里几条超长样本的attention部分出了问题。你可以先过滤一下序列长度接近512的样本,或者按token数做个分布看看是不是有极端值。另外qlora的scale参数确实可能是个坑,尤其4bit下,试着把lora_alpha调小点或者换成paged_adamw优化器,有时候能缓解。还有个小技巧,loss爆掉之前先存个checkpoint,然后拿那几条可疑样本单独跑一下前向,看能不能复现。
我之前跑bloom-7b也撞过一模一样的鬼,后来发现是数据集里几条编码错乱的样本,token id直接炸出词表边界。建议你先写个脚本扫一遍token长度分布和特殊字符,把超过512截断后仍异常的样本单独拎出来看看。另外qlora的scale我习惯跟着target modules走,你要是用了默认值,试试调到16或者32,有时候低秩矩阵初始化太激进也会在某个step突然爆loss。还有个骚操作是加载时把model.config里的use_cache关了,能省点显存,虽然不解决nan但至少能多撑几步观察。
我之前跑别的模型也撞到过这种前100步正常然后突然nan的情况,后来发现是数据里混了几条超长重复片段,tokenizer切完会爆出极端值。建议先扫一遍数据,看看有没有超过序列长度80%以上的异常样本,直接过滤掉再试。另外qlora的scale可以试着从默认值往下调个0.5倍,有时候4bit量化加double quant会让某些层的梯度范数突然炸掉,跟lr关系不大。你检查一下是不是loss spike之前有某个batch的梯度norm先飙到上百了,如果是的话,把梯度裁剪阈值降到0.5甚至0.3试试。
我之前跑bloom-7b也撞过一模一样的鬼打墙,最后发现是数据里混了几条超长重复字符的脏样本,tokenize完直接让某个位置的激活值爆了。你可以先按loss变nan那一步的样本id倒查一下,看是不是集中在特定领域关键词上。另外qlora的scale参数确实值得查,默认值在4bit下有时候太激进,我调到16之后稳定性明显好了。还有个小技巧,把optimizer换成adamw的bfloat16版本,某些极端情况下能兜住溢出。
前100步正常后面突然nan,我赌五毛是数据里有极端长序列或者某个样本的label有问题,先写个脚本扫一下token长度分布和特殊字符,特别是看看有没有单个样本loss贡献异常大的。qlora的scale我一般不动,但你要是用了新版transformers,检查下rope theta和flash attention的版本兼容性,我上次就是升级库之后莫名其妙爆nan。另外A100 40G跑8b qlora按理说很宽裕,OOM大概率是nan之后优化器状态炸了,可以先catch异常跳过坏样本试试。
我遇到过类似的,最后查出来是某个样本里有连续几千个重复的token,直接把logits干溢出了。你可以先用fp16跑一遍,不开混合精度,看loss是不是还爆,能定位是不是精度问题。梯度裁剪对nan没用,得调min_loss或者用torch.nan_to_num兜底。还有个小技巧,把数据集按loss排序,先训简单的再逐步混入难的,能避免一开始就撞上毒样本。
我之前跑bloom-7b也遇到过一模一样的鬼情况,也是200步左右突然loss飞升然后显存爆掉,后来发现是数据里混了几条超长重复文本,tokenizer把某些罕见字切成了超长sequence,直接让激活值爆炸。建议你先把数据按token长度重新统计一下,看看有没有异常长的样本,另外qlora的alpha可以试着跟r的比例调小点,比如4bit下alpha=16配r=8可能太激进了。还有检查一下是不是有个别样本里带特殊符号导致embedding算出来特别大,可以写个脚本把loss超过阈值时对应的数据抓出来看看。
我之前跑bloom-7b也撞过一模一样的墙,loss突然飙nan然后显存跟着炸,排查半天发现是数据里几条超长样本的attention部分数值异常,跟token长尾关系不大。你可以试着把序列长度砍到256跑几百步看看还炸不炸,如果稳定了基本就是样本长度分布的问题。另外qlora的scale参数确实是个坑,默认值在4bit下有时候会放大异常梯度,建议直接降到原来的四分之一试试,反正微调效果影响不大。还有个小技巧,用torch.cuda.amp的grad_scaler把scale下限设高点,能拦住一部分溢出。要是还不行,就把数据集按loss排序,把那几个异常样本单独拎出来看,八成是标签错误或者特殊字符没清洗干净。别急着调参,先花半小时把数据清洗逻辑捋一遍,这种问题多半是脏数据惹的祸。