TensorFlow GPU训练显存不足,指定部分操作跑CPU无效求解决
问题
在GPU上训练TensorFlow模型时始终遇到显存不足(OOM)问题,尝试将部分操作放在CPU上运行但未解决问题,相关代码如下:
主函数中的GPU配置:
gpus = tf.config.list_physical_devices(device_type='GPU') tf.config.experimental.set_visible_devices(gpus[0],'GPU') tf.config.experimental.set_memory_growth(gpus[0], enable=True)
继承tf.keras.Model的类中,在__init__调用的函数里设置CPU设备:
def _setup_C(self, double_length=False): """ Construct C~ from C double_length: current C is for length L, convert it to length 2L """ with tf.device('/cpu:0'): C = _r2c(self.C) self._setup_state() dA_L = power(self.L, self.dA) # Multiply C by I - dA_L C_ = _conj(C) prod = contract("h m n, c h n -> c h m", tf.transpose(dA_L,perm = [0,2,1]), C_) if double_length: prod = -prod # Multiply by I + dA_L instead C_ = C_ - prod C_ = C_[..., :self.N] # Take conjugate pairs again self.C = tf.identity(_c2r(C_)) if double_length: self.L *= 2 self._omega(self.L, dtype=C.dtype, cache=True)
解决思路
- 验证CPU设备上下文的实际生效:TensorFlow的设备上下文可能被函数内的显式设备设置覆盖,比如
_setup_state()、power()、contract()这些自定义函数如果内部指定了GPU,会优先于外部的tf.device。可以在这些函数内部也添加with tf.device('/cpu:0'),或者用tf.debugging.assert_on_cpu()在代码块内验证执行设备。 - 临时迁移核心张量到CPU:即使操作放在CPU,
self.C、self.dA等核心张量可能仍驻留在GPU显存。可以在_setup_C执行前把这些张量移到CPU,操作完成后按需移回GPU:# 临时移到CPU处理 self.C = tf.identity(self.C, device='/cpu:0') self.dA = tf.identity(self.dA, device='/cpu:0') with tf.device('/cpu:0'): # 原有操作逻辑 ... # 后续需在GPU使用时移回 self.C = tf.identity(self.C, device='/GPU:0') self.dA = tf.identity(self.dA, device='/GPU:0') - 优化运算的内存占用:
contract()的张量收缩可能产生大中间张量,即使在CPU执行,GPU上的训练参数、输入数据仍可能占满显存。可以拆分contract的运算步骤,避免一次性生成大张量;同时直接降低批量大小,这是缓解OOM最直接的方法。 - 清理显存碎片与无用张量:在训练循环间隙调用
tf.keras.backend.clear_session()(不要在模型构建阶段调用),或者用TensorFlow调试工具查看显存占用的具体张量,定位未及时回收的大张量。开启内存增长后仍可能存在碎片,可尝试固定显存分配(用tf.config.set_logical_device_configuration设置显存上限)。 - 定位模型显存大户:
_setup_C可能只是小部分操作,训练时的前向/反向传播(尤其是梯度张量)才是显存占用核心。用TensorBoard Profiler分析显存时间线,找到真正的高占用部分,再针对性迁移到CPU或优化。
内容的提问来源于stack exchange,提问作者jessie
相关产品推荐
相关产品推荐

