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!
- Features
- Model Architecture
- Project Structure
- Getting Started
- Quick, Draw! Generation Details
- References
- License
- 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.
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]
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
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 --releaseTo train on your GPU (Wgpu backend):
cargo run --release -- --gpu--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 (noisefor DDPM/DDIM,velocityfor Flow Matching).
(Note: These flags can be combined, for example: cargo run --release -- --epochs 40 --lr 0.0003 --prediction-type velocity --gpu)
| 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 --gpuTo train on the Google Quick, Draw! dataset (25 classes):
cargo run --release -- --dataset quickdraw --gpuDataset Cache: Dataset files are downloaded once and cached at
target/emnist_dataset/andtarget/quickdraw_dataset/.
Run the mathematical tests verifying the forward scheduling, time embeddings, and U-Net blocks:
cargo testGenerate 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!
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).
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 between1and128).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".
The models compile to WebAssembly for fully client-side generation. Model weights (*-model.bin) are downloaded on-demand in the browser and cached locally.
Make sure you have trained the models first, then:
-
Install wasm-pack:
cargo install wasm-pack
-
Build the WebAssembly module:
cargo run --bin build_web
-
Install a local static file server:
cargo install basic-http-server
-
Serve locally:
basic-http-server docs
Navigate to http://localhost:4000.
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:
-
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 -
Commit and push the updated version files:
git add docs/weights-version.txt git commit -m "Update model weights version" git push origin master
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:
16or32
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.
This project supports doodle generation using a subset of the Quick, Draw! Dataset.
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
- Burn β Deep Learning Framework for Rust
- The Burn Book
- wasm-pack β Rust WebAssembly Packager
- Rust Drawing Classifier Web (Original Fork Base)
- Denoising Diffusion Probabilistic Models (DDPM)
- Denoising Diffusion Implicit Models (DDIM)
- Classifier-Free Diffusion Guidance (Ho & Salimans, 2022)
- Flow Matching for Generative Modeling (Lipman et al., 2022)
- EMNIST Dataset (NIST Special Database 19)
- Google Quick, Draw! Dataset
This project is licensed under the MIT License.
