Skip to content

Latest commit

Β 

History

130 Commits

Folders and files

NameName
Last commit message
Last commit date
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 

Repository files navigation

πŸ”’ MNIST, EMNIST & Quick, Draw! Classifier β€” Burn (Rust)

An interactive handwritten digit, letter, and doodle classifier built with the Burn deep learning framework in Rust. Train a CNN model, run inference from the CLI, or draw in the browser!

πŸš€ Try the Live WebAssembly Demo!

image


πŸ“‘ Table of Contents


✨ Features

  • MobileNet-style CNN Architecture β€” Depthwise Separable Convolutions, 1x1 projection and identity residual shortcuts, BatchNorm, GAP, and Dropout
  • Interactive Web Demo β€” Draw on a canvas and get real-time predictions
  • WebAssembly Client-Side Inference β€” Runs entirely in the browser via WASM, no backend required
  • CLI Inference β€” Predict with ASCII art visualization
  • Fully in Rust β€” Training, inference, and web frontend in a unified workspace
  • Data Augmentations β€” Random spatial translation, scale/zoom shifts, and horizontal flips (QuickDraw) applied dynamically during batch collation for robust canvas prediction

πŸ—οΈ Model Architecture

Input [1Γ—28Γ—28]
  β†’ Stem: Conv2d(1β†’24, 3Γ—3, stride=1) β†’ BatchNorm β†’ ReLU                 β†’ [24Γ—28Γ—28]
  β†’ Block 1: SeparableConv(24β†’48, stride=2) + Proj Shortcut(24β†’48)
             β†’ BatchNorm β†’ ReLU                                          β†’ [48Γ—14Γ—14]
  β†’ Block 2: SeparableConv(48β†’96, stride=2) + Proj Shortcut(48β†’96)
             β†’ BatchNorm β†’ ReLU                                          β†’ [96Γ—7Γ—7]
  β†’ Block 3: SeparableConv(96β†’96, stride=1) + Identity Shortcut
             β†’ BatchNorm β†’ ReLU                                          β†’ [96Γ—7Γ—7]
  β†’ Classifier:
      β†’ Global Average Pooling (GAP)                                     β†’ [96Γ—1Γ—1]
      β†’ Flatten                                                          β†’ [96]
      β†’ Dropout(0.5)
      β†’ Linear(96β†’num_classes)

πŸ“ Project Structure

burn-drawing-classifier/
β”œβ”€β”€ model_shared/           # Shared library workspace crate
β”‚   └── src/lib.rs          # CNN model definition & LayerNorm layers
β”œβ”€β”€ web/                    # Rust WASM crate (wasm-pack entry point)
β”œβ”€β”€ src/                    # Training & CLI inference (Burn backend)
β”‚   β”œβ”€β”€ main.rs
β”‚   β”œβ”€β”€ model.rs            # Re-exports shared model definition
β”‚   β”œβ”€β”€ training.rs
β”‚   └── ...
β”œβ”€β”€ docs/                   # Static web frontend (served by GitHub Pages)
β”‚   β”œβ”€β”€ index.html          # Single-page drawing app with Developer Console
β”‚   └── pkg/                # Compiled WASM output (gitignored, built by CI)
β”œβ”€β”€ assets/                 # README images and training curves
β”œβ”€β”€ build.rs                # Copies model weights at build time
β”œβ”€β”€ publish-weights.ps1     # Helper script to upload weights to GitHub Releases
└── .github/workflows/
    └── deploy.yml          # CI: build WASM β†’ assert β†’ deploy β†’ verify

Git Branches:

  • master β€” The only branch you need. All development happens here. Compiled binaries are gitignored and built fresh by CI on every deploy.

πŸš€ Getting Started

Train the Model

By default, training runs on the CPU (NdArray backend):

cargo run --release

To train on your GPU (Wgpu backend):

cargo run --release -- --gpu

To train on the EMNIST Letters dataset (26 classes, downloads dynamically):

# CPU
cargo run --release -- --dataset emnist

# GPU
cargo run --release -- --dataset emnist --gpu

To train on the Google Quick, Draw! dataset (25 classes, downloads dynamically):

# CPU
cargo run --release -- --dataset quickdraw

# GPU
cargo run --release -- --dataset quickdraw --gpu

Dataset Cache: Dataset files are downloaded once and cached at target/emnist_dataset/ and target/quickdraw_dataset/. If you change the configurations, delete the caches first so they re-download:

rm -rf target/quickdraw_dataset   # Linux / macOS
Remove-Item -Recurse -Force target\quickdraw_dataset  # Windows PowerShell

Note: Always use --release for optimized tensor math performance.

πŸ“Š Results

After 5 epochs of training:

Dataset Validation Accuracy Validation Loss
MNIST (10 classes) ~98%+ ~0.05
EMNIST Letters (26 classes) ~90%+ ~0.30
Quick, Draw! (25 classes) ~86%+ ~0.51
πŸ“ˆ View MNIST Training Progress Curve

MNIST Training Curve


Run Tests

cargo test

CLI Inference

Once trained, predict from the MNIST test set:

cargo run --release -- --predict

Predict from the EMNIST Letters test set:

cargo run --release -- --predict --dataset emnist

Predict from the Quick, Draw! test set:

cargo run --release -- --predict --dataset quickdraw
πŸ“ Example Output (MNIST)
Loading model for inference...

Input Image:
      ######                
      ################      
      ################      
           ###########      
                  ####      
                 ####       
                 ####       
                ####        
                ####        
               ####         
               ###          
              ####          
             ####           
            #####           
            ####            
           #####            
           ####             
          #####             
          #####             
          ####              
                            
Target Label (Ground Truth): 7
Top Predictions:
  1. 7            : 99.42%
  2. 9            : 0.35%
  3. 2            : 0.11%
πŸ“ Example Output (EMNIST Letters)
Loading model for inference (dataset: emnist)...
Loading and parsing EMNIST Letters test data...

Input Image:






           ...
         .#####
      ..######.
     ..#######.
    .########....
   .######...####.
   .#####...#####.
   ###############
  .###############
   .##############.
   .###############.
    ..####..   .####..
      ....     .######..
       ..       ...####..
                   ......








Target Label (Ground Truth): A
Top Predictions:
  1. A            : 88.57%
  2. R            : 3.80%
  3. D            : 1.84%
πŸ“ Example Output (Quick, Draw!)
Loading model for inference (dataset: quickdraw)...

Input Image:
         #          
        ###         
       #####        
      #######       
     #########      
    ###########     
   #############    
  ###############   
 #################  

Target Label (Ground Truth): triangle
Top Predictions:
  1. triangle     : 96.81%
  2. mountain     : 2.14%
  3. house        : 0.45%

Interactive Web Server (Axum backend)

Start the browser-based drawing pad backed by the Rust Axum server:

  • MNIST Digits (default):

    cargo run --release -- --serve

    Open http://127.0.0.1:3000 to draw digits (0–9).

  • EMNIST Letters:

    cargo run --release -- --serve --dataset emnist

    Open http://127.0.0.1:3000 to draw letters (A–Z).

  • Quick, Draw! Doodles:

    cargo run --release -- --serve --dataset quickdraw

    Open http://127.0.0.1:3000 to draw and predict doodles (25 classes).


Client-Side WebAssembly App (WASM)

The trained models compile to WebAssembly for fully client-side inference. Model weights are decoupled from Git history to avoid binary bloat:

  • Locally: The weights (*-model.bin) are copied to the docs/ folder (which is ignored by Git).
  • In CI: The GitHub Actions deployment workflow automatically downloads the weights from GitHub Releases during build.

1. Build the WASM bundle locally

Make sure you have trained the models first, then:

  1. Install wasm-pack:

    cargo install wasm-pack
  2. Build the WebAssembly module:

    • Using the helper binary (Recommended, cross-platform):
      cargo run --bin build_web
    • Or run the command manually (does not automatically clean up duplicate/redundant folders):
      wasm-pack build web --target web --out-dir ../docs/pkg
  3. Install a local static file server (needed for the preview):

    cargo install basic-http-server
  4. Serve locally:

    basic-http-server docs

    Navigate to http://localhost:4000.

2. Automatic Deployments & Release Management

The CI workflow (.github/workflows/deploy.yml) automatically builds and deploys to GitHub Pages on every push to master. It also verifies the live WASM URL after deploying.

To update the model weights used by the CI runner:

  1. Ensure the GitHub CLI (gh) is installed and authenticated:

    • Install (Windows): winget install --id GitHub.cli (restart VS Code after)
    • Authenticate: gh auth login β†’ GitHub.com β†’ HTTPS β†’ browser
  2. Upload your local weights to a GitHub Release:

    # Default v1.0.0
    cargo run --bin publish_weights
    
    # Custom version tag
    cargo run --bin publish_weights -- v2.0.0

    The script now updates web/weights-version.txt and docs/weights-version.txt automatically, so you no longer need to edit those files by hand.

  3. Commit and push the updated version files along with any code changes:

    git add web/weights-version.txt docs/weights-version.txt
    git commit -m "Update model weights version"
    git push origin master
  4. Trigger the deployment:

    • Code changes: git push origin master
    • Weights only: Go to Actions β†’ Deploy WebAssembly to GitHub Pages β†’ Run workflow
  5. Verify your repository settings under Settings β†’ Pages β†’ Build and deployment:

    • Source: GitHub Actions

    ℹ️ The workflow uses the official GitHub Pages API (upload-pages-artifact + deploy-pages), so no gh-pages branch is needed.


🎨 Quick, Draw! Classification Details

This project supports doodle classification using the public Google Quick, Draw! Dataset.

Selected Categories (25 classes)

Rather than training on all 345 categories (39 GB of raw data), we train on a curated subset of 25 diverse and sketchable classes:

Group Classes
Nature / Weather sun, moon, star, tree, flower
Animals cat, dog, fish, butterfly
Common Objects cup, key, umbrella, hat, clock, envelope, toothbrush
Structures / Vehicles house, car
Shapes circle, triangle, square, smiley face
Clothing pants, t-shirt
Food apple

Key Design Considerations

  1. Compute & Storage: 25 classes β†’ ~250 MB footprint, fits in memory, trains in minutes on a consumer GPU.
  2. Model Capacity: A simple CNN easily achieves high accuracy on 25 distinct shapes vs. struggling with all 345.
  3. Canvas Drawability: These shapes are simple and iconic enough to draw clearly on a 28Γ—28 canvas.
  4. License & Privacy: CC BY 4.0, no personally identifiable information (PII).

πŸ“š References


πŸ“„ License

This project is licensed under the MIT License.

About

An interactive, browser-based drawing classifier using MobileNet-style CNNs. Built with the Burn framework in Rust, featuring training and client-side inference via WebAssembly. πŸ”’πŸ¦€

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages