最近在跑一个文本分类任务,模型用的是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?
全部回复
共 100 条40G的A100跑BERT-base batch16就爆,感觉有点不对劲,你检查下是不是序列长度太长或者参数量没冻结?我之前用类似配置batch能到32。ZeRO确实有效但别一上来就上3,ZeRO-2够用了,配置也不复杂,官方文档照着改几行就行。另外可以试试把优化器换成AdamW的8-bit版,能省不少显存。
40G还爆的话,多半是padding太长或激活没开gradient_checkpointing,先试这个,比折腾DeepSpeed快多了。
ZeRO-2够用了,ZeRO-3慢还容易踩坑,单卡真没必要上。
40G的A100跑BERT-base把batch压到16还OOM,我怀疑你seq length是不是拉太长了,或者attention mask没处理好。先试试把序列长度砍到128或者用dynamic padding,很多情况下比上DeepSpeed立竿见影。ZeRO-2和ZeRO-3我实际用下来,单卡场景下两者差别不大,ZeRO-3主要是为了多机扩展,你这种单卡任务真没必要折腾,而且DeepSpeed跟原生循环混在一起调试起来挺折磨的。另外可以看下是不是优化器状态占了大头,换AdamW加bitsandbytes的8bit版能省不少。
40G跑BERT-base才16就爆?先查下是不是序列长度和pad的锅,梯度累积要配amp用。
ZeRO-2够用了,别直接上3,通信开销大。
先试试梯度检查点,能省不少显存,Deepspeed配置确实费劲但ZeRO-2性价比挺高的。
40G的A100跑BERT-base batch16就爆,你这序列长度是不是拉太长了?先check一下input长度,能截断到512以内就尽量截,实在不行再考虑DeepSpeed。ZeRO-2对你这个单卡场景基本没用,ZeRO-3倒是能省但配置麻烦,收益不如直接砍batch加梯度累积,慢就慢点反正几十万条数据也就多跑几小时。
另外可以试试torch.utils.checkpoint,激活重计算能省不少显存,代价是大概20%的算力换显存,你这个场景应该挺划算的。amp你已经试过了,那再检查下dataloader有没有把整个batch的tensor都搬到GPU上,有时候一些小细节能省出一大块。
先试gradient_checkpointing,能省一半多显存,比动DeepSpeed快多了。
40G跑BERT-base还爆,多半是序列长度或数据加载的锅,先查查这两个。ZeRO-2够用了,别直接上3,配置麻烦收益还不大。
40G的A100跑BERT-base到16就爆,多半是序列长度和attention的锅,先看看max_len是不是设太长了,能砍到128的话显存直接减半。DeepSpeed配置确实烦,但ZeRO-2对单卡场景其实没啥用,它主要省的是多卡通信显存,单卡该爆还是爆。我建议你先试试torch.utils.checkpoint,把bert的梯度检查点打开,显存能降30%以上,速度牺牲比梯度累积小得多。另外你数据几十万条,分类任务其实可以试试用更大的batch但截断到256长度,很多情况准确率不会掉太多。
40G的A100跑BERT-base batch16就爆,有点不对啊,你检查下是不是max_len设太长或者dataloader没开pin_memory?我之前用类似配置batch32都没问题。真要省显存,DeepSpeed的ZeRO-2其实配置没那么吓人,改几行就行,主要省的是优化器状态,比砍batch划算多了。ZeRO-3是连参数和梯度都分片,但通信开销大,单机单卡反而可能更慢,你这种情况ZeRO-2足够了。另外可以试试gradient_checkpointing,虽然会慢点但能省一半显存,比梯度累积强。
ZeRO-2够用了,先开AMP再上DeepSpeed,别纠结batch大小,速度慢多半是梯度累积步数没调好。
amp + 梯度累积还不够的话,直接上DeepSpeed吧,ZeRO-2配置不复杂,省显存效果立竿见影。
你这情况我太熟了,BERT-base单卡40G只吃16的batch确实不正常,先别急着上DeepSpeed,八成是序列长度或者激活值没处理好。建议你先开torch.utils.checkpoint,也就是梯度检查点,能省一半左右显存,代价就是多20%算力,但比砍batch强多了。另外你把max_length从512砍到128试试,文本分类任务很多用不到那么长,显存直接掉一大截。至于DeepSpeed,如果你不想动训练循环,ZeRO-2其实改动挺小的,官方有现成wrapper,但说实话单卡上ZeRO收益不大,它主要是为多卡设计的。ZeRO-3就更重了,还会把参数切到CPU,单卡反而可能变慢,不建议碰。我自己的经验是,先检查一下dataloader里有没有把padding搞到最长,用动态padding或者bucket sampler能省不少。还有你说amp效果有限,检查一下是不是loss scaling出了问题,或者干脆换成bf16,A100支持得更好。如果这些都试完还爆,那再考虑DeepSpeed也不迟,但大概率你砍完序列长度就能跑32的batch了。
40G的A100跑BERT-base batch16就爆,大概率是序列长度和注意力矩阵吃满了,先看看max_len能不能砍,很多文本分类任务根本用不到512。DeepSpeed配置确实烦,但ZeRO-2对单卡也有offload优化,比硬砍batch强,至少能保住梯度质量。ZeRO-3主要是为了多卡分参数,单卡收益没那么明显,别被网上吹晕了。另外试下torch.utils.checkpoint,用计算换显存,配合amp可能就够你用了。
40G的A100跑BERT-base batch16就爆,可能是序列长度或者padding没处理好,建议先看看有没有无效token占显存。DeepSpeed确实值得搞,但ZeRO-2就够你用了,ZeRO-3主要是为超大模型跨节点设计的,单卡上收益不大。另外可以试试torch.utils.checkpoint,用一点点计算换显存,配合AMP基本能解决你的问题。梯度累积慢是因为你batch太小,checkpoint加上之后batch能翻倍,速度反而更快。
40G的A100跑BERT-base单卡16都OOM,这不太正常,你是不是忘了关梯度检查点或者序列长度没设上限?我建议先看一眼显存到底被啥吃了,有时候是激活值在作怪,不是参数。DeepSpeed迁移成本确实高,但ZeRO-2其实挺好上手的,比想象中简单,至少比砍batch强,砍到8以下收敛效果会明显变差。ZeRO-3的话单卡基本用不上,那是给多卡通信设计的,别被名字唬住。你这种情况我猜是激活值占大头,试试开个gradient_checkpointing,配合amp,可能直接就能塞下32的batch,速度比梯度累积快多了。
40G单卡跑BERT-base还OOM,batch16确实有点反常,先检查下是不是序列长度没padding到统一阈值,或者DataLoader里num_workers开太多导致显存碎片化。DeepSpeed迁移成本其实没那么高,ZeRO-2基本够用,ZeRO-3主要省在超长序列和超大模型上,你这种任务用ZeRO-2加offload optimizer就够了。另外可以试试把优化器换成Adafactor,省显存效果比砍batch明显,就是收敛得调下学习率。
你试试gradient_checkpointing,BERT-base开这个能省一半显存,代价就是多20%左右的训练时间,但比梯度累积那种慢法强多了。DeepSpeed说实话配置起来有点折腾,如果只是单卡的话,torch.compile加混合精度再加上checkpointing,应该能稳住batch32。ZeRO-3在单卡上没啥优势,多卡才值得上。
我上次跑类似任务直接换成了PagedAdamW,显存占用直接掉了30%,配合amp基本没再OOM。ZeRO-2和ZeRO-3实际差别主要在通信开销,单卡场景下ZeRO-2就够,别折腾ZeRO-3了。另外你确认下是不是用了动态padding,固定长度的话浪费太多显存,改成动态batch能省不少。
建议先看一眼是不是模型没设eval模式
先试试gradient checkpointing,能省不少显存,比换DeepSpeed简单多了。
40G跑BERT-base还爆显存,先查查是不是max_len设太长或数据没padding到统一长度。
ZeRO最香的是能开超大batch,但单卡上收益有限,不如试试gradient checkpointing,代码改动就两行。
40G跑BERT-base还OOM?先查查是不是数据加载或padding的锅,AMP加梯度累积够用了。