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

TensorFlow 2.8中model.fit正常但model.predict内存不足问题及解决

TensorFlow 2.8中model.fit正常但model.predict触发OOM错误问题及解决方法

问题背景

基于TensorFlow 2.8(使用谷歌AI平台自定义镜像,基础镜像为gcr.io/deeplearning-platform-release/tf-gpu.2-8)训练Keras模型时遇到如下矛盾现象:

  • 训练阶段用batch size=128执行model.fit(dataset, epochs=10)完全正常
  • 训练完成后调用model.predict(dataset)时,哪怕把batch size降到1,依然会触发**内存不足(Out of Memory)**错误

数据集加载代码如下:

options = tf.data.Options()
options.experimental_distribute.auto_shard_policy = tf.data.experimental.AutoShardPolicy.OFF
dataset = (
            tf.data.Dataset.list_files(f'{dataset_location}/*.csv')
            .flat_map(tf.data.TextLineDataset)
            .skip(1)
            .map(decode_line) # 自定义CSV解析函数
        )
dataset = dataset.apply(tf.data.experimental.ignore_errors())
dataset = dataset.batch(batch_size).with_options(options)

疑问:为什么内存消耗更高的fit调用能正常运行,predict却出现OOM?

解决方法及推测

推测问题根源是TensorFlow 2.8的predict函数存在内存泄漏。通过修改数据集加载逻辑,改为逐个加载单份CSV文件并在小数据集上执行预测,而非一次性加载所有CSV文件,可有效规避OOM问题。示例代码如下:

predictions = []
for file_name in tf.io.gfile.listdir(dataset_location):
    file_path = f'{dataset_location}/{file_name}'
    dataset = tf.data.experimental.make_csv_dataset(
        file_pattern=file_path,
        batch_size=10,
        shuffle=False,
        label_name="label",
        field_delim=';',
        column_defaults=[[""], [""], [0.0]],
        num_epochs=1
    )
    pred = model.predict(dataset)
    predictions.append(pred)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 04:51:12