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

如何清除GPU内存中的tf.data.Dataset以解决内存不足问题?

如何清除GPU内存中的tf.data.Dataset以解决内存不足问题?

我碰到过类似的GPU内存占用困扰,核心问题在于tf.data.Dataset.from_tensor_slices处理numpy数组时,会自动把数据转换成TensorFlow张量;如果GPU处于可用状态,这些张量大概率会被默认放到GPU内存里,后续的批处理、预取操作还会让更多数据驻留GPU,导致训练完成后内存无法自动释放,进而影响后续推理任务。下面给你几个实用的解决办法:

一、从根源避免数据一次性加载到GPU

最有效的思路是不让数据集一开始就跑到GPU上,而是留在CPU内存中,训练时再按需传输到GPU,这样能从源头控制GPU内存占用:

1. 用CPU上下文管理器创建Dataset

创建Dataset时,显式指定使用CPU设备,确保所有张量都留在CPU,不会占用GPU内存:

# 创建训练数据集时绑定CPU设备
with tf.device('/CPU:0'):
    train_X = {'data': train_data, 'index': train_index}
    train_dataset = tf.data.Dataset.from_tensor_slices((train_X, train_y))
    train_dataset = train_dataset.batch(256).prefetch(tf.data.AUTOTUNE)

# 验证数据集同理操作
with tf.device('/CPU:0'):
    val_X = {'data': val_data, 'index': val_index}
    val_dataset = tf.data.Dataset.from_tensor_slices((val_X, val_y))
    val_dataset = val_dataset.batch(256).prefetch(tf.data.AUTOTUNE)

2. 改用流式加载的Dataset格式

如果数据量极大,建议把数据保存为TFRecord格式,再用tf.data.TFRecordDataset流式读取。这种方式会按需加载数据,不会一次性占用大量内存(不管是CPU还是GPU):

# 先将数据写入TFRecord文件(示例代码)
def write_tfrecord(data_dict, labels, filename):
    with tf.io.TFRecordWriter(filename) as writer:
        for idx in range(len(labels)):
            # 构造TFRecord特征
            data_feature = tf.train.Feature(float_list=tf.train.FloatList(value=data_dict['data'][idx].flatten()))
            index_feature = tf.train.Feature(int64_list=tf.train.Int64List(value=data_dict['index'][idx]))
            label_feature = tf.train.Feature(int64_list=tf.train.Int64List(value=labels[idx]))
            
            feature_dict = {'data': data_feature, 'index': index_feature, 'label': label_feature}
            example = tf.train.Example(features=tf.train.Features(feature=feature_dict))
            writer.write(example.SerializeToString())

# 写入训练和验证数据
write_tfrecord({'data': train_data, 'index': train_index}, train_y, 'train.tfrecord')
write_tfrecord({'data': val_data, 'index': val_index}, val_y, 'val.tfrecord')

# 读取TFRecord并解析
def parse_tfexample(example_proto):
    feature_desc = {
        'data': tf.io.FixedLenFeature([15000], tf.float32),
        'index': tf.io.FixedLenFeature([2], tf.int64),
        'label': tf.io.FixedLenFeature([1], tf.int64)
    }
    parsed = tf.io.parse_single_example(example_proto, feature_desc)
    # 恢复数据形状
    data = tf.reshape(parsed['data'], (15000, 1))
    x_input = {'data': data, 'index': parsed['index']}
    y_label = tf.cast(parsed['label'], tf.int32)
    return x_input, y_label

# 构建流式数据集
train_dataset = tf.data.TFRecordDataset('train.tfrecord').map(parse_tfexample).batch(256).prefetch(tf.data.AUTOTUNE)
val_dataset = tf.data.TFRecordDataset('val.tfrecord').map(parse_tfexample).batch(256).prefetch(tf.data.AUTOTUNE)

二、强制释放已占用的GPU内存

如果已经创建了占GPU内存的Dataset,可以尝试以下步骤强制清理资源:

  1. 先删除Dataset对象及相关的张量引用:
    del train_dataset, val_dataset, train_X, val_X
    
  2. 调用Python垃圾回收,同时触发TensorFlow清理未使用的GPU内存:
    import gc
    gc.collect()
    # 重新开启GPU内存增长会触发内存池清理
    gpus = tf.config.list_physical_devices('GPU')
    tf.config.experimental.set_memory_growth(gpus[0], True)
    
  3. 可以结合keras.backend.clear_session()一起使用,它能清除模型相关的GPU资源,辅助释放内存。

三、排查隐藏的内存占用点

  • 如果你在Dataset上使用了cache()方法,默认会把数据缓存到GPU内存,建议改成缓存到磁盘:dataset.cache('/tmp/dataset_cache')
  • 预取缓冲大小不要手动设置过大,prefetch(tf.data.AUTOTUNE)是最优选择,它会根据系统资源自动调整缓冲规模。

备注:内容来源于stack exchange,提问作者Alb

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 09:27:58