拓冰建站拓冰建站
首页 / 资讯中心 / 正文

PyTorch显存不足怎么办?从定位到优化,一套完整排查方案

用 PyTorch 训练模型最让人上头的时刻不是 loss 曲线突然上扬而是刚好跑完一个 epoch正准备保存 checkpoints屏幕“啪”地弹出一行大红字CUDA out of memory。很多同学遇到这个问题的第一反应是调小 batch size再不行就换显卡。结果 batch size 调到 1 还是爆换卡又没钱最后卡在内存不足这个坑里好几天。说实话PyTorch 里的内存不足尤其是显存不足绝大多数情况不是“显存真的不够”而是我们没搞清楚内存被谁吃掉了、为什么 PyTorch 的显存占用看起来只增不减。这篇内容我会把我在实际项目中排查和解决 PyTorch 内存不足的整套思路完整写下来包括怎么定位是显存还是系统内存的问题、PyTorch 显存管理机制是怎么回事、梯度累积和混合精度这些降低占用的实操方法以及我踩过的一系列坑。不管你是刚配好 PyTorch 环境准备跑第一个模型还是已经在调大规模 Transformer这套排查逻辑应该都能直接用上。1. 定位内存瓶颈先搞清楚是显存不够还是内存不够很多人一看到“内存不足”四个字就开始无脑调 batch size其实这是最没有效率的做法。PyTorch 运行过程中涉及两种完全不同的内存资源GPU 显存VRAM和系统内存RAM。这两种资源不足的时候报错、表现、排查手段都不一样第一步必须先分清楚到底是谁不够。1.1 两种“内存不足”的报错长什么样GPU 显存不足时最常见的报错是RuntimeError: CUDA out of memory. Tried to allocate 512.00 MiB (GPU 0; 8.00 GiB total capacity; 7.20 GiB already allocated; 62.36 MiB free; 6.89 GiB reserved in total by PyTorch)这个报错信息量非常大后面我会拆开讲。系统内存不足时通常不会直接给你一个 PyTorch 异常而是整个进程被操作系统杀掉常见表现有Linux 下出现Killed提示命令直接退出。训练到一半整个电脑卡死鼠标动不了。Windows 下弹窗提示“内存不足”或者进程直接闪退。报RuntimeError: DataLoader worker (pid xxx) is killed by signal: Killed。这两种情况的处理思路完全不同。显存不够主要靠模型侧优化比如混合精度、梯度检查点、减小 batch size内存不够则要重点查 DataLoader 的 worker 配置、数据集加载方式、是否有内存泄漏甚至可能需要用memmap方式加载大文件数据。注意如果是服务器上的多卡训练还要考虑其他用户/进程是否占用了显存。比如nvidia-smi显示显存剩余不少但 PyTorch 就是分配不到这种大概率不是显存不够而是驱动或权限层面的问题。1.2 先用这几条命令判断占用情况定位的第一步是看“还剩多少内存”。打开终端执行nvidia-smi输出里的Memory-Usage列会显示每个 GPU 的显存总容量、已用容量和剩余容量。如果你发现显存明明还剩 4GBPyTorch 却报 out of memory那问题基本出在 PyTorch 的缓存分配器上而不是物理显存不够。nvidia-smi显示的是整张卡的占用情况但它看不到 PyTorch 内部是怎么分配显存的。想查 PyTorch 进程自己的显存使用可以在代码里加一段import torch print(allocated: %.2f GB % (torch.cuda.memory_allocated() / 1024**3)) print(reserved: %.2f GB % (torch.cuda.memory_reserved() / 1024**3))allocated是 PyTorch 实际使用的显存量reserved是 PyTorch 从显卡上预留下来的量。大多数时候reserved会明显大于allocated因为 PyTorch 为了加速会多囤一些显存这个属于正常行为。如果两者差距特别大比如allocated只有 2GBreserved却到了 7GB那就需要检查代码里是否创建了大量临时张量或者显存碎片化已经很严重了。更完整的查看方式是用 PyTorch 自带的显存诊断接口torch.cuda.memory_summary(deviceNone, abbreviatedFalse)它会把每一步的显存占用、缓存块数量、碎片化程度全部列出来是排查显存问题的第一利器。我一般在代码里加一个命令行参数比如--debug_memory开启后在第 1 个 batch 跑完就输出memory_summary()方便快速判断模型和数据默认吃掉了多少显存。1.3 CPU 系统内存的排查方法系统内存不足没那么好定位但思路也很直接看进程的内存占用曲线是否持续上涨。最简单的方法是用topLinux/Mac或任务管理器Windows观察。如果内存占用随着训练步数上涨并且不回落很可能存在内存泄漏如果在固定某个位置突然飙升则大概率是 DataLoader 或者某个预处理函数把数据一次性加载进来了。想更精确地排查 Python 代码里的内存占用可以用标准库的tracemallocimport tracemalloc tracemalloc.start() # 你的训练代码 ... current, peak tracemalloc.get_traced_memory() print(f当前内存: {current / 1024**2:.1f} MB) print(f峰值内存: {peak / 1024**2:.1f} MB)peak会告诉你程序运行过程中的最高内存占用。如果峰值非常高就去搜代码里哪一步瞬间构建了超大对象比如一次性把整个数据集读成 list、把图片批量转成 numpy 数组但没有释放等。我自己的习惯是如果怀疑系统内存问题先把num_workers调成 0DataLoader 不新开子进程跑一遍。如果内存问题消失基本可以确定是 DataLoader 多进程导致的。如果num_workers0时依然内存暴涨再往数据集加载或模型本身查。2. 理解 PyTorch 显存管理机制为什么“释放”了不够用很多人在 PyTorch 里写了del tensor甚至调用了torch.cuda.empty_cache()然后发现nvidia-smi里显存占用还是老样子就以为程序泄露了。其实这背后是 PyTorch 显存缓存分配器的工作机制。2.1 requested 和 reserved 是两个不同的概念用银行取钱来打比方torch.cuda.empty_cache()相当于你把钱包里的零钱整理了一下但银行卡里的余额nvidia-smi看到的显存占用不会因为你整理钱包就变少。PyTorch 默认的分配策略是当你在 GPU 上创建一个 Tensor它一次性向驱动申请一大块显存比如 1GB这块显存成了 PyTorch 的“储备池”。之后你再创建小 TensorPyTorch 直接从储备池里划一块出去不需要每次都向驱动“要钱”。这样做的好处是快坏处就是你看到的显存占用会被放大。举个例子import torch # 只创建一个很小的 tensor x torch.zeros(1024, devicecuda) print(torch.cuda.memory_allocated()) # 实际占用可能是 4KB 左右 print(torch.cuda.memory_reserved()) # 预留显存可能是 2GB 甚至更多第一次运行torch.cuda.memory_reserved()很可能返回一个比较大的值比如 1GB 或 2GB这个不是 bug而是 PyTorch 启动时会预先囤显存。2.2 为什么 nvidia-smi 显示高占用但 allocated 不高如果你发现这种情况最常见的原因是PyTorch 的缓存分配器预留了大量显存但里面大部分是空闲块。其他进程比如另一个 Python 脚本、浏览器、ComfyUI 等也占用了显存。显存碎片化严重空闲块分散在小角落无法满足一个大块的分配请求。显存碎片化是训练过程中特别容易遇到的隐性杀手。模型前向传播、反向传播会产生大量不同尺寸的临时张量。假设显存里已经存在很多不连续的空闲块每块只有 100MB这时候你要分配一个 200MB 的矩阵尽管总的空闲显存可能超过 1GB但分配器找不到连续的 200MB 空闲块于是直接报 out of memory。2.3 empty_cache 的正确使用误区torch.cuda.empty_cache()确实会清掉 PyTorch 缓存池里“未使用”的显存块把它还给驱动。但我实际看很多人滥用它在训练循环的每一步都调用结果显存是被释放了但运行速度慢了好几倍。因为清空缓存池之后下一步训练又得重新向驱动申请显存重复申请-释放这个昂贵操作白白浪费大量时间。我的建议是训练循环内不要用empty_cache()。在验证阶段开始前或者保存完 checkpoint、加载新模型时如果显存紧张可以调用一次。多个模型交替推理时用完的模型显存需要释放可以del model后调用empty_cache()。3. 降低训练显存占用的可落地方法一从模型和数据维度下手定位完问题如果发现确实是自己模型的显存峰值太高接下来就是实打实地降低占用。这一部分我会按“从简单到复杂”的顺序列出亲测有效的方法并给出关键实现思路。3.1 梯度累积用时间换空间梯度累积是最简单、最不容易踩坑的降显存手段。做法是不在每个 batch 后更新参数而是累积多个 batch 的梯度后再统一更新。accumulation_steps 4 optimizer.zero_grad() for i, (inputs, labels) in enumerate(train_loader): outputs model(inputs) loss criterion(outputs, labels) # 这里除以 accumulation_steps相当于取平均 loss loss / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()从显存角度看梯度累积把“一次计算 N 个 batch 的梯度”变成“分 N 次计算”每次只把一个 batch 的中间结果放在显存里所以峰值占用明显降低。需要注意的是loss.item()看的是单个 batch 的平均 loss如果你用累积方式监控日志里可以同时保留一个 batch 的 loss 和整体的平滑 loss避免训练曲线抖动太大不好判断收敛。3.2 梯度检查点拿计算换显存如果梯度累积已经不够用可以上激活重计算也就是 PyTorch 里的torch.utils.checkpoint。它的核心思想是前向传播时只保留少数关键激活值不保存中间层的所有激活反向传播需要某个中间张量时再重新做一次前向去计算出来。from torch.utils.checkpoint import checkpoint def forward_with_checkpoint(module, x): return checkpoint(module, x, use_reentrantFalse)在实际的 Transformer 代码里通常在残差分支前插入 checkpointclass TransformerLayer(nn.Module): def forward(self, x): # nn.MultiHeadAttention 这一类重模块用 checkpoint 包裹 attn_out checkpoint(self_attn, x, x, x, use_reentrantFalse) x x attn_out x x checkpoint(ffn, x, use_reentrantFalse) return x显存能省多少如果模型有 12 层 Transformer每层都做 checkpoint显存开销大约从“所有层激活值总和”降到“单层激活值加上少量保存的输入”。代价是训练时间可能增加 20% 到 40%因为反向传播时要重复计算前向。经验是如果你的模型深度很大、显存差一点点就够用时优先对最重的注意力层做 checkpoint而不是全模型无脑包一层。3.3 混合精度训练现代显卡必须开启混合精度AMP已经是 PyTorch 训练的基本操作了。它的原理很简单FP32 的权重和梯度保持不动计算过程中的中间激活值用 FP16 保存显存占用直接减半。同时借助 Tensor Core计算速度还能提升。我现在习惯的写法from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for inputs, labels in train_loader: inputs, labels inputs.cuda(), labels.cuda() optimizer.zero_grad() with autocast(): outputs model(inputs) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()需要注意的坑是某些操作对精度特别敏感比如 log softmax、某些归一化autocast 会自动跳过不需要手动干预。如果梯度出现 NaN/Inf可以先检查scaler是否正常更新不要急着改模型结构。FP16 对某些老显卡如 GTX 10 系效果有限性能提升不明显但显存降低是实打实的。3.4 DataLoader 参数调优内存与速度的平衡除了模型本身数据加载也是吃内存的大户。很多人的显存明明还够但内存先爆了就是 DataLoader 配置不当。num_workers不是越大越好。每个 worker 会复制一部分数据到自己的内存空间比如num_workers8数据集预处理时每个 worker 额外占用 1GB 内存8 个 worker 就多占 8GB。如果prefetch_factor再设高内存占用会更夸张。我的建议配置是物理内存 16GB 的机器num_workers设 2 到 4。物理内存 32GB 以上的机器num_workers可以设 8但prefetch_factor默认 2 就好。数据集里的图片或文本如果预处理后体积大尽量在Dataset.__getitem__里做轻量处理不要把所有预处理结果缓存到内存。pin_memoryTrue会把数据放到锁页内存pinned memory传输到 GPU 的速度更快但它同样会占用系统内存。如果你的系统内存本身紧张可以关掉。3.5 输入数据布局与内存复用有时候显存并不是被模型吃掉的而是被输入数据吃掉的。如果你在训练循环里写了类似inputs inputs.cuda(non_blockingTrue)每行都会重新分配显存。如果transform或预处理在 GPU 上进行频繁创建新的中间 Tensor也会快速塞满缓存池。更好的做法是尽量不直接给模型喂原尺寸大图先做 resize。在 batch 的__getitem__里确认返回的 Tensor 是连续的tensor.contiguous()避免后续操作产生大量非连续内存拷贝。复用一个固定大小的输入 Tensor靠切片赋值更新内容而不是每次new_tensor。4. 常见内存报错场景与实操排查清单下面是我这几年跑 PyTorch 项目中真正遇到的几个高频内存报错场景每条都给了定位思路和解决办法可以直接对照查。4.1 初始 batch 就爆显存报错发生在训练刚开始、第一个 batch 前向传播的时候。原因通常是batch size 设置过大即便是初始测试也会直接超出显存容量。模型输入尺寸过大尤其是(N, C, H, W)的图片H 和 W 大一个维度显存消耗是平方级增长。加载了预训练模型但忘记切换到eval()或者no_grad()比如在推理阶段还挂着梯度计算。解决思路先打印每层输入输出的形状用torch.cuda.memory_summary()看是哪一步分配的显存最多。如果 batch size 降到 1 仍然爆检查模型本身是否填入了过大的输入尺寸比如max_length4096的 Transformer 在单卡 8GB 上很容易爆。4.2 训练中途突然报 fragmented memory报错信息里通常会出现RuntimeError: CUDA error: out of memory ... Tried to allocate 2.00 GiB (GPU 0; 8.00 GiB total capacity; 6.99 GiB already allocated; 0 bytes free; 7.13 GiB reserved in total by PyTorch)这里的0 bytes free是整卡真正剩余的空闲空间很少但6.99 GiB already allocated是 PyTorch 实际用掉的。如果reserved接近上限但free为 0说明显存碎片化严重虽然 PyTorch 缓存池里有很多小块空闲但没法拼出一块连续的 2GB。处理方法将PYTORCH_CUDA_ALLOC_CONF设置为max_split_size_mb:128减少大块分配请求被碎片化影响。检查代码里有没有在循环里创建大量不同 shape 的临时张量比如动态 mask、动态长度 padding这类情况尽量统一 shape。用torch.cuda.empty_cache()在验证阶段前释放缓存。4.3 Dataloader 内存翻倍系统卡死表现是训练开始后系统内存占用一路飙升最后电脑卡死或进程被 kill。排查思路把num_workers降为 0 试一下如果恢复正常说明是 worker 数量太高。检查__getitem__里有没有返回共享的大对象。某些实现会在__getitem__里读取整张图片的原始数据然后用PIL处理每个 worker 都会复制一份。用tracemalloc在训练循环里打印内存峰值找瞬间暴涨的代码位置。有一种常见情况是collate_fn里做了torch.stack把所有样本堆叠成一个超大 Tensor。如果每个样本是(3, 512, 512)的 float32 图片128 个样本一次 stack 就是(128, 3, 512, 512) * 4B ≈ 384MB多几个 worker 同时处理内存直接爆炸。这时可以考虑降低 batch size 或者用更小的输入分辨率。4.4 模型加载或推理时显存不足训练没问题但一加载大模型做推理就报 out of memory。原因通常是推理时没有关闭梯度模型中仍然保存了中间激活值。用model.to(cuda)后忘记model.eval()BN/Dropout 还在训练模式。多卡环境中每张卡都加载了一份完整权重。正确做法model.eval() with torch.no_grad(): output model(input)如果模型特别大比如几十亿参数单卡推理确实放不下这个时候只能做模型并行、量化或使用更节省显存精度的加载方式比如torch_dtypetorch.float16的 HuggingFace 模型加载。4.5 多次运行后显存不释放通常发生在 Jupyter Notebook 或交互式环境里。你跑了一个训练脚本结束时报错退出了但nvidia-smi里显存还是占着。原因可能是Python 进程没有退出显存还被进程持有。某些后台线程比如 DataLoader worker没有正常关闭。解决方式用kill -9 pid强制结束残留进程。在代码退出前显式del model、torch.cuda.empty_cache()。长期运行的脚本建议每隔一定步数用nvidia-smi确认显存占用是否随训练推进不断增长如果一直涨优先怀疑缓存分配器碎片化或临时张量没释放而不是简单地认为是卡有问题。4.6 排查速查表现象大概率原因优先检查项第一个 batch 就爆batch size 或输入尺寸过大torch.cuda.memory_summary()、打印 Tensor shape训练中途突然爆动态 shape/碎片化nvidia-smifree 显存、max_split_size_mb配置系统内存一路涨DataLoader worker 过多或数据集缓存过大num_workers0对比测试、tracemalloc推理时爆显存没开no_grad()确认model.eval()和torch.no_grad()进程结束后显存不释放进程未退出/缓存未清kill -9、empty_cache()多卡训练总显存不足每卡加载完整模型检查model nn.DataParallel(model)或DDP的配置5. 环境配置与长期维护建议排查和优化的技巧说完了再聊一点容易被忽略但影响极大的事情环境配置。很多时候内存不足问题根本不是代码问题而是 PyTorch 环境本身没配好或者版本行为差异导致的。我在实际项目里见过不少人把CUDA out of memory归咎于显卡不够结果换了个版本问题就消失了。5.1 PyTorch 版本和 CUDA 版本的影响PyTorch 在不同版本里的显存管理策略有所变化。比如较新版本改进了缓存分配器对碎片化的处理更好某些版本在autocast下对 FP16 的分配更积极。如果你跑的是老项目从 PyTorch 1.13 升到 2.x 时很可能会发现显存占用有变化这时候不要慌先重新测量基线。CUDA 版本也有影响。建议安装时直接参考 PyTorch 官网给出的匹配版本不要自己用最新版 CUDA。实测下来CUDA 11.8 和 CUDA 12.1 在稳定性和显存管理上都有明显差异某些显卡在老 CUDA 驱动下反而更好用。5.2 用 Anaconda 创建独立环境很多内存问题和环境混乱有关。比如系统里同时装了 CPU 版和 GPU 版的 PyTorch代码 import 到了 CPU 版模型只能跑 CPU数据却拼命运到内存里最后内存被吃光。我的习惯是用 Anaconda 建独立环境conda create -n torch_gpu python3.10 conda activate torch_gpu pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118创建环境之前先确认显卡支持的 CUDA 版本用nvidia-smi看右上角的CUDA Version。它只要大于等于 PyTorch 所需的 CUDA 版本就行不需要完全一致。安装完成后在 Python 里确认import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.version.cuda)如果torch.cuda.is_available()返回 False大概率是驱动太老或者装成了 CPU 版这会直接影响显存能否被使用进而导致“内存不足”的错觉。5.3 调整缓存分配策略PYTORCH_CUDA_ALLOC_CONFPyTorch 从很早就支持通过环境变量控制显存分配行为。几个常用的配置export PYTORCH_CUDA_ALLOC_CONFmax_split_size_mb:128max_split_size_mb指定缓存块拆分的最大单位。设小一点可以减少碎片化但可能略微增加分配开销。如果遇到 fragmented memory可以尝试从 128 调整到 64 或 256 观察效果。还有一个配置是garbage_collection_threshold比如export PYTORCH_CUDA_ALLOC_CONFgarbage_collection_threshold:0.8,max_split_size_mb:128它控制当缓存池使用率达到 80% 时触发内部垃圾回收。对长时间训练的脚本有一定帮助但也要实测调太激进会拖慢训练速度。在 Windows 上环境变量可以通过系统设置里添加也可以在代码开头设置import os os.environ[PYTORCH_CUDA_ALLOC_CONF] max_split_size_mb:128 import torch注意os.environ设置必须在import torch之前否则不会生效。5.4 真不够用时怎么兜底如果代码优化做完了、环境也调了但模型确实超过显存容量那就只能上更重的方案模型并行把模型的不同层放到不同 GPU 上比如model.layers[0:2].to(cuda:0)但这个需要手动管理数据流动。DeepSpeed ZeRO主要针对大规模分布式训练可以显存节省到非常夸张的程度但对单卡小规模任务收益有限配置成本也高。torch.compilePyTorch 2.x 自带的编译优化在部分模型上能降低显存峰值和提升速度但首次编译时间长并且和某些动态图代码不兼容。FlashAttention如果你的模型里有大量注意力计算换用 FlashAttention 能大幅减少显存占用并且带来一定的速度提升。这在开源模型的推理和训练里都是实践验证过的方向。这些并不是首选项而是当你把所有常规手段都用完之后再考虑的“硬核方案”。我个人的观点是先把自己代码里的临时张量减少到极致再考虑上外部框架否则就算换了 DeepSpeed混乱的代码依然会出现莫名其妙的内存问题。最后再分享一个小技巧。我每次开始一个新的 PyTorch 项目都会先在代码里加一个--profile_memory参数训练的前 3 个 batch 结束后自动打印一次torch.cuda.memory_summary()并把峰值显存写入日志。这样无论后面怎么改模型、换数据集都有基线数据可以对照。排查内存不足问题最怕的就是没有基线、到处乱试。先把复盘数据做起来大多数问题一眼就能看到答案。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门