PyTorch implementation of Flow Matching for unconditional and conditional image generation on MNIST and CIFAR-10.
Flow Matching learns continuous normalizing flows by predicting velocity fields that transform noise into data. This implementation includes:
- Unconditional generation: Random sampling from learned data distributions
- Conditional generation: Generate specific classes/digits on command
- Optimized training: Straightened path training with endpoint emphasis for fast generation
- Modular architecture: Shared base classes for easy extension
# Install dependencies
pip install torch torchvision matplotlib
# Unconditional generation
python unconditional/fm_mnist.py # Generates random MNIST digits
python unconditional/fm_cfortan.py # Generates random CIFAR-10 images
# Conditional generation
python conditional/fm_mnist_conditional.py # Train conditional MNIST model
python conditional/fm_cifar_conditional.py # Train conditional CIFAR-10 modelAfter training, generate specific content:
# Generate specific MNIST digits
from conditional.fm_mnist_conditional import ConditionalFlowMatchingNet, generate_specific_digits
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = ConditionalFlowMatchingNet().to(device)
# ... load trained weights ...
generate_specific_digits(model, device, [1, 3, 7, 9])
# Generate specific CIFAR-10 classes
from conditional.fm_cifar_conditional import ConditionalFlowMatchingNetCIFAR, generate_specific_classes
model = ConditionalFlowMatchingNetCIFAR().to(device)
# ... load trained weights ...
generate_specific_classes(model, device, [0, 3, 5]) # airplane, cat, dogBoth models use encoder-decoder architectures with time and class conditioning:
- Encoder: Downsamples images to latent representations
- Time Embedding: Sinusoidal embeddings for flow time t ∈ [0,1]
- Class Embedding: Learned embeddings for conditional generation
- Fusion: Concatenates spatial, temporal, and class features
- Decoder: Upsamples to predict velocity fields
Flow Matching with optimizations:
- Sample noise x₀ ~ N(0,I) and data x₁ from dataset
- Sample time t ~ Beta(0.5,0.5) for endpoint emphasis
- Straightened path: x_t = (1-t²)x₀ + t²x₁
- Predict velocity: v_θ(x_t, t, c) ≈ 2t(x₁ - x₀)
- Minimize: ||v_θ(x_t, t, c) - 2t(x₁ - x₀)||²
Generation uses Heun integration with cosine time grid for efficient few-step sampling.
Generated samples are saved in the results/ folder:
15steps.png- Main result: 15-step generation100steps.png- Baseline: 100-step generation12steps.png,20steps.png- Additional speed comparisonsconditional_samples.png- All CIFAR-10 classes
Performance: 15 steps vs 100 steps = 6.7x speedup
conditional_samples.png- All MNIST digitsspecific_digits.png- Targeted digit generation