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
相关产品推荐
相关产品推荐

