YAML Metadata Warning:empty or missing yaml metadata in repo card
Check out the documentation for more information.
- End-to-End CNN Model for Image Classification
End-to-End CNN Model for Image Classification
A complete implementation of a Convolutional Neural Network (CNN) for image classification using PyTorch. This project includes data loading, model training, evaluation, and inference capabilities with comprehensive metrics tracking and visualization.
π Features
- Modular Architecture: Clean separation of data, models, training, and utilities
- CIFAR-10 Dataset: Built-in support for CIFAR-10 with automatic downloading
- Data Augmentation: Configurable augmentation pipeline (random crops, flips, color jitter)
- Training Infrastructure:
- Early stopping and learning rate scheduling
- Model checkpointing (best and latest)
- Comprehensive metrics tracking (accuracy, precision, recall, F1)
- Training history visualization
- Evaluation Tools:
- Confusion matrix generation
- Per-class metrics analysis
- Sample prediction visualization
- Inference: Easy-to-use script for making predictions on new images
π Project Structure
CNN1/
βββ config.py # Configuration and hyperparameters
βββ train.py # Main training script
βββ evaluate.py # Model evaluation script
βββ inference.py # Inference on new images
βββ requirements.txt # Python dependencies
βββ data/ # Data loading and preprocessing
β βββ __init__.py
β βββ dataset.py # Dataset loaders
β βββ transforms.py # Data transformations
βββ models/ # Model architectures
β βββ __init__.py
β βββ cnn.py # CNN model definition
β βββ utils.py # Model utilities
βββ utils/ # Training and evaluation utilities
β βββ __init__.py
β βββ trainer.py # Trainer class
β βββ metrics.py # Metrics tracking
β βββ visualization.py # Visualization functions
βββ checkpoints/ # Saved model checkpoints
βββ logs/ # Training logs
βββ results/ # Evaluation results and plots
π Installation
- Clone the repository (or navigate to the project directory):
cd c:\CNN\CNN1
- Install dependencies:
pip install -r requirements.txt
π Dataset
The project uses the CIFAR-10 dataset, which consists of:
- 60,000 32x32 color images in 10 classes
- 50,000 training images
- 10,000 test images
- Classes: airplane, automobile, bird, cat, deer, dog, frog, horse, ship, truck
The dataset will be automatically downloaded on first run.
π― Usage
Training
Train the model with default settings:
python train.py
Train with custom parameters:
python train.py --epochs 50 --batch-size 64 --lr 0.0001
Training outputs:
- Model checkpoints saved to
checkpoints/ - Training history saved to
results/training_history.json - Training plots saved to
results/training_history.png
Evaluation
Evaluate the trained model on the test set:
python evaluate.py
Evaluate a specific checkpoint:
python evaluate.py --checkpoint checkpoints/best_model.pth
Evaluation outputs:
- Metrics saved to
results/evaluation_metrics.json - Confusion matrix saved to
results/confusion_matrix.png - Per-class metrics plot saved to
results/per_class_metrics.png - Sample predictions saved to
results/sample_predictions.png
Inference
Make predictions on a new image:
python inference.py --image path/to/image.jpg --visualize
With custom parameters:
python inference.py --image path/to/image.jpg --checkpoint checkpoints/best_model.pth --top-k 3 --visualize --output results/prediction.png
βοΈ Configuration
All hyperparameters and settings can be modified in config.py:
Model Architecture
conv_channels: [32, 64, 128, 256] - Channels for each conv blockfc_hidden: 512 - Hidden units in FC layerdropout_rate: 0.5 - Dropout probability
Training Hyperparameters
NUM_EPOCHS: 100BATCH_SIZE: 128LEARNING_RATE: 0.001OPTIMIZER: 'Adam' (options: Adam, SGD, AdamW)SCHEDULER: 'ReduceLROnPlateau'
Data Augmentation
- Random crop with padding
- Random horizontal flip
- Color jitter (brightness, contrast, saturation, hue)
- Optional random rotation
Early Stopping
EARLY_STOPPING_PATIENCE: 15 epochsMIN_DELTA: 0.001
ποΈ Model Architecture
The CNN model consists of:
Convolutional Blocks (4 blocks):
- Each block: Conv2D β BatchNorm β ReLU β Conv2D β BatchNorm β ReLU β MaxPool
- Progressive channel increase: 3 β 32 β 64 β 128 β 256
Fully Connected Layers:
- FC1: feature_size β 512 (with dropout)
- FC2: 512 β 256 (with dropout)
- FC3: 256 β 10 (output)
Regularization:
- Batch normalization after each conv layer
- Dropout (0.5) in FC layers
- Weight decay (L2 regularization)
π Expected Results
With default configuration, you can expect:
- Test Accuracy: ~75-85% on CIFAR-10
- Training Time: ~2-3 hours on GPU, longer on CPU
- Model Size: ~5-10M parameters
π§ Advanced Usage
Custom Dataset
To use a custom dataset, modify data/dataset.py:
def get_custom_loaders():
# Implement your custom dataset loading
pass
Model Customization
Modify the model architecture in models/cnn.py:
model = CNN(
num_classes=10,
conv_channels=[64, 128, 256, 512], # Deeper network
fc_hidden=1024,
dropout_rate=0.3
)
Training Callbacks
The trainer supports:
- Early stopping
- Learning rate scheduling (ReduceLROnPlateau, StepLR, CosineAnnealing)
- Model checkpointing
- Metrics logging
π Metrics Tracked
During training and evaluation:
- Loss: Cross-entropy loss
- Accuracy: Overall classification accuracy
- Precision: Weighted average precision
- Recall: Weighted average recall
- F1 Score: Weighted average F1
- Per-class metrics: Individual metrics for each class
- Confusion Matrix: Detailed classification matrix
π Troubleshooting
CUDA Out of Memory
Reduce batch size in config.py:
BATCH_SIZE = 64 # or 32
Slow Training
- Reduce
NUM_WORKERSif CPU is bottleneck - Use GPU if available
- Reduce model size (fewer channels/layers)
Poor Accuracy
- Increase training epochs
- Adjust learning rate
- Enable more data augmentation
- Try different optimizer/scheduler
π Dependencies
- Python 3.7+
- PyTorch 2.0+
- torchvision
- numpy
- matplotlib
- scikit-learn
- tqdm
- Pillow
- seaborn
π€ Contributing
Feel free to:
- Report bugs
- Suggest features
- Submit pull requests
- Improve documentation
π License
This project is open source and available for educational purposes.
π Acknowledgments
- CIFAR-10 dataset by Alex Krizhevsky
- PyTorch framework
- Community contributions and feedback
Happy Training! π