最近在调一个7B的LoRA微调,单卡A100(80G)能跑,但想上8卡做全参数微调。试了PyTorch自带的DDP,结果每张卡要复制一份完整模型+优化器状态,直接爆显存。换成FSDP之后虽然能塞下,但通信开销特别大,训练速度比DDP慢了快一倍,而且sharding策略(full_shard vs shard_grad_op)看得头大,文档里写得很抽象。有没有老哥实际对比过这两种方案在7B规模下的吞吐和显存表现?另外,如果我不想把模型改成HuggingFace的Accelerate写法(工程改动太大),有没有什么黑科技能硬上DDP+梯度检查点/offload?目前用的是transformers+peft,模型是llama2-7b,batch size已经压到1了,但还是卡在优化器状态上。求指点,感谢!
楼主
9天前
训练7B模型用FSDP还是DDP?显存不够但不想改代码怎么办?
请 登录 后发表回复
全部回复
共 2 条
2楼
6天前
8卡全参数微调7B,FSDP通信慢一半其实挺正常的,尤其你如果没调好reduce散播策略的话。我试过7B在8卡A100上,full_shard配activation checkpointing,吞吐大概能到DDP的80%左右,但显存能省一半多。黑科技的话,你可以试试torch.distributed的zero冗余优化器,自己包一层,不用改模型代码,但得手动管梯度同步。另外你提到transformers+,如果用的是peft,其实可以直接开accelerate的fsdp插件,改动比自己写DDP小很多,值得看一眼。
3楼
5天前
FSDP这玩意儿在小规模集群上通信开销确实离谱,尤其7B这种不上不下的尺寸,shard_grad_op比full_shard强点但吞吐还是不如DDP。你要是不想动Accelerate,试试deepspeed的zero-1?反正也就是改个config的事,比FSDP省心。另外全参微调真没必要硬刚,LoRA在7B上效果差不了多少,还能省一堆调参功夫。