YAML Metadata Warning:empty or missing yaml metadata in repo card
Check out the documentation for more information.
MNIST GAN - Handwritten Digit Generator
An end-to-end Deep Convolutional GAN (DCGAN) implementation for generating handwritten digits using PyTorch, trained on the MNIST dataset and deployed with Flask.
π Features
- DCGAN Architecture: Deep Convolutional Generative Adversarial Network optimized for MNIST
- Training Pipeline: Complete training script with checkpointing and visualization
- Web Interface: Flask-based web application for easy digit generation
- Interactive Generation: Command-line interface for quick digit generation
- API Endpoints: RESTful API for programmatic access
π Project Structure
gan/
βββ app.py # Flask web application
βββ train.py # GAN training script
βββ generate.py # Inference script
βββ gan_model.py # GAN model architecture
βββ requirements.txt # Python dependencies
βββ README.md # This file
βββ checkpoints/ # Saved model checkpoints
β βββ generator_final.pth
βββ samples/ # Training sample outputs
β βββ samples_epoch_*.png
βββ generated/ # Generated digit outputs
βββ templates/ # HTML templates
β βββ index.html
β βββ error.html
βββ static/ # Static assets
βββ css/
β βββ style.css
βββ js/
βββ app.js
π Quick Start
1. Install Dependencies
pip install -r requirements.txt
2. Train the Model
python train.py --epochs 100 --batch-size 128
Training parameters:
--epochs: Number of training epochs (default: 100)--batch-size: Batch size (default: 128)--lr: Learning rate (default: 0.0002)--latent-dim: Latent space dimension (default: 100)
3. Run the Web Application
python app.py
Then open http://localhost:5000 in your browser.
4. Generate Digits (CLI)
# Generate 16 digits and save as grid
python generate.py --num-samples 16 --save-grid --show
# Generate individual images
python generate.py --num-samples 64 --save-individual --output-dir ./my_digits
# Interactive mode
python generate.py
π§ Model Architecture
Generator
- Input: Random noise vector (100 dimensions)
- Architecture: Transposed convolutions with BatchNorm and ReLU
- Output: 28Γ28 grayscale image with tanh activation
Discriminator
- Input: 28Γ28 grayscale image
- Architecture: Convolutions with BatchNorm and LeakyReLU
- Output: Probability (real/fake classification)
π API Endpoints
| Endpoint | Method | Description |
|---|---|---|
/ |
GET | Main web interface |
/generate?num=N |
GET | Generate N digits (1-64) |
/generate-single |
GET | Generate single digit |
/api/status |
GET | Check model status |
/health |
GET | Health check |
Example API Usage
# Generate 9 digits
curl "http://localhost:5000/generate?num=9"
# Generate single digit
curl "http://localhost:5000/generate-single"
# Check status
curl "http://localhost:5000/api/status"
π Training Visualization
The training script saves:
- Sample images every 10 epochs (
samples/samples_epoch_*.png) - Checkpoints every 10 epochs (
checkpoints/gan_checkpoint_epoch_*.pth) - Final generator model (
checkpoints/generator_final.pth) - Training history plot (
checkpoints/training_history.png)
π οΈ Customization
Modify Model Architecture
Edit gan_model.py to change:
- Generator layers (
self.mainin Generator class) - Discriminator layers (
self.mainin Discriminator class) - Latent dimension size (
latent_dimparameter)
Training Hyperparameters
Edit train.py or use command-line arguments:
python train.py --epochs 200 --batch-size 64 --lr 0.0001
π¦ Dependencies
- torch >= 1.9.0
- torchvision >= 0.10.0
- flask >= 2.0.0
- matplotlib >= 3.4.0
- numpy >= 1.20.0
- Pillow >= 8.0.0
π― Expected Results
After 50-100 epochs, the generator should produce:
- Clear, recognizable handwritten digits (0-9)
- Varied handwriting styles
- Consistent image quality
π Notes
- Training on GPU is recommended for faster results
- MNIST dataset will be automatically downloaded
- Model files are saved in the
checkpoints/directory - Web application loads the pre-trained generator automatically
π License
MIT License
Inference Providers NEW
This model isn't deployed by any Inference Provider. π Ask for provider support