AI百事通

从零开始:用PyTorch实现一个轻量级图像分类模型

📅 2026-07-27📰 ai_generated👁 1 次阅读
从零开始:用PyTorch实现一个轻量级图像分类模型
PyTorch图像分类深度学习卷积神经网络教程

前言

图像分类是计算机视觉的基础任务。本教程将使用PyTorch框架,在CIFAR-10数据集上训练一个轻量级卷积神经网络(CNN),模型参数量小于1M,适合在CPU上快速运行。

环境准备

pip install torch torchvision matplotlib

步骤一:数据加载

使用torchvision下载CIFAR-10,并进行数据增强。

import torch
import torchvision
import torchvision.transforms as transforms

transform_train = transforms.Compose([
    transforms.RandomHorizontalFlip(),
    transforms.RandomCrop(32, padding=4),
    transforms.ToTensor(),
    transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010))
])

trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform_train)
trainloader = torch.utils.data.DataLoader(trainset, batch_size=128, shuffle=True, num_workers=2)

步骤二:定义模型

构建一个简单的CNN:

import torch.nn as nn
import torch.nn.functional as F

class LightCNN(nn.Module):
    def __init__(self):
        super(LightCNN, self).__init__()
        self.conv1 = nn.Conv2d(3, 32, 3, padding=1)
        self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
        self.pool = nn.MaxPool2d(2, 2)
        self.fc1 = nn.Linear(64 * 8 * 8, 256)
        self.fc2 = nn.Linear(256, 10)
        self.dropout = nn.Dropout(0.25)

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = x.view(-1, 64 * 8 * 8)
        x = F.relu(self.fc1(x))
        x = self.dropout(x)
        x = self.fc2(x)
        return x

net = LightCNN()

步骤三:训练模型

定义损失函数和优化器,训练10个epoch:

import torch.optim as optim

criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(net.parameters(), lr=0.001)

for epoch in range(10):
    running_loss = 0.0
    for i, data in enumerate(trainloader, 0):
        inputs, labels = data
        optimizer.zero_grad()
        outputs = net(inputs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
        running_loss += loss.item()
        if i % 100 == 99:
            print(f'[Epoch {epoch+1}, Batch {i+1}] loss: {running_loss/100:.3f}')
            running_loss = 0.0
print('Finished Training')

步骤四:评估模型

在测试集上计算准确率:

testset = torchvision.datasets.CIFAR10(root='./data', train=False, download=True, transform=transforms.ToTensor())
testloader = torch.utils.data.DataLoader(testset, batch_size=100, shuffle=False)

correct = 0
total = 0
with torch.no_grad():
    for data in testloader:
        images, labels = data
        outputs = net(images)
        _, predicted = torch.max(outputs.data, 1)
        total += labels.size(0)
        correct += (predicted == labels).sum().item()

print(f'Accuracy of the network on the 10000 test images: {100 * correct / total}%')

预期准确率约75%。

进阶优化

  • 使用学习率调度器(如StepLR)
  • 添加Batch Normalization
  • 尝试更深的网络(如ResNet-18)

总结

本教程展示了用PyTorch实现图像分类的基本流程。你可以将模型保存并部署到移动设备或Web端。代码已上传至GitHub,欢迎Star。