This project is partially supported by Google TPU Research Cloud. I would like to thank the Google Cloud TPU team for providing me with the resources to train the bigger text-conditional models in multi-host distributed settings.
In recent years, diffusion and score-based multi-step models have revolutionized the generative AI domain. However, the latest research in this field has become highly math-intensive, making it challenging to understand how state-of-the-art diffusion models work and generate such impressive images. Replicating this research in code can be daunting.
FlaxDiff is a library of tools (schedulers, samplers, models, etc.) designed and implemented in an easy-to-understand way. The focus is on understandability and readability over performance. I started this project as a hobby to familiarize myself with Flax and Jax and to learn about diffusion and the latest research in generative AI.
I initially started this project in Keras, being familiar with TensorFlow 2.0, but transitioned to Flax, powered by Jax, for its performance and ease of use. The old notebooks and models, including my first Flax models, are also provided.
The Diffusion_flax_linen.ipynb notebook is my main workspace for experiments. Several checkpoints are uploaded to the pretrained folder along with a copy of the working notebook associated with each checkpoint. You may need to copy the notebook to the working root for it to function properly.
The library has since grown beyond images: it now also covers video diffusion, flow matching, latent diffusion with a VAE, and I-JEPA/V-JEPA self-supervised training, all running through the same trainer with data-parallel and FSDP sharding. The bigger text-conditional models in the gallery were trained on a TPU-v4-32 pod.
In the tutorial notebooks folder, you will find comprehensive notebooks for various diffusion techniques, written entirely from scratch and are independent of the FlaxDiff library. Each notebook includes detailed explanations of the underlying mathematics and concepts, making them invaluable resources for learning and understanding diffusion models.
-
Diffusion explained (nbviewer link) (local link)
- WORK IN PROGRESS An in-depth exploration of the concept of Diffusion based generative models, DDPM (Denoising Diffusion Probabilistic Models), DDIM (Denoising Diffusion Implicit Models), and the SDE/ODE generalizations of diffusion, with step-by-step explainations and code.
-
EDM (Elucidating the Design Space of Diffusion-based Generative Models)
- TODO A thorough guide to EDM, discussing the innovative approaches and techniques used in this advanced diffusion model.
These notebooks aim to provide a very easy to understand and step-by-step guide to the various diffusion models and techniques. They are designed to be beginner-friendly, and thus although they may not adhere to the exact formulations and implementations of the original papers to make them more understandable and generalizable, I have tried my best to keep them as accurate as possible. If you find any mistakes or have any suggestions, please feel free to open an issue or a pull request.
-
Multi-host Data parallel training script in JAX
- Training script for multi-host data parallel training in JAX, to serve as a reference for training large models on multiple GPUs/TPUs across multiple hosts. A full-fledged tutorial notebook is in the works.
-
TPU utilities for making life easier
- A collection of utilities and scripts to make working with TPUs easier, such as cli to create/start/stop/setup TPUs, script to setup TPU VMs (install everything you need), mounting gcs datasets etc.
I worked as a Machine Learning Researcher at Hyperverge from 2019-2021, focusing on computer vision, specifically facial anti-spoofing and facial detection & recognition. Since switching to my current job in 2021, I haven't engaged in as much R&D work, leading me to start this pet project to revisit and relearn the fundamentals and get familiar with the state-of-the-art. My current role involves primarily Golang system engineering with some applied ML work just sprinkled in. Therefore, the code may reflect my learning journey. Please forgive any mistakes and do open an issue to let me know.
Also, few of the text may be generated with help of github copilot, so please excuse any mistakes in the text.
- A Versatile and simple Diffusion Library
- Disclaimer (and About Me)
- Features
- Installation of FlaxDiff
- Getting Started with FlaxDiff
- References and Acknowledgements
- Pending things to do list
- Gallery
- Contribution
- License
Implemented in flaxdiff.schedulers:
- LinearNoiseScheduler (
flaxdiff.schedulers.LinearNoiseScheduler): A beta-parameterized discrete scheduler. - CosineNoiseScheduler (
flaxdiff.schedulers.CosineNoiseScheduler): A beta-parameterized discrete scheduler. - ExpNoiseScheduler (
flaxdiff.schedulers.ExpNoiseScheduler): A beta-parameterized discrete scheduler. - CosineContinuousNoiseScheduler (
flaxdiff.schedulers.CosineContinuousNoiseScheduler): A continuous scheduler. - CosineGeneralNoiseScheduler (
flaxdiff.schedulers.CosineGeneralNoiseScheduler): A continuous sigma parameterized cosine scheduler. - SqrtContinuousNoiseScheduler (
flaxdiff.schedulers.SqrtContinuousNoiseScheduler): A continuous scheduler using the sqrt schedule proposed for diffusion language models. - KarrasVENoiseScheduler (
flaxdiff.schedulers.KarrasVENoiseScheduler): A sigma-parameterized continuous scheduler proposed by Karras et al. 2022, best suited for inference. - EDMNoiseScheduler (
flaxdiff.schedulers.EDMNoiseScheduler): A sigma-parameterized continuous scheduler based on the EDM paper, best suited for training with the KarrasVENoiseScheduler. - FlowMatchingScheduler (
flaxdiff.schedulers.FlowMatchingScheduler): A rectified-flow scheduler with logit-normal timestep sampling and resolution-dependent shifting, as used in Stable Diffusion 3.
Implemented in flaxdiff.predictors:
- EpsilonPredictionTransform (
flaxdiff.predictors.EpsilonPredictionTransform): The model predicts the noise in the data. - DirectPredictionTransform (
flaxdiff.predictors.DirectPredictionTransform): The model predicts the original data from the noisy data. - VPredictionTransform (
flaxdiff.predictors.VPredictionTransform): The model predicts a linear combination of the data and noise. - FlowMatchPredictionTransform (
flaxdiff.predictors.FlowMatchPredictionTransform): The model predicts the flow velocity. - KarrasPredictionTransform (
flaxdiff.predictors.KarrasPredictionTransform): A generalized transform for the EDM, integrating various parameterizations. - get_diffusion_preset (
flaxdiff.predictors.get_diffusion_preset): One call which pairs a training schedule, a sampling schedule and a transform for the"edm","karras","cosine"and"flow"setups.
Implemented in flaxdiff.samplers:
- DDPMSampler (
flaxdiff.samplers.DDPMSampler): Implements the Denoising Diffusion Probabilistic Model (DDPM) sampling process. - DDIMSampler (
flaxdiff.samplers.DDIMSampler): Implements the Denoising Diffusion Implicit Model (DDIM) sampling process. - EulerSampler (
flaxdiff.samplers.EulerSampler): An ODE solver sampler using Euler's method. - EulerAncestralSampler (
flaxdiff.samplers.EulerAncestralSampler): Euler sampling with ancestral noise injection. - HeunSampler (
flaxdiff.samplers.HeunSampler): An ODE solver sampler using Heun's method. - RK4Sampler (
flaxdiff.samplers.RK4Sampler): An ODE solver sampler using the Runge-Kutta method. - MultiStepDPM (
flaxdiff.samplers.MultiStepDPM): Implements a multi-step sampling method inspired by the Multistep DPM solver as presented here: tonyduan/diffusion
All samplers support classifier-free guidance, including interval-limited guidance.
Implemented in flaxdiff.trainer:
- GeneralDiffusionTrainer (
flaxdiff.trainer.GeneralDiffusionTrainer): Manages the training loop, loss calculation, EMA, gradient accumulation, checkpointing and wandb logging, for both image and video data. It runs data-parallel and FSDP sharded training throughjax.jitwithNamedShardingon a(data, fsdp)mesh. - Objectives (
flaxdiff.trainer.objectives): What to optimize is pluggable.DiffusionObjectiveis the default;JepaObjective(flaxdiff.jepa) trains I-JEPA/V-JEPA encoders on the same trainer.
Implemented in flaxdiff.models and constructed via flaxdiff.models.registry.build_model:
- Unet: A classic convolutional UNet.
- UNet3D: A video UNet which can inflate 2D Unet checkpoints.
- UViT / SimpleUDiT: U-shaped transformers.
- SimpleDiT / SimpleMMDiT / HierarchicalMMDiT: DiT and SD3-style multi-modal DiT variants.
- HybridSSMAttentionDiT: Interleaves S5 state-space blocks with attention.
- VideoDiT: A factorized spatial-temporal DiT for video.
- Hilbert and zigzag patch scan orders are available via the
+hilbertand+zigzagarchitecture suffixes. - Autoencoders (
flaxdiff.models.autoencoder):StableDiffusionVAE(vendored Flax VAE, loads Hugging Face hub weights) andSimpleAutoEncoderfor latent diffusion without any external weights.
Implemented in flaxdiff.metrics: FID (vendored InceptionV3), CLIP score, PSNR, SSIM, and linear/kNN probes for JEPA.
To install FlaxDiff, you need to have Python 3.11 or higher:
pip install flaxdiffOptional extras pull in the heavier dependencies only when you need them:
flaxdiff[av]: video/audio sources and readers (OpenCV, decord, moviepy, PyAV)flaxdiff[metrics]: FID (scipy) and Inception weight downloadflaxdiff[streaming]: online URL-streaming loader (Hugging Facedatasets)flaxdiff[tfds]: TFDS-backed dataset sources
Or for development, clone the repo and install in editable mode with the test dependencies:
pip install -e .[test]
JAX_PLATFORMS=cpu pytest -m "not network" -qThe test suite covers model forward passes for every architecture, scheduler and transform invariants, sampler convergence against an analytic denoiser, trainer smoke runs for images and videos, FSDP and data-parallel parity on a simulated 8-device mesh, sharded checkpoint round-trips with mid-epoch data resume, and the JEPA objectives. Tests marked network download pretrained weights and are excluded by default.
Here is a simplified example to get you started with training a diffusion model using FlaxDiff:
from datetime import datetime
import jax, optax
from flaxdiff.data.dataloaders import get_dataset_grain
from flaxdiff.inputs import DiffusionInputConfig, ConditionalInputConfig
from flaxdiff.inputs.encoders import CLIPTextEncoder
from flaxdiff.models.registry import build_model
from flaxdiff.predictors import get_diffusion_preset
from flaxdiff.trainer import GeneralDiffusionTrainer
from flaxdiff.samplers.euler import EulerAncestralSampler
BATCH_SIZE, IMAGE_SIZE = 16, 128
data = get_dataset_grain("oxford_flowers102", batch_size=BATCH_SIZE, image_scale=IMAGE_SIZE)
text_encoder = CLIPTextEncoder.from_modelname("openai/clip-vit-large-patch14")
input_config = DiffusionInputConfig(
sample_data_key="image",
sample_data_shape=(IMAGE_SIZE, IMAGE_SIZE, 3),
conditions=[ConditionalInputConfig(encoder=text_encoder)],
)
train_schedule, sample_schedule, transform = get_diffusion_preset("edm")
model = build_model("simple_dit", dict(
emb_features=512, num_layers=8, num_heads=8, patch_size=8,
))
trainer = GeneralDiffusionTrainer(
model=model,
optimizer=optax.adamw(2e-4),
input_config=input_config,
noise_schedule=train_schedule,
model_output_transform=transform,
rngs=jax.random.PRNGKey(4),
name=f"flowers-edm-{datetime.now():%Y-%m-%d_%H%M}",
distributed_training=True,
checkpoint_base_path="./checkpoints",
)
trainer.fit(
data,
training_steps_per_epoch=data["train_len"] // BATCH_SIZE,
epochs=100,
sampler_class=EulerAncestralSampler,
sampling_noise_schedule=sample_schedule,
)The full-featured entry points are training.py for diffusion and training_jepa.py for I-JEPA/V-JEPA.
Here is a simplified example for generating images using a trained model:
from flaxdiff.inference.pipeline import DiffusionInferencePipeline
pipeline = DiffusionInferencePipeline.from_wandb_registry(
modelname="diffusion-model", project="my-project",
)
images = pipeline.generate_samples(
num_samples=16, resolution=128, diffusion_steps=50,
guidance_scale=3.0, conditioning_data=["a water lily", "a rose"],
)- The Original Denoising Diffusion Probabilistic Models (DDPM) paper
- Denoising Diffusion Implicit Models (DDIM) paper
- Improved Denoising Diffusion Probabilistic Models paper
- Diffusion Models beat GANs on image synthesis paper
- Score-Based Generative Modeling through Stochastic Differential Equations paper
- Elucidating the design space of Diffusion-based generative models (EDM) paper
- Perception Prioritized Training of Diffusion Models (P2 Weighting) paper
- Pseudo Numerical Methods for Diffusion Models on Manifolds (PNMDM) paper
- The DPM-Solver: A Fast ODE Solver for Diffusion Probabilistic Model Sampling in Around 10 Steps paper
- Scalable Diffusion Models with Transformers (DiT) paper
- Scaling Rectified Flow Transformers for High-Resolution Image Synthesis (SD3) paper
- Flow Matching for Generative Modeling paper
- Self-Supervised Learning from Images with a Joint-Embedding Predictive Architecture (I-JEPA) paper
- Simplified State Space Layers for Sequence Modeling (S5) paper
- Applying Guidance in a Limited Interval Improves Sample and Distribution Quality (interval-limited CFG) paper
- Diffusion-LM Improves Controllable Text Generation (sqrt schedule) paper
- An incredible series of blogs on various diffusion related topics by Sander Dieleman. The posts particularly on diffusion models, Typicality, Geometry of Diffusion Guidance and Noise Schedules are a must read
- An awesome blog series by Tony Duan on Diffusion models from scratch. Although it trains models for MNIST and the implementations are a bit basic, the maths is explained in a very nice way. The codebase is here
- The k-diffusion codebase by Katherine Crowson, which hosts an exhaustive implementation of the EDM paper (Karras et al) along with the DPM-Solver, DPM-Solver++ (both 2S and 2M) in pytorch. Most other diffusion libraries borrow from this.
- The Official EDM implementation by Tero Karras, in pytorch. Really neat code and the reference implementation for all the karras based samplers/schedules.
- The Hugging Face Diffusers Library. The vendored Flax VAE and parts of the attention module derive from it (Apache-2.0, attribution headers preserved).
- jax-fid, the origin of the vendored InceptionV3 used for FID.
- The Keras DDPM Tutorial by A_K Nain, and the Keras DDIM implementation by András Béres, which are great starting points for beginners to understand the basics of diffusion models. I started my journey by trying to implement the concepts introduced in these tutorials from scratch.
- Multi-host validation of the revamped trainer on an actual TPU pod
- A proper precision policy (dtype/param_dtype are still threaded ad-hoc)
- Full FID-50k evaluation (the current FID metric is per-validation-batch)
- Autoregressive LM and diffusion-LM objectives on the same trainer
Model trained on Laion-Aesthetics 12M + CC12M + MS COCO + 1M aesthetic 6+ subset of COYO-700M on TPU-v4-32:
a beautiful landscape with a river with mountains, a beautiful landscape with a river with mountains, ...
Params:
Dataset: Laion-Aesthetics 12M + CC12M + MS COCO + 1M aesthetic 6+ subset of COYO-700M
Batch size: 256
Image Size: 128
Training Epochs: 5
Steps per epoch: 74573
Model Configurations: feature_depths=[128, 256, 512, 1024]
Training Noise Schedule: EDMNoiseScheduler
Inference Noise Schedule: KarrasVENoiseScheduler
Images generated by the following prompts using classifier free guidance with guidance factor = 2:
'water tulip, a water lily, a water lily, a water lily, a photo of a marigold, a water lily, a water lily, a photo of a lotus, a photo of a lotus, a photo of a lotus, a photo of a rose, a photo of a rose, a photo of a rose, a photo of a rose, a photo of a rose'
Params:
Dataset: oxford_flowers102
Batch size: 16
Image Size: 128
Training Epochs: 1000
Steps per epoch: 511
Training Noise Schedule: EDMNoiseScheduler
Inference Noise Schedule: KarrasVENoiseScheduler
Params:
Dataset: oxford_flowers102
Batch size: 16
Image Size: 64
Training Epochs: 1000
Steps per epoch: 511
Training Noise Schedule: CosineNoiseScheduler
Inference Noise Schedule: CosineNoiseScheduler
Model: UNet(emb_features=256, feature_depths=[64, 128, 256, 512], attention_configs=[{"heads":4}, {"heads":4}, {"heads":4}, {"heads":4}, {"heads":4}], num_res_blocks=2, num_middle_res_blocks=1)
Images generated by Heun Sampler in 10 steps (20 model inferences as Heun takes 2x inference steps) [Unconditional]
Params:
Dataset: oxford_flowers102
Batch size: 16
Image Size: 64
Training Epochs: 1000
Steps per epoch: 511
Training Noise Schedule: EDMNoiseScheduler
Inference Noise Schedule: KarrasVENoiseScheduler
Feel free to contribute by opening issues or submitting pull requests. Let's make FlaxDiff better together!
This project is licensed under the MIT License.




