A minimalist PyTorch implementation of CoFlow, which adjusts flow matching trajectories through contrastive repulsion to improve few-step generation.
CoFlow training uses --flow_type neg. Use --flow_type standard for the standard flow matching baseline or --flow_type delta for the DeltaFM baseline.
conda create -n negfm python=3.10.0
conda activate negfm
pip install torch==2.1.0 torchvision==0.16.0 --index-url https://download.pytorch.org/whl/cu121
pip install -r requirements.txtPlace ImageNet training images in an ImageFolder directory, with one subdirectory per class:
data/
└── imagenet/
└── images/
└── train/
├── n01440764/
│ └── example.JPEG
└── ...
Caching is needed only for training with --use_latent. First convert the images to LMDB:
(cd preprocess && python image2lmdb.py)Next, cache the VAE posterior parameters for both original and flipped images:
torchrun --standalone --nnodes=1 --nproc_per_node=1 preprocess/main_cache.py \
--source_lmdb ./data/imagenet/imagenet_train_lmdb \
--target_lmdb ./data/imagenet/train_vae_latents_lmdb \
--img_size 256 \
--batch_size 32 \
--lmdb_size_gb 400Train CoFlow directly from images:
accelerate launch --mixed_precision=bf16 --num_processes=4 train.py \
--model_type 'SiT-S/2' \
--flow_type neg \
--neg_type 'res xt' \
--neg_policy random \
--neg_lam 0.20To train from cached VAE inputs, add --use_latent. For example, use four GPUs:
accelerate launch --mixed_precision=bf16 --num_processes=4 train.py \
--model_type 'SiT-S/2' \
--flow_type neg \
--neg_type 'res xt' \
--neg_policy random \
--neg_lam 0.20 \
--use_latentTrain the standard flow matching baseline with:
accelerate launch --mixed_precision=bf16 --num_processes=4 train.py \
--model_type 'SiT-S/2' \
--flow_type standardAdd --use_repa for representation alignment with DINOv2 ViT-B/14, or --use_ot for exact optimal transport coupling within each local batch.
Training writes an experiment directory:
exp/<YYYYMMDD_HHMMSS>/
├── config.json
├── log.txt
└── steps_<NNNNNNN>/
├── custom_checkpoint_0.pkl # EMA weights used by sample.py
└── ...
Checkpoints are saved every 50,000 optimizer steps and at the final training step.
Before sampling, prepare the ImageNet 256 reference batch at ./data/VIRTUAL_imagenet256_labeled.npz. Sampling automatically runs evaluation using this fixed path.
Replace YYYYMMDD_HHMMSS and 50000 below with an experiment ID and a saved checkpoint step under exp/:
accelerate launch --mixed_precision=bf16 --num_processes=4 sample.py \
--exp YYYYMMDD_HHMMSS \
--ckpt_step 50000 \sample.py loads the EMA checkpoint and uses an Euler–Maruyama sampler, with a deterministic final step.