PyTorch torchvision Computer Vision Module
torchvision is an extension library in the PyTorch ecosystem specifically for computer vision tasks. It provides the following core features:
- Pre-trained models: Includes classic CNN architecture implementations (such as ResNet, VGG, AlexNet, etc.)
- Dataset tools: Built-in common visual datasets (such as CIFAR10, MNIST, ImageNet, etc.)
- Image transforms: Provides various image preprocessing and data augmentation methods
- Utility tools: Includes auxiliary functions such as video processing, image operations, etc.
# 安装 torchvision(通常与 PyTorch 一起安装) pip install torch torchvision
Core Components Analysis
1. torchvision.models
Provides pre-trained computer vision models that can be directly used for transfer learning:
Example
import torchvision.models as models
# Load pre-trained model
resnet18 = models.resnet18(pretrained=True)
alexnet = models.alexnet(pretrained=True)
vgg16 = models.vgg16(pretrained=True)
# Load pre-trained model
resnet18 = models.resnet18(pretrained=True)
alexnet = models.alexnet(pretrained=True)
vgg16 = models.vgg16(pretrained=True)
Common model list:
| Model Name | Applicable Scenario | Parameter Count | Top-1 Accuracy |
|---|---|---|---|
| ResNet | General image classification | 11M-60M | 69%-80% |
| VGG | Feature extraction | 138M | 71.3% |
| MobileNet | Mobile application | 3.4M | 70.6% |
| EfficientNet | Efficient model | 5M-66M | 77%-84% |
2. torchvision.datasets
Built-in common computer vision datasets to simplify the data loading process:
Example
from torchvision import datasets
# Load CIFAR10 dataset
train_data = datasets.CIFAR10(
root='data',
train=True,
download=True,
transform=transforms.ToTensor()
)
# Load MNIST dataset
test_data = datasets.MNIST(
root='data',
train=False,
download=True
)
# Load CIFAR10 dataset
train_data = datasets.CIFAR10(
root='data',
train=True,
download=True,
transform=transforms.ToTensor()
)
# Load MNIST dataset
test_data = datasets.MNIST(
root='data',
train=False,
download=True
)
Supported dataset types:
Example
graph TD
A[torchvision.datasets] --> B[Classification datasets]
A --> C[Detection datasets]
A --> D[Segmentation datasets]
B --> B1[CIFAR10/100]
B --> B2[MNIST/FashionMNIST]
B --> B3[ImageNet]
C --> C1[COCO]
C --> C2[VOC]
D --> D1[Cityscapes]
A[torchvision.datasets] --> B[Classification datasets]
A --> C[Detection datasets]
A --> D[Segmentation datasets]
B --> B1[CIFAR10/100]
B --> B2[MNIST/FashionMNIST]
B --> B3[ImageNet]
C --> C1[COCO]
C --> C2[VOC]
D --> D1[Cityscapes]
3. torchvision.transforms
Core tools for image preprocessing and data augmentation:
Example
from torchvision import transforms
# Define image transform pipeline
transform = transforms.Compose([
transforms.Resize(256), # Resize
transforms.CenterCrop(224), # Center crop
transforms.ToTensor(), # Convert to tensor
transforms.Normalize( # Normalize
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]
)
])
# Define image transform pipeline
transform = transforms.Compose([
transforms.Resize(256), # Resize
transforms.CenterCrop(224), # Center crop
transforms.ToTensor(), # Convert to tensor
transforms.Normalize( # Normalize
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]
)
])
Classification of common transform methods:
| Category | Example Methods | Purpose |
|---|---|---|
| Geometric transforms | RandomRotation, RandomResizedCrop | Increase position invariance |
| Color transforms | ColorJitter, Grayscale | Enhance color robustness |
| Blur/Noise | GaussianBlur, RandomErasing | Prevent overfitting |
| Combined transforms | RandomApply, RandomChoice | Flexible combination strategy |
Practical Example: Image Classification Pipeline
1. Data Preparation
Example
import torch
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
# Define data transforms
train_transform = transforms.Compose([
transforms.RandomHorizontalFlip(),
transforms.RandomRotation(10),
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
# Load dataset
train_set = datasets.CIFAR10(
root='./data',
train=True,
download=True,
transform=train_transform
)
# Create data loader
train_loader = DataLoader(
train_set,
batch_size=32,
shuffle=True
)
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
# Define data transforms
train_transform = transforms.Compose([
transforms.RandomHorizontalFlip(),
transforms.RandomRotation(10),
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
# Load dataset
train_set = datasets.CIFAR10(
root='./data',
train=True,
download=True,
transform=train_transform
)
# Create data loader
train_loader = DataLoader(
train_set,
batch_size=32,
shuffle=True
)
2. Model Training
Example
import torch.nn as nn
import torch.optim as optim
# Use pre-trained model
model = models.resnet18(pretrained=True)
# Modify the last layer (adapt to CIFAR10's 10 classes)
num_ftrs = model.fc.in_features
model.fc = nn.Linear(num_ftrs, 10)
# Define loss function and optimizer
criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(model.parameters(), lr=0.001, momentum=0.9)
# Training loop
for epoch in range(10):
for images, labels in train_loader:
outputs = model(images)
loss = criterion(outputs, labels)
optimizer.zero_grad()
loss.backward()
optimizer.step()
import torch.optim as optim
# Use pre-trained model
model = models.resnet18(pretrained=True)
# Modify the last layer (adapt to CIFAR10's 10 classes)
num_ftrs = model.fc.in_features
model.fc = nn.Linear(num_ftrs, 10)
# Define loss function and optimizer
criterion = nn.CrossEntropyLoss()
optimizer = optim.SGD(model.parameters(), lr=0.001, momentum=0.9)
# Training loop
for epoch in range(10):
for images, labels in train_loader:
outputs = model(images)
loss = criterion(outputs, labels)
optimizer.zero_grad()
loss.backward()
optimizer.step()
Advanced Features
1. Custom Dataset
Example
from torchvision.datasets import VisionDataset
class CustomDataset(VisionDataset):
def __init__(self, root, transform=None):
super().__init__(root, transform=transform)
# Implement __getitem__ and __len__
def __getitem__(self, index):
# Return (image, target) tuple
pass
def __len__(self):
# Return dataset size
pass
class CustomDataset(VisionDataset):
def __init__(self, root, transform=None):
super().__init__(root, transform=transform)
# Implement __getitem__ and __len__
def __getitem__(self, index):
# Return (image, target) tuple
pass
def __len__(self):
# Return dataset size
pass
2. Model Export and Deployment
Example
# Export to ONNX format
dummy_input = torch.randn(1, 3, 224, 224)
torch.onnx.export(
model,
dummy_input,
"model.onnx",
input_names=["input"],
output_names=["output"]
)
dummy_input = torch.randn(1, 3, 224, 224)
torch.onnx.export(
model,
dummy_input,
"model.onnx",
input_names=["input"],
output_names=["output"]
)
Best Practice Recommendations
1. Data augmentation strategy:
- Use random transforms to augment data during training
- Use only deterministic transforms during validation/testing
2. Transfer learning tips:
Example
# Freeze all parameters except the last layer
for param in model.parameters():
param.requires_grad = False
model.fc.requires_grad = True
for param in model.parameters():
param.requires_grad = False
model.fc.requires_grad = True
3. Performance optimization:
- Use
num_workersparameter to accelerate data loading - For large datasets, consider using
Datasetsubset of
4. Common mistakes:
- Forgetting to call
zero_grad() - Confused
train()andeval()mode - Image tensor shape does not match model requirements (should be C×H×W)