This repository contains a high-performance implementation of a Convolutional Neural Network (CNN) using PyTorch to solve the handwritten digit classification problem (MNIST) from the Kaggle Digit Recognizer competition.
The project focuses on building an automated system capable of recognizing handwritten digits (0-9). As a Computer and Control Engineer, the implementation emphasizes not just accuracy, but also the stability of the training process and the robustness of the model architecture.
The model was designed using a Sequential structure with a focus on feature extraction and regularization:
- Feature Extraction: Utilized
Conv2dlayers to capture spatial hierarchies in the images. - Stability: Integrated
BatchNorm2dlayers after convolutions to stabilize the hidden state distributions and accelerate convergence. - Non-Linearity: Used
ReLUactivation functions to enable the model to learn complex patterns. - Dimensionality Reduction: Employed
MaxPool2dfor spatial downsampling while retaining essential features. - Generalization (Regularization): Applied
Dropoutlayers to mitigate overfitting, ensuring the model performs well on unseen test data.
The project leverages the full power of the PyTorch ecosystem and Python data science stack:
- Core Framework:
torch&torch.nnfor model building. - Data Pipeline:
DataLoaderandTensorDatasetfor efficient batching and memory management. - Preprocessing:
pandasandsklearnfor data splitting and normalization. - Visualization:
seabornandmatplotlibfor analyzing the training curves and error distribution.
In line with control engineering principles, the training process was treated as an optimization problem:
- Optimizer: Used the Adam optimizer for its adaptive learning rate capabilities.
- Loss Function: CrossEntropyLoss was chosen as the objective function for multi-class classification.
- Dynamic Feedback: Implemented a Learning Rate Scheduler (
ReduceLROnPlateau). This mimics a closed-loop system where the "system" (the model) monitors the validation loss and automatically reduces the learning rate when the improvement stalls (plateaus), ensuring fine-tuned convergence.
- Data Source: Kaggle Digit Recognizer.
- Model Weights: The final trained state of the model is saved using
torch.save(model.state_dict(), 'model_weights.pth'), allowing for easy deployment or further fine-tuning. - Evaluation: Detailed analysis was performed using a
confusion_matrixto identify specific digit-class confusion (e.g., distinguishing between 4 and 9).
- Engineering Approach: The project doesn't just provide code; it explains the Optimization Strategy (Adam + ReduceLROnPlateau) and why these specific control mechanisms were chosen.
- Professional Documentation: By including a
requirements.txt(torch, pandas, seaborn, etc.), the project follows professional software development standards. - Data Ethics: The README explicitly mentions that large CSV data files are ignored (via
.gitignore) and provides links to the official source, respecting storage best practices. - Visual Insights: Emphasis is placed on the Confusion Matrix to demonstrate a deep understanding of model behavior beyond just "total accuracy" percentages.