Skip to content
AI.info

Research

gfnx: Fast and Scalable Library for Generative Flow Networks in JAX

gfnx: Fast and Scalable Library for Generative Flow Networks in JAX Overview Research area: Machine learning software infrastructure — specifically Generative Flow Networks (GFlowNets), a family of ge

gfnx: Fast and Scalable Library for Generative Flow Networks in JAX
arXiv
2511.16592
Published
2025-11-20
Authors
Daniil Tiapkin, Artem Agarkov, Nikita Morozov, Ian Maksimov, Askar Tsyganov, Timofei Gritsaev, Sergey Samsonov

AI summary

gfnx: Fast and Scalable Library for Generative Flow Networks in JAX

Overview

  • Research area: Machine learning software infrastructure — specifically Generative Flow Networks (GFlowNets), a family of generative models for sampling from unnormalized discrete distributions.
  • Technical level: Intermediate. Readers should know what GFlowNets are and be comfortable with JAX/PyTorch-style training loops, though the paper is written as a library and benchmarking report rather than a theory paper.
  • Scope: A single paper introducing gfnx, a JAX-native library of JIT-compiled environments, reward modules, and metrics for training and benchmarking GFlowNets, together with wall-clock comparisons against PyTorch-based baselines.

What This Paper Is About

GFlowNets need to be trained by repeatedly rolling out trajectories through an environment, which makes environment speed a practical bottleneck for research. Existing toolkits such as torchgfn execute environment logic on the CPU host, forcing data transfers to and from GPUs or TPUs during training. The authors build gfnx, a library in which environments, reward modules, and metrics are all written in JAX so that the entire training loop can be just-in-time (JIT) compiled and run on-device.

Key Contributions

  1. A JAX-native GFlowNet library. gfnx implements vectorized, JIT-able environments, reward modules, and success metrics entirely in JAX, together with single-file baseline training scripts in the spirit of CleanRL. The library is available on GitHub (https://github.com/d-tiapkin/gfnx), on PyPI (https://pypi.org/project/gfnx/), and documented at https://gfnx.readthedocs.io.
  2. Eight implemented environments. Hypergrids, Bit sequences, TFBind8, QM9, AMP, Phylogenetic trees, Bayesian structure learning, and the Ising model — covering synthetic, molecular, biological-sequence, phylogenetic, causal-discovery, and energy-based tasks.
  3. Modular decoupling of reward from dynamics. The library separates environment dynamics from reward computation, so reward families can be swapped or learned during training without recompiling environment logic.
  4. Large wall-clock speedups over baselines. Reported speedups include up to 55× for the QM9 environment on CPU and up to 80× for the Bayesian structure learning setup on GPU.

Main Findings

  • Speedups on CPU sequence environments: On TFBind8 the TB objective rises from 230.4 ± 2.7 it/s (baseline, Shen et al. [2023]) to 6929.3 ± 26.0 it/s with gfnx. On QM9, TB rises from 162.3 ± 1.5 it/s to 9061.6 ± 28.7 it/s, corresponding to the reported 55× speedup.
  • Speedup on GPU structure learning: Bayesian structure learning with the MDB objective rises from 0.73 ± 0.03 it/s (Lahlou et al. [2023]) to 58.0 ± 1.0 it/s with gfnx — the 80× case.
  • Hypergrid on CPU: With DB, throughput rises from 178.3 ± 2.0 it/s to 1560.0 ± 3.6 it/s; with TB from 219.9 ± 1.1 it/s to 1463.3 ± 1.4 it/s; with SubTB from 121.0 ± 0.7 it/s to 596.0 ± 0.3 it/s. The paper states gfnx is at least about five times faster than torchgfn in this setting.
  • Bit sequences on GPU: With DB, 52.4 ± 0.8 it/s becomes 1666.4 ± 8.0 it/s; with TB, 54.0 ± 0.2 it/s becomes 2433.6 ± 67.1 it/s. The paper reports gfnx is at least 30× faster than the PyTorch implementation here.
  • AMP on GPU: TB throughput rises from 21.2 ± 1.7 it/s to 413.1 ± 3.4 it/s.
  • Phylogenetic trees on GPU (FLDB objective, eight datasets): DS-1: 13.3 ± 0.7 → 263.6 ± 2.1 it/s; DS-2: 9.9 ± 0.1 → 208.4 ± 4.0 it/s; DS-3: 8.6 ± 0.1 → 149.7 ± 2.4 it/s; DS-4: 10.0 ± 0.1 → 122.0 ± 1.0 it/s; DS-5: 11.7 ± 0.7 → 76.1 ± 0.9 it/s; DS-6: 6.3 ± 0.3 → 66.3 ± 1.4 it/s; DS-7: 2.9 ± 0.1 → 36.7 ± 0.4 it/s; DS-8: 4.1 ± 0.2 → 39.1 ± 0.6 it/s.
  • Ising model (no baseline available): gfnx reaches 27.8 ± 0.2 it/s for N = 9 and 26.4 ± 0.6 it/s for N = 10 on GPU with the TB objective. The paper notes no open-source implementation of this EBM setting from Zhang et al. [2022] was available for comparison.
  • Small and large hypergrids (Table 2, CPU): 2-dimensional hypergrid with side 20 — DB 321.8 ± 2.5 → 3235.0 ± 86.7 it/s, TB 359.8 ± 13.1 → 1853.3 ± 14.9 it/s, SubTB 222.1 ± 1.4 → 2258.0 ± 53.6 it/s. 8-dimensional hypergrid with side 10 — DB 209.6 ± 1.8 → 1453.6 ± 2.9 it/s, TB 225.3 ± 3.7 → 1443.8 ± 4.0 it/s, SubTB 143.6 ± 0.7 → 598.1 ± 3.4 it/s.
  • Sampling quality is preserved. On the 4-dimensional hypergrid with side length 20, both gfnx and torchgfn converge to the same total variation metric against the true reward distribution for DB, TB, and SubTB, with the paper reporting that a perfect sampler also has a nonzero metric because the metric is computed from a finite sample of terminal states.
  • Bit-sequence correlation is comparable. For n = 120, k = 8, gfnx matches the metric values of the PyTorch implementation from Tiapkin et al. [2024] while running faster, measured as Pearson correlation between terminating-state log-probability and log-reward over 7200 randomly sampled bit sequences.

Methodology in Plain English

The authors reimplement GFlowNet environments, reward functions, and evaluation metrics as JAX programs so that the environment step, reward evaluation, and the training update can all be fused into a single compiled computation. Three design choices drive the speedups:

  1. Stateless environments. All mutable data lives in an EnvState object returned by env.reset and modified explicitly by env.step, matching the JAX functional paradigm and making vectorization straightforward.
  2. Fused reward evaluation. Instead of naively mapping a terminal check over the batch with jax.vmap, the library wraps reward computation in jax.lax.cond and only evaluates rewards when at least one element of the batch is terminal, avoiding redundant work.
  3. Log-scale rewards. Environments return log_reward rather than raw reward — terminal transitions yield their log-reward, non-terminal steps return zero — which matches how most GFlowNet algorithms consume the training signal.

The library is organized into base.py, environment/, reward/, metrics/, and utils/ modules, with companion baselines/ single-file training scripts (using Equinox for neural networks) and proxy/ utilities for training dataset-driven proxy reward models. Backward actions are abstracted over structural choices rather than exact symbol inverses, which lets a backward rollout be implemented simply by replacing initial states with terminal ones and env.step with env.backward_step.

Benchmarks compare gfnx against torchgfn and against author implementations from prior work, reporting iterations per second with a 3-sigma standard error interval, averaged over at least 3 random seeds. Hypergrid experiments ran on an Intel Core i7-10700F CPU with 32 GB RAM; bit-sequence experiments ran on a single NVIDIA V100 GPU.

Why This Matters

Impact on research. The paper argues that runtime efficiency is critical for GFlowNet research because fast environments enable rapid iteration, large-scale hyperparameter sweeps, and more reliable statistical comparisons across random seeds. By standardizing environments and metrics and releasing compiled single-file baselines, gfnx aims to lower the barrier to reproducing and extending GFlowNet methods.

Real-world applications (as represented by the implemented environments):

  • Molecular generation — the QM9 environment rewards molecules by a proxy model predicting the HOMO-LUMO gap, an important molecular property.
  • Antimicrobial peptide design — the AMP environment generates variable-length peptides (up to 60 tokens over a 20-amino-acid vocabulary) rewarded by a proxy trained on the DBAASP database.
  • DNA sequence design — TFBind8 generates nucleotide sequences of length 8 (vocabulary size 4) scored by wet-lab measured binding activity to the human transcription factor SIX6.
  • Phylogenetic tree construction and causal discovery — the phylogenetic environment builds rooted binary trees from n singleton species over n − 1 merge steps, and the Bayesian structure learning environment builds a DAG edge by edge with reward equal to a log-posterior under linear-Gaussian or BGe scores.

Industry relevance. The paper itself does not report industrial deployments or case studies; its stated relevance is infrastructure — providing differentiable, GPU/TPU-native environments that avoid CPU–accelerator data transfers, and releasing the package on GitHub and PyPI with public documentation.

Future Directions

  1. Continuous action spaces. gfnx currently supports only discrete action spaces, but the authors note that many problems have continuous components, including a more complete version of the phylogenetic tree generation environment.
  2. Non-acyclic environments. Many natural problems, such as permutation generation, might be more clearly written without an acyclicity constraint.
  3. Multi-objective support. The need to generate diverse Pareto-optimal solutions arises in many problems and is not yet covered.
  4. Additional baselines and trainer vectorization. The authors propose implementing backward policy optimization algorithms and exploration techniques, adding entropy-regularized RL training baselines, and batching over seeds and hyperparameters in the style of purejaxrl to improve reproducibility and hyperparameter sweeps on small environments.

Target Audience

GFlowNet researchers and practitioners who need fast, reproducible training and evaluation loops; reinforcement-learning engineers already working in JAX who want compiled environments in the style of purejaxrl; and applied scientists in molecular generation, biological sequence design, phylogenetics, or causal discovery who want ready-made benchmark environments with documented reward modules and metrics. Readers looking for new GFlowNet theory, convergence guarantees, or industrial deployment results will not find them here — the paper is a library and benchmarking contribution.

Authors’ abstract

In this paper, we present gfnx, a fast and scalable package for training and evaluating Generative Flow Networks (GFlowNets) written in JAX. gfnx provides an extensive set of environments and metrics for benchmarking, accompanied with single-file implementations of core objectives for training GFlowNets. We include synthetic hypergrids, multiple sequence generation environments with various editing regimes and particular reward designs for molecular generation, phylogenetic tree construction, Bayesian structure learning, and sampling from the Ising model energy. Across different tasks, gfnx achieves significant wall-clock speedups compared to Pytorch-based benchmarks (such as torchgfn library) and author implementations. For example, gfnx achieves up to 55 times speedup on CPU-based sequence generation environments, and up to 80 times speedup with the GPU-based Bayesian network structure learning setup. Our package provides a diverse set of benchmarks and aims to standardize empirical evaluation and accelerate research and applications of GFlowNets. The library is available on GitHub (https://github.com/d-tiapkin/gfnx) and on pypi (https://pypi.org/project/gfnx/). Documentation is available on https://gfnx.readthedocs.io.

Read the original paper