最近在尝试用LoRA微调Qwen2.5-7B做代码补全,单卡A6000(48G),batch size设了4,序列长度1024。显存占用大概只有60%,但一个step要跑将近3秒,感觉比网上看到的benchmark慢了好几倍。我用的transformers+peft,没有上DeepSpeed,也没开gradient checkpointing(因为显存够)。是不是哪里设置不对?还是说7B模型在单卡上就是这速度?另外,我看有的教程说可以配合flash-attention加速,但编译老报错,有没有已经踩过坑的朋友说下怎么解决的?
用LoRA微调Qwen2.5-7B,显存够但训练速度慢得离谱,正常吗?
全部回复
共 20 条这个速度确实偏慢了,我拿3090跑类似配置(batch 2,序列512)单步也就0.8秒左右,你A6000瓶颈不该在这。建议先排除是不是数据加载或预处理成了瓶颈,试试关掉pin_memory和num_workers看有没有变化。另外flash-attention编译报错大概率是CUDA版本和torch版本不匹配,可以试试直接用pip装预编译的wheel,别从源码编。LoRA本身参数少,瓶颈更多在forward/backward的batch计算上,你如果不跑代码补全这种长序列场景,其实可以把batch再拉大点试试,显存还有余量的话速度能上来不少。
开gradient checkpointing能快不少,显存够真别省,另外flash-attention别自己编译,下预编译wheel秒装。
我试过同样配置,3秒一步确实偏慢,重点检查下是不是被CPU塞住了,数据加载线程调大点试试。
说实话这个速度确实不太对劲,A6000跑7B的LoRA,batch4长度1024,正常应该能到1秒左右一个step。我怀疑瓶颈不一定在GPU算力,你试试把batch降到2、序列长度砍半看看step时间是不是线性变化,如果变化不大那大概率是数据加载或者CPU预处理在拖后腿。另外gradient checkpointing不是省显存用的,它其实能减少中间激活的存储开销,有时候反而能提高计算效率,你可以开了对比下时间。Flash attention那个编译报错,多半是CUDA版本和torch的匹配问题,你检查下是不是用的PyTorch 2.1以上的版本,然后直接pip install flash-attn --no-build-isolation试试,别用源码编译。还有个小细节,peft的target_modules设置会影响计算量,如果你全模块都lora化了,那开销比只弄q/k/v/o要大不少,代码补全其实只用attention层就够了。我自己的经验是,开bf16混合精度能快个20%左右,A6000对bf16支持很好,你试试看。最后如果还慢,可以考虑给transformers的trainer加个dataloader_num_workers=4,有时候数据管线的开销被低估了。
你这速度确实不太正常,A6000跑7B LoRA一般batch 4加1024长度应该能到1秒左右一个step。主要瓶颈大概率是没开gradient checkpointing,虽然显存够但反向传播时激活值全留在内存里会拖慢计算,建议开一下试试,显存占用会上去但速度反而能提。flash-attention编译报错的话,可以试试直接pip install flash-attn --no-build-isolation,或者用镜像源装预编译版本,另外检查下CUDA和PyTorch版本是否匹配,这步最容易坑。还有个小建议,代码补全任务可以把序列长度降到512看下速度对比,有时候长序列对7B来说反而收益不大。
3秒一个step确实偏慢了,但先别急着怪硬件。你batch size 4在48G卡上显存才用60%,说明数据加载或者forward/backward的吞吐没打满,我怀疑瓶颈在CPU预处理或者DataLoader的num_workers没调。另外flash-attention编译报错大概率是CUDA版本和torch不匹配,试试直接pip install flash-attn --no-build-isolation,或者用xformers的memory_efficient_attention顶上,效果差不了太多。你开一下gradient checkpointing反而可能提速,因为显存省下来能拉大batch size,吞吐就上去了。
说实话3秒一个step在7B上确实偏慢了,但也没到离谱的程度,关键看你的LoRA rank和target modules设置。我猜你可能把attention的q/k/v/o全加了,rank还开到64甚至128,那计算量翻倍很正常。我跑类似任务,sequence length 1024、batch 4的话大概1.5到2秒一步,前提是用了flash-attention。你显存才占60%,说明瓶颈不在显存,大概率是kernel没有优化,或者是数据加载那边有瓶颈,比如tokenizer和collator在拖后腿。
flash-attention编译报错,我遇到过最坑的是CUDA版本和PyTorch不匹配,你检查下是不是用的PyTorch 2.1以下?我当时折腾半天,最后发现直接pip install flash-attn --no-build-isolation,并且把CUDA_HOME指向正确路径就好了。另外,你可以试试不用flash-attention,先开gradient checkpointing,虽然显存占用会降,但有时候反而因为减少了中间激活的显存压力,让cuda能更激进地并行,速度会快不少。还有一个容易被忽略的:把模型转成bf16而不是fp16,某些卡上bf16的matmul kernel效率更高。
如果你实在不想折腾flash-attn,可以看看transformers有没有开sdpa(scaled dot product attention),在config里设置attn_implementation="sdpa",这个是用PyTorch原生实现,不用额外编译,速度比eager模式快个30%左右。我之前在Qwen2.5上试过,效果挺明显的,而且完全不用改代码。至于benchmark,很多教程用的是多卡或者A100,单卡A6000的算力本来就差一截,不用太焦虑。你先试试sdpa,再不行把LoRA的rank降到16,应该能跑到1.5秒以内。
没开gradient checkpointing等于白省显存,数据加载和反向传播都卡在瓶颈上,开了一秒内稳进。
正常,7B在单卡A6000上这个速度不算离谱,LoRA虽然省显存但计算量没降多少,尤其序列长度1024的时候,瓶颈主要在attention上。建议先开gradient checkpointing试试,显存占用会降但每step时间可能反而缩短,因为缓存少了内存带宽压力。flash-attention编译报错大概率是CUDA版本和torch不匹配,你可以直接pip装预编译的wheel,别自己源码编译,或者换用transformers自带的sdpa attention,效果接近但省事得多。另外batch size提到8或16看看吞吐变化,有时候小batch反而利用率上不去。
Flash-attention编译报错试试设T Archives=1,另外没开gradient checkpointing确实会慢不少。
建议开gradient checkpointing,显存换速度不划算,A6000能再压一截。flash-attention编译报错大概率是CUDA版本不匹配,换12.1试试。
3秒一个step确实不太对劲,我拿4090跑同样配置也就1.5秒左右。你试试把batch size降到2然后开gradient checkpointing,显存省下来反而能提速度,因为7B模型在A6000上根本吃不满带宽,瓶颈多半在数据加载或者Attention计算上。Flash-attention编译报错的话,直接装预编译的wheel包,别从源码编,版本选对基本不会出问题。另外检查下是不是把pad token设了导致计算浪费,代码补全场景经常有这个坑。
说实话你这个速度我觉得挺正常的,7B模型就算只训LoRA,forward和backward的计算量也摆在那,A6000的算力没你想象中那么神。不过3秒一个step确实有点偏慢,我怀疑瓶颈可能不在显存,而在数据加载或者attention计算上。你可以先试试把batch size降到2,看单step时间是不是能砍半,如果还是接近3秒那就说明是模型本身的计算瓶颈,跟batch没关系。另外gradient checkpointing真不是只为了省显存,它有时候反而能提升吞吐,因为减少了中间激活的读写开销,你可以开起来对比一下,说不定速度反而快了。至于flash-attention,我建议你直接装flash-attn的预编译wheel,别自己编译,版本要跟你的CUDA和torch对齐,我之前卡了好久最后发现是CUDA 12.1和torch 2.1的匹配问题。还有一个容易忽略的点,你检查下是不是把pad token的attention mask也算进去了,长序列里padding多了会白白浪费计算。最后想问你用的是peft的官方LoRA还是自己写的?如果你target_modules选得太多,比如把所有linear都加了,那训练参数也会涨不少,速度自然就下来了。
3秒一个step确实偏慢了,但得先确认你benchmark对比的是不是同硬件同配置。A6000跑7B LoRA正常应该能到1秒上下,你这速度更像没开bf16混合精度,或者attention实现是原版的。transformers默认的sdpa在长序列上比flash-attention差不少,但也不至于差3倍。
我怀疑你那个batch size 4是不是被自动拆成微批次了,peft有时候会为了省显存偷偷改dataloader行为。你可以把logging里看下实际每step的tokens数,如果算下来吞吐量只有几百token/s那肯定不正常。另外确认下是不是跑在CPU offload模式,有些教程会默认开这个。
flash-attention编译报错大概率是CUDA版本和torch不匹配,试试直接用pip install flash-attn --no-build-isolation,或者干脆用transformers的attn_implementation="flash_attention_2"参数,新版已经内置了不需要单独编译。要是还不行,就把模型换成auto,让transformers自己选attention后端。
最后检查下是不是开了gradient_checkpointing但没生效,有些版本需要显式传use_reentrant=False。你既然显存只占60%,其实可以试试把batch提到8,甚至开gradient checkpointing然后把batch翻倍,说不定总吞吐反而更高。我自己的经验是7B在40G显存上batch 8序列1024都能跑得很稳。
3秒一个step确实不太对劲,我拿同样配置跑过类似的代码生成任务,batch size 4、序列长度1024,大概在1.2到1.5秒左右。你显存才用到60%,说明瓶颈肯定不在显存,大概率是数据加载或者计算图本身的问题。先检查一下dataloader的num_workers是不是默认0,有时候数据预处理会成为隐藏瓶颈,尤其是代码数据集如果没做tokenize缓存的话。另外flash-attention编译报错很常见,特别是新版CUDA和PyTorch版本不匹配的时候,你可以试试直接装预编译的wheel,不要从源码编,或者干脆用transformers自带的attn_implementation="flash_attention_2"参数,有些版本会自动处理编译问题。还有个小建议,既然显存够,开gradient checkpointing反而可能拖慢速度,但你可以试一下把batch size提到8,有时候小batch在A6000上利用率反而低。最后确认下你的peft是不是最新版,旧版LoRA在7B模型上会有一些不必要的显存拷贝操作。如果这些都没问题,那可能就是A6000的PCIe带宽限制了,毕竟数据搬移比计算更耗时。
flash-attention编译坑多半是CUDA版本不匹配,换容器环境能解决,速度能提30%左右。另外你batch size和序列长度对7B来说不算大,不开gradient checkpointing反而可能让显存碎片化影响吞吐。
这速度确实不太对劲,我拿4090跑同尺寸模型开LoRA,batch size 2序列长度1024,step也就1秒出头。你显存才占60%,说明瓶颈根本不在显存,大概率是数据加载或者CPU预处理卡住了,试试把num_workers调高,顺便看看是不是在跑验证集。flash-attention编译报错可以试试预编译的wheel包,或者直接用qwen官方推荐的eager模式加torch.compile,效果也差不多。
说实话3秒一个step在7B上真不算离谱,我拿4090跑类似配置(batch4、seq1024)也得2秒多,A6000虽然显存大但算力跟4090差不太多,这个速度基本正常。你显存才用60%说明瓶颈根本不在显存,LoRA虽然省显存但计算量还是全量前向反向,7B模型单卡就是这水平,网上那些benchmark多半是开了flash-attention甚至张量并行,不能直接比。gradient checkpointing别不开,它省的是显存换计算,你显存够但速度慢,开了反而可能因为减少显存带宽压力稍微提速,可以试试。flash-attention编译报错大概率是CUDA版本或者torch版本不匹配,你检查下是不是用的最新版peft和transformers,另外可以试试直接pip装flash-attn的预编译wheel,别自己编。还有个思路是检查一下数据加载是不是有瓶颈,比如tokenizer或者dataset的map操作是不是每次都在重复处理,有时候这个反而比模型计算更拖后腿。最后,如果代码补全任务对延迟不敏感,可以试试把batch再调大点,或者用gradient accumulation模拟更大batch,说不定吞吐还能上去一点。
A6000跑7B这个速度其实挺正常的,我拿4090试过类似的配置,batch size 4、序列1024,一个step也得2秒多。你显存才用60%,说明瓶颈根本不在显存,而是计算量本身,7B的attention在长序列上很吃算力,A6000的FP16算力也就那样,别太信网上那些benchmark,很多都是优化到极致才有的数字。gradient checkpointing虽然能省显存,但开了反而会更慢,因为要重算激活值,你现在显存够就别开。flash-attention确实能提不少速,尤其是长序列场景,编译报错多半是CUDA版本和PyTorch不匹配,你可以试试直接用pip装预编译的wheel,别自己编,或者看下是不是flash-attn版本太新,换2.3.x试试。另外你如果只做代码补全,序列长度是不是一定要1024?降到512可能速度直接翻倍,效果未必差很多。最后提一句,可以试试torch.compile,有时候能白嫖15%-20%提速,就是第一次编译会卡一会儿。
没开gradient checkpointing确实会慢,但3秒/step还是偏高了,建议先看一眼是不是数据加载卡IO瓶颈了。
flash-attention编译报错大概率是CUDA版本不匹配,直接上预编译wheel包试试,能省不少事。
3秒一个step确实偏慢了,我拿3090跑类似配置差不多1.5秒左右。你试试开gradient checkpointing,虽然显存够但能省不少计算,反而可能更快。flash-attention编译报错大概率是CUDA版本和torch不匹配,直接去装预编译的wheel包能省事很多。另外检查下是不是被CPU吃住了,比如tokenizer或数据加载成了瓶颈。