Skip to content

About

An interactive, browser-based drawing generator (DDPM/DDIM, and Flow Matching). Built with the Burn framework in Rust, featuring client-side inference via WebAssembly. πŸŽ¨πŸ¦€

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Latest commit

Β 

History

197 Commits

Folders and files

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

Repository files navigation

🎨 Rust Drawing Generator β€” Burn (Rust)

An interactive drawing generator utilizing a Denoising Diffusion Probabilistic Model (DDPM/DDIM) built with the Burn deep learning framework in Rust. Generate handwritten digits, letters, or doodles right in your browser!

πŸš€ Try the Live WebAssembly Demo!

image


πŸ“‘ Table of Contents


✨ Features

  • Conditional U-Net Architecture β€” Sinusoidal time embedding module, class conditioning embedding module, self-attention modules, skip connections, and residual blocks.
  • DDPM & Flow Matching (Rectified Flow) Support β€” Dual support for both standard noise-prediction (DDPM/DDIM) and straight-line velocity-prediction (Flow Matching) generative paradigms.
  • DDIM/Euler & Heun Scheduler Denoising β€” Accelerated reverse sampling supporting both DDIM/Euler (1st-Order) and Heun (2nd-Order) samplers, with customizable Polynomial Spacing Schedules (exponents 0.25–7.0) to achieve superior drawing quality in as few as 5–10 steps.
  • Classifier-Free Guidance (CFG) with Dynamic Linear Decay β€” Extrapolates conditional predictions from an unconditional baseline to enhance semantic shape accuracy, utilizing a dynamic decay schedule to eliminate late-stage drawing artifacts.
  • Interactive Educational Tooltips β€” Sleek inline SVG information toggles that explain the parameter spaces on click or hover to guide user selections.
  • Cosine Annealing Learning Rate Scheduler β€” Smoothly decays the learning rate to zero during training to optimize final output convergence.
  • Progressive Denoising Animation β€” Visual progressive rendering showing the drawing emerge frame-by-frame from pure Gaussian noise.
  • Dual Inference Modes β€” Server-side Axum streaming via Server-Sent Events (SSE) and client-side WebAssembly local browser execution.
  • Fully in Rust β€” Training, scheduler, inference engine, and web frontend in a unified workspace.

πŸ—οΈ Model Architecture

The generator relies on a lightweight Denoising Diffusion Probabilistic Model (DDPM/DDIM) with a conditional U-Net structure:

Inputs: Latent State x_t [BΓ—1Γ—28Γ—28], Timestep t [B], Class ID c [B]
  β†’ Time Embedding: Sinusoidal Positional Encoding (48) + MLP (48β†’192)
  β†’ Class Embedding: Embedding Lookup (48) + Linear Projection (48β†’192)
  β†’ Merged Embedding: Addition of Time & Class Embeddings [BΓ—192]
  β†’ U-Net Encoder:
      β†’ Stem: Conv2d(1β†’48 channels) [BΓ—48Γ—28Γ—28]
      β†’ Down Block 1: UNetBlock + Time/Class Injection [BΓ—48Γ—28Γ—28]
      β†’ Downsample 1: Conv2d(48β†’96 channels, Stride 2) [BΓ—96Γ—14Γ—14]
      β†’ Down Block 2: UNetBlock + Time/Class Injection [BΓ—96Γ—14Γ—14]
      β†’ Downsample 2: Conv2d(96β†’192 channels, Stride 2) [BΓ—192Γ—7Γ—7]
  β†’ Bottleneck Middle Block: UNetBlock + Time/Class Injection + Self-Attention [BΓ—192Γ—7Γ—7]
  β†’ U-Net Decoder:
      β†’ Upsample 1: ConvTranspose2d(192β†’96 channels, Stride 2) [BΓ—96Γ—14Γ—14]
      β†’ Up Block 1: Concatenate(Upsample 1, Skip 2) β†’ UNetBlock(192β†’96 channels) + Self-Attention [BΓ—96Γ—14Γ—14]
      β†’ Upsample 2: ConvTranspose2d(96β†’48 channels, Stride 2) [BΓ—48Γ—28Γ—28]
      β†’ Up Block 2: Concatenate(Upsample 2, Skip 1) β†’ UNetBlock(96β†’48 channels) [BΓ—48Γ—28Γ—28]
      β†’ Output Layer: Conv2d(48β†’1 channels) [BΓ—1Γ—28Γ—28]

πŸ“ Project Structure

burn-drawing-generator/
β”œβ”€β”€ model_shared/           # Shared library workspace crate
β”‚   β”œβ”€β”€ src/lib.rs          # Model architecture & DDIMScheduler definition
β”‚   β”œβ”€β”€ src/unet.rs         # Conditional U-Net blocks and modules
β”‚   └── src/scheduler.rs    # DDIM forward process and reverse sampling math
β”œβ”€β”€ web/                    # Rust WASM crate (wasm-pack entry point)
β”‚   └── src/lib.rs          # Stateful GeneratorWasm wrapper exposing .step()
β”œβ”€β”€ src/                    # Training, CLI generation, & serving (Burn backend)
β”‚   β”œβ”€β”€ main.rs             # CLI router & Axum API serve endpoint
β”‚   β”œβ”€β”€ model.rs            # Re-exports shared Model wrapper
β”‚   β”œβ”€β”€ training.rs         # Autodiff training loop and MSE loss wrapper
β”‚   β”œβ”€β”€ inference.rs        # Iterative DDIM/Heun progressive sampling & ASCII art
β”‚   β”œβ”€β”€ data.rs             # Noise collator & normalized dataset batcher
β”‚   β”œβ”€β”€ emnist.rs           # EMNIST classes mapping helper
β”‚   β”œβ”€β”€ quickdraw.rs        # QuickDraw dataset loading and class mappings
β”‚   └── bin/                # Executable runner scripts
β”‚       β”œβ”€β”€ build_web.rs    # Script to build, optimize, and bundle WASM
β”‚       β”œβ”€β”€ convert.rs      # Utility to convert trained weights to binary format
β”‚       └── publish_weights.rs # Utility to bundle and publish weights to release targets
β”œβ”€β”€ docs/                   # Static web frontend (served by GitHub Pages)
β”‚   β”œβ”€β”€ index.html          # Drawing generator UI with Developer Console
β”‚   └── pkg/                # Compiled WASM output (gitignored, built by CI)
└── assets/                 # README showcase demo assets

πŸš€ Getting Started

Train the Model

Warning

Migration Warning: If you are migrating from an older checkpoint model (without Classifier-Free Guidance support), you must delete the existing checkpoint directory (e.g., ./target/mnist-model/) before starting training. Otherwise, the model loading step will crash due to class embedding size shape mismatch (num_classes + 1 rows).

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

cargo run --release

To train on your GPU (Wgpu backend):

cargo run --release -- --gpu

Training Configuration Flags:

  • --epochs <num>: Customize the number of epochs to train (defaults to 5).
  • --lr <num>: Customize the peak learning rate (defaults to 2e-4, decayed to 0 via Cosine Annealing).
  • --prediction-type <noise|velocity>: Choose the target type (noise for DDPM/DDIM, velocity for Flow Matching).

(Note: These flags can be combined, for example: cargo run --release -- --epochs 40 --lr 0.0003 --prediction-type velocity --gpu)


πŸ’‘ Hyperparameter Recommendations & Training Guides

Objective Paradigm Recommended Command
Quick Run (~10 Mins) DDPM (Noise) cargo run --release -- --epochs 6 --lr 2.5e-4 --prediction-type noise --gpu
Quick Run (~10 Mins) Flow Matching cargo run --release -- --epochs 6 --lr 8e-4 --prediction-type velocity --gpu
Fully Converged DDPM (Noise) cargo run --release -- --epochs 25 --lr 2e-4 --prediction-type noise --gpu
Fully Converged Flow Matching cargo run --release -- --epochs 30 --lr 5e-4 --prediction-type velocity --gpu

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

cargo run --release -- --dataset emnist --gpu

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

cargo run --release -- --dataset quickdraw --gpu

Dataset Cache: Dataset files are downloaded once and cached at target/emnist_dataset/ and target/quickdraw_dataset/.


Run Tests

Run the mathematical tests verifying the forward scheduling, time embeddings, and U-Net blocks:

cargo test

CLI Generation

Generate drawings from random Gaussian noise and watch the progressive ASCII rendering directly in your terminal:

# Generate MNIST digits
cargo run --release -- --predict

# Generate EMNIST Letters
cargo run --release -- --predict --dataset emnist

# Generate Quick, Draw! doodles
cargo run --release -- --predict --dataset quickdraw
πŸ“ Example Terminal Output (MNIST digit 4 progress)
Loading model for generation (dataset: mnist)...
Generating drawing for class: '4' (class ID: 4) using 20 DDIM steps...

Generated Output:
                            
                            
               ..           
               ##.          
              .###          
              ####          
             #####          
            .#####          
            #####           
           .####.           
           #####            
          ######            
         #######            
         ####.##            
        ##### ##            
       .#### .##            
       ####  .##            
      ######.###.           
      ##########.           
      ##########            
          .  ##             
             ##.            
             ##.            
             ##.            
                            
                            

Generation complete!

Interactive Web Server (Axum backend)

Start the browser-based generator UI backed by the Rust Axum server streaming denoising steps via Server-Sent Events (SSE):

  • MNIST Digits (default):

    cargo run --release -- --serve

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

  • EMNIST Letters:

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

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

  • Quick, Draw! Doodles:

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

    Open http://127.0.0.1:3000 to generate doodles (25 classes).

GET /api/generate SSE Stream

The Axum server exposes a GET endpoint /api/generate that streams the progressive denoising steps using Server-Sent Events (SSE):

  • Query Parameters:
    • class_id (integer, required): The target class ID to generate.
    • steps (integer, optional, default: 16): Number of denoising steps (clamped between 1 and 128).
    • schedule (string, optional, default: "linear"): Denoising schedule type or power exponent (e.g. "linear", "quadratic", or a custom float string like "3.0", "7.0").
    • sampler (string, optional, default: "ddim"): Denoising sampler algorithm: "ddim" or "heun".

Client-Side WebAssembly App (WASM)

The models compile to WebAssembly for fully client-side generation. Model weights (*-model.bin) are downloaded on-demand in the browser and cached locally.

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:

    cargo run --bin build_web
  3. Install a local static file server:

    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.

To update the model weights used by the CI runner:

  1. Upload your local weights to a GitHub Release:

    cargo run --bin publish_weights
    # or with custom tag
    cargo run --bin publish_weights -- v2.0.0
  2. Commit and push the updated version files:

    git add docs/weights-version.txt
    git commit -m "Update model weights version"
    git push origin master

πŸ’‘ Recommended Parameters Setup

To get the highest quality digit and doodle drawings at the best inference speeds, we recommend using the following parameters:

  • Classifier-Free Guidance (CFG): On (3.0) (Amplifies class details during early stages)
  • Denoising Schedule (Power): Early Focus (0.5) (Spends 70% of the step budget on global shape layouts)
  • Sampler Algorithm: Euler / DDIM (1st-Order) (Fastest inference mode, requiring only 2 forward passes per step)
  • Denoising Steps: 16 or 32

🧠 Dynamic Classifier-Free Guidance Decay

This workspace implements a Linear CFG Decay Schedule in both the backend server (src/inference.rs) and local browser WebAssembly crate (web/src/lib.rs): $$\text{Scale}{\text{step}} = 1.0 + (\text{Scale}{\text{initial}} - 1.0) \times (1.0 - \text{progress})$$

  • At early steps (noise to outline): CFG scale is at its maximum (e.g. 3.0), strongly forcing the model to establish correct digit and category structures.
  • At late steps (smoothing & sharpening): CFG decays to 1.0 (conditional-only), completely eliminating high-frequency stroke distortions and artifacts commonly caused by constant guidance.

🎨 Quick, Draw! Generation Details

This project supports doodle generation using a subset of the Quick, Draw! Dataset.

Selected Categories (25 classes)

We train on a curated subset of 25 doodle 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

πŸ“š References


πŸ“„ License

This project is licensed under the MIT License.

About

An interactive, browser-based drawing generator (DDPM/DDIM, and Flow Matching). Built with the Burn framework in Rust, featuring client-side inference via WebAssembly. πŸŽ¨πŸ¦€

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages