最近在尝试用DeepSpeed的ZeRO-3微调一个7B的LLaMA模型,单卡A100 80G,batch size已经降到1了,结果训练到一半还是OOM。我用了deepspeed的auto offload,也试过把cpu_offload和pin_memory打开,但感觉没太大改善。是不是我理解错了ZeRO-3的原理?还是说7B模型本来就不太可能在单卡上微调?另外,我看到的显存占用大部分是optimizer states和activation,有没有什么trick能省一省?比如用gradient checkpointing或者混合精度?但我试了bf16好像效果不稳定。求大佬指点,先谢过了。
用DeepSpeed跑LLaMA微调,ZeRO-3内存还是爆了,是我设置有问题吗?
全部回复
共 2 条说实话,7B模型单卡A100 80G跑ZeRO-3微调确实挺极限的,但理论上不至于完全跑不动。我觉得问题可能出在activation memory上,ZeRO-3只分摊了模型参数、梯度和优化器状态,但激活值还是每份都存着,batch size=1下如果序列长度太长,比如超过2k,那激活值照样能吃掉大几十G。你试过gradient checkpointing吗?这个能显著降低激活内存,但会慢一点,通常配合bf16用。说到bf16不稳定,你检查过模型本身的数值范围吗?有些LLaMA变体用bf16会有loss spike,可以试试用torch.cuda.amp的混合精度,手动控制哪些层用fp16哪些用bf16。另外,auto offload我建议慎用,CPU offload本身会有大量数据传输,反而可能因为内存碎片或带宽瓶颈导致隐性OOM,不如手动把optimizer states切到CPU上,只保留模型参数和梯度在GPU。最后一个小trick:如果序列长度不固定,可以试试动态padding或者用flash attention,能省不少显存。
刚看到这帖子,感觉你遇到的情况挺典型的。ZeRO-3本身是把模型参数、梯度和优化器状态分散到多卡或者CPU上,但单卡A100跑7B模型,其实理论上不是完全没戏,问题往往出在activation和临时buffer上。你batch size都降到1了还爆,大概率是activation占据了大量显存,ZeRO-3优化不了这部分,因为它是按层计算的。gradient checkpointing确实是个好方向,它会在反向传播时重新计算中间激活,而不是全部存储,基本能省下30%-50%的显存,代价是增加一点点计算时间。混合精度的话,bf16在A100上通常比fp16稳定,但如果你发现loss震荡,可以试试把一些敏感层(比如embedding)保留在fp32,或者加上loss scaling。还有个小trick是调整deepspeed的pin_memory参数,它有时候反而会占用额外内存,不如手动控制CPU offload的阈值。另外检查一下是否不小心加载了额外的cache或eval模式下的缓存,比如transformers库的past_key_values。总的来说,7B单卡微调是可行的,我见过有人用gradient checkpointing加ZeRO-3把batch size开到2甚至4,你试试把offload只对optimizer state和gradient生效,activation部分留给显存,可能会平衡一些。