最近在跑一个文本分类任务,模型用的是BERT-base,数据大概几十万条。单卡A100(40G),batch size调到16就OOM了,试了梯度累积但训练速度慢得离谱。看网上说DeepSpeed的ZeRO能省显存,但配置起来有点复杂,而且我现在用的是PyTorch原生训练循环,不知道值不值得花时间迁过去。另外也试过torch.cuda.amp混合精度,效果有限。想问下各位大佬,这种情况是直接砍batch size硬扛,还是上DeepSpeed?或者有没有其他更简单的显存优化技巧?顺便问下,ZeRO-2和ZeRO-3在实际使用中差别大吗?谢谢!
用PyTorch训练Transformer时显存总爆,换DeepSpeed还是直接砍batch?
全部回复
共 112 条40G都爆说明不是batch大小问题,先查下是不是数据加载或模型本身有冗余,比如max_len设太长。
ZeRO-2够用了,ZeRO-3是给上百亿参数准备的,你这规模迁移成本不划算。
讲真,40G卡跑BERT-base才16就爆有点反常,先查查是不是max_len设太长或者data loader里没开pin_memory,另外把attention的显存占用打一下,可能你压根没到必须上ZeRO的程度。真要省心就先把batch砍到8配合梯度累积,虽然慢但至少能跑,DeepSpeed那套配置对原生循环来说光改分布式钩子就够你折腾半天。ZeRO-2和3在单卡场景下差别不大,3主要是为了跨节点分片参数,你这规模用不上,别被忽悠着上重武器。
40G单卡跑BERT-base才16的batch就爆,感觉不太正常,你是不是把序列长度拉太长了?我一般先用AMP加gradient checkpointing,这俩加起来基本能翻倍,DeepSpeed那个配置对新手确实不友好,尤其ZeRO-3还得改分布式逻辑。ZeRO-2和3实际差距主要看模型规模,你这种单卡场景其实ZeRO-2就够了,但真不如先查查是不是数据加载或者padding策略有问题。另外你可以试试把优化器换成Adafactor,省显存效果比砍batch明显,速度影响也小。
40G的A100跑BERT-base batch16就爆,这有点不对劲啊,你序列长度是不是特别长?我之前跑BERT-large也就这配置,batch32还能凑合。先确认下是不是数据加载或者padding没优化,试试dynamic padding和sort by length,能省不少显存。至于DeepSpeed,说实话ZeRO-2配你这种单卡场景意义不大,它主要解决多卡通信冗余,单卡上省的那点内存还不够折腾配置的成本。ZeRO-3倒是能把参数分片到CPU,但速度会掉得让你怀疑人生,文本分类任务真没必要。我建议你先砍到batch8,配合梯度累积到等效32,用amp加上gradient checkpointing,bert-base的话显存能压到12G左右,训练速度其实比OOM反复调参快多了。另外你试试torch.utils.checkpoint,对bert这种深层模型效果很猛,虽然多了点重计算,但比砍batch带来的收敛变稳要划算。几十万条数据单卡大概也就几小时,真不值得为这上分布式那套。
先试试gradient_checkpointing,能省一半显存,比你折腾ZeRO划算多了。
40G跑BERT-base batch16就爆有点离谱,你是不是把序列长度搞太长了或者忘了关梯度checkpoint?先试试把max_len砍到128或者256,再开torch.compile,说不定直接就能跑起来。DeepSpeed配置确实费劲,但ZeRO-2这种offload优化对单卡也有用,不过迁移成本你得掂量下,如果只是临时任务真不如直接砍batch配梯度累积,虽然慢点但稳。ZeRO-3主要是为多卡跨节点设计的,单卡上跟ZeRO-2差别不大,别被网上教程忽悠了。
40G的卡跑BERT-base batch 16就爆,感觉有点不太正常啊,你检查下是不是序列长度或者padding那边有冗余开销。ZeRO确实管用但为了这个任务迁移有点重,我建议先试试gradient checkpointing,开起来显存能省一大截,速度损失比梯度累积小多了。ZeRO-2和ZeRO-3主要差别在参数分片粒度,单机单卡的情况下ZeRO-2就够用,ZeRO-3反而是为了多机超大规模模型准备的,你这场景大概率用不上。
另外你提到amp效果有限,要不要看一眼是不是把优化器状态也切成fp32存储了,有些情况下混合精度配合torch.compile能再挤出点空间。我自己的经验是,先把batch压到8,加checkpointing,再用amp,基本能稳住,速度比梯度累积快不少。
40G显存跑BERT-base到16就OOM,这不太正常啊,你序列长度是不是给得太长了?先检查下max_len,说不定砍到128就能解决大半问题。DeepSpeed迁移成本确实高,但ZeRO-2配置其实不难,主要是offload那步容易踩坑,为了省显存值得折腾一次。
我个人经验是梯度累积加AMP基本够用,但慢的话试试只看梯度累积步数能不能降下来,毕竟batch16累积4步等效64,显存占用不变但吞吐会好很多。ZeRO-3就别考虑了,那是给上百亿参数模型用的,你这种规模用ZeRO-2纯属杀鸡用牛刀,反而可能因为通信开销变慢。
实在不想动代码的话,还有个偏方:把BERT换成DistilBERT,精度损失不大但显存直接砍半,训练速度还能翻倍,比你纠结优化省心多了。我上次跑类似任务就是这么干的,效果挺稳。
先看看你max_len是不是设太长了,BERT默认512很吃显存,砍到128试试,比折腾DeepSpeed快多了。
BERT-base几十万条数据,batch 16就OOM确实有点怪,先确认下是不是max_length拉太长了,512降到128能省一大截显存。要我选的话会先试ZeRO-2,改造成本比ZeRO-3低不少,原生训练循环加个deepspeed.initialize就行,不用大改。ZeRO-3参数分片更狠但通信开销明显,单卡场景其实收益有限,还不如省下时间调调sequence length和checkpointing。
BERT-base才16就OOM有点离谱,先确认下是不是max_length拉太长了,512和128的显存差好几倍。如果长度没法砍,上DeepSpeed ZeRO-2其实挺香的,配置没想象中那么麻烦,官方示例改改就能跑,单卡也能用。ZeRO-3对单卡收益不大还容易拖速度,你这情况ZeRO-2够用了。实在不想折腾框架,先试试checkpoint梯度或者换8bit Adam,改动小见效快。
几十万条数据用BERT-base,A100 40G跑到batch 16就OOM有点不太正常,先确认下是不是max_length设太长了,512砍到256显存能省一大截。DeepSpeed ZeRO-2迁起来其实没那么吓人,配置文件改改就行,原生循环也能接,比ZeRO-3省心。ZeRO-3适合模型大到单卡放不下,你这个规模ZeRO-2基本够用,别一上来就上3。