VIOT pretrained checkpoints

Pretrained weights for A Variational Optimal Transport Operator on Incompressible Flow (VIOT), arXiv:2609.13729.

Code: github.com/jinjinhe2001/VIOT · Project page: jinjinhe2001.github.io/viot · Live demo in the browser

VIOT transports in 2D and 3D Transports generated by the released operators (paper teaser): 2D MNIST digits and MPEG-7 silhouettes, a 3D "SMOKE" sequence, a 3D human-pose sequence, and a sphere-to-airplane transport. Red boxes mark target keyframes; insets show the ground truth.

VIOT is a neural operator. Given the current density, the target density and a time t, it predicts a velocity field that is divergence-free by construction:

  • in 2D a Fourier neural operator (FNO) outputs a stream function;
  • in 3D it outputs a vector potential.

The velocity is the spectral curl of that potential, truncated at |k| <= k_max. A source density is carried to a target by advecting it with this velocity for 50 steps.

The operator is trained by unrolling this transport. The loss is the terminal density mismatch plus a kinetic-energy (Benamou-Brenier) term and a viscosity (enstrophy) term. The weights of these terms are in each config.json; the 3D models use a kinetic-energy weight of 0.

These are the models from the paper, for our method only (no baselines, no ablations).

Gallery

Sample transports from these checkpoints

Generated with the files in this repository and the examples/ of the code repository. Each segment is one 50-step rollout, and the final density of a segment is the source of the next one.

Sample transports

From left to right: MNIST test digits 2→0→2→6 (mnist_2d), 春→江→潮→水 rendered with SimHei (cjk_2d), V→I→O→T in DejaVu Sans (latin_font_2d), apple→heart→horse→turtle from the MPEG-7 pool (mpeg7_2d), and extruded 3D glyphs V→I→O→T seen at 45° (font_3d).

Figures from the paper

sphere2airplane_3d: seven sphere-to-airplane rollouts; each column is one rollout, top to bottom, with the target airplane inset at the bottom.

Sphere to airplane

humman_3d_figures: five chains, each through five randomly selected human poses, with two intermediate frames between adjacent target keyframes.

Human poses

font_3d: seven chains of volumetric glyphs; red boxes mark target keyframes.

3D fonts

mpeg7_2d: one hundred random MPEG-7 transitions. Each tile is the final density after five consecutive transport segments (dark inset: the target).

latin_font_2d: one operator successively forming the glyphs of "FLUIDOT OPERATOR".

Glyph chain

Method. The operator receives the current density, the target density and the time, predicts a stream function (2D) or a vector potential (3D), and its curl is the divergence-free velocity that advects the density.

Pipeline

Models

name task resolution architecture params size advection (training) used in the paper for
mnist_2d MNIST digit to digit 256² FNO2D w64 / m32 / L8, k_max 0.25 134.5 M 538 MB MacCormack teaser (MNIST row), interactive demo figure, appendix MNIST comparison, held-out / momentum / timing tables
cjk_2d Chinese character to character 256² FNO2D w64 / m32 / L8, k_max 0.25 134.5 M 538 MB MacCormack held-out and momentum tables, poem video
mpeg7_2d MPEG-7 silhouette to silhouette 256² FNO2D w64 / m32 / L8, k_max 0.25 134.5 M 538 MB MacCormack teaser (MPEG-7 row), 10×10 grid figure, appendix MPEG-7 comparison, step-count study, held-out / momentum / timing tables
latin_font_2d Latin glyph (A-Z, a-z, 0-9) to glyph 256² FNO2D w64 / m16 / L8, k_max 0.0625 33.8 M 135 MB first-order semi-Lagrangian appendix glyph chains, held-out / momentum / timing tables
sphere2airplane_3d sphere to ShapeNet airplane 128³ FNO3D w32 / m16 / L6, k_max 0.25 201.5 M 810 MB first-order semi-Lagrangian teaser (airplane row), appendix airplane chains, held-out and timing tables
humman_3d human pose to pose (HuMMan) 128³ FNO3D w32 / m16 / L6, k_max 0.25 201.5 M 810 MB first-order semi-Lagrangian all HuMMan metrics (held-out and momentum tables)
font_3d extruded glyph to glyph (DejaVu) 128³ FNO3D w32 / m16 / L6, k_max 0.25 201.5 M 810 MB first-order semi-Lagrangian teaser (3D glyph row), appendix 3D font chains, held-out / momentum / timing tables
humman_3d_figures human pose to pose (HuMMan) 128³ FNO3D w32 / m16 / L6, k_max 0.25 201.5 M 810 MB first-order semi-Lagrangian teaser (HuMMan row), appendix HuMMan chains, HuMMan timing row
  • Architecture column. w = channel width, m = Fourier modes per dimension, L = number of FNO layers.
  • Headline table. The first seven models are the runs listed in the paper's headline table.
  • humman_3d_figures continues humman_3d: 5000 more steps at viscosity weight 0.007 instead of 0.005.
    • The paper's HuMMan figures and HuMMan timing row were made with it.
    • It is shipped only so that those figures can be reproduced.
    • It does not underlie any reported metric.
  • Advection at inference. The advection scheme is a rollout choice; the network is the same either way.
    • The paper's MPEG-7 teaser row, MPEG-7 grid and glyph-chain figures used WENO advection at inference.
    • Some 3D figures used MacCormack at inference.
    • config.json records the scheme each model was trained with.

Each folder <name>/ contains two files:

  • model.safetensors: the fp32 state_dict.
    • Keys and tensors are identical to the original PyTorch checkpoint, including the trunc_mask buffer.
    • The saved trunc_mask is used at the native resolution. At other resolutions the mask is rebuilt from k_max.
  • config.json, which records:
    • the constructor arguments (arch) and the rollout settings;
    • the data description;
    • the exact training command (training.args, training.command);
    • the checkpoint step, SHA256 of the original .pt and of model.safetensors;
    • what the checkpoint was used for in the paper.

MANIFEST.json lists every file with its size and SHA256.

Browser weights (web/)

web/viot_mnist256_q4.safetensors (65 MB) and web/viot_mnist256_q8.safetensors (124 MB) are mnist_2d exported for the WebGPU demo (web/export_web.py in the code repository):

  • the lift is folded exactly into the first spectral layer;
  • the spectral weights of layers 1-7 are stored in 4 or 8 bits, with an f16 scale per input channel and mode.

On 32 MNIST test pairs, int4 changes a 50-step rollout by 4.7% and int8 by 0.4%, relative to strict fp32. For comparison, PyTorch's default TF32 convolutions change it by 2.9%. The terminal L2 error is unchanged in all cases (0.00110-0.00111). These files are only for the browser demo; use mnist_2d/ for anything else.

Usage

pip install torch safetensors huggingface_hub
pip install git+https://github.com/jinjinhe2001/VIOT
from viot import load_pretrained

model, cfg = load_pretrained("mnist_2d", device="cuda")   # fetches mnist_2d/ from this repository
print(cfg["arch"], cfg["rollout"])

The viot package provides rollouts (viot.ops_2d.rollout_2d, viot.ops_3d.rollout_3d) with the advection scheme from cfg["rollout"]["advection"].

Loading without the helper

import json
import torch
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
from viot.model_2d import FNO2D          # FNO3D lives in viot.model_3d

repo = "jinjinhe2001/VIOT"
cfg = json.load(open(hf_hub_download(repo, "mnist_2d/config.json")))
state = load_file(hf_hub_download(repo, "mnist_2d/model.safetensors"))

arch = {k: v for k, v in cfg["arch"].items() if k != "class"}
model = FNO2D(**arch)
model.load_state_dict(state, strict=True)
model.eval()

# rho_t, rho_1: [B, 1, 256, 256] densities, each summing to 1 (2D models); t: [B] in [0, 1)
with torch.no_grad():
    v = model.forward_velocity_only(rho_t, t, rho_1)   # [B, 2, 256, 256], divergence-free

Rollout convention

The trainer's evaluation and the paper use this rollout:

  • n_steps = 50, dt = 1 / n_steps, and the model is queried at t_i = i / n_steps.
  • After each advection step the density is clamped at 0 and rescaled to its initial total mass.

Input normalisation differs by dimension:

  • 2D models take densities normalised to unit sum.
  • 3D models take densities with peak value 1.

Checkpoint selection

The shipped checkpoints are model_best. Every training run saved a checkpoint every 1000 steps and kept as model_best the one with the lowest terminal loss on the current training batch (this is not a validation-based selection). The paper's figures, its held-out, momentum and timing tables, and the interactive demo use these checkpoints. The final weights of the runs are not included.

name shipped step (model_best) total steps
mnist_2d 64000 80000
cjk_2d 71000 80000
mpeg7_2d 79000 80000
latin_font_2d 35000 40000
sphere2airplane_3d 4000 5000
humman_3d 3000 5000
font_3d 1000 1500
humman_3d_figures 3000 5000

Reproducibility notes

  • latin_font_2d

    • It was trained with first-order semi-Lagrangian advection.
    • Its training pool is 62 glyphs × 8 font faces. Neither the build script nor the font identities were recovered, so the pool cannot be regenerated exactly.
    • Its launch command was reconstructed from the training log and the checkpoint shapes.
  • 3D models. All four were trained, and evaluated by the trainer, with first-order semi-Lagrangian advection.

  • Continuation runs.

    • humman_3d and font_3d were initialised from intermediate checkpoints that are not released.
    • humman_3d_figures was initialised from humman_3d.
    • The lineage is described in each config.json (training.init).
  • MPEG-7 chain number in the text. The paper's Sec. 5.2 figure of 0.0023 (rel. 0.289) came from a different, unreleased model. It cannot be reproduced with mpeg7_2d.

  • Which checkpoint the figures used. It was not recorded for:

    • the 3D font chains (font_3d);
    • the HuMMan chains (humman_3d_figures).

    Those figures may have been made with the final weights of the same runs.

Data and licence notes

Licence of the weights: CC BY-NC 4.0 (non-commercial use, with attribution). The dataset terms below apply in addition.

model data terms
mnist_2d MNIST training split (torchvision) MNIST is distributed under CC BY-SA 3.0
cjk_2d 3000 CJK characters (U+4E00-U+59B7) rasterised from SimHei, Microsoft YaHei and SimSun Windows-bundled commercial fonts; neither the fonts nor the rasterised glyphs are redistributed. Check the font licences for your use.
latin_font_2d 62 Latin glyphs × 8 font faces Source fonts unidentified (see above)
mpeg7_2d MPEG-7 CE-Shape-1 silhouettes Public academic benchmark distributed without an explicit licence
sphere2airplane_3d Synthetic sphere; ShapeNetCore.v2 airplanes (watertight, voxelized) ShapeNet terms of use (non-commercial research and educational use) apply to these weights
humman_3d, humman_3d_figures HuMMan SMPL poses (voxelized) HuMMan licence (non-commercial research) and the SMPL model licence (non-commercial) apply to these weights
font_3d 62 glyphs × 22 DejaVu faces, extruded and voxelized DejaVu fonts are freely redistributable

No training data is included in this repository. Data-preparation scripts and ID lists are in the code repository.

Citation

@misc{he2026variational,
  title={A Variational Optimal Transport Operator on Incompressible Flow},
  author={Jinjin He and Shenyifan Lu and Sinan Wang and Zhiqi Li and Duowen Chen and Bo Zhu},
  year={2026}, eprint={2609.13729}, archivePrefix={arXiv}, primaryClass={cs.LG},
  url={https://arxiv.org/abs/2609.13729}
}
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Paper for jinjinhe2001/VIOT