Skip to content

Latest commit

 

History

4 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 

Repository files navigation

CNN 图像分类项目

项目简介

这是一个使用 PyTorch 实现的卷积神经网络(CNN)图像分类项目。模型基于 CIFAR-10 数据集进行训练,能够对 10 类图像(如飞机、汽车、鸟类等)进行准确分类。

该项目的功能包括:

  • 数据预处理(数据增强、归一化)
  • CNN 模型构建(3 个卷积层 + 全连接层)
  • 模型训练和评估
  • 使用国内镜像下载 CIFAR-10 数据集

项目目录结构


项目依赖

在运行此项目之前,请确保安装以下依赖库:

  • Python >= 3.7
  • torch >= 1.11.0
  • torchvision >= 0.12.0
  • numpy

您可以通过以下命令快速安装所需依赖:

pip install torch torchvision numpy

数据集

本项目使用的是 CIFAR-10 数据集,包含 10 个类别的 60,000 张 32x32 彩色图像:

  • 训练集:50,000 张
  • 测试集:10,000 张

项目中采用了国内镜像源(百度云镜像)来下载 CIFAR-10 数据集,以加快下载速度。


如何运行项目

1. 克隆仓库

git clone https://github.com/YourUsername/YourRepoName.git
cd YourRepoName

2. 运行主程序

确保安装依赖后,运行以下命令开始训练和评估模型:

python main.py

3. 修改参数

您可以通过修改 main.py 文件中的以下内容调整训练参数:

  • 学习率lr
  • 批量大小batch_size
  • 训练轮次epochs

模型架构

本项目的 CNN 模型包含以下部分:

卷积层块 1:

  • 卷积层:输入通道数为 3,输出通道数为 32,核大小为 3x3
  • Batch Normalization:批量规范化
  • ReLU 激活
  • 最大池化层:2x2,步幅 2

卷积层块 2:

  • 类似卷积层块 1,输出通道数为 64

卷积层块 3:

  • 类似卷积层块 2,输出通道数为 128

全连接层:

  1. 扁平化:将 128 个特征图(4x4)拉平成一维向量
  2. 全连接层 1:输入 2048,输出 512,Dropout = 0.5
  3. 全连接层 2:输出 10(对应 CIFAR-10 的 10 类)

文件说明

1. main.py

  • 负责数据加载、训练、评估流程的主脚本。
  • 数据增强使用了随机水平翻转和随机裁剪。

2. model.py

  • 定义 CNN 模型架构,包含 3 个卷积层块和 1 个全连接层块。

结果展示

在训练 50 个 epoch 后,模型可以达到约 75%-80% 的测试集准确率(具体取决于超参数设置和训练环境)。

作者信息

欢迎提交 issue 或 PR 来改进本项目!

About

No description, website, or topics provided.

Resources

Stars

1 star

Watchers

1 watching

Forks

Releases

Packages

Contributors

Languages