最近在跑一个BERT微调任务,batch size设了8,序列长度256,刚开始loss下降挺正常,但跑到第3个epoch时突然CUDA out of memory。诡异的是,同样的代码和数据,前两个epoch显存占用一直是稳定的(大概11G左右),到了第三个epoch就飙到14G+,直接爆了。
PyTorch训练到一半显存爆炸但loss正常,是代码问题还是正常现象?
全部回复
共 2 条这种情况我也踩过坑,大概率不是代码逻辑的问题,而是PyTorch的缓存分配器在搞鬼。前两个epoch显存看着稳定,其实碎片已经攒得差不多了,第三个epoch某个特殊长度的张量一申请,就触发了重新向CUDA申请大块内存,直接爆掉。
你可以试试在epoch开始时加个torch.cuda.empty_cache(),或者用torch.cuda.set_per_process_memory_fraction限制一下上限,看能不能缓解。另外检查下是不是有某个batch的样本长度刚好踩到padding的临界值,导致激活值突然变大。
如果loss一直正常,基本可以排除梯度爆炸或者数据异常,不用太慌。我之前跑GPT微调也遇到过一模一样的现象,最后用梯度累积把batch size降下来,显存反而稳住了。
我之前也踩过类似的坑,不过是在训练Transformer做生成任务的时候。你这种情况大概率不是代码逻辑问题,而是PyTorch的缓存分配器在作祟——前两个epoch显存看着稳定,其实可能已经留了一些碎片化的缓存块,第三个epoch正好碰上某个特殊长度的中间张量,触发了重新分配,显存就一下子上去了。你可以试试在epoch之间手动调一下torch.cuda.empty_cache(),虽然不一定根治,但至少能看出是不是缓存碎片的问题。还有个细节,如果你的DataLoader里做了动态padding或者有样本长度差异很大的情况,第三个epoch可能正好抽到了更长的那批数据,导致激活内存峰值暴涨,这个用torch.profiler看每步的显存占用曲线就能验证。另外,如果loss正常但显存爆,大概率不是梯度问题,因为梯度爆炸一般会伴随loss抖动或NaN。要是实在排查不出来,可以把batch size降到4,或者用gradient checkpointing,虽然慢点但稳定很多。最后想问一下,你是用的Adam还是AdamW?有些优化器的状态缓存会在特定step才膨胀,我之前遇到过类似情况。