Official PyTorch implementation of the main 2D, PDE, and CIFAR-10 experiments in:
VTV-FM: Flow Matching through Variational Terminal-Velocity Closure
Haoyang Jiang*, Yuheng Li*, Di Yang*, Yanhai Xiong, Haipeng Chen, Yi He
NeurIPS 2026 (* equal contribution)
This repository contains only the VTV-FM method, its three principal experiment families, and the minimal third-party components required to run them. It intentionally excludes comparison baselines, exploratory sweeps, obsolete variants, and figure-only scripts.
See experiments/README.md for the exact inclusion and exclusion policy.
VTV-FM learns second-order phase-space dynamics
dx/dt = v, dv/dt = a_theta(x, v, t/T)
from a source prior x0 ~ N(0, I), v0 ~ N(0, sigma_v0^2 I). CIFAR-10 is a
missing-terminal-velocity problem, so VTV-FM uses one variational construction
for both source-target pairing and trajectory supervision:
- Compute the projected minimum-acceleration cost (Eq. 6) and solve minibatch entropic OT.
- Complete the missing terminal velocity with
v1* = 3(x1 - x0)/(2T) - v0/2(Eq. 5). - Evaluate the induced cubic Hermite bridge at a random time.
- Regress its acceleration, optionally with local consistency and acceleration-divergence matching (ADM).
The training objective is
L_total = L_acc + lambda_va L_va + lambda_xv L_xv + lambda_adm L_adm.
L_acc is the core method. L_va, L_xv, and L_adm are the three optional
regularizers reported in the paper's objective ablation. All can be enabled or
disabled independently from the command line.
| Component | Implementation |
|---|---|
| Terminal-velocity closure, Eq. 5 | vtvfm.matcher.terminal_velocity_closure |
| Projected pairing cost, Eq. 6 | vtvfm.matcher.projected_cost_matrix |
| Minimum-acceleration bridge, Proposition 3.1 | vtvfm.matcher.bridge_hermite_cubic |
| Pairing and supervision, Algorithm 2 | vtvfm.matcher.VTVFMatcher |
| Local consistency, Eqs. 10a-10b | vtvfm.losses.local_consistency_losses |
| ADM, Eq. 11 | vtvfm.losses.acceleration_divergence_loss |
| Minibatch Sinkhorn OT | vtvfm.optimal_transport.OTPlanSampler |
| Euler, Heun, Taylor2, Verlet, Dopri5 | vtvfm.samplers |
| CIFAR-10 acceleration network | vtvfm.models.build_acceleration_unet |
git clone https://github.com/HaoyangJiang-WM/VTV-FM.git
cd VTV-FM
conda create -n vtvfm python=3.11 -y
conda activate vtvfm
# Choose the PyTorch command matching your CUDA installation.
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu128
pip install -e .The low-dimensional release contains the seven datasets reported in the paper: Dinosaur, Swiss Roll, Checkerboard, Two Moons, Parallel Lines, Tree, and Star.
python -m experiments.toy_2d.train --dataset allUse --dataset <name> for a single dataset. The script trains VTV-FM only,
evaluates sample-based W2 with 2,000 samples, and writes one checkpoint per
dataset plus summary.csv.
PDE experiments use the observed/derivable-terminal-velocity regime and the
full-boundary action cost. A dataset file must contain snapshot and
terminal_tangent_phys tensors with matching [N, H, W] or [N, 1, H, W]
shapes.
python -m experiments.pde.train \
--system ns2d \
--dataset path/to/ns2d_snapshots.pt
python -m experiments.pde.evaluate \
--checkpoint results/pde/ns2d/vtvfm_step400000.pt \
--dataset path/to/ns2d_snapshots.pt \
--sampler heun --steps 100Replace ns2d with ks2d for the Kuramoto-Sivashinsky experiment. Evaluation
reports MMD, PSD log-L2, and the derivative-sensitive gap used in the paper.
Run the full VTV-FM objective:
bash experiments/cifar10/scripts/train.shReproduce the objective ablations:
OBJECTIVE=acc bash experiments/cifar10/scripts/train.sh # L_acc only
OBJECTIVE=acc_cons bash experiments/cifar10/scripts/train.sh # L_acc + L_va + L_xv
OBJECTIVE=acc_adm bash experiments/cifar10/scripts/train.sh # L_acc + L_adm
OBJECTIVE=full bash experiments/cifar10/scripts/train.sh # all termsThe default image configuration uses a 192-channel U-Net, batch size 256,
400,000 optimization steps, EMA decay 0.99995, and Sinkhorn pairing. CIFAR-10
is downloaded to ./data automatically. Checkpoints and preview grids are
written below ./results.
To resume training or inspect all options:
python -m experiments.cifar10.train --resume results/cifar10/vtvfm_full/vtvfm_cifar10_step400000.pt
python -m experiments.cifar10.train --helpfullGenerate candidate 8x8 sample grids with 100-step Velocity Verlet:
bash experiments/cifar10/scripts/sample_grid.sh path/to/checkpoint.ptThe script saves original and 4x-upscaled grids under ./samples/grids.
Evaluate 50,000 generated samples against CIFAR-10 training statistics using
clean-fid in legacy_tensorflow mode:
python -m experiments.cifar10.eval_fid --ckpt path/to/checkpoint.pt --sampler verlet --steps 100
bash experiments/cifar10/scripts/eval_fid.sh path/to/checkpoint.ptThe shell script evaluates Euler, Heun, Taylor2, Velocity Verlet, and adaptive
Dopri5. Fixed-step Euler and Taylor2 use N function evaluations; Heun and
Verlet use 2N.
.
├── experiments/
│ ├── toy_2d/ # seven low-dimensional datasets
│ ├── pde/ # NS2D and KS2D training/evaluation
│ └── cifar10/ # training, FID, grids, shell entry points
├── tests/ # mathematical and solver unit tests
└── vtvfm/
├── matcher.py # closure, cost, bridge, and pairing
├── losses.py # consistency and ADM regularizers
├── optimal_transport.py # minibatch OT
├── samplers.py # phase-space ODE solvers
├── checkpoint.py
└── models/unet/ # acceleration U-Net
- The controlled cross-method results in Table 3 use acceleration matching only: set all three regularization weights to zero.
- Table 4 uses the four objective settings shown above.
- The latent-terminal-velocity ADM target is
-3; Appendix E.5 derives-4for an observed fixed terminal velocity. - Class-conditional training is available through
--class_cond, but is not used for the unconditional CIFAR-10 results in the paper.
If you use this implementation, please cite:
@inproceedings{jiang2026vtvfm,
title = {{VTV-FM}: Flow Matching through Variational Terminal-Velocity Closure},
author = {Jiang, Haoyang and Li, Yuheng and Yang, Di and Xiong, Yanhai and Chen, Haipeng and He, Yi},
booktitle = {Advances in Neural Information Processing Systems},
year = {2026}
}The minibatch OT utilities and U-Net implementation are adapted from TorchCFM, which is released under the MIT License. The U-Net implementation originates from OpenAI guided-diffusion, also released under the MIT License. See NOTICE for details.
This repository is released under the MIT License.