显存溢出是显存分配策略、数据流和计算图共同作用的结果,需综合排查memory_growth设置时机、输入尺寸、混合精度配置、训练/推理模式等多因素。

显存溢出不是“模型太大”这么简单,而是显存分配策略、数据流和计算图共同作用的结果。只要模型参数+梯度+激活值+优化器状态的总和超过GPU当前可用显存,ResourceExhaustedError: OOM when allocating tensor 就会立刻报出——哪怕只差几MB。
tf.config.experimental.set_memory_growth没生效的典型原因
这个设置必须在 import tensorflow as tf 之后、**任何模型构建或 tf.device 上下文创建之前**执行,否则无效。
- 常见错误:先写了
model = tf.keras.Sequential(...),再调用set_memory_growth—— 此时TensorFlow已按默认策略预分配显存,增长模式被忽略 - 另一个坑:用了
tf.config.experimental.set_memory_limit(gpu, ...)后又试图启用set_memory_growth,两者互斥,后者直接失效 - 验证是否生效:运行中用
nvidia-smi -l 1观察Memory-Usage是否缓慢爬升;如果一启动就打满,说明没生效
batch_size调小却仍OOM,问题可能不在batch上
单样本前向/反向所需的显存峰值,由输入尺寸、网络深度、激活函数类型决定。比如一张 2048x2048 图像过 ResNet50,光中间激活值就可能吃掉 6GB+ 显存,此时把 batch_size 从 32 降到 1 毫无意义。
- 检查输入 shape:用
print(x_train.shape)确认是否意外加载了超高分辨率图像或长序列 - BN 层在
batch_size=1下会因统计不稳定而触发额外缓存(尤其在训练模式),建议推理时固定training=False - 某些自定义层(如带内部状态的 RNN、Attention)会在 batch 维度外隐式放大显存占用,需逐层注释排查
混合精度没起效,反而让loss爆炸
mixed_float16 策略要求输出层保持 float32,否则 softmax 或 loss 计算容易下溢或 NaN。TensorFlow 不会自动帮你兜底。
- 必须显式指定关键层 dtype:
tf.keras.layers.Dense(10, dtype='float32'),尤其是最后一层 - 优化器要包装成
LossScaleOptimizer,否则梯度更新阶段仍用 float32,失去显存收益 - 如果用了自定义 loss 函数,需确保其内部运算支持 float16,例如避免直接用
tf.log而不用tf.math.log(后者有 dtype 推导)
推理时显存越用越多,大概率是没关 training 模式
model(batch, training=False) 和 model.predict() 行为不同:后者不保证关闭 BN/ Dropout 的状态更新逻辑,尤其在自定义模型中。
- 务必在推理循环中显式传入
training=False,哪怕模型里没写 BN 层——某些 ops(如tf.nn.l2_normalize)内部也有训练/推理分支 - 避免在循环内反复调用
model(batch)而不释放引用;Python 的del pred+gc.collect()对 GPU 张量无效,得靠上下文管理或重用变量名 - 使用
tf.data.Dataset流式推理时,prefetch和cache会常驻显存,调试阶段可临时禁用
真正卡住人的从来不是某一行代码,而是多个机制叠加后的隐性资源竞争:比如 set_memory_growth 开了但混合精度没配对、batch_size 调小了但输入分辨率没降、推理关了 training 却忘了清理 tf.function 缓存图——这些点单独看都合理,合起来就爆了。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











