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

如何用PyTorch/TensorFlow准备转换数据集?附自定义Dataset代码求验证

自定义图像数据集准备:PyTorch & TensorFlow 实用指南

Hey there! Let's walk through your 80-image small dataset preparation step by step, fix your PyTorch code, and cover TensorFlow implementation too—we'll also focus on key best practices tailored for small datasets to avoid common pitfalls.


一、Your Current PyTorch Dataset Code: Key Issues to Fix

Your core idea is on the right track, but there are a few critical details that'll cause errors or suboptimal results:

  • Undefined variable path2: You're using path2 in __getitem__ but never defined it—you should use the path parameter passed to __init__ instead.
  • Wrong transform order: transforms.Normalize requires tensor input, so it must come after transforms.ToTensor() (your current order is reversed).
  • Redundant code: self.transformations is defined but never used, and you duplicate ToTensor() with self.to_tensor—we can clean this up.
  • Poor label handling: Using filenames directly as labels rarely works unless your filenames explicitly encode class info. For most cases, you'll want to extract labels from folder structure (the standard approach).
  • Overly large batch size: With only 80 images, a batch_size=100 means you'll only get one batch—decrease it to 8 or 16 for meaningful batch training.
  • Unnecessary np.size: len(self.name) does the same job more simply than np.size(self.name).

二、Fixed PyTorch Custom Dataset Code

Let's assume you're using the standard folder structure (highly recommended):

your_data/
├── class_a/
│   ├── img1.jpg
│   └── img2.jpg
└── class_b/
    ├── img3.jpg
    └── img4.jpg

Here's the revised code that handles this structure correctly:

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

class MyCustomDataset(Dataset):
    def __init__(self, root_dir, transform=None):
        # Collect all image paths and map classes to numeric labels
        self.image_paths = []
        self.labels = []
        self.classes = sorted(os.listdir(root_dir))
        self.class_to_idx = {cls_name: idx for idx, cls_name in enumerate(self.classes)}

        # Traverse each class folder
        for cls_name in self.classes:
            cls_folder = os.path.join(root_dir, cls_name)
            for img_filename in os.listdir(cls_folder):
                self.image_paths.append(os.path.join(cls_folder, img_filename))
                self.labels.append(self.class_to_idx[cls_name])

        # Set default transforms if none are provided
        self.transform = transform or transforms.Compose([
            transforms.Resize((224, 224)),  # Standardize image size
            transforms.ToTensor(),
            transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))  # Normalize after ToTensor
        ])

    def __getitem__(self, index):
        # Load image and ensure it's RGB (avoids grayscale issues)
        img = Image.open(self.image_paths[index]).convert('RGB')
        label = self.labels[index]

        # Apply transforms
        if self.transform:
            img = self.transform(img)

        return img, label

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

if __name__ == '__main__':
    # Initialize dataset (replace './your_data' with your actual path)
    dataset = MyCustomDataset(root_dir='./your_data')
    # Use a small batch size for your 80-image dataset
    data_loader = torch.utils.data.DataLoader(dataset, batch_size=8, shuffle=True)

    # Test the loader
    for batch_imgs, batch_labels in data_loader:
        print(f"Batch shape: {batch_imgs.shape}, Labels: {batch_labels}")
        break

Quick Notes on the Fixed Code:

  • Automatically extracts labels from folder names (no manual label mapping needed)
  • Standardizes image size to avoid shape mismatches in your model
  • Fixes the transform order to ensure normalization works correctly
  • Handles grayscale images by converting them to RGB
  • Uses a batch size that makes sense for your small dataset

If your images are all in a single folder (labels encoded in filenames, e.g., class_a_img001.jpg), adjust the __init__ to parse labels from filenames:

# Example for single-folder datasets
def __init__(self, img_dir, transform=None):
    self.image_paths = [os.path.join(img_dir, fname) for fname in os.listdir(img_dir)]
    self.labels = []
    # Parse class from filename (adjust this logic to match your naming scheme)
    for fname in os.listdir(img_dir):
        cls_name = fname.split('_')[0]
        self.labels.append(0 if cls_name == 'class_a' else 1)  # Map to numeric label
    # Rest of the __init__ remains the same

三、TensorFlow Implementation for Small Datasets

TensorFlow has built-in tools that make small dataset loading even easier. Here are two common approaches:

Use image_dataset_from_directory to auto-load and split your data:

import tensorflow as tf

# Load training and validation sets (20% of data for validation)
train_ds = tf.keras.utils.image_dataset_from_directory(
    './your_data',
    validation_split=0.2,
    subset="training",
    seed=123,
    image_size=(224, 224),
    batch_size=8
)

val_ds = tf.keras.utils.image_dataset_from_directory(
    './your_data',
    validation_split=0.2,
    subset="validation",
    seed=123,
    image_size=(224, 224),
    batch_size=8
)

# Add data augmentation (critical for small datasets to prevent overfitting)
data_augmentation = tf.keras.Sequential([
    tf.keras.layers.RandomFlip("horizontal"),
    tf.keras.layers.RandomRotation(0.1),
    tf.keras.layers.RandomZoom(0.1)
])

# Apply augmentation to training data
train_ds = train_ds.map(lambda x, y: (data_augmentation(x, training=True), y))

# Normalize pixel values to [-1, 1] (matches PyTorch's normalization)
normalizer = tf.keras.layers.Rescaling(1./127.5, offset=-1)
train_ds = train_ds.map(lambda x, y: (normalizer(x), y))
val_ds = val_ds.map(lambda x, y: (normalizer(x), y))

2. Single-Folder Dataset

If you need to load from a single folder, use from_generator:

import os
import tensorflow as tf

def load_single_folder_data(img_dir):
    img_paths = [os.path.join(img_dir, fname) for fname in os.listdir(img_dir)]
    for path in img_paths:
        # Load and preprocess image
        img = tf.io.read_file(path)
        img = tf.image.decode_jpeg(img, channels=3)
        img = tf.image.resize(img, (224, 224))
        # Parse label from filename (adjust to your naming scheme)
        cls_name = os.path.basename(path).split('_')[0]
        label = 0 if cls_name == 'class_a' else 1
        yield img, label

# Create dataset
dataset = tf.data.Dataset.from_generator(
    lambda: load_single_folder_data('./single_folder'),
    output_signature=(
        tf.TensorSpec(shape=(224, 224, 3), dtype=tf.float32),
        tf.TensorSpec(shape=(), dtype=tf.int32)
    )
).batch(8).shuffle(80)

四、Best Practices for Small Image Datasets

  • Use a standard folder structure: It simplifies label handling and works with most built-in dataset loaders.
  • Add data augmentation: Random flips, rotations, and zooms help your model generalize better (critical for 80 images).
  • Split into training/validation sets: Reserve 10-20% of your data for validation to catch overfitting early.
  • Standardize image size: Ensure all images have the same dimensions to avoid model input errors.
  • Normalize pixel values: Scale values to a consistent range (e.g., [-1,1] or [0,1]) to speed up model convergence.

内容的提问来源于stack exchange,提问作者Mahmoud Al zaitoun

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 09:01:25