最近在拿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 条这问题我上周刚踩过一模一样的坑,最后查出来是数据里有几条超长重复片段,tokenizer切出来带了一串异常id,loss直接炸。你先别动qlora的scale,把loss为nan那一步的batch单独dump出来看看里面有没有特别长的样本,或者用torch.autograd.detect_anomaly()定位一下。另外4bit下bf16的精度有时候也会出这种诡异问题,试试把qlora的compute_dtype换成fp32跑几十步看看还炸不炸。
上个月我微调mistral也遇到过,最后发现是dataset里混了几条全角符号的脏数据,embedding算出来inf,清洗完就好了。你的数据是清洗过的吗?可以先统计一下token长度分布,看看有没有超过512截断后还残留特殊字符的。另外梯度检查点对显存帮助有限,核心还是qlora的4bit反传精度问题,可以考虑把quant_type改成nf4再试试。
我这边之前用8b做领域微调也炸过一次,后来发现是某个样本的label里有nan值,loss算出来直接飘了。你检查一下数据预处理那步,特别是有没有做masking或者padding时不小心把attention_mask搞成全0的行。还有个小技巧,把整个数据集按token数排序,从短到长喂进去,能有效避免突然的显存尖峰。
跑200步才炸大概率是数据里有脏样本,先筛一下loss突增那几步的batch看看。
我上次是调低qlora的alpha到16解决的,你可以试试。
这情况我上周刚踩过一模一样的坑,也是llama-3-8b,qlora跑到300步直接nan然后OOM。后来查了半天,问题出在数据里有个样本的label包含了一个超长的重复片段,导致loss计算时softmax溢出。建议你先写个脚本扫一遍数据,看看有没有极端长度的样本或者某个token出现频率异常高的,特别是那种重复几百次的连续token,很可能是元数据混进去了。另外qlora的scale参数我试过默认的确实在长序列下容易炸,改成0.5之后稳定了不少,你可以试试看。还有个小技巧,把loss改成bf16计算,能扛住一部分溢出,虽然治标不治本。梯度裁剪你试了是好事,但nan出现后裁剪其实没用了,因为梯度已经变成inf了,得在forward阶段就防住。要是实在找不到坏样本,可以考虑加个loss的nan检测,跳过异常步,至少能保住前面100步的训练成果。
试试把qlora的alpha调成r的两倍看看,之前我遇到过类似问题这么解决的。
我之前跑llama-2的时候也碰到过一模一样的状况,loss降到一半突然nan,紧接着显存就炸了。后来查下来是数据集里有几条超长重复文本,tokenizer把某些罕见字符合并成了超长序列,导致attention矩阵那一步直接溢出。你试试把数据里长度超过512的样本单独筛出来看下,尤其是那些包含特殊符号或者emoji的,很可能就是它们触发的。另外qlora的scale参数确实值得检查,特别是你用double quant之后,如果alpha和r的比例没配好,某些层的梯度范数会异常放大,你可以把lora的初始化改成高斯分布试试,默认的kaiming有时候在4bit下不太稳。还有个小技巧,把优化器换成adamw with foreach=True,有时候能缓解部分数值问题,但根治还得靠清洗数据。你那个20w条数据量其实不算大,建议先跑个数据分布统计,看看有没有那种token长度超过700的极端样本,直接截断到512或者干脆删掉。如果实在排查不出来,可以试试在loss计算前加个logit clamp,把输出限制在-30到30之间,能挡住大部分溢出路径。
我之前也踩过类似的坑,最后发现是dataset里混了几条超长文本,虽然截断到512但某些token的id异常大,embedding直接爆了。你可以先跑个脚本统计一下输入ids的最大值和分布,看看有没有离群值,顺便查一下label里有没有-100以外的异常值。另外qlora的scale确实要留意,你试过把lora_alpha调低到8或者16吗?有时候默认32配4bit会把梯度放大得比较厉害。还有个笨办法,就是按loss升序把样本排序,先喂正常数据让模型稳下来,再慢慢混入难样本,能避开前期崩溃。
这情况我也踩过坑,八成是数据里混了超长token或脏值,先跑个token分布统计,异常样本直接过滤掉试试。
我之前跑别的模型也遇到过类似情况,后来发现是数据里有个别样本的token长度超出了预期,导致embedding层梯度爆炸。你可以先写个脚本把数据里超过512的样本筛出来看看,或者统计下loss变nan前那几步的输入是不是有异常。另外qlora的scale参数确实可能跟基座模型不匹配,试试把lora_alpha降到16或者8,有时候比调学习率管用。如果还不行,可以看看是不是某个样本的label里有特殊字符,我之前就栽在了一个全角空格上。
我之前也遇到过类似情况,最后发现是数据里混了几条超长文本,虽然截断到512但某些token的embedding值特别极端,直接让中间激活炸了。建议你先单独跑一下数据集的loss分布,把那些loss异常高的样本筛出来看看。另外QLoRA的scale参数确实有影响,我试过把4bit的scale调小一个量级,数值稳定性好了不少。还有个小技巧,可以把优化器换成Adafactor,对显存和数值波动都更友好。
前100步正常后面突然nan,我赌五毛是数据里混了超长token或者脏文本,建议先跑个token长度分布统计,把异常长的样本单独拎出来看。另外qlora的scale其实不用动,但你可以试试把4bit的double quant关掉,有时候量化误差会在某些极端激活值上爆掉。还有个小技巧,loss飘了之后先别急着重跑,load最近的checkpoint把batch size再砍半试试,我之前碰到过类似情况这样救回来的。
我之前跑别的模型也撞过这种前100步正常然后突然nan的鬼事,最后查出来是数据里几条超长样本把embedding的梯度搞炸了,尽管序列截断到512,但tokenizer的attention_mask没对齐。你可以先单独跑一下数据集的loss,把那些loss异常高的样本筛出来看看,大概率是脏数据。qlora的scale参数一般不用动,但4bit下如果用了double quant,建议把trust_remote_code关掉试试,有时候是反量化精度在特定值上溢出。另外OOM可能是nan之后优化器状态崩了导致的假象,先解决nan源头再说,别急着调显存。
看到这个情况我第一反应是想起之前调LLaMA-2时踩过的坑,loss突然飙nan大概率不是lr的问题,更像是数据里混了特别长的重复片段或者特殊unicode字符,建议你写个脚本扫一下tokenizer前后的长度分布,尤其看看有没有超过512截断后依然残留的异常token。另外qlora的scale参数确实值得怀疑,4bit下如果某个奇异值被放得太大,反向传播时梯度范数会突然爆炸,即使有grad clip也可能在clip之前就把中间激活值冲爆了,可以试试把lora的r从默认的64降到8或者16,同时把alpha调成跟r一样大,看会不会稳定一些。还有个小细节,你用的是double quant,但要注意量化后的常数是否在transformers版本里被正确传递,我之前升级库之后出现过量化参数没生效导致显存暴增的怪事。最笨但有效的办法是,先用一个小批量比如100条数据反复跑几十步,确认能收敛,再加数据量,这样能快速定位是不是某个样本的问题。如果实在找不到是哪个样本,就写个hook把loss异常的batch index打印出来,直接筛掉那几条数据,虽然治标不治本但能先跑通流程。对了,你用的什么框架,peft和bitsandbytes的版本最好也报一下,版本不匹配经常出这种玄学问题。
我之前跑bert-large也撞过类似的,loss突然变nan八成不是lr的问题,你试试把qlora的lora_alpha往低调,scale跟着r的比例走,别用默认值。另外强烈建议扫一遍数据,看有没有长度接近512但padding特别多的样本,那种容易让某些token的梯度爆炸。还有个偏方,把optimizer换成adamw with eps=1e-8,有时候能救回来。你检查下是不是在200步正好碰到某个batch里有重复的极端值,单独拎出来跑一下验证很快的。
查一下是不是有样本标签错了,我之前遇到loss飙nan就是脏数据搞的鬼。
你试试按loss排序筛掉异常样本,或者把qlora的alpha调低点看稳不稳。
我之前跑LLaMA-2-7B也撞过一模一样的墙,loss突然跳nan然后显存跟着爆,排查到最后发现是数据里有一条样本的input_id全是0,embedding算出来直接inf。你那个怀疑长尾token的方向我觉得挺靠谱,可以先写个脚本扫一遍数据,看看有没有极端长度的样本,或者某个token出现频率特别低导致embedding没学好。另外qlora的scale参数确实值得怀疑,默认值在4bit下有时候会放大异常梯度,我之前从默认的16调到8之后稳定了不少,你可以试试看。还有个小技巧,把optimizer换成adamw的bfloat16版本,有时候能躲过数值溢出,虽然治标不治本但至少能帮你定位是不是数据问题。梯度检查点救不回显存的话,试试看把activation checkpointing和梯度累积结合用,batch size=1的时候梯度累积步数调大点,虽然慢但至少不会爆。最后建议你把loss飙nan那一步的样本单独dump出来,看看到底是什么内容,比瞎猜快多了。
我之前跑别的模型也撞过类似情况,后来发现是数据里混了几条超长重复片段,tokenizer没处理干净导致embedding直接炸了。建议你先跑一遍数据统计,看看max token长度有没有异常,顺便检查label里有没有极端值。qlora的scale其实不太会引发nan,优先怀疑数据问题,可以写个脚本单独喂那些可疑样本试试。另外A100 40G跑8b qlora应该够,OOM大概率是梯度爆了之后显存没释放,把torch的缓存清理加上,或者换个策略定期重启训练。
我之前跑llama-2的时候也撞过这个鬼情况,最后查出来是数据里几条超长文本没清洗干净,tokenizer硬切出来一堆异常id,loss直接飘了。你那个长尾token的思路我觉得靠谱,可以先写个脚本统计下每条样本的token长度分布和特殊字符,把top几的极端值单独拎出来看看。另外qlora的scale参数一般不用动,但你可以试着把4bit的nf4换成fp4,或者干脆降到3bit,有时候量化粒度太粗也会在反向传播时搞出inf。如果排查完数据还不行,试试在dataloader里加个梯度累积步数,顺便把loss打印改成每步都输出,定位到具体哪一步爆的。
我之前跑llama-2的时候也撞过一模一样的nan,查了半天发现是数据里有个超长样本,虽然序列截断到512但某些token的embedding异常大,直接爆了loss。你试试用dataloader加个异常值过滤,或者把qlora的alpha调低到16看看,scale参数确实有影响,我之前从32降到16就稳了。另外A100 40G跑8b还OOM有点怪,确认下是不是显存碎片化,可以开个环境变量试试。
跑200步才炸八成是数据里混进了脏样本,建议先定位loss飙的那一步的输入,别急着调参。
我之前跑别的模型也遇到过类似情况,后来发现是数据里有个别样本的token长度虽然没超限,但特殊字符把embedding搞出了极值,导致loss爆掉。你可以先跑个脚本统计下所有样本的loss,把异常高的那几条单独拎出来看看是不是集中在某个领域或格式上。另外qlora的scale参数确实值得查,特别是你用double quant的时候,默认值不一定适合你的任务,试着按rank比例缩一下可能就稳了。还有个小技巧,把优化器换成adamw带权重衰减,有时候对数值稳定性有帮助。