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

使用Keras DataGenerator,如何在拆分前打乱数据解决模型过拟合?

Hey there! Let's break this down step by step since you're dealing with overfitting and having trouble getting shuffle=True to work as expected. Here's how to properly load, shuffle, and split your data with your image folder + CSV setup:

Key Reason Why shuffle=True Might Not Be Working

Chances are your original CSV has samples grouped by label (e.g., all cat images first, then all dog images). If you split the data before shuffling, your training set might be dominated by one label, and your validation set by another. Even if you set shuffle=True in your data loader, this lopsided label distribution will still lead to poor generalization (and overfitting). The fix? Shuffle the entire dataset first, then split.

Step-by-Step Solution

We'll use pandas to handle the CSV, sklearn to split, and then show examples for both TensorFlow/Keras and PyTorch (two common frameworks for beginners).

1. Load & Shuffle the CSV Data

First, we'll load your label CSV and completely randomize the row order to mix up labels:

import pandas as pd
from sklearn.model_selection import train_test_split

# Load your CSV file (replace with your file path)
df = pd.read_csv("your_labels.csv")

# Shuffle the entire dataset (frac=1 means use all rows)
# random_state ensures reproducibility (you can use any number)
shuffled_df = df.sample(frac=1, random_state=42).reset_index(drop=True)

2. Split into Train/Validation Sets

Use train_test_split with the stratify parameter to guarantee the label distribution is the same in both sets (critical for fair evaluation):

# Split into 80% training, 20% validation
train_df, val_df = train_test_split(
    shuffled_df,
    test_size=0.2,
    random_state=42,
    stratify=shuffled_df["your_label_column_name"]  # Replace with your label column name
)

3. Load Images with Labels

For TensorFlow/Keras Users

Use flow_from_dataframe to link your CSV labels to the image files:

from tensorflow.keras.preprocessing.image import ImageDataGenerator

# Normalize pixel values to 0-1
train_datagen = ImageDataGenerator(rescale=1./255)
val_datagen = ImageDataGenerator(rescale=1./255)

# Training data generator
train_generator = train_datagen.flow_from_dataframe(
    dataframe=train_df,
    directory="train/",  # Path to your image folder
    x_col="your_filename_column_name",  # CSV column with image filenames
    y_col="your_label_column_name",
    target_size=(224, 224),  # Adjust to match your model's input size
    batch_size=32,
    class_mode="categorical"  # Use "binary" for 2-class problems
)

# Validation data generator
val_generator = val_datagen.flow_from_dataframe(
    dataframe=val_df,
    directory="train/",
    x_col="your_filename_column_name",
    y_col="your_label_column_name",
    target_size=(224, 224),
    batch_size=32,
    class_mode="categorical"
)

For PyTorch Users

Create a custom Dataset class to load images and labels, then use DataLoader:

import torch
from torch.utils.data import Dataset, DataLoader
from PIL import Image
import os
from torchvision import transforms

# Custom dataset class
class ImageLabelDataset(Dataset):
    def __init__(self, df, img_dir, transform=None):
        self.df = df
        self.img_dir = img_dir
        self.transform = transform

    def __len__(self):
        return len(self.df)

    def __getitem__(self, idx):
        # Get image path
        img_path = os.path.join(self.img_dir, self.df.iloc[idx]["your_filename_column_name"])
        # Load image (convert to RGB to handle grayscale images)
        image = Image.open(img_path).convert("RGB")
        # Get label
        label = self.df.iloc[idx]["your_label_column_name"]
        
        if self.transform:
            image = self.transform(image)
        return image, label

# Define transforms (add data augmentation for training!)
train_transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.RandomHorizontalFlip(),  # Random flip to augment data
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

val_transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

# Create datasets and loaders
train_dataset = ImageLabelDataset(train_df, img_dir="train/", transform=train_transform)
val_dataset = ImageLabelDataset(val_df, img_dir="train/", transform=val_transform)

train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4)
val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4)
Bonus Tips to Fight Overfitting

Since you're dealing with overfitting, here are a few extra steps to help:

  • Add more data augmentation: For example, random rotations, zoom, or brightness adjustments (look up ImageDataGenerator for Keras or torchvision.transforms for PyTorch).
  • Use Dropout layers: Add Dropout(0.2) between dense layers in your model to randomly disable neurons during training.
  • Weight decay (L2 regularization): Add weight_decay=1e-4 to your optimizer to penalize large weights.
  • Early stopping: Monitor your validation loss, and stop training when it stops improving (use EarlyStopping callback in Keras or implement it manually in PyTorch).

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.12 04:30:50