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/s,PCIe 5.0 翻倍到 64 GB/s。注意,这个数字和显存带宽(1000 GB/s 级别)差了一个数量级。
这意味着什么?意味着数据搬运本身可能是瓶颈。如果你在训练循环里频繁地把数据从 CPU 搬到 GPU,或者把结果从 GPU 搬回 CPU,GPU 的计算单元就会经常处于“等数据”的状态,利用率上不去。这就是为什么 DataLoader 要用多进程预读取、为什么要用pin_memory=True、为什么要把数据提前放到 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_memory=True配合non_blocking=True:
dataloader = DataLoader(dataset, batch_size=32, pin_memory=True, num_workers=4) for data, target in dataloader: data = data.to(device, non_blocking=True) target = target.to(device, non_blocking=True) # ...pin_memory=True会把数据放在锁页内存(pinned memory)里,这种内存不会被操作系统换出,GPU 可以直接通过 DMA 访问,传输速度更快。non_blocking=True则允许传输和计算重叠,进一步压榨性能。
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(f"Step {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 是 14GB,INT4 是 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 参数的模型,用 LoRA,rank=8,batch size=4,序列长度=512,FP16 混合精度。来估算一下:
- 基础模型参数:FP16 下 2.6GB。LoRA 只训练一小部分参数,假设新增参数 10M,可以忽略。
- 梯度:只对 LoRA 参数算梯度,很小。
- 优化器状态:只对 LoRA 参数,也很小。
- 激活值:这个是大头。1.3B 模型大约 24 层,每层激活值估算下来,batch=4、seq=512 的情况下,可能在 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和FSDP(Fully 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 常见问题速查表
| 问题现象 | 可能原因 | 排查方法 | 解决手段 |
|---|---|---|---|
| 训练时 OOM | batch 太大、激活值太高 | 看 memory_summary | 减小 batch、梯度检查点 |
| 推理时 OOM | 模型太大、精度太高 | 算参数大小 | 量化、CPU 卸载 |
| 显存占用高但利用率低 | 数据传输瓶颈 | 看 GPU 利用率 | pin_memory、异步传输 |
| 多卡显存不均 | batch 分配不均 | 逐卡查看 | 调整分配、检查设备 |
| 内存爆了 | DataLoader worker 太多 | 看内存占用 | 减少 worker、流式读取 |
| 加载模型时 OOM | 权重加载占内存 | 看加载过程 | 分片加载、mmap |
7. 我踩过的坑和几条实在建议
第一个坑是以为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(f'cuda:{local_rank}')可以避免。
几条实在建议:先估算再动手,别上来就训,先算算参数、梯度、优化器状态和激活值大概多少,心里有数;监控要常态化,在训练脚本里加显存打印,每隔几步输出一次,出问题能快速定位;量化优先于卸载,量化实现简单、效果直接,卸载调优复杂、容易引入新问题;batch size 不是越大越好,有时候小 batch 配合梯度累积,效果一样但显存更省。
最后分享一个小技巧:如果你不确定某个操作会不会爆显存,可以先用一个很小的输入跑一遍,看峰值显存是多少,再按比例放大估算。比如用 batch=1、seq=128 跑一遍,记录max_memory_allocated(),然后按 batch 和 seq 的倍数估算实际需求。这个方法虽然不精确,但能帮你快速判断方案是否可行,省去很多试错时间。