如何在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.
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" )
- 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_sizeorimage_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

