最近在调一个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了,但还是卡在优化器状态上。求指点,感谢!