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:

  1. Pre-trained models: Includes classic CNN architecture implementations (such as ResNet, VGG, AlexNet, etc.)
  2. Dataset tools: Built-in common visual datasets (such as CIFAR10, MNIST, ImageNet, etc.)
  3. Image transforms: Provides various image preprocessing and data augmentation methods
  4. 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)

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
)

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]

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]
    )
])

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
)

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()

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

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"]
)

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

3. Performance optimization:

  • Usenum_workersparameter to accelerate data loading
  • For large datasets, consider usingDatasetsubset of

4. Common mistakes:

  • Forgetting to callzero_grad()
  • Confusedtrain()andeval()mode
  • Image tensor shape does not match model requirements (should be C×H×W)
Other Extensions