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.main in Generator class)
  • Discriminator layers (self.main in Discriminator class)
  • Latent dimension size (latent_dim parameter)

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

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. πŸ™‹ Ask for provider support