如何在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.
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')
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

