You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.10 12:05:22