Skip to content

Latest commit

 

History

2 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

PyTorch CNN Models

使用 PyTorch 从零搭建并训练经典卷积神经网络的学习项目,包含 LeNet、AlexNet、VGG16、GoogLeNet、ResNet,以及面向自定义水果图片数据集的 GoogLeNet 分类器。

项目内容

目录 模型 数据集 输入尺寸 类别数
LeNet/ LeNet FashionMNIST 1 × 28 × 28 10
ALexNet/ AlexNet(适配小尺寸灰度图) FashionMNIST 1 × 28 × 28 10
VGG/ VGG16(适配 FashionMNIST) FashionMNIST 1 × 64 × 64 10
GoogLeNet/ GoogLeNet / Inception FashionMNIST 1 × 224 × 224 10
ResNet/ ResNet(残差网络) FashionMNIST 1 × 28 × 28 10
Fruit_GoogLeNet/ GoogLeNet 自定义水果图片 3 × 224 × 224 由数据集自动确定

每个模型目录主要包含:

  • model.py:网络结构定义;
  • model_train.py:数据加载、训练、验证、早停及权重保存;
  • model_test.py:加载权重并评估测试集;
  • data_partitioning.py:仅水果分类项目使用,用于划分训练集与测试集。

环境要求

  • Python 3.9+
  • PyTorch 2.0+

建议使用虚拟环境:

git clone https://github.com/lbw-work/pytorch-cnn-models.git
cd pytorch-cnn-models

python -m venv .venv
source .venv/bin/activate  # Windows: .venv\Scripts\activate
python -m pip install --upgrade pip
pip install -r requirements.txt

训练脚本会自动选择可用设备,优先级为 CUDA、Apple Silicon MPS、CPU(具体顺序以各脚本实现为准)。

FashionMNIST 模型

FashionMNIST 数据会在首次训练或测试时自动下载到对应模型目录的 data/ 中。请进入目标目录再运行脚本,例如:

cd LeNet

# 查看模型结构
python model.py

# 训练并生成 best_model.pth
python model_train.py

# 使用已生成的权重评估测试集
python model_test.py

其他 FashionMNIST 模型使用相同方式:

cd ALexNet   # 或 VGG、GoogLeNet、ResNet
python model.py
python model_train.py
python model_test.py

LeNet/plot.py 可用于查看一个批次的 FashionMNIST 样本:

cd LeNet
python plot.py

水果分类 GoogLeNet

将原始图片按类别文件夹放入 Fruit_GoogLeNet/data/:

Fruit_GoogLeNet/data/
├── apple/
│   ├── image_001.jpg
│   └── image_002.jpg
├── banana/
│   └── image_001.jpg
└── orange/
    └── image_001.jpg

然后依次执行:

cd Fruit_GoogLeNet

# 按类别将数据划分到 data/train 和 data/test
python data_partitioning.py

# 训练模型,类别数会根据文件夹自动确定
python model_train.py

# 评估测试集
python model_test.py

若要对单张图片推理,可将图片放在 Fruit_GoogLeNet/ 根目录后运行 python model_test.py。脚本会优先检测项目目录中的图片;未发现图片时自动评估测试集。

data_partitioning.py 会移动原始图片并重新组织目录。请先备份唯一的数据副本。

数据与模型权重

数据集、训练权重和生成图片未提交到仓库:

  • FashionMNIST 可由 torchvision 自动下载;
  • 自定义水果数据需要自行准备;
  • *.pth、*.pt、*.ckpt 等权重文件需要运行训练脚本生成。

这些内容通常体积较大,并可能受数据来源或模型分发许可约束,因此统一由 .gitignore 排除。

说明

本项目用于理解经典 CNN 架构和完整的训练、验证、测试流程。部分结构针对 FashionMNIST 的单通道、小尺寸输入做了适配,并非论文原始配置的逐参数复现。

About

Classic CNN implementations and training workflows built with PyTorch

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages