Skip to content
HaoyangJiang-WMPublic

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Latest commit

 

History

2 Commits

Folders and files

Repository files navigation

VTV-FM: Flow Matching through Variational Terminal-Velocity Closure

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.

Description

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:

  1. Compute the projected minimum-acceleration cost (Eq. 6) and solve minibatch entropic OT.
  2. Complete the missing terminal velocity with v1* = 3(x1 - x0)/(2T) - v0/2 (Eq. 5).
  3. Evaluate the induced cubic Hermite bridge at a random time.
  4. 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.

VTV-FM package

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

Installation

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 .

Experiments

Two-dimensional datasets

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 all

Use --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-governed fields

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 100

Replace ns2d with ks2d for the Kuramoto-Sivashinsky experiment. Evaluation reports MMD, PSD log-L2, and the derivative-sensitive gap used in the paper.

CIFAR-10

Run the full VTV-FM objective:

bash experiments/cifar10/scripts/train.sh

Reproduce 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 terms

The 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 --helpfull

Sampling

Generate candidate 8x8 sample grids with 100-step Velocity Verlet:

bash experiments/cifar10/scripts/sample_grid.sh path/to/checkpoint.pt

The script saves original and 4x-upscaled grids under ./samples/grids.

FID evaluation

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.pt

The 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.

Project structure

.
├── 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

Reproducibility notes

  • 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 -4 for 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.

How to cite

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}
}

Acknowledgements

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.

License

This repository is released under the MIT License.

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages