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

如何在Keras中加载自定义本地数据集替代cifar10?

Hey there! Great question—making the jump from pre-built datasets like CIFAR-10 to your own custom images is such an important step for real-world computer vision work. Let me walk you through exactly how to set this up.

Keras (and TensorFlow under the hood) works best with a directory-based structure where each class has its own subfolder. This lets the library automatically assign labels to your images without extra manual work. Here's a standard setup for classification tasks:

my_custom_dataset/
├── train/
│   ├── cat/
│   │   ├── fluffy_cat.jpg
│   │   ├── tabby_cat.png
│   │   └── ...
│   ├── dog/
│   │   ├── golden_retriever.jpg
│   │   ├── poodle.png
│   │   └── ...
│   └── bird/
│       └── ...
└── validation/
    ├── cat/
    ├── dog/
    └── bird/
  • train/: Holds your training images, split into subfolders named after each class (e.g., "cat", "dog"). Every image in the "cat" folder will automatically get a "cat" label.
  • validation/: Contains images for evaluating your model during training—use the same class subfolder names as the train directory.
  • If you have a separate test set, add a test/ directory with the same structure too.

Pro tip: Stick to common image formats (JPG, PNG, BMP) — Keras handles these without extra setup.

2. Loading Your Dataset with Keras

There are two main ways to load your data, depending on whether you need data augmentation (to boost model generalization) or not.

Option 1: tf.keras.utils.image_dataset_from_directory (Simplest, TensorFlow 2.3+)

This is the modern, hassle-free method. It returns a tf.data.Dataset object that integrates smoothly with Keras models.

import tensorflow as tf

# Set your image dimensions and batch size
img_height = 128  # Adjust to match your images' actual size (or what your model expects)
img_width = 128
batch_size = 32

# Load training data
train_ds = tf.keras.utils.image_dataset_from_directory(
  "path/to/my_custom_dataset/train",
  image_size=(img_height, img_width),
  batch_size=batch_size
)

# Load validation data
val_ds = tf.keras.utils.image_dataset_from_directory(
  "path/to/my_custom_dataset/validation",
  image_size=(img_height, img_width),
  batch_size=batch_size
)

# Get your class names automatically (super handy!)
class_names = train_ds.class_names
print("Classes found:", class_names)  # Outputs ['cat', 'dog', 'bird']

Option 2: ImageDataGenerator (For Data Augmentation)

If you want to apply transformations like rotation, flipping, or zooming to your training images (to prevent overfitting), use ImageDataGenerator with flow_from_directory.

from tensorflow.keras.preprocessing.image import ImageDataGenerator

# Set up training generator with augmentation
train_datagen = ImageDataGenerator(
    rescale=1./255,  # Normalize pixel values to [0, 1] (critical for stable training!)
    rotation_range=20,  # Randomly rotate images up to 20 degrees
    width_shift_range=0.2,  # Randomly shift images horizontally
    height_shift_range=0.2,  # Randomly shift images vertically
    horizontal_flip=True  # Randomly flip images horizontally
)

# Validation generator only needs rescaling (no augmentation here!)
val_datagen = ImageDataGenerator(rescale=1./255)

# Load training data
train_generator = train_datagen.flow_from_directory(
    "path/to/my_custom_dataset/train",
    target_size=(img_height, img_width),
    batch_size=batch_size,
    class_mode="categorical"  # Use 'binary' if you only have 2 classes
)

# Load validation data
val_generator = val_datagen.flow_from_directory(
    "path/to/my_custom_dataset/validation",
    target_size=(img_height, img_width),
    batch_size=batch_size,
    class_mode="categorical"
)
3. Quick Tips to Avoid Headaches
  • Rescaling: Always normalize pixel values by dividing by 255—this helps your model train faster and more reliably.
  • Class Mode: Use class_mode="categorical" for multi-class tasks (one-hot encoded labels), 'binary' for binary classification, or 'sparse' if you want integer labels.
  • Image Sizes: Make sure all images are resized to the same dimensions (using target_size or image_size)—neural networks need fixed input shapes.
  • File Paths: If you're getting errors about missing files, use absolute paths (like /home/user/my_custom_dataset/train) instead of relative paths to avoid confusion.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 18:27:56