LDMAX is an educational, modular, and high-performance research codebase for training and sampling from Diffusion Transformers (DiT) from scratch in the modern JAX ecosystem.
LDMAX is built on a Clean Architecture with strict separation of concerns across dataset pipelines, model architectures, training orchestration, and evaluation services:
flowchart TD
Config["YAML Experiment Configs\n(configs/*.yaml)"] --> CLI["CLI Dispatchers\n(scripts/train.py, scripts/sample.py)"]
CLI --> Trainer["Unified Trainer Engine\n(src/training/trainer.py)"]
subgraph Data Layer
Trainer --> DataFactory["Data Factory\n(src/data/factory.py)"]
DataFactory --> Loaders["Grain / Hugging Face Loaders\n(CIFAR-10, Fashion-MNIST, CelebA)"]
end
subgraph Model Layer
Trainer --> ModelFactory["Model Factory\n(src/models/factory.py)"]
ModelFactory --> DiTModel["Diffusion Transformer (DiT)\n(Flax NNX)"]
end
subgraph Services & Infrastructure
Trainer --> CkptService["Checkpoint Service\n(src/training/checkpointing.py)"]
Trainer --> Evaluator["Evaluation & Visualizer\n(src/training/evaluator.py)"]
Evaluator --> VAE["VAEManager\n(Latent Space Decoding)"]
Evaluator --> PixelNorm["Pixel Unnormalizer\n(Raw Pixel Space)"]
end
- Core Framework: JAX (
jax,jax.numpy), JIT compilation (jax.jit). - Model Definition: Flax NNX (
flax.nnx), the reference-based object-oriented API for Flax. - Optimization: Optax (
optax) with AdamW and EMA weight tracking. - Checkpointing: Orbax (
orbax.checkpoint) with async saving and GCS synchronization. - Data Loading: Google Grain (
grain.python) for high-throughput sharded pipelines (CIFAR-10, CelebA), alongside Hugging Facedatasets/ NumPy in-memory iterators (Fashion-MNIST). - Precision: Mixed-precision support with
bfloat16and explicit FP32 parameter and schedule management for TPU/GPU execution.
βββ app/ # Full-stack web application
β βββ frontend/ # Static web app (GitHub Pages: ghif.github.io/ldmax)
β βββ backend/ # FastAPI backend server (Google Cloud Run)
βββ configs/ # YAML experiment configurations
β βββ celeba*.yaml # CelebA latent-space configurations (256x256 -> 32x32 latents)
β βββ cifar10*.yaml # CIFAR-10 latent & native-pixel configurations
β βββ fashion_mnist*.yaml# Fashion-MNIST raw-pixel configurations (28x28 grayscale)
βββ docs/ # Specs, design documents, and research notes
βββ scripts/ # Thin CLI entry points
β βββ train.py # Unified training CLI
β βββ sample.py # Unified standalone sampling CLI
β βββ train_cifar10.py # CIFAR-10 training launcher
β βββ train_fashion_mnist.py # Fashion-MNIST training launcher
β βββ train_celeba.py # CelebA training launcher
β βββ demo.py # Unified multi-dataset tabbed Gradio demo
βββ src/ # Core library code
β βββ data/ # Dataset sources and factory
β β βββ celeba.py # CelebA Grain pipeline
β β βββ cifar.py # CIFAR-10 Grain pipeline
β β βββ fashion_mnist.py # Fashion-MNIST pipeline
β β βββ factory.py # Unified DataLoaderBundle & metadata factory
β βββ models/ # Model architectures and factory
β β βββ dit/ # DiT & AdaLN-Zero blocks
β β βββ factory.py # Unified create_model factory
β βββ sampling/ # Sampling utilities & offline generation
β β βββ generator.py # Unified standalone image generator
β βββ training/ # Training infrastructure & services
β β βββ checkpointing.py # State validation & Orbax restore services
β β βββ ema.py # Exponential Moving Average manager
β β βββ evaluator.py # Sampling evaluator & grid visualizer
β β βββ sampler.py # DDIM sampling engine
β β βββ step.py # JIT training step & MSE loss computation
β β βββ trainer.py # Unified config-driven Trainer engine
β βββ utils/ # Utilities (checkpoint, config, logging, RNG, VAE)
βββ tests/ # Unit & integration tests
β βββ unit/ # Unit tests for models, factories, checkpointing, trainer
β βββ integration/ # End-to-end integration tests
βββ pyproject.toml # Project configuration and linter settings
# 1. Create and activate conda environment
conda create -n ldmax python=3.11 -y
conda activate ldmax
# 2. Install dependencies for your target hardware:
# CPU:
pip install -r requirements_cpu.txt
# GPU (CUDA):
# pip install -r requirements_gpu.txt
# TPU:
# pip install -r requirements_tpu.txtRun training across any dataset using the unified entry point:
# Unified Training CLI
PYTHONPATH=. python scripts/train.py \
--config configs/cifar10_pixel.yaml \
--output_dir outputs/cifar10_run
# Fashion-MNIST Raw-Pixel Diffusion (28x28 Grayscale)
PYTHONPATH=. python scripts/train_fashion_mnist.py \
--config configs/fashion_mnist.yaml \
--output_dir outputs/fashion_mnist_run
# CIFAR-10 Native-Pixel Diffusion (32x32 RGB)
PYTHONPATH=. python scripts/train_cifar10.py \
--config configs/cifar10_pixel.yaml \
--output_dir outputs/cifar10_pixel_run
# CelebA Latent Diffusion (256x256 -> 32x32 Latents with VAE)
PYTHONPATH=. python scripts/train_celeba.py \
--config configs/celeba.yaml \
--output_dir outputs/celeba_runTo resume training seamlessly from an existing run or specific checkpoint step:
PYTHONPATH=. python scripts/train.py \
--config configs/cifar10_pixel.yaml \
--resume_from outputs/cifar10_run \
--output_dir outputs/cifar10_resumedGenerate visual sample grids from a trained checkpoint or EMA weights:
# Unified Sampling CLI
PYTHONPATH=. python scripts/sample.py \
--config configs/cifar10_pixel.yaml \
--checkpoint outputs/cifar10_pixel_run/checkpoints/5000 \
--num_samples 16 \
--class_id 3 \
--output_path samples/cifar10_class3.png
# Attribute-Conditioned CelebA Sampling
PYTHONPATH=. python scripts/sample.py \
--config configs/celeba.yaml \
--checkpoint outputs/celeba_run/checkpoints/50000 \
--num_samples 16 \
--attribute_names "Smiling,Eyeglasses" \
--output_path samples/celeba_custom.pngLaunch TensorBoard to monitor live loss curves and generated sample grids:
tensorboard --logdir outputsLaunch the interactive Gradio browser demo to generate and blend classes or facial attributes across datasets in dedicated tabs:
# Launch unified multi-dataset tabbed demo (CIFAR-10, Fashion-MNIST, CelebA)
PYTHONPATH=. python scripts/demo.py \
--cifar10-config configs/cifar10_pixel.yaml \
--fashion-config configs/fashion_mnist_tpu_v4.yaml \
--celeba-config configs/celeba.yaml \
--celeba-checkpoint gs://diffjax/models/celeba_ldm_ccond_tpu-v6e-1_18-08-2026/checkpoints/270000 \
--port 7860Run the production full-stack web application located in app/:
# 1. Start FastAPI backend (Cloud Run target)
PYTHONPATH=. uvicorn app.backend.main:app --host 127.0.0.1 --port 8000 --reload
# 2. Serve static frontend (GitHub Pages target: ghif.github.io/ldmax)
python3 -m http.server 3000 --directory app/frontendOpen http://localhost:3000 in your browser.
LDMAX includes a comprehensive test suite verifying factories, checkpoint serialization, state management, model outputs, and end-to-end training runs:
# Run unit and integration tests
JAX_PLATFORMS=cpu PYTHONPATH=. pytest tests/
# Run code linter
ruff check .
# Run code formatter check
ruff format --check .- Scalable Diffusion Models with Transformers (DiT) (Peebles & Xie, 2023)
- Flax NNX (Documentation)
- Google Grain (Repository)
- Orbax Checkpoint (Documentation)