最近在拿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 条先查数据里有没有全角空格或特殊字符,我上次就是被一个异常token搞崩的。
之前跑别的模型也遇到过这种前段正常后段突然nan的情况,最后查出来是数据里有一条超长文本截断后留下了半个unicode字符,embedding直接炸了。你可以先做个极端值筛查,看看有没有token id特别诡异的样本,顺便把qlora的target_modules换一下试试,有时候某些模块对量化误差更敏感。另外A100 40G跑8b按理说很宽裕,如果梯度检查点开着还OOM,可能显存碎片化了,试试torch.cuda.empty_cache()放step循环里。
之前跑llama2也炸过,查了半天是脏数据里有超长连续token,过滤掉就好了,你可以先扫一遍。
我之前跑bloom-7b也撞过一模一样的墙,loss飙升和OOM几乎是同时来的。后来发现是数据里有几条超长重复片段,把tokenizer的max_length撑爆导致embedding层溢出,清洗后就好了。你可以先扫一遍数据,看看有没有长度接近512的极端样本,或者试试把序列长度砍到384跑200步验证一下。qlora的scale一般不用动,倒是建议把4bit的嵌套量化关掉,有时候double quant在极端数值下会放大误差。
跑200步才nan而且前100步正常,大概率不是全局超参问题,我也在8b上遇过类似情况,最后定位到是数据里几条超长文本的尾部生成了异常ID,虽然截断到512但分词后某些稀有token的embedding在反向传播时爆了。你可以先按batch逐条算一下loss,把nan样本单独拎出来看看,或者把qlora的lora_alpha从16降到8试试,scale参数确实会影响数值稳定性。另外A100 40G跑8b qlora按理说绰绰有余,OOM可能是nan后显存状态没被清干净,建议加个异常检测直接跳过坏batch而不是中断训练。
这情况我也踩过坑,20w条数据里混着超长重复片段或者特殊符号时,qlora的4bit下特别容易炸。你先别急着调scale,用grep把tokenize后长度接近512的样本筛出来看看,八成是某条数据里有连续数字或base64字符串。另外OOM是nan之后的连锁反应,建议把qlora的alpha从16降到8试下,我上次这么干直接稳住了。如果还不行就开trust_remote_code慢慢跑,把max_grad_norm设成0.3。
这情况我也踩过,而且跟你几乎一模一样的配置,最后查出来是数据里有一条样本的label全是-100,mask掉之后loss直接变0,然后反向传播梯度就炸了。建议你先跑个脚本扫一遍数据,看看有没有极端长的target序列或者全padding的样本,这种在qlora下特别容易触发nan。另外你说的scale参数,我试过把lora_alpha从16降到8,确实能稳一点,但代价是收敛变慢,不太像根因。还有个思路是换用bf16混合精度,A100对bf16支持很好,能避免fp16下某些小数值溢出,我之前用bf16之后就没再出现过这种突然爆loss的情况。如果还不行,就试着把qlora换成普通的lora(不用4bit),显存多占一点但稳定性会好很多,毕竟20w条数据量不算小,没必要在量化上死磕。最后建议开个wandb盯一下每层的梯度范数,看看是哪一层先爆的,能帮你定位是embedding还是输出层的问题。
跑200步才炸大概率是数据里有极端值,先洗一下长尾token再试试,lc调0.5看看。
查一下是不是有样本标签错位,之前我遇到过脏数据直接让loss炸掉。
试试把qlora的alpha调小点,跟r保持2:1比例,别用默认值。
我之前跑bloom-7b也撞过一模一样的墙,后来发现是数据里有几条超长重复片段把embedding的梯度搞炸了。你可以先写个脚本扫一下token长度分布,把超过512的样本单独拎出来看,或者干脆截断到480试试。另外qlora的alpha别动,先检查target_modules是不是全量注入了,有时候只注入了部分层反而更容易爆。实在不行就换paged_adamw优化器,显存能省不少。
我之前跑别的模型也遇到过这种前100步正常然后突然nan的情况,最后查出来是数据里几条特别长的样本把激活值顶爆了,虽然序列长度设了512但得看看是不是有样本没截断干净。你可以先加个try-except或者统计一下每步的输入tensor最大值,定位到具体是哪一步出的问题。另外qlora的scale确实可能是个坑,默认值在8b上不一定合适,试试调小到4或者2,我上次这么干直接稳住了。还有个小建议,开了梯度检查点的话把激活检查点策略改成选择性开启,别全开,能省不少显存。
之前跑llama-2遇到过类似情况,最后发现是数据里混了几条超长重复片段,tokenizer编码后直接爆了embedding的scale。你可以先单独抽那几步的loss曲线,配合gradient accumulation的step数看是不是固定间隔。qlora的scale我一般不动,但会把4bit的normalize改成False试试,有些异常值会在这放大。另外检查下有没有空行或特殊unicode字符,那种很容易让loss飞掉。
跑微调时我也被nan坑过一回,最后定位到是某个样本里有个超长的数字串,分词后产生了几千个重复token,梯度直接炸了。你先写个脚本把所有样本的token长度分布打出来,看有没有远超512的尾巴,把那些截断或者过滤掉。qlora的scale通常不用管,但你可以顺手把double quant关掉对比一下,说不定是量化误差累积的锅。
之前用qlora微调llama-3-7b也遇到loss中途跳nan,排查了一圈是数据里有一批样本带了不可见字符,编码后变成异常embedding。建议你先把数据清洗一遍,过滤掉token长度超过480的,再跑个50步看看。如果还炸,可以试试把qlora的r从64降到32,有时低秩矩阵的scale对数值敏感,稍微调小点能稳定。另外确认下是不是某个特定batch导致的,
查一下是不是有样本包含超长重复片段,之前我碰到过类似情况,清洗掉就稳了。
跑的时候盯一下loss曲线,nan前有没有突然冲高,可能是某个batch的梯度爆炸,qlora的scale调小点试试。
试过把qlora的alpha调成和rank一致吗,之前遇到过类似问题,这样能稳住loss。
跑一下数据集的token分布,重点看是不是有超长或重复的异常片段,直接过滤掉试试。
我之前跑bloom-7b也撞过一模一样的情况,后来查出来是数据里有一条超长样本,tokenizer截断后反而把中间某段异常字符拼成了特别极端的id序列,直接让embedding爆炸。建议你先做个极端值扫描,看看loss变nan前那几个batch的输入是不是有规律,或者干脆把每个样本单独跑一遍前向看哪个触发溢出。qlora的scale参数一般不用动,倒是可以试试把4bit的quant_type从nf4改成fp4,有时候数值分布不一样能避开这种坑。另外实在不行就加个loss clamp,虽然丑但能先跑通流程。
我之前跑bloom的时候也遇到过这种前100步正常然后突然nan的,后来发现是数据集里藏着几条超长文本,虽然截断到512但某些token的embedding直接爆了。你可以写个脚本扫一下input_ids的分布,看看有没有极端值,或者把qlora的alpha调低点试试,有时候4bit量化下scale太激进也会这样。另外建议你在dataloader里加个异常检测,遇到loss异常就跳过当前batch,先定位是不是数据问题,别急着调模型。
我倒是觉得不一定是数据的问题,A100 40G跑qlora按理说绰绰有余,突然OOM更像是某个时刻显存碎片化或者梯度累积爆了。你可以试试把optimizer换成adafactor,省显存效果很明显,或者干脆在loss变nan前手动保存checkpoint,然后二分法找触发样本。要是排查数据太费时间,直接按loss值过滤掉那些异常高的batch,先跑完再说。
个人经验哈,这种nan十有八九是fp16精度下某个样本的logits溢出,跟学习率真没关系。你可以开bf16试试,A100支持,数值稳定性会好很多。另外qlora的target_modules默认可能没覆盖全部线性层,导致某些层还是全精度微调,显存压力全堆在那几个模块上了,改成全部attention加mlp试试。我之前就是这么解决的,直接稳
我之前跑7B也遇到过一模一样的状况,后来发现是数据里混了几条超长且重复token的脏数据,把输入清洗一下就好了。另外你试试把qlora的r从8降到4,同时把alpha调成r的两倍,我这么改完loss稳定多了。还有个笨办法,用gradient accumulation配合微batch排查,定位到具体哪条样本炸的,直接扔掉比调参省时间。
nan突然出现大概率是数据里有脏样本,先跑个脚本检查下token长度和特殊字符分布。
我之前遇到过类似情况,把qlora的alpha调成跟rank一样大就稳了。
这题我好像也踩过,20w条数据量不算大,但qlora的scale参数确实容易被忽略,默认值在长序列上偶尔会爆。你可以试试把loss打印成float16看是不是inf,或者直接抓一下哪几条样本的attention输出异常大,我之前排查就是定位到某条重复样本。另外建议把优化器换成adamw+eps调大点,比如1e-8改到1e-6,有时候数值稳定性差就差在这。梯度检查点救不了本质问题,先确认是不是某个batch的输入带特殊字符导致embedding溢出吧。
查一下loss爆炸那几步的输入样本,八成是超长token序列触发了embedding溢出。