如何用PyTorch/TensorFlow准备转换数据集?附自定义Dataset代码求验证
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 usingpath2in__getitem__but never defined it—you should use thepathparameter passed to__init__instead. - Wrong transform order:
transforms.Normalizerequires tensor input, so it must come aftertransforms.ToTensor()(your current order is reversed). - Redundant code:
self.transformationsis defined but never used, and you duplicateToTensor()withself.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=100means 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 thannp.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:
1. Folder-Structured Dataset (Recommended)
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

