题目:卷积神经网络图像分类系统设计与实现——以CIFAR-10/猫狗分类为例

技术栈:Python · PyTorch · Flask · Vue3

核心工作


三、核心实验:数据增强对过拟合的抑制效果
3.1 实验设计(对照实验)
变量实验组(Aug)对照组(No-Aug)
随机水平翻转
随机裁剪(padding=4)
随机旋转(15°)
颜色抖动
归一化✓(相同)✓(相同)

控制变量:相同模型结构、相同训练轮次(50 epoch)、相同优化器参数(SGD, lr=0.1, momentum=0.9)

3.2 关键代码实现
# 数据增强策略定义
def get_transforms(use_augment=True):
    if use_augment:
        train_transform = transforms.Compose([
            transforms.RandomHorizontalFlip(p=0.5),
            transforms.RandomRotation(15),
            transforms.RandomCrop(32, padding=4),
            transforms.ColorJitter(brightness=0.2, contrast=0.2),
            transforms.ToTensor(),
            transforms.Normalize((0.4914, 0.4822, 0.4465), 
                               (0.2023, 0.1994, 0.2010)),
        ])
    else:
        train_transform = transforms.Compose([
            transforms.ToTensor(),
            transforms.Normalize((0.4914, 0.4822, 0.4465), 
                               (0.2023, 0.1994, 0.2010)),
        ])
    return train_transform, test_transform
3.3 实验结果
指标无数据增强有数据增强提升
训练集准确率94.2%91.8%-2.4%
验证集准确率78.5%85.3%+6.8%
过拟合差距15.7%6.5%-9.2%
最终测试准确率76.8%84.6%+7.8%

关键发现

3.4 可视化分析

training_curves.png

左图:对照组(蓝线训练loss持续下降,橙线验证loss反弹);右图:实验组双曲线同步收敛


四、工程实现要点
4.1 模型结构适配

CIFAR-32×32图像远小于ImageNet-224×224,直接使用标准ResNet-18会导致:

model = resnet18(num_classes=10)
model.conv1 = nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1, bias=False)
model.maxpool = nn.Identity()  # 恒等映射替代下采样
4.2 Mac MPS加速踩坑

问题:PyTorch 2.0+自动检测MPS后端,但DataLoader多进程报错 解决:设置num_workers=0(MPS暂不支持多进程数据加载)

trainloader = DataLoader(trainset, batch_size=128, 
                        shuffle=True, num_workers=0)  # MPS必须单进程
4.3 跨域与部署

Flask后端配置CORS:

from flask_cors import CORS
app = Flask(__name__)
CORS(app)  # 开发环境全开放,生产环境需配置白名单

Vue3前端代理配置(vite.config.js):

server: {
  proxy: {
    '/api': {
      target: 'http://localhost:5001',
      changeOrigin: true,
      rewrite: (path) => path.replace(/^\/api/, '')
    }
  }
}

五、项目结构
cnn_project/
├── models/
│   └── resnet18_best.pth          # 最佳模型权重(85.3%验证准确率)
├── data/
│   └── cifar-10-batches-py/     # 数据集(自动下载)
├── train.py                       # 训练脚本(含对照实验)
├── app.py                         # Flask推理服务
├── frontend/                      # Vue3前端
│   ├── src/
│   │   └── App.vue               # 主页面(拖拽上传+结果可视化)
│   └── vite.config.js            # 代理配置
└── README.md                      # 详细使用文档

六、快速开始

完整代码https://github.com/yourname/cnn-cifar10-classification

  • 实现ResNet-18/VGG网络,在CIFAR-10上达到85%+准确率

  • 设计数据增强策略(旋转、裁剪、归一化),分析过拟合与正则化效果

  • 开发Web交互系统:前端上传图片,后端调用模型推理返回结果

  • cnn_project/                    # PyCharm 打开此文件夹
    ├── .idea/                      # PyCharm 配置(自动生成)
    ├── venv/                       # Python 虚拟环境
    ├── models/
    │   └── resnet18_best.pth       # 训练好的模型
    ├── data/                       # CIFAR-10 数据集(自动下载)
    ├── train.py                    # 训练脚本
    ├── app.py                      # Flask 后端
    └── frontend/                   # WebStorm 打开此文件夹
        ├── .idea/                  # WebStorm 配置(自动生成)
        ├── node_modules/
        ├── src/
        │   └── App.vue             # 主页面
        ├── vite.config.ts          # 代理配置
        └── package.json

    核心目标:建立"数据→模型→训练→评估→部署"的完整认知闭环。


    二、技术架构
    ┌─────────────┐     ┌─────────────┐     ┌─────────────┐
    │   数据采集   │ --> │   模型训练   │ --> │   服务部署   │
    │  CIFAR-10   │     │  ResNet-18  │     │ Flask API   │
    └─────────────┘     └─────────────┘     └─────────────┘
                                                  │
                                                  v
                                           ┌─────────────┐
                                           │   前端交互   │
                                           │   Vue 3     │
                                           └─────────────┘

    技术栈

  • 深度学习:Python 3.12 · PyTorch 2.1 · torchvision

  • 模型架构:ResNet-18(适配修改)

  • 后端服务:Flask · Flask-CORS · Gunicorn(生产)

  • 前端交互:Vue 3 · Vite · Axios

  • 开发环境:PyCharm · WebStorm · Mac M3(MPS加速)

  • 数据增强降低了训练集准确率(更难拟合),但显著提升验证集表现

  • 对照组出现典型过拟合:训练loss持续下降,验证loss在epoch 30后反弹

  • 实验组loss曲线同步收敛,泛化性能稳定

  • 首层卷积(kernel=7, stride=2)+ MaxPool后,特征图仅剩4×4

  • 解决方案:修改首层为kernel=3, stride=1,移除MaxPool


  •  

更多推荐