Skip to content

Latest commit

 

History

History

Folders and files

NameName
Last commit message
Last commit date

parent directory

..
 
 
 
 
 
 
 
 
 
 
 
 
 
 

README.md

猫狗分类项目代码说明

项目简介

基于 PyTorch 和 ResNet-101 的猫狗图像分类项目,使用迁移学习进行二分类任务。

文件说明

核心文件

  • train.py: 模型训练脚本,包含训练、验证、评估和可视化功能
  • test.py: 测试集预测脚本,对测试集进行预测并生成可视化结果
  • dog_rename.py: 对训练集中狗的图片进行重命名,避免编号重复引起歧义和冲突,仅在开始时运行一次即可
  • dataset.py: 自定义数据集类 DogCat,用于加载和处理图像数据
  • models.py: ResNet 模型定义,包含 ResNet-18/34/50/101/152 的实现

依赖文件

  • requirements.txt: Python 依赖包列表

使用方法

1. 安装依赖

pip install -r requirements.txt

2. 训练模型

cd code
python train.py --num_workers 0 --nepoch 10

主要参数:

  • --num_workers: 数据加载进程数(Windows 建议设为 0)
  • --nepoch: 训练轮数
  • --batchSize: 批次大小(默认 64)
  • --lr: 学习率(默认 0.001)
  • --early_stop: 启用早停(默认启用)
  • --patience: 早停耐心值(默认 2)

3. 测试集预测

cd code
python test.py --num_workers 0 --num_samples 8

主要参数:

  • --model_path: 模型权重路径(默认 ../results/best_model.pth
  • --num_samples: 可视化样本数量(默认 8)
  • --test_dir: 测试集目录(默认 ../data/test1

输出结果

训练完成后,结果保存在 ../results/ 目录:

  • best_model.pth: 最佳模型权重
  • logs/tensorboard/train/: TensorBoard 训练日志
  • logs/test_predictions.csv: 测试集预测结果
  • figures/test_predictions.png: 测试集预测可视化拼图

注意事项

  1. 数据路径:代码中使用相对路径 ../data/,请确保数据目录结构正确
  2. Windows 系统:建议将 --num_workers 设置为 0 以避免多进程问题
  3. GPU 支持:代码会自动检测 CUDA,如无 GPU 会自动使用 CPU