使用VGG16模型加载CelebA数据集时遇加载失败问题求助
解决CelebA数据集加载错误的方案
错误原因
flow_from_directory API的设计逻辑是:目标目录下必须包含对应类别的子文件夹,每个子文件夹存放该类别的图片。而你的img_align_celeba文件夹直接存放所有图片,没有分类子目录,导致API无法识别类别和图片,因此出现Found 0 images belonging to 0 classes的错误。
正确处理CelebA数据集的方法
CelebA的标签存储在单独的CSV文件(list_attr_celeba.csv)中,推荐用flow_from_dataframe替代flow_from_directory来加载数据,具体步骤如下:
步骤1:导入库并加载标签文件
确保你已下载标签文件,和图片目录路径对应:
import tensorflow as tf import pandas as pd from tensorflow.keras.preprocessing.image import ImageDataGenerator from tensorflow.keras.applications.vgg16 import VGG16, preprocess_input # 路径配置 img_dir = '/content/drive/MyDrive/Datasets/img_align_celeba' label_path = '/content/drive/MyDrive/Datasets/list_attr_celeba.csv' # 替换为你的标签文件路径 # 加载标签,以"Male"属性为例,可替换为其他属性 df = pd.read_csv(label_path) # 将原标签的-1/1转换为0/1,适配二分类任务 df['Male'] = df['Male'].apply(lambda x: 1 if x == 1 else 0)
步骤2:配置ImageDataGenerator
注意:preprocess_input已包含归一化逻辑,无需重复设置rescale=1./255:
datagen = ImageDataGenerator( preprocessing_function=preprocess_input, validation_split=0.2 # 按需划分训练/验证集比例 )
步骤3:生成训练/验证数据
# 训练集生成器 train_generator = datagen.flow_from_dataframe( dataframe=df, directory=img_dir, x_col='image_id', # CSV中存图片文件名的列 y_col='Male', # 训练目标属性列 target_size=(224, 224), batch_size=32, class_mode='binary', # 二分类用binary,多分类用categorical subset='training' ) # 验证集生成器 val_generator = datagen.flow_from_dataframe( dataframe=df, directory=img_dir, x_col='image_id', y_col='Male', target_size=(224, 224), batch_size=32, class_mode='binary', subset='validation' )
备选方案:仅加载图片(无监督/特征提取场景)
如果不需要标签,只是加载图片做特征提取或无监督任务,可手动构造路径列表后用tf.data.Dataset加载:
import os # 获取所有图片路径 img_paths = [os.path.join(img_dir, fname) for fname in os.listdir(img_dir) if fname.endswith('.jpg')] # 定义图片加载函数 def load_image(path): img = tf.io.read_file(path) img = tf.image.decode_jpeg(img, channels=3) img = tf.image.resize(img, (224, 224)) img = preprocess_input(img) return img # 构建数据集 dataset = tf.data.Dataset.from_tensor_slices(img_paths) dataset = dataset.map(load_image, num_parallel_calls=tf.data.AUTOTUNE) dataset = dataset.batch(32) # 获取验证批次 X_val = next(iter(dataset))
注意事项
- 确保标签文件的
image_id列与图片文件名完全对应(比如000001.jpg) - 不要重复设置
rescale和preprocess_input,避免预处理冲突 - 多分类任务需将
class_mode改为categorical,并确保标签格式适配
内容的提问来源于stack exchange,提问作者Jasmine
相关产品推荐
相关产品推荐

