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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 03:50:01