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

tensorflow.keras同训练配置下数据量增大触发GPU显存溢出问题

问题原因与解决方案

一、认知误区修正

你之前的认知仅适用于数据按需加载到GPU的场景,实际当你用from_tensor_slices直接传入内存中的numpy数组、或者直接把numpy数组传给fit接口时,高版本TensorFlow存在多个会把全量训练数据提前加载到GPU显存的默认行为/逻辑漏洞:

  • 当使用MirroredStrategy分布式策略时,TensorFlow 2.3及以上版本会默认尝试将整份输入张量预先拷贝到所有可用GPU的显存中做分片,而非仅在每个训练step拷贝当前batch
  • Dataset.from_tensor_slices会将传入的numpy数组直接转为TensorFlow常量,默认存储在显存而非内存,总数据量越大占用的固定显存越高,会直接挤占模型参数、梯度、batch数据的可用显存空间

二、TensorFlow 2.6+版本出现该问题的核心原因

对比TensorFlow 2.1,2.6版本有3个直接导致显存占用升高的改动:

  1. 分布式训练默认数据预拷贝逻辑:2.3版本之后为了提升分布式训练性能,默认会把全量输入张量预先分发到各个GPU显存,你总数据量2196的情况下,两份(x+y)全量数据会直接占掉Tesla M60单卡8G显存的2-3G,叠加模型参数、梯度、batch数据的占用就会触发OOM
  2. tensorflow.data默认预取行为:2.4版本之后Dataset会默认开启prefetch(tf.data.AUTOTUNE),你没有显式关闭的话,会提前加载多batch数据到显存,进一步抬高显存占用
  3. Windows平台分布式适配bug:Windows系统下的MirroredStrategy比Linux平台多了冗余的张量拷贝逻辑,你遇到的AUTO sharding警告本身就是bug的一部分,TensorFlow在找不到文件源的情况下会强制做全量数据分片缓存,额外占用显存

三、具体解决方法

按优先级依次尝试:

1. 禁止TensorFlow预先把全量数据加载到显存

把from_tensor_slices的输入替换为tf.data.Dataset.from_generator,避免直接把numpy数组转为显存常量:

def train_generator():
    for x, y in zip(x_train, y_train):
        yield x, y
def val_generator():
    for x, y in zip(x_test, y_test):
        yield x, y

train_data = tf.data.Dataset.from_generator(
    train_generator,
    output_signature=(
        tf.TensorSpec(shape=x_train.shape[1:], dtype=x_train.dtype),
        tf.TensorSpec(shape=y_train.shape[1:], dtype=y_train.dtype)
    )
)
val_data = tf.data.Dataset.from_generator(
    val_generator,
    output_signature=(
        tf.TensorSpec(shape=x_test.shape[1:], dtype=x_test.dtype),
        tf.TensorSpec(shape=y_test.shape[1:], dtype=y_test.dtype)
    )
)
# 后续的shuffle、batch、options逻辑保持不变

2. 显式关闭不必要的缓存与预取

构造Dataset之后追加如下配置:

# 仅预取1个batch,不要使用AUTOTUNE
train_data = train_data.prefetch(1)
# 缓存到磁盘而非显存,Windows环境下修改为本地磁盘路径
train_data = train_data.cache("C:/tmp/train_cache")

3. 开启TensorFlow显存动态分配

在代码最开头加入如下配置,禁止TensorFlow一次性占用全部显存:

import tensorflow as tf
gpus = tf.config.list_physical_devices('GPU')
for gpu in gpus:
    tf.config.experimental.set_memory_growth(gpu, True)

4. 替换跨设备通信策略

Windows环境下HierarchicalCopyAllReduce显存占用高且性能差,替换为tf.distribute.ReductionToOneDevice:

strategy = tf.distribute.MirroredStrategy(
    cross_device_ops=tf.distribute.ReductionToOneDevice()
)

内容的提问来源于stack exchange,提问作者Romaxx

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 22:48:00