深度学习入门-基于CNN的图像分类系统
题目:卷积神经网络图像分类系统设计与实现——以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 可视化分析
左图:对照组(蓝线训练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
更多推荐


所有评论(0)