PyTorch Dataset

In deep learning tasks, data loading and processing are a crucial part.

PyTorch provides powerful data loading and processing tools, mainly including:

  • torch.utils.data.Dataset: An abstract class for datasets, which needs to be customized and implemented__len__(dataset size) and__getitem__(get samples by index).

  • torch.utils.data.TensorDataset: A tensor-based dataset, suitable for handling data-label pairs, directly supports batching and iteration.

  • torch.utils.data.DataLoader: An iterator that wraps Dataset, providing batching, data shuffling, multi-threaded loading, etc., to facilitate data input for model training.

  • torchvision.datasets.ImageFolder: Loads image data from folders, each subfolder represents a class, suitable for image classification tasks.

PyTorch Built-in Datasets

PyTorch provides many commonly used datasets through the torchvision.datasets module, for example:

  • MNIST: Handwritten digit image dataset, used for image classification tasks.
  • CIFAR: A dataset of 60,000 32x32 color images in 10 classes, used for image classification tasks.
  • COCO: A large-scale dataset for general object detection, segmentation, and keypoint detection, containing over 330k images and 2.5M object instances.
  • ImageNet: Contains over 14 million images, used for tasks such as image classification and object detection.
  • STL-10: A dataset containing 100k 96x96 color images, used for image classification tasks.
  • Cityscapes: Contains 5,000 finely annotated city street scene images, used for semantic segmentation tasks.
  • SQUAD: A dataset used for machine reading comprehension tasks.

The above datasets can be loaded through functions in the torchvision.datasets module, and other datasets can also be loaded through custom methods.

torchvision and torchtext

  • torchvision: A graphics library that provides APIs and dataset interfaces related to image data processing, including dataset loading functions and common image transformations.
  • torchtext: A natural language processing toolkit that provides tools for text data processing and modeling, including data preprocessing and data loading methods.

torch.utils.data.Dataset

Dataset is a class in PyTorch used for dataset abstraction.

To create a custom dataset, you need to inherit torch.utils.data.Dataset and override the following two methods:

  • __len__: Returns the size of the dataset.
  • __getitem__: Gets a data sample and its label by index.

Example

import torch
from torch.utils.data import Dataset

# Custom dataset
class MyDataset(Dataset):
    def __init__(self, data, labels):
        # Data initialization
        self.data = data
        self.labels = labels

    def __len__(self):
        # Return dataset size
        return len(self.data)

    def __getitem__(self, idx):
        # Return data and label by index
        sample = self.data[idx]
        label = self.labels[idx]
        return sample, label

# Generate sample data
data = torch.randn(100, 5)  # 100 samples, each with 5 features
labels = torch.randint(0, 2, (100,))  # 100 labels, values are 0 or 1

# Instantiate dataset
dataset = MyDataset(data, labels)

# Test dataset
print("Dataset size:", len(dataset))
print("Sample 0:", dataset[0])

The output result is as follows:

数据集大小: 100
第 0 个样本: (tensor([-0.2006,  0.7304, -1.3911, -0.4408,  1.1447]), tensor(0))

torch.utils.data.DataLoader

DataLoader is a data loader provided by PyTorch, used to load datasets in batches.

It provides the following features:

  • Batch loading: By settingbatch_size。
  • Data shuffling: By settingshuffle=True。
  • Multi-thread acceleration: By settingnum_workers。
  • Iterative access: Conveniently access data batch by batch.

Example

import torch
from torch.utils.data import Dataset
from torch.utils.data import DataLoader

# Custom dataset
class MyDataset(Dataset):
    def __init__(self, data, labels):
        # Data initialization
        self.data = data
        self.labels = labels

    def __len__(self):
        # Return dataset size
        return len(self.data)

    def __getitem__(self, idx):
        # Return data and label by index
        sample = self.data[idx]
        label = self.labels[idx]
        return sample, label

# Generate sample data
data = torch.randn(100, 5)  # 100 samples, each with 5 features
labels = torch.randint(0, 2, (100,))  # 100 labels, values are 0 or 1

# Instantiate dataset
dataset = MyDataset(data, labels)
# Instantiate DataLoader
dataloader = DataLoader(dataset, batch_size=10, shuffle=True, num_workers=0)

# Iterate over DataLoader
for batch_idx, (batch_data, batch_labels) in enumerate(dataloader):
    print(f"Batch {batch_idx + 1}")
    print("Data:", batch_data)
    print("Labels:", batch_labels)
    if batch_idx == 2:  # Only display the first 3 batches
        break

The output result is as follows:

批次 1
数据: tensor([[ 0.4689,  0.6666, -1.0234,  0.8948,  0.4503],
        [ 0.0273, -0.4684, -0.7762,  0.7963,  0.2168],
        [ 1.0677, -0.3502, -0.9594, -1.1318, -0.2196],
        [-1.4989,  0.0267,  1.0405, -0.7284,  0.2335],
        [-0.5887, -0.4934,  1.6283,  1.4638,  0.0157],
        [-1.1047, -0.6550, -0.0381,  0.3617, -1.2792],
        [ 0.3592, -0.8264,  0.0231, -1.5508,  0.6833],
        [-0.6835,  0.6979,  0.9048, -0.4756,  0.3003],
        [ 1.1562, -0.4516, -1.2415,  0.2859,  0.5837],
        [ 0.7937,  1.5316, -0.6139,  0.7999,  0.5506]])
标签: tensor([0, 1, 1, 1, 1, 0, 1, 1, 0, 0])
批次 2
数据: tensor([[-0.0388, -0.3658,  0.8993, -1.5027,  1.0738],
        [-0.6182,  1.0684, -2.3049,  0.8338,  0.1363],
        [-0.5289,  0.1661, -0.0349,  0.2112,  1.4745],
        [-0.3304, -1.2114, -0.2982, -0.3006,  0.5252],
        [-1.4394, -0.3732,  1.0281,  0.5754,  1.0081],
        [ 0.8714, -0.1945, -0.2451, -0.2879, -2.0520],
        [ 0.0235,  0.4360,  0.1233,  0.0504,  0.5908],
        [ 0.5927,  0.1785, -0.9052, -0.9012,  0.8914],
        [ 0.4693,  0.5533, -0.1903,  0.0267,  0.4077],
        [-1.1683,  1.6699, -0.4846, -0.7404,  0.3370]])
标签: tensor([1, 1, 0, 1, 0, 1, 1, 0, 1, 1])
批次 3
数据: tensor([[ 0.2103, -0.7839,  1.4899,  2.2749, -0.7548],
        [-1.2836,  1.0025, -1.1162, -0.4261,  1.0690],
        [-0.7969,  1.0418, -0.7405,  0.8766,  0.2347],
        [-1.1071,  1.8560, -1.2979, -0.8364, -0.2925],
        [-1.0488,  0.4802, -0.6453,  0.2009,  0.5693],
        [ 0.8883,  0.4619, -0.2087,  0.2189, -0.3708],
        [-1.4578,  0.3629,  1.8282,  0.5353, -1.1783],
        [-1.2813,  0.5129, -0.4598, -0.2131, -1.2804],
        [ 1.7831,  1.1730, -0.2305, -0.6550,  0.1197],
        [-0.9384, -0.0483,  1.9626,  0.3342,  0.1700]])
标签: tensor([0, 0, 0, 1, 0, 1, 1, 1, 0, 1])

Using Built-in Datasets

PyTorch provides several commonly used datasets in torchvision, especially suitable for image tasks.

Load the MNIST dataset:

Example

import torchvision
import torchvision.transforms as transforms
from torch.utils.data import DataLoader

# Define data preprocessing
transform = transforms.Compose([
    transforms.ToTensor(),  # Convert to tensor
    transforms.Normalize((0.5,), (0.5,))  # Standardize
])

# Load training dataset
train_dataset = torchvision.datasets.MNIST(
    root='./data', train=True, transform=transform, download=True)

# Use DataLoader to load data
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)

# View a batch of data
data_iter = iter(train_loader)
images, labels = next(data_iter)
print(f"Batch image size: {images.shape}")  # Output shape is [batch_size, 1, 28, 28]
print(f"Batch labels: {labels}")

The output result is:

批次图像大小: torch.Size([32, 1, 28, 28])
批次标签: tensor([0, 4, 9, 8, 1, 3, 8, 1, 7, 2, 1, 1, 1, 2, 6, 3, 9, 7, 6, 9, 4, 9, 7, 1,
        3, 7, 3, 0, 7, 7, 6, 7])

Custom Application of Dataset and DataLoader

The following is an example of using a CSV file as a data source and reading data through a custom Dataset and DataLoader.

The CSV file content is as follows (downloadexample_pytorch_data.csv):

Example

import torch
import pandas as pd
from torch.utils.data import Dataset, DataLoader

# Custom CSV dataset
class CSVDataset(Dataset):
    def __init__(self, file_path):
        # Read CSV file
        self.data = pd.read_csv(file_path)

    def __len__(self):
        # Return dataset size
        return len(self.data)

    def __getitem__(self, idx):
        # Use .iloc to explicitly use position-based indexing
        row = self.data.iloc[idx]
        # Separate features and labels
        features = torch.tensor(row.iloc[:-1].to_numpy(), dtype=torch.float32)  # Features
        label = torch.tensor(row.iloc[-1], dtype=torch.float32)  # Labels
        return features, label

# Instantiate dataset and DataLoader
dataset = CSVDataset("example_pytorch_data.csv")
dataloader = DataLoader(dataset, batch_size=4, shuffle=True)

# Iterate over DataLoader
for features, label in dataloader:
    print("Features:", features)
    print("Labels:", label)
    break

The output result is:

特征: tensor([[ 1.2000,  2.1000, -3.0000],
        [ 1.0000,  1.1000, -2.0000],
        [ 0.5000, -1.2000,  3.3000],
        [-0.3000,  0.8000,  1.2000]])
标签: tensor([1., 0., 1., 0.])
tianqixin@Mac-mini example-test % python3 test.py
特征: tensor([[ 1.5000,  2.2000, -1.1000],
        [ 2.1000, -3.3000,  0.0000],
        [-2.3000,  0.4000,  0.7000],
        [-0.3000,  0.8000,  1.2000]])
标签: tensor([0., 1., 0., 0.])
Other Extensions