ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

PyTorch显存管理实战:从CUDA OOM到模型部署优化

PyTorch显存管理实战:从CUDA OOM到模型部署优化 1. 从一个真实场景说起模型加载时的那声“CUDA out of memory”如果你在算法团队待过大概率见过这样的画面同事兴冲冲地跑过来说“模型训崩了”你凑过去一看终端里赫然一行红字——RuntimeError: CUDA out of memory. Tried to allocate 2.00 GiB。然后就是熟悉的操作把 batch size 从 32 调到 16再调到 8还是崩最后干脆把模型从 GPU 上挪到 CPU 上跑速度慢得像蜗牛爬但至少不报错了。这个场景背后其实藏着一个算法工程师绕不开的基本功搞清楚模型到底放在哪里。是放在 CPU 的内存里还是放在 GPU 的显存里两者之间怎么搬运为什么有时候内存够、显存不够有时候反过来这些问题看起来像是“运维的事”但实际上一旦你开始做模型训练、微调、推理部署它们就会变成每天都要面对的现实问题。我写这篇东西的出发点很简单网上讲 CPU 内存和 GPU 显存的文章要么是硬件科普讲一堆 DDR5、HBM、PCIe 带宽的参数看完还是不知道怎么用要么是框架文档直接甩给你torch.cuda.empty_cache()和model.to(cuda)但不告诉你为什么这么写、什么时候不该这么写。我想做的是把这两端接起来——从算法工程师的实际工作流出发把“模型放在哪里”这件事讲透。这篇文章适合几类人看刚入行、第一次接触 GPU 训练的算法新人做过一些训练但总是被 OOM 卡住、想系统理解显存管理的工程师还有那些需要在有限硬件上部署模型、天天琢磨怎么省显存的老手。我会从存储器的基本分工讲起然后落到 PyTorch 的实际操作再讲显存估算、常见坑和排查方法。全程不堆公式尽量用你能直接上手的方式来说。2. 先把地图画清楚CPU 内存和 GPU 显存到底分工是什么2.1 两种存储器两种性格CPU 内存和 GPU 显存本质上都是“存储器”但它们的性格完全不同。你可以把 CPU 内存想象成一个大仓库容量大现在工作站动辄 64GB、128GB服务器上 512GB 也不稀奇什么都能放但搬运速度相对慢而且离计算单元CPU 核心比较远。GPU 显存则像一个工作台容量小消费级显卡 8GB、12GB、24GB专业卡能到 48GB、80GB但离计算单元CUDA 核心、Tensor Core极近带宽高得离谱。这个“近”和“带宽高”有多重要举个直观的例子。一块 RTX 4090 的显存带宽大约是 1008 GB/s而一套双通道 DDR5-5600 的内存带宽大约是 89.6 GB/s。也就是说显存的数据吞吐能力是内存的十倍以上。GPU 之所以能在矩阵乘法、卷积这类操作上碾压 CPU很大程度上就是因为计算单元不用等数据——数据就在旁边而且来得飞快。但代价也很明显显存贵、容量小、扩展性差。你没法像插内存条那样给显卡加显存买回来是多少就是多少。这就决定了算法工程师的核心矛盾计算要快就得把数据放显存显存不够就得想办法省着用或者往内存里挪。2.2 模型参数、梯度、优化器状态、激活值显存里到底住了谁很多人以为“模型占显存”就是参数大小比如一个 7B 模型FP16 精度下参数占 14GB那 24GB 显存应该够了吧结果一跑就 OOM。原因是显存里住的不只是参数。在训练场景下显存的主要住户有这么几位模型参数Parameters这是模型的权重FP16 下每个参数占 2 字节FP32 下占 4 字节。7B 模型 FP16 就是 14GB。梯度Gradients反向传播算出来的梯度通常和参数同精度同大小又是 14GB。优化器状态Optimizer States如果你用 Adam每个参数要存一阶矩和二阶矩FP32 下就是 8 字节/参数7B 模型就是 56GB。这就是为什么全量微调大模型这么吃显存。激活值Activations前向传播过程中每一层的输出需要保留到反向传播用。这部分和 batch size、序列长度强相关往往是大头。临时缓冲区CUDA 算子执行时的中间结果、通信缓冲区等。所以一个 7B 模型全量微调显存需求轻松超过 100GB单卡根本放不下。这也是为什么现在流行 LoRA、QLoRA 这些参数高效微调方法——它们把大部分参数冻住只训练一小部分优化器状态和梯度都大幅减少。推理场景就简单多了只需要参数和激活值不需要梯度和优化器状态。所以 7B 模型 FP16 推理14GB 参数加上一些激活值24GB 显存基本够用。如果再量化到 INT8 或 INT4显存需求还能砍半甚至更多。2.3 数据是怎么在内存和显存之间流动的理解了“谁住在哪”接下来要理解“怎么搬家”。CPU 和 GPU 之间的数据传输走的是 PCIe 总线或者 NVLink如果你有的话。PCIe 4.0 x16 的带宽大约是 32 GB/sPCIe 5.0 翻倍到 64 GB/s。注意这个数字和显存带宽1000 GB/s 级别差了一个数量级。这意味着什么意味着数据搬运本身可能是瓶颈。如果你在训练循环里频繁地把数据从 CPU 搬到 GPU或者把结果从 GPU 搬回 CPUGPU 的计算单元就会经常处于“等数据”的状态利用率上不去。这就是为什么 DataLoader 要用多进程预读取、为什么要用pin_memoryTrue、为什么要把数据提前放到 GPU 上。PyTorch 里的.to(cuda)或.cuda()就是触发这个搬运的操作。它会把张量从内存复制到显存。反过来.cpu()会把张量从显存复制回内存。每次调用都有开销所以能批量搬就不要零散搬能提前搬就不要在循环里搬。3. PyTorch 里的“放哪里”从 to(device) 到显存分配器3.1 device 对象和模型迁移的基本操作在 PyTorch 里“模型放在哪里”是通过device来控制的。最基础的写法是这样import torch device torch.device(cuda if torch.cuda.is_available() else cpu) model MyModel().to(device) data data.to(device) output model(data)这几行代码看起来简单但每一行背后都有讲究。torch.cuda.is_available()检查的是当前环境有没有可用的 CUDA 设备包括驱动、运行时和显卡是否匹配。如果返回 False可能是驱动没装好、CUDA 版本不匹配或者显卡被其他进程占满了。model.to(device)会把模型的所有参数和缓冲区搬到指定设备。注意这个操作是原地修改的也就是说model本身的参数会被替换成 GPU 上的张量。如果你有多个模型或者多个设备要小心别把不该搬的搬了。data.to(device)同理把输入数据搬到 GPU。这里有个常见的性能陷阱如果你在训练循环里对每个 batch 都调用.to(device)而且没有用pin_memory和异步传输那么数据传输会和计算串行GPU 利用率会很低。正确的做法是用DataLoader的pin_memoryTrue配合non_blockingTruedataloader DataLoader(dataset, batch_size32, pin_memoryTrue, num_workers4) for data, target in dataloader: data data.to(device, non_blockingTrue) target target.to(device, non_blockingTrue) # ...pin_memoryTrue会把数据放在锁页内存pinned memory里这种内存不会被操作系统换出GPU 可以直接通过 DMA 访问传输速度更快。non_blockingTrue则允许传输和计算重叠进一步压榨性能。3.2 显存分配器PyTorch 不是每次都找系统要显存很多人有个误解以为 PyTorch 每次创建张量都会向系统申请显存。实际上PyTorch 有一个缓存分配器Caching Allocator。它会在第一次需要显存时向系统申请一大块然后自己管理这块显存后续的张量创建和释放都在这个池子里进行。这个设计的好处是避免频繁的系统调用因为每次向 CUDA 申请显存都有开销。坏处是当你看到nvidia-smi显示显存占用很高时可能其中一部分是被缓存占着、但实际上没在用的。这就是为什么有时候你删了模型、调了torch.cuda.empty_cache()显存占用才降下来。torch.cuda.empty_cache()的作用是释放缓存分配器里那些没有被引用的显存块还给系统。注意它不会释放还在被张量引用的显存。所以如果你有一个大张量还活着调多少次empty_cache()都没用。这里有个实操心得不要在训练循环里频繁调用empty_cache()。因为它会清空缓存导致下一次分配又要向系统申请反而拖慢速度。正确的时机是在你确实不再需要某些大张量之后比如一个 epoch 结束、或者切换模型之前。3.3 查看显存占用的几种方式想知道模型到底占了多少显存有几种方式nvidia-smi最直接能看到每块卡的总体占用但看不到具体是哪个张量占的。torch.cuda.memory_allocated()返回当前被张量占用的显存字节数。torch.cuda.memory_reserved()返回缓存分配器保留的显存字节数通常大于memory_allocated()。torch.cuda.max_memory_allocated()返回峰值占用排查 OOM 时很有用。torch.cuda.memory_summary()打印详细的显存分配报告包括各个分配块的大小和状态。我一般在训练脚本里加这么一段方便随时监控def print_memory(step): allocated torch.cuda.memory_allocated() / 1024**3 reserved torch.cuda.memory_reserved() / 1024**3 print(fStep {step}: allocated{allocated:.2f}GB, reserved{reserved:.2f}GB)如果发现reserved远大于allocated说明缓存里有不少空闲块可以考虑在合适的时候empty_cache()。如果allocated本身就很高那就是真的有那么多张量活着得从模型结构或 batch size 上想办法。4. 显存估算动手算一遍比拍脑袋靠谱4.1 参数、梯度、优化器状态的显存公式前面说了显存里住着谁现在来算具体数字。假设模型有 P 个参数训练时用混合精度AMP优化器用 Adam模型参数FP16 下 2P 字节FP32 下 4P 字节。混合精度通常保留一份 FP32 主权重和一份 FP16 计算权重所以是 6P 字节。梯度FP16 下 2P 字节。Adam 优化器状态FP32 的一阶矩和二阶矩共 8P 字节。激活值这个最难估和网络结构、batch size、序列长度都有关通常需要实测。所以一个 P 参数的模型混合精度 Adam 全量微调光是参数、梯度和优化器状态就是 16P 字节。7B 模型就是 112GB单张 80GB 的 A100 都放不下。这就是为什么全量微调大模型需要多卡并行或者用 ZeRO 这类优化技术。推理就简单了FP16 下 2P 字节INT8 下 1P 字节INT4 下 0.5P 字节。7B 模型 FP16 是 14GBINT4 是 3.5GB。所以如果你只是推理量化能省很多显存。4.2 激活值那个容易被忽略的大头激活值的显存占用经常被低估。以 Transformer 为例每一层的激活值包括注意力矩阵、前馈网络的中间结果等。注意力矩阵的大小是batch_size × num_heads × seq_len × seq_len序列长度翻倍这部分显存翻四倍。这就是为什么长序列训练特别吃显存。如果你要训 8K 甚至 32K 序列激活值可能比参数还大。解决办法包括梯度检查点Gradient Checkpointing用计算换显存、Flash Attention优化注意力计算减少中间矩阵、序列并行等。梯度检查点的原理是前向传播时不保存所有激活值只保存部分检查点反向传播时重新计算需要的激活值。这样显存占用从 O(n) 降到 O(sqrt(n))代价是多了大约 30% 的计算量。PyTorch 里可以用torch.utils.checkpoint.checkpoint来包装需要检查的层。4.3 一个具体的估算例子假设你要微调一个 1.3B 参数的模型用 LoRArank8batch size4序列长度512FP16 混合精度。来估算一下基础模型参数FP16 下 2.6GB。LoRA 只训练一小部分参数假设新增参数 10M可以忽略。梯度只对 LoRA 参数算梯度很小。优化器状态只对 LoRA 参数也很小。激活值这个是大头。1.3B 模型大约 24 层每层激活值估算下来batch4、seq512 的情况下可能在 2-4GB 左右。临时缓冲区1-2GB。总计大约 6-9GB一张 12GB 的 RTX 3060 就能跑。这也解释了为什么热词里有人问“minimaxh3 用 rtx3060 的 12g 显存能跑吗”——如果是 LoRA 微调或者量化推理12GB 是有希望的如果是全量微调那肯定不够。5. 省显存的实战手段从量化到卸载5.1 量化用精度换显存量化是最直接的省显存手段。FP32 转 FP16 省一半转 INT8 再省一半转 INT4 再省一半。7B 模型从 FP16 的 14GB 降到 INT4 的 3.5GB一张 6GB 显存的卡都能跑。但量化不是免费的。INT8 和 INT4 会带来精度损失尤其是对数值敏感的层。实践中常用的方法包括训练后量化PTQ模型训练完后直接量化简单但精度损失可能较大。量化感知训练QAT训练时就模拟量化误差精度更好但需要重新训练。GPTQ、AWQ、GGUF这些是针对大语言模型的量化格式各有优劣。GPTQ 适合 GPU 推理GGUF 适合 CPU 推理。热词里提到的“6g 显存”“低显存运行模型”基本都要靠量化来实现。我实测下来7B 模型 INT4 量化后6GB 显存跑推理是可行的但速度会受限于显存带宽和计算能力。5.2 梯度检查点和激活值重计算前面提过梯度检查点这里补充实操细节。在 PyTorch 里你可以这样用from torch.utils.checkpoint import checkpoint class MyModel(nn.Module): def forward(self, x): x checkpoint(self.layer1, x) x checkpoint(self.layer2, x) return x注意checkpoint要求被包装的函数是纯函数不能有副作用。另外它和torch.no_grad()不兼容因为反向传播时需要重新计算。梯度检查点的显存节省效果很明显但会增加计算时间。我的经验是如果显存是瓶颈、计算资源相对充裕那就值得用如果计算本身就是瓶颈那要权衡一下。5.3 CPU 卸载把暂时不用的挪回内存CPU 卸载Offloading的思路是把暂时不用的参数或优化器状态放到内存里需要时再搬到显存。DeepSpeed 的 ZeRO-Offload 和 PyTorch 的CPUOffload都支持这个。代价是数据传输开销。前面说过PCIe 带宽比显存带宽低一个数量级所以频繁卸载会导致 GPU 等数据。适合的场景是显存极度紧张、但内存充裕而且模型的计算密度不是特别高。实操中我一般会先尝试量化和梯度检查点如果还不够再考虑卸载。因为卸载的调优比较复杂容易引入新的性能问题。5.4 模型并行和分布式训练如果单卡怎么都放不下那就只能多卡了。模型并行把模型的不同层放到不同卡上数据并行把不同 batch 放到不同卡上。PyTorch 的DistributedDataParallel和FSDPFully Sharded Data Parallel是常用方案。FSDP 的思路是把参数、梯度、优化器状态都分片到多张卡上每张卡只存一部分需要时通过通信拼起来。这样显存占用随卡数线性下降。代价是通信开销所以对网络带宽要求高。热词里提到的“gpu 租用”“gpu 配额”反映的就是多卡训练的现实需求。如果自己买不起多卡租云上的 GPU 集群是常见选择。6. 常见问题与排查技巧实录6.1 OOM 了怎么办一套排查流程遇到CUDA out of memory别急着调 batch size先按这个流程走一遍确认是不是真的显存不够用nvidia-smi看当前占用用torch.cuda.memory_summary()看分配详情。有时候是其他进程占着显存没释放。定位是哪一步 OOM是在模型加载时、前向传播时、反向传播时还是优化器更新时不同阶段 OOM 的原因不同。检查 batch size 和序列长度这两个是最直接的影响因素。先减半试试。检查是否有张量泄漏比如在循环里不断创建新张量而不释放或者把计算图保留了下来比如没加torch.no_grad()。考虑量化和梯度检查点如果模型本身太大就得从这些手段入手。最后才考虑换卡或分布式这是成本最高的方案。6.2 内存够但显存不够或者反过来有时候会遇到“内存还剩很多但显存爆了”或者“显存够用但内存爆了”。前者通常是因为模型和数据都在 GPU 上CPU 那边没什么负担解决办法就是量化、卸载、或者减小 batch。后者通常是因为 DataLoader 的num_workers太多、或者数据预处理太占内存解决办法是减少 worker 数量、用更高效的数据格式、或者流式读取。还有一种情况是“显存够但内存不够”这通常发生在模型加载阶段你需要先把权重加载到内存再搬到显存。如果内存不够加载就会失败。解决办法是用mmap方式加载、或者分片加载。6.3 多卡训练时的显存不均衡多卡训练时有时候会发现某张卡的显存占用明显高于其他卡。原因可能是数据并行时每张卡的 batch 大小不一样。模型并行时不同层的参数量不一样。通信缓冲区分配不均。解决办法包括调整 batch 分配、用torch.cuda.set_device()明确指定设备、检查是否有张量默认创建在了错误的设备上。6.4 常见问题速查表问题现象可能原因排查方法解决手段训练时 OOMbatch 太大、激活值太高看 memory_summary减小 batch、梯度检查点推理时 OOM模型太大、精度太高算参数大小量化、CPU 卸载显存占用高但利用率低数据传输瓶颈看 GPU 利用率pin_memory、异步传输多卡显存不均batch 分配不均逐卡查看调整分配、检查设备内存爆了DataLoader worker 太多看内存占用减少 worker、流式读取加载模型时 OOM权重加载占内存看加载过程分片加载、mmap7. 我踩过的坑和几条实在建议第一个坑是以为empty_cache()是万能药。刚入行时一遇到 OOM 就调empty_cache()结果发现没什么用。后来才明白它只释放缓存不释放活着的张量。真正有用的是找到那些不该活着的张量比如忘记detach()的计算图、或者循环里累积的列表。第二个坑是忽略激活值的显存占用。有次微调一个模型参数才 2GB但 batch size 开到 16 就 OOM。后来用memory_summary()一看激活值占了 8GB。把 batch 降到 4再开梯度检查点就稳了。第三个坑是在多卡环境里没指定 device。PyTorch 默认用cuda:0如果你不显式指定所有张量都会往第一张卡上挤其他卡闲着。用torch.cuda.set_device(local_rank)和device torch.device(fcuda:{local_rank})可以避免。几条实在建议先估算再动手别上来就训先算算参数、梯度、优化器状态和激活值大概多少心里有数监控要常态化在训练脚本里加显存打印每隔几步输出一次出问题能快速定位量化优先于卸载量化实现简单、效果直接卸载调优复杂、容易引入新问题batch size 不是越大越好有时候小 batch 配合梯度累积效果一样但显存更省。最后分享一个小技巧如果你不确定某个操作会不会爆显存可以先用一个很小的输入跑一遍看峰值显存是多少再按比例放大估算。比如用 batch1、seq128 跑一遍记录max_memory_allocated()然后按 batch 和 seq 的倍数估算实际需求。这个方法虽然不精确但能帮你快速判断方案是否可行省去很多试错时间。
返回列表