最近在试着微调一个7B的LLaMA模型做垂直领域任务,单卡A100 80G跑了一轮就OOM了。我知道可以用LoRA或者QLoRA,但这次想试试全量微调看看效果上限。
用PyTorch微调LLaMA时显存总爆,大家用DeepSpeed还是自己写梯度检查点?
全部回复
共 168 条我之前也卡在过这个坎上,80G跑全量7B确实紧巴巴的。DeepSpeed ZeRO-3配合offload能撑住,但速度慢得让人怀疑人生,而且CPU内存得够大。我自己后来是梯度检查点+手动控制激活值显存,勉强跑起来,但batch size缩到1了,loss还容易震荡。你试过把序列长度砍到512或者用flash-attention吗?有时候瓶颈不在模型本身,在输入长度上。
说实话我之前也卡在这,全量微调7B确实吃显存,光优化器状态就够呛。我后来是DeepSpeed ZeRO-3加手动梯度检查点混着用的,把checkpointing放在每个Transformer层中间,峰值能压到60G左右。不过你既然想试上限,建议先看看是不是激活值占大头,用activation offload到CPU能再省不少。另外batch size调到1试试,梯度累积开大点,别让OOM卡死整个实验。
80G都OOM的话,全量微调7B确实得把优化器状态和激活值都算进去,光靠梯度检查点不够使。我试过DeepSpeed ZeRO-3配offload,能跑起来但速度慢得怀疑人生,后来发现把batch拆小点配合梯度累积,再手动把attention的激活值缓存砍掉,反而比硬上框架省心。你检查过是不是序列长度太长导致的峰值暴涨?有时候把max_len从2048降到1024,显存直接砍半。另外如果非要用全量,建议看看torch的activation checkpointing能不能单独用在某些层上,别全开。
80G都爆的话,多半是激活值没管住,先试试梯度检查点加mixed precision,能撑住再上DeepSpeed。
说实话全量微调7B在单卡A100上就是地狱难度,80G看着大但光优化器状态就能吃掉快30G。建议先试DeepSpeed ZeRO-2加offload,比手写梯度检查点省心太多,至少不用自己管内存碎片。另外你确定需要全量吗,我试过用LoRA在同样数据上只差2-3个点,但训练时间能缩短一半以上。如果非要全量,可以看看activation checkpointing配DeepSpeed的partition activation,能再省一波显存。
全量微调7B确实挺吃显存的,光靠梯度检查点也就省个30%左右,optimizer states那部分才是大头。我建议直接上DeepSpeed ZeRO-2,省心不少,offload到CPU虽然慢点但至少不炸。自己写检查点容易踩坑,除非你有特殊需求,不然没必要重复造轮子。
全量微调7B单卡80G确实紧,我用的DeepSpeed ZeRO-2加梯度检查点才勉强跑起来,你可以先试试这个组合。
7B全量微调单卡80G确实够呛,optimizer states加上梯度激活值随便就上百G了。我之前试过DeepSpeed ZeRO-2,能跑但通信开销不小,单卡的话其实没啥优势。自己写梯度检查点更省显存,但训练速度会慢个20%左右,得看你任务能不能接受。