最近开始尝试用Llama 3.1 8B做领域微调,跟着教程写了LoRA,batch size设到1,gradient checkpointing也开了,结果3090(24G)还是OOM。我的输入长度大概2k tokens,是不是跟序列长度有关?看到有人说用DeepSpeed ZeRO-3或者Flash Attention能省显存,但配置起来有点复杂,不太确定是哪里出了问题。另外,torch.compile会有帮助吗?现在用的是PyTorch 2.1,cuda 12.1。求有经验的大佬指点下常见坑,谢谢!
刚转大模型方向,用PyTorch跑LLM微调,显存总爆掉怎么优化?
全部回复
共 178 条跟序列长度关系很大,2k tokens对8B模型来说,光激活值就很吃显存,你试试把max_length砍到512或者1k,应该立竿见影。Flash Attention确实能省不少,不过3090上直接开sdp_a maybe就够了,不用上FA2那么麻烦。DeepSpeed ZeRO-3对单卡其实帮助有限,反而可能拖慢速度,建议先把offload关掉,用ZeRO-2配合gradient checkpointing试试。torch.compile在2.1版本对LLM微调支持一般,容易编译报错,先别折腾,等跑通再优化。
另外检查下是不是LoRA只冻住了transformer层但没冻embedding,常用做法是把embedding和lm_head也一起冻掉。我踩过最坑的是dataloader里没设pin_memory=False,结果pinned memory把显存挤爆了,你排查下这个。
说实话你这个配置我太熟了,之前用3090跑7B模型也踩过一模一样的坑。24G显存跑8B LoRA按理说是够的,但2k序列长度确实是个大头,因为attention的计算和显存占用是跟序列长度平方相关的,所以问题大概率出在这。建议你先别急着上DeepSpeed,那个配置复杂度对新手不太友好,先把Flash Attention加上试试,它能显著降低attention部分的显存峰值,而且现在transformers库支持起来很简单。另外torch.compile可以开,但别指望它省显存,它主要是加速,有时候反而会因为编译过程多占一点显存。还有个容易被忽略的点,检查下你的LoRA是不是只作用在attention层,如果也加了MLP层,可以试着减少rank值或者只微调部分模块。我当时的解法是开了gradient checkpointing之后,再把batch size调到2,同时用unsloth的优化版模型,显存直接降了差不多30%,你可以参考下。另外注意下PyTorch 2.1和CUDA 12.1的版本搭配,有些操作符在特定版本下会有额外的显存开销,如果方便的话可以升到2.2以上。最后如果实在不行,就考虑用4bit量化加载模型,QLoRA方式微调,效果损失很小但显存能省一大截。
24G跑8B LoRA按理说够用,但2k序列长度确实是个坎,attention的显存占用是平方增长的。你试试把LoRA的rank降到8,同时用unsloth优化过的模型加载,能省不少。DeepSpeed ZeRO-3配置确实麻烦,但可以用deepspeed的zero.Init直接包model,配合flash attention基本能解决。torch.compile建议先别开,它跟gradient checkpointing有时候有兼容问题,反而更吃显存。你检查下是不是输入了没必要的padding,把max_seq_len硬设成2048也会浪费。
我最近也在搞类似的事情,8B模型2k长度在24G上确实很极限。你开了gradient checkpointing但没开Flash Attention的话,attention部分还是会把激活值撑爆,建议先把flash attn加上,能省不少。另外LoRA是不是把target modules设全了?比如q,k,v,o,gate,up,down这些全加上,显存占用会明显比只微调部分层高。ZeRO-3配置起来确实麻烦,但如果你只是单卡跑,其实用ZeRO-2或者干脆offload到CPU就够了,ZeRO-3主要解决多卡分片问题,单卡上收益不大。torch.compile我试过,对显存帮助有限,主要是加速,而且第一次编译很慢,建议先把其他优化做完再考虑。还有个容易忽略的点,你的输入长度2k是tokens还是字符?如果是字符的话,tokenizer切出来可能远不止2k,建议打印一下实际序列长度确认。我之前就栽过这个坑,以为长度可控,结果每个样本都塞满了padding到最大长度,显存全浪费在padding上了。你可以试试把max_length设成训练集实际最大长度,别用模型默认的4k或8k,能省一大块。
你这个配置跑Llama 3.1 8B全量微调确实很极限,我怀疑问题出在输入长度上——2k tokens的激活值在24G里本来就吃紧,LoRA只能省可训练参数,省不了激活内存。建议先试下Flash Attention,配合gradient checkpointing能把激活占用砍掉一大截,我之前用类似配置从OOM降到能跑。torch.compile对显存帮助不大,主要是提速,但可能跟某些库不兼容,先别折腾。另外DeepSpeed ZeRO-3对单卡场景其实提升有限,反而容易引入通信开销,不如先把attention优化做了。如果还爆,把LoRA的rank降到8,或者用QLoRA(4bit量化)直接换显存空间,效果立竿见影。
8B塞24G还得靠flash attention,顺手把rope改成动态缩放,能再省一截。
唉,这配置单看真不怪你,8B模型塞2k序列,24G确实卡在临界点上。你开LoRA但别忽略base model的权重占用,试试把gradient checkpointing换成offload到CPU,或者把batch size压到1但用梯度累积模拟更大batch。Flash Attention能上就上,序列长度一长省显存效果立竿见影。torch.compile对显存帮助不大,主要是提速,别指望它救OOM。还有个小坑,检查下是不是把优化器状态也塞GPU了,换AdamW的8bit版能省不少。
3090跑8B全参微调本来就极限,你开LoRA还爆的话大概率不是显存不够,是激活值峰值太高了——2k序列长度其实不算长,但Llama的注意力计算很吃显存,建议先试试把flash attention打开,对显存占用立竿见影,而且代码改动就几行。DeepSpeed ZeRO-3在单卡上收益不大,反而增加通信开销,不如把优化器状态分片到CPU(ZeRO-Offload)来得实在,能省下几个G。torch.compile目前对训练加速有限,但对显存优化基本没帮助,别指望它救火。最后确认下你的LoRA是不是只冻结了原模型但没把梯度设成False,很多人漏了这一步,导致反向传播时还保留全部参数的梯度,那24G肯定不够。
我之前跑7B也遇到过这问题,2k长度加上8B确实容易爆,Flash Attention能直接省一半左右的激活内存,值得先配上。DeepSpeed ZeRO-3对单卡其实帮助不大,主要是多卡场景,单卡更建议把offload打开试试。torch.compile我这版用下来加速还行但省显存有限,不如先把attention的kernel换了。另外检查下是不是把eval也包在gradient checkpointing里了,那个经常被忽略。
试试把输入截断到1k,再开gradient checkpointing加bf16,24G跑8B LoRA应该够,Flash Attention提升不大但能省点。
你input长度确实是主要瓶颈,2k对8B来说太吃显存了,换成4bit量化加载模型,或者把batch再拆小点。
24G跑8B LoRA还爆显存,2k序列长度确实是关键,你试试把max_length临时砍到1k看能不能跑通,能跑就是长度问题。Flash Attention值得折腾一下,配置其实不难,主要是注意和你的PyTorch版本匹配,torch.compile对显存优化不大但能提速。另外检查下是不是把LoRA的target modules设太多了,只冻结全部原参数、只训attention层的话能省不少。
24G跑8B LoRA按理说够用了,你开gradient checkpointing还爆大概率是序列长度在作祟,2k tokens的激活值很吃显存,可以试试把max_length先砍到1k看还爆不爆。Flash Attention确实能省不少,而且现在transformers里直接调attn_implementation="flash_attention_2"就行,没你想的那么复杂。DeepSpeed ZeRO-3对这种单卡场景其实帮助不大,主要是省下的是参数和优化器状态,你LoRA本身就只训练少量参数,不如先把attention的显存优化搞定。torch.compile对显存没啥帮助,顶多加速,别指望它能救OOM。另外确认下你是不是把输入padding到固定长度了,用动态padding能省不少。
同款配置踩过坑,2k长度确实吃显存,但24G爆掉大概率是padding没处理干净,试试把attention mask和pad token统一到左边,能省不少。Flash Attention值得上,不用DeepSpeed也能缓解,但记得先升级到PyTorch 2.2,torch.compile对显存优化不大,主要提速度。另外检查下是不是梯度过大导致峰值涨了,加个gradient clipping能稳很多。
24G跑8B LoRA还爆显存,大概率不是显存不够,而是你踩了Llama 3.1的坑——它的rope默认用fp32计算,会额外吃一块显存。你可以试试在config里把rope_scaling的type改成"linear",或者直接换用bitsandbytes的4bit量化加载base model,LoRA只训练adapter,这样显存占用能直接砍半。我自己的经验是,2k长度输入下,8B模型用4bit+LoRA,batch size 1在3090上大概占用14-16G,完全能跑。Flash Attention确实很有用,但注意它跟gradient checkpointing配合时,有的版本会有bug导致反向传播显存不降反升,建议先单独开FA试试。torch.compile在微调场景收益不大,尤其在动态shape下反而可能触发重编译卡顿,不如把精力放在优化输入padding上——你检查下是不是用了动态padding,如果每个batch都按最长序列padding,2k长度但实际很多样本可能只有几百token,浪费巨大。我一般用unpad或者把数据集按长度分组排序,能省出20-30%显存。另外DeepSpeed ZeRO-3对单卡场景其实没必要,那是多卡才需要的,单卡开ZeRO-2配合offload optimizer就够了,配置也就十几行。最后建议你装个nvidia-smi监控实时显存,看看到底是激活值还是优化器状态在涨,这样定位问题会快很多。
你这个问题我太熟了,2k长度加上8B模型,24G确实卡在临界点上。先别急着上DeepSpeed,把flash attention开了能省不少,而且torch.compile在2.1上对LLM微调收益挺明显的,尤其长序列。另外检查下是不是把optimizer状态也塞进显存了,用8bit adam或者offload optimizer能救回来。
8B全参微调本来就不是给单卡玩的,算了下你这配置光激活值就超了,试着把max_len砍到1k或者直接上seq parallel吧。
你试试把max_length从2k砍到1k,8B在24G上跑2k序列本来就紧,开flash attention比DeepSpeed省事多了。
显存爆大概率是序列长度和attention的锅,先试Flash Attention,torch.compile对显存帮助不大。
24G跑8B LoRA按理说应该够的,你试下把max_length从2k砍到1k看看,很多情况下OOM就是序列长度直接拉爆了激活显存。DeepSpeed ZeRO-3在单卡上其实没啥必要,反而可能拖慢速度,Flash Attention倒是值得装一下,能省不少显存。torch.compile对显存优化帮助不大,主要提升的是计算速度,你现在瓶颈在显存不在算力。另外检查下是不是优化器状态或者梯度累积没关,有时候这些细节比LoRA配置更吃显存。
我之前也卡在这过,2k序列长度对8B来说确实挺吃显存的,LoRA虽然省了训练参数但激活值还是大头。建议先试试torch.compile,配合最新版flash-attention能省不少,我这边开了之后显存占用直接降了三分之一。ZeRO-3配置确实麻烦,但比FSDP更容易踩坑,如果只是单卡的话不如先检查下是不是padding策略没做对,把动态padding加上可能比上那些大件更立竿见影。