使用tfds加载places365_small时model.fit()致VSCode崩溃求助
问题分析与解决方案
你遇到的内存耗尽问题,核心原因是数据集处理流程的内存占用设计严重超出了机器承载能力,和CIFAR10的小数据规模不同,places365_small的大尺寸、大样本量需要针对性调整处理策略,具体问题和解决方法如下:
核心问题点
Shuffle缓冲区过大
你设置了shuffle(ds_info.splits['train'].num_examples),这会尝试把整个训练集(约180万张256×256图像)加载到内存做全局打乱,直接触发内存溢出。CIFAR10只有5万张32×32小图,内存可以容纳,但places365_small的规模完全不是一个量级。全量内存缓存
cache()在batch()之前调用,会把未分批的原始图像数据全部缓存到内存,同样远超内存容量。Batch Size过高
256×256的图像用128的batch size,单批次数据加上卷积计算的中间特征图,会占用极高的内存/显存,直接拉满资源。
修复步骤与代码示例
1. 调整Shuffle缓冲区大小
把全局打乱改成小缓冲区打乱,既保证数据随机性,又避免内存过载:
ds_train = ds_train.shuffle(1000) # 用1000的缓冲区,而非整个训练集大小
2. 优化缓存策略
要么将cache()移到batch()之后(缓存分批后的数据),要么改用磁盘缓存(指定路径,避免占用内存):
# 先分批再缓存,或指定磁盘缓存路径 ds_train = ds_train.batch(32) ds_train = ds_train.cache('./places365_train_cache') # 磁盘缓存
3. 减小Batch Size
根据你的内存/显存情况,把batch size从128降到32或16:
ds_train = ds_train.batch(32) ds_test = ds_test.batch(32)
4. (可选)缩小图像分辨率
如果任务不需要256×256的高分辨率,可以将图像缩小到128×128,大幅降低内存占用:
def normalize_img(image, label): image = tf.image.resize(image, (128, 128)) # 缩小图像尺寸 return tf.cast(image, tf.float32) / 255., label
同时对应调整模型输入形状:
model.add(tf.keras.layers.Conv2D(32, (3, 3), activation='relu', input_shape=(128, 128, 3)))
完整修复后代码
import tensorflow as tf import tensorflow_datasets as tfds from tensorflow.keras import datasets, layers, models import matplotlib.pyplot as plt import numpy as np (ds_train, ds_test), ds_info = tfds.load( 'places365_small', split=['train', 'test'], shuffle_files=True, as_supervised=True, with_info=True, ) def normalize_img(image, label): """Normalizes images: `uint8` -> `float32`.""" # 可选:缩小图像尺寸进一步降低内存占用 # image = tf.image.resize(image, (128, 128)) return tf.cast(image, tf.float32) / 255., label ds_train = ds_train.map( normalize_img, num_parallel_calls=tf.data.AUTOTUNE) ds_train = ds_train.shuffle(1000) # 小缓冲区打乱 ds_train = ds_train.batch(32) # 减小batch size ds_train = ds_train.cache('./places365_train_cache') # 磁盘缓存 ds_train = ds_train.prefetch(tf.data.AUTOTUNE) ds_test = ds_test.map( normalize_img, num_parallel_calls=tf.data.AUTOTUNE) ds_test = ds_test.batch(32) ds_test = ds_test.cache('./places365_test_cache') ds_test = ds_test.prefetch(tf.data.AUTOTUNE) model = tf.keras.models.Sequential() # 若使用resize,对应修改input_shape为(128,128,3) model.add(tf.keras.layers.Conv2D(32, (3, 3), activation='relu', input_shape=(256, 256, 3))) model.add(tf.keras.layers.MaxPooling2D((2, 2))) model.add(tf.keras.layers.Conv2D(64, (3, 3), activation='relu')) model.add(tf.keras.layers.MaxPooling2D((2, 2))) model.add(tf.keras.layers.Conv2D(64, (3, 3), activation='relu')) model.add(tf.keras.layers.Flatten()) model.add(tf.keras.layers.Dense(64, activation='relu')) model.add(tf.keras.layers.Dense(365)) model.compile( optimizer=tf.keras.optimizers.Adam(0.001), loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=[tf.keras.metrics.SparseCategoricalAccuracy()], ) model.fit( ds_train, epochs=6, validation_data=ds_test, )
内容的提问来源于stack exchange,提问作者user21122083
相关产品推荐
相关产品推荐

