pytorch 2.x中torch.compile默认不启用显存分页,需手动配置torch.cuda.set_per_process_memory_fraction(0.7–0.85)并禁用cudnn,配合mode="max-autotune"和fsdp的use_orig_params=true才能触发inductor按页粒度(如2mb)动态分配与复用显存。

PyTorch 2.x 的 torch.compile 默认不启用显存分页,需手动配置 torch.cuda.set_per_process_memory_fraction
显存分页(Paged Memory)在 PyTorch 2.x 中并非独立开关,而是通过 torch._C._cuda_setMemoryFraction 和底层 CUDA 上下文管理协同实现的。它实际依赖于 torch.compile 启用的 AOTInductor 后端 + CUDA Graph 捕获时的内存预留策略。直接调用 torch.cuda.empty_cache() 或设置 max_split_size_mb 不触发分页——这些只是常规缓存清理或分配限制。
真正起作用的是:在启动训练前,用 torch.cuda.set_per_process_memory_fraction(0.8) 限定进程可见显存上限,再配合 torch.compile(..., mode="max-autotune"),让 Inductor 在图捕获阶段主动按分页粒度(通常 2MB 对齐)申请和复用显存块。
- 不设
set_per_process_memory_fraction→ Inductor 认为显存充足,倾向于大块连续分配,无法触发分页调度逻辑 - 设为
1.0→ 等同于未设,效果消失 - 建议值:0.7–0.85,留出约 10% 显存给 CUDA Graph 元数据和梯度同步缓冲区
使用 torch.compile 时必须关闭 torch.backends.cudnn.enabled = False
CuDNN 的默认卷积/归一化算子会绕过 AOTInductor 的内存调度器,导致分页机制失效。即使启用了 torch.compile,只要 CuDNN 参与计算,Inductor 就无法接管显存生命周期管理。
实测发现:开启 CuDNN(默认 True)时,nvtop 显示显存占用呈锯齿状上升且不回落;关闭后,同一模型在相同 batch size 下显存峰值下降 23%~31%,且波动平缓——这是分页复用生效的典型特征。
- 在
import torch后立即执行:torch.backends.cudnn.enabled = False - 若模型含
torch.nn.Conv2d或torch.nn.BatchNorm2d,还需额外禁用 CuDNN 确定性模式:torch.backends.cudnn.benchmark = False、torch.backends.cudnn.deterministic = True - 注意:禁用 CuDNN 可能轻微降低单卡吞吐(约 5%~8%),但对多卡 DDP 场景影响更小
torch.distributed.fsdp 与显存分页共存时,必须用 use_orig_params=True
FSDP 默认启用 use_orig_params=False,会将参数 flatten 成超大张量并集中管理,这与分页机制的“按模块粒度动态分配/释放”原则冲突。此时 Inductor 观察到的是一个巨型 flat param,无法按层切分页块。
启用 use_orig_params=True 后,FSDP 保留原始参数结构,torch.compile 才能在每个 nn.Module 子图中独立应用分页策略——例如,只对 Transformer 的 SelfAttention 子图启用高频率页交换,而对 MLP 子图保持常驻。
- FSDP 初始化时传参:
fsdp_config = {"use_orig_params": True} - 对应地,
optimizer必须迭代model.parameters()而非model.named_parameters(),否则可能漏参 - 该配置会略微增加 FSDP 的通信开销(因参数未 flatten),但换来显存可预测性提升
验证分页是否生效:看 nvidia-smi 的 Used 波动 + torch.cuda.memory_snapshot() 中 segment 数量
仅靠总显存占用下降不足以确认分页启用。关键指标是:同一训练 step 内,显存 Used 值是否呈现高频小幅波动(±50–200MB),而非缓慢爬升后陡降——前者说明页块正在被快速复用。
更准确的方式是插入检查点:
if step % 100 == 0:
torch.cuda.memory._dump_snapshot(f"mem_{step}.pickle")
然后用 torch.cuda.memory.trace_plot(snapshot) 查看输出 HTML,重点关注 segments 列表长度是否稳定在数百量级(分页模式)而非个位数(传统大块模式)。
- 常见误判:看到
torch.cuda.memory_allocated()下降就认为分页生效 —— 实际可能是empty_cache()清理了碎片 - 真实分页特征:多个
segment的size集中在 2MB / 4MB / 8MB,且state在active和inactive间频繁切换 - 如果 snapshot 中大量
segment的traceback指向inductor/kernel.py,基本可确认生效
batch_size 和 seq_len,再逐项开关配置验证。Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











