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个直接导致显存占用升高的改动:
- 分布式训练默认数据预拷贝逻辑:2.3版本之后为了提升分布式训练性能,默认会把全量输入张量预先分发到各个GPU显存,你总数据量2196的情况下,两份(x+y)全量数据会直接占掉Tesla M60单卡8G显存的2-3G,叠加模型参数、梯度、batch数据的占用就会触发OOM
- tensorflow.data默认预取行为:2.4版本之后
Dataset会默认开启prefetch(tf.data.AUTOTUNE),你没有显式关闭的话,会提前加载多batch数据到显存,进一步抬高显存占用 - 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
相关产品推荐
相关产品推荐

