CIFAR-10图像分类实战:从零设计深度神经网络与SVM分类器

1次阅读
没有评论

共计 2356 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

CIFAR-10 数据集简介

CIFAR-10 是计算机视觉领域经典的基准数据集,包含 60000 张 32×32 像素的彩色图像,分为 10 个类别(飞机、汽车、鸟等),每个类别 6000 张。数据集已预先分为 50000 张训练图像和 10000 张测试图像。

CIFAR-10 图像分类实战:从零设计深度神经网络与 SVM 分类器

主要挑战包括:

  • 小尺寸图像导致特征提取困难
  • 类别间相似度高(如猫 / 狗、卡车 / 汽车)
  • 光照、视角等干扰因素

技术方案对比

1. CNN 方案实现

采用轻量化的 LeNet- 5 架构进行演示:

import torch
import torch.nn as nn

class LeNet5(nn.Module):
    def __init__(self):
        super(LeNet5, self).__init__()
        self.conv1 = nn.Conv2d(3, 6, 5)  # 3 输入通道,6 输出通道,5x5 卷积核
        self.pool = nn.MaxPool2d(2, 2)   # 2x2 最大池化
        self.conv2 = nn.Conv2d(6, 16, 5)
        self.fc1 = nn.Linear(16*5*5, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)

    def forward(self, x):
        x = self.pool(F.relu(self.conv1(x)))  # 卷积→ReLU→池化
        x = self.pool(F.relu(self.conv2(x)))
        x = x.view(-1, 16*5*5)  # 展平特征图
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        return x

关键组件说明:

  • 卷积层:局部感受野提取空间特征
  • ReLU 激活:引入非线性
  • 最大池化:降维并保持平移不变性
  • 全连接层:整合全局信息

2. SVM 方案实现

使用 HOG 特征 +SVM 的经典流程:

from skimage.feature import hog
from sklearn.svm import SVC

# 提取 HOG 特征
def extract_hog(images):
    features = []
    for img in images:
        fd = hog(img, orientations=9, pixels_per_cell=(8,8),
                cells_per_block=(2,2), channel_axis=-1)
        features.append(fd)
    return np.array(features)

# 训练 SVM
svm = SVC(C=1.0, kernel='rbf', gamma='scale')
svm.fit(train_features, train_labels)

完整代码实现

CNN 完整训练流程(PyTorch)

# 数据增强
transform_train = transforms.Compose([transforms.RandomHorizontalFlip(),
    transforms.RandomCrop(32, padding=4),
    transforms.ToTensor(),
    transforms.Normalize((0.5,0.5,0.5), (0.5,0.5,0.5))
])

# 学习率调度
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)

# 训练循环
for epoch in range(100):
    model.train()
    for inputs, labels in train_loader:
        optimizer.zero_grad()
        outputs = model(inputs)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
    scheduler.step()

SVM 参数调优(scikit-learn)

from sklearn.model_selection import GridSearchCV

params = {'C': [0.1, 1, 10],
    'gamma': [0.01, 0.1, 1]
}

grid = GridSearchCV(SVC(), params, cv=3)
grid.fit(train_features, train_labels)
print(f"最佳参数: {grid.best_params_}")

性能对比

指标 CNN(LeNet-5) SVM(RBF 内核)
测试准确率 72.3% 58.1%
训练时间 25min(GPU) 8min(CPU)
内存占用 1.2GB 650MB

常见问题解决方案

  1. 像素值归一化错误
  2. 错误做法:直接使用 0 -255 整数值
  3. 正确做法:标准化到 [-1,1] 或[0,1]范围

  4. 过拟合处理

  5. 添加 Dropout 层(p=0.5)
  6. 使用 L2 正则化(weight_decay=1e-4)
  7. 早停机制(patience=10)

  8. 类别不平衡

  9. 采样策略:过采样少数类 / 欠采样多数类
  10. 损失函数加权:nn.CrossEntropyLoss(weight=class_weights)

延伸思考

  1. 迁移到 CIFAR-100
  2. 使用更深的网络(如 ResNet-34)
  3. 添加中间分类头(分组类别)
  4. 引入知识蒸馏

  5. SVM 适用场景

  6. 计算资源受限环境
  7. 小规模数据集(<10k 样本)
  8. 需要模型可解释性的场合

参考文献

  • [CIFAR-10 官方说明]https://www.cs.toronto.edu/~kriz/cifar.html
  • [PyTorch 图像分类教程]https://pytorch.org/tutorials/beginner/blitz/cifar10_tutorial.html
  • [HOG 特征原论文]https://lear.inrialpes.fr/people/triggs/pubs/Dalal-cvpr05.pdf
正文完
 0
评论(没有评论)