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

如何在PyTorch中创建数据集并按规则导出CIFAR10训练/测试集?

Hey there! Let's tackle your two PyTorch dataset questions one by one—they're both common use cases, so I'll break down practical, actionable steps with code examples.


Question 1: Creating a PyTorch Dataset and Saving It for Later Use

First, you'll want to wrap your data in a custom Dataset class (inheriting from torch.utils.data.Dataset) to make it compatible with PyTorch's DataLoader. Once you have that, saving it depends on your dataset size and needs. Here are the most common approaches:

1. Save the entire Dataset object (simple for small/medium data)

If your dataset is picklable (most are, as long as you don't hold non-picklable objects like open file handles), you can directly save it with torch.save:

import torch
from torch.utils.data import Dataset

class CustomDataset(Dataset):
    def __init__(self, data, labels):
        self.data = data
        self.labels = labels
    
    def __len__(self):
        return len(self.data)
    
    def __getitem__(self, idx):
        return self.data[idx], self.labels[idx]

# Example data (replace with your actual data)
data = torch.randn(100, 3, 32, 32)  # 100 RGB images of 32x32
labels = torch.randint(0, 10, (100,))  # 100 class labels

# Create and save the dataset
my_dataset = CustomDataset(data, labels)
torch.save(my_dataset, 'my_custom_dataset.pt')

# Load it later
loaded_dataset = torch.load('my_custom_dataset.pt')

2. Save data/labels separately (better for large datasets)

For bigger datasets, saving the underlying tensors/arrays individually is more efficient. You can then recreate the dataset when loading:

# Save components separately
torch.save(my_dataset.data, 'dataset_data.pt')
torch.save(my_dataset.labels, 'dataset_labels.pt')

# Load and rebuild the dataset
loaded_data = torch.load('dataset_data.pt')
loaded_labels = torch.load('dataset_labels.pt')
loaded_dataset = CustomDataset(loaded_data, loaded_labels)

3. Use HDF5 for out-of-memory datasets

If your data is too large to fit in RAM, use h5py to save it to disk and create a dataset that loads samples on-the-fly:

import h5py

# Save to HDF5
with h5py.File('large_dataset.h5', 'w') as f:
    f.create_dataset('data', data=my_dataset.data.numpy())
    f.create_dataset('labels', data=my_dataset.labels.numpy())

# Custom Dataset for HDF5 loading
class HDF5Dataset(Dataset):
    def __init__(self, file_path):
        self.file_path = file_path
        with h5py.File(file_path, 'r') as f:
            self.length = len(f['labels'])
    
    def __len__(self):
        return self.length
    
    def __getitem__(self, idx):
        with h5py.File(self.file_path, 'r') as f:
            data = torch.tensor(f['data'][idx])
            label = torch.tensor(f['labels'][idx])
        return data, label

# Load later
large_dataset = HDF5Dataset('large_dataset.h5')

Question 2: Extracting CIFAR10 Data in Custom Order and Saving Train/Test Sets

Your initial thought of using a non-shuffled DataLoader to track indices is valid, but we can optimize this by computing your custom score first, sorting indices based on that score, then creating sorted subsets. Let's use your pixel sum example to walk through this:

Step 1: Load the original CIFAR10 dataset

First, load CIFAR10 (skip download=True if you already have it):

from torchvision.datasets import CIFAR10
from torch.utils.data import DataLoader, Subset
import torchvision.transforms as transforms

cifar10_train = CIFAR10(root='./data', train=True, download=False)
cifar10_test = CIFAR10(root='./data', train=False, download=False)

Step 2: Define and compute your custom score function

Let's implement f(Image, label)—here, we'll calculate the sum of all pixels in the image:

def compute_pixel_sum_score(image, label):
    # Convert PIL image to tensor (values 0-255)
    img_tensor = transforms.ToTensor()(image) * 255
    # Sum all pixels across channels, height, and width
    return img_tensor.sum().item()

# Compute scores for every sample in train and test sets
train_scores = [compute_pixel_sum_score(img, lbl) for img, lbl in cifar10_train]
test_scores = [compute_pixel_sum_score(img, lbl) for img, lbl in cifar10_test]

Step 3: Sort indices based on the scores

Decide if you want ascending or descending order, then get the sorted indices:

# Sort training indices by score (ascending; use reverse=True for descending)
sorted_train_indices = sorted(range(len(train_scores)), key=lambda i: train_scores[i])
# Sort test indices the same way
sorted_test_indices = sorted(range(len(test_scores)), key=lambda i: test_scores[i])

Step 4: Create sorted subsets

Use PyTorch's Subset class to wrap the original dataset with your sorted indices:

sorted_train_dataset = Subset(cifar10_train, sorted_train_indices)
sorted_test_dataset = Subset(cifar10_test, sorted_test_indices)

Step 5: Save for later use

You have two options here:

Option A: Save the entire Subset objects

This is straightforward, but it saves a reference to the original dataset (so keep your CIFAR10 data in the same location):

torch.save(sorted_train_dataset, 'sorted_cifar10_train.pt')
torch.save(sorted_test_dataset, 'sorted_cifar10_test.pt')

# Load later
loaded_sorted_train = torch.load('sorted_cifar10_train.pt')
loaded_sorted_test = torch.load('sorted_cifar10_test.pt')

Option B: Save just the sorted indices (lightweight)

This is better if you want to reuse the original dataset later (e.g., with different transforms):

torch.save(sorted_train_indices, 'sorted_train_indices.pt')
torch.save(sorted_test_indices, 'sorted_test_indices.pt')

# Load later and recreate subsets
loaded_train_indices = torch.load('sorted_train_indices.pt')
loaded_test_indices = torch.load('sorted_test_indices.pt')

# Recreate the original CIFAR10 dataset (add transforms if needed)
cifar10_train = CIFAR10(root='./data', train=True, download=False)
cifar10_test = CIFAR10(root='./data', train=False, download=False)

sorted_train_dataset = Subset(cifar10_train, loaded_train_indices)
sorted_test_dataset = Subset(cifar10_test, loaded_test_indices)

Step 6: Use with DataLoader

Now you can use these sorted datasets with DataLoader like any other dataset:

# No shuffle if you want to keep the custom order; set shuffle=True for training shuffling
train_loader = DataLoader(sorted_train_dataset, batch_size=32, shuffle=False)
test_loader = DataLoader(sorted_test_dataset, batch_size=32, shuffle=False)

If you want to shuffle during training but retain the custom order for initial loading, just set shuffle=True—this will shuffle the sorted batches each epoch without altering the underlying sorted dataset.


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 08:16:57