Riemannian Flow Models with Reinforcement Learning for Molecular Crystal Structure Prediction
Authors: Thomas Egg, Harry Winston Sullivan, Maya M. Martirossyan, Philipp Höllmer, Cheng Zeng, Adrian Roitberg, Mingjie Liu, Richard Hennig, +3 more
Organizations: Center for Soft Matter Research, Department of Physics, New York University, New York 10003, USA · Simons Center for Computational Physical Chemistry, Department of Chemistry, New York University, New York 10003, USA · Department of Chemical Engineering and Materials Science, University of Minnesota, Minneapolis, MN 55455, USA · Department of Chemistry, University of Florida, Gainesville, FL 32611, USA · Quantum Theory Project, University of Florida, Gainesville, FL 32611, USA · Department of Materials Science & Engineering, University of Florida, Gainesville, FL 32611, USA · Department of Chemistry, University of Minnesota, Minneapolis, MN 55455, USA · Department of Aerospace Engineering and Mechanics, University of Minnesota, Minneapolis, MN 55455, USA · Courant Institute of Mathematical Sciences, New York University, New York 10003, USA · Center for Neural Science, New York University, New York 10003, USA
Crystal structure governs material properties, making crystal structure prediction (CSP) a fundamental problem in materials science. Generative models are a promising approach for solving this problem, but the prevalence of polymorphism, coupled with large unit cells and complex packing geometry, makes the molecular CSP task challenging for existing models. To address this, we introduce Coarse-Grained Open Materials Generation (CG-OMatG), an equivariant Riemannian flow-based generative model. CG-OMatG predicts molecular crystal structures \textit{via} a coarse-grained, hierarchical representation. CG-OMatG treats molecules as rigid bodies---performing both inter- and intra-molecular message passing to construct a geometric representation for molecular packings---and learns to reconstruct molecule centroid positions, orientations, and lattice parameters, conditioned on chemical species and conformer geometry. We train the model on subsets of the Open Molecular Crystals (OMC25) and Cambridge Structural Database (CSD) datasets. Further, we fine-tune the model \textit{via} policy gradient reinforcement learning to steer the model towards generating low-energy candidate structures. We validate the generated structures on the CSP blind test benchmark, assessing agreement with experimentally determined crystals using COMPACK packing-similarity analysis. CG-OMatG exhibits strong performance for generative molecular crystal structure prediction, paving the way for accelerated polymorph screening and organic solid-state materials discovery.
Figures & tables
Figure 1: Comparison of generated and ground-truth crystal packings for three OMC25-MCF targets. COMPACK alignments are shown with RDKit molecular diagrams and computed RMSDNMatches . CSD refcodes, left to right: NIYDOO, PEMWOT, JURRET. No relaxation was performed.
Metric
MCF †
Base
Orb-IRL
UMA-IRL
Solved ↑
0.0391
0.0742±0.0039
0.1008±0.0046
0.1273±0.0081
Solved (coll. allowed) ↑
0.1797
0.2109±0.0064
0.2227±0.0055
0.2594±0.0064
Packing match ↑
0.4062
0.4680±0.0115
0.5031±0.0085
0.5766±0.0079
Packing match/draw ↑
0.0534
0.0679±0.0013
0.0749±0.0009
0.0865±0.0009
Clash ↓
0.6182
0.3896±0.0022
0.2716±0.0016
0.1732±0.0015
Table 1: CCDC packing-similarity metrics on the first 128 structures of the OMC25-MCF test set ( k=30 inference). Values are rates; ↑ higher is better and ↓ lower is better.
Figure 2: UMA-driven relaxation for the three homomolecular CSP blind-test 6 targets (NACJAF, XAFPAY, XAFQIH), showing the ten lowest-energy draws per generator. Top: energy deviation from the UMA-relaxed ground truth across the three-stage BFGS relaxation; bottom: final deviations sorted by energy. Energies are in eV/atom. Relaxation details and complementary energy–density analyses are in Appendix I and Appendix K .
Appendix figures & tables19 assets
Supplementary material from the paper’s appendix.
Appendix
Figure 3: A cartoon of the molecular crystal manifold M . The blue curve represents a geodesic {\color[rgb]{0.2695,0.3789,0.9883}x_{t}}=\exp_{{\color[rgb]{0.1172,0.7305,0.1094}x_{0}}}(t\log_{{\color[rgb]{0.1172,0.7305,0.1094}x_{0}}}\left({\color[rgb]{0.5,0,0.4961}x_{1}}\right)) connecting an initially random rigid body positions {{\color[rgb]{0.1172,0.7305,0.1094}x_{0}}} to an optimal set of positions {\color[rgb]{0.5,0,0.4961}x_{1}} . The initial velocity {{\color[rgb]{0.957,0.6172,0.0234}{v}}} is an element of the tangent space {\color[rgb]{0.9883,0.2695,0.2695}T_{x_{0}}\mathcal{M}} at the point {{\color[rgb]{0.1172,0.7305,0.1094}x_{0}}} . The abstract blobs are rigid bodies with PCA frames attached to them which represent molecules being crystallized.
Figure 4: Geometric molecule embedding. Rotated atomic coordinates define a radius graph; edge lengths are expanded with a Gaussian radial basis and edge directions with spherical harmonics, then combined with species and time embeddings to form equivariant node features to be fed into a deep message passing network whose final hidden states are averaged to produce a molecule embedding.
Figure 5: Overall architecture. Geometric inputs are embedded and propagated on a periodic radius graph over centroid coordinates (via ghost centroids), yielding per-centroid hidden states for lattice, fractional-coordinate, and rotational-velocity prediction.
Figure 6: Prediction heads. Per-centroid hidden states feed four heads: fractional translation, cell rotation, cell stretch, and molecular orientation. Outputs are mapped to tangent vectors on each factor of M via Riemannian logarithm or Lie-algebra left-translation.
Figure 7: Deep message-passing network. Left: two message-update blocks with a residual-sum skip and a final concatenation skip. Right: the message layer uses an RBF-conditioned weighted tensor product with spherical harmonics; the update layer applies tensor augmentation, scalar gating, and geometric layer normalization.
Parameter
Value
Inference
config-name
omc25_inference.yaml
ckpt_path
model-checkpoints/omc25-mcf/best.ckpt
num_samples
30
Interpolant — sampling
num_timesteps
50
Appendix
Table 2: MolCrystalFlow inference hyperparameters used with the OMC25-MCF checkpoint, following the values recommended in the project README.
Parameter
Value
Master
Z
Per-crystal (#unique bb_indices )
MPI ranks
8
Workflow
tasks
generation , symm_rigid_press
Generation
Appendix
Table 3: Genarris hyperparameters used for the molecular-crystal generation baseline.
Parameter
Value
Training
Batch size (global)
4×320=1280
Optimizer
AdamW, lr =5×10−4 , weight decay =0.00818
LR schedule
Cosine annealing, 1500 epochs, ηmin=10−7
Gradient clipping
0.5, per-element
Time sampling
Logit-normal: t=σ(1.7Z+0.8) , Z∼N(0,1)
Appendix
Table 4: CG-OMatG pre-training hyperparameters. Each MLP has a single hidden layer with the listed width.
Parameter
Value
Training
Optimizer
Adam, lr =10−4
Max steps
5000
Gradient clipping
1.0, global norm
GRPO / PPO
Group size / num. groups
64 / 5
Appendix
Table 5: CG-OMatG RL fine-tuning hyperparameters (GRPO with PPO clipping).
Target
Method
Solved ↑
Solved (collisions allowed) ↑
Packing match ↑
Packing match (per draw) ↑
Clash ↓
NACJAF
OXtal
0.30±0.15
0.30±0.15
1.00±0.00
0.08±0.01
0.00±0.00
MCF
0.00
0.00
0.00
0.00
0.93
CG-OMatG
0.10±0.10
0.10±0.10
0.40±0.16
0.02±0.01
0.20±0.02
CG-OMatG-IRL
0.00±0.00
0.00±0.00
0.50±0.17
0.03±0.01
0.07±0.01
CG-OMatG (relaxed)
0.60±0.16
0.60±0.16
0.80±0.13
0.05±0.01
0.02±0.01
CG-OMatG-IRL (relaxed)
0.10±0.10
0.10±0.10
0.50±0.17
0.02±0.01
0.00±0.00
Appendix
Table 6: CCDC packing-similarity metrics for the three homomolecular sixth CSD blind-test targets ( k=30 inference). Values are rates; ↑ higher is better and ↓ lower is better.
Method
Solved
Collisions allowed
Packing match
Per draw
Clash
OXtal
0.20±0.02
0.20±0.02
0.50±0.02
0.10±0.00
0.00±0.00
CG-OMatG
0.05±0.00
0.08±0.00
0.39±0.01
0.05±0.00
0.43±0.00
CG-OMatG-IRL
0.06±0.00
0.10±0.00
0.49±0.01
0.06±0.00
0.25±0.00
Appendix
Table 7: CCDC packing-similarity metrics on the 37 OMC targets absent from both models’ training data. Values are mean ± SEM over ten K=30 blocks.
Figure 8: Energy versus density for generated sixth CSD blind-test structures before relaxation (top) and after UMA relaxation (bottom). Energies are reported relative to the UMA-relaxed experimental target. Relaxation sharpens the energy–density distributions for NACJAF, XAFPAY, and XAFQIH; for all three targets, the best CG-OMatG sample lies within 0.01 eV/atom of the lowest-energy relaxed ground truth with low density error.
Figure 9: Full OMC128 velocity-annealing sweep for the base and reinforced models at K=30 . Rows vary positional annealing and columns vary rotational annealing. Shown are solved rate, clash rate per draw, and target-level packing similarity.
Figure 10: Reinforcement-learning trajectories for CSD with a UMA reward (left), OMC with a UMA reward (center), and OMC with an Orb reward (right). Top: per-step reward −E/N and its 50-step rolling mean. Bottom: clashes per generated structure. Reward increases and clashes decrease across all three runs.
CSD ID
Rotatable bonds
Energy rank
RMSD after MMFF94s
RMSD after UMA
XAFQIH
5
10
0.571 Å
0.648 Å
XAFPAY
6
24
0.401 Å
0.348 Å
NACJAF
0
1
0.057 Å
0.040 Å
Appendix
Table 8: Recovery of experimental blind-test conformers from ETKDGv3/MMFF94s sampling.
Target
Pipeline
Solved
Collisions allowed
Packing match
Per draw
Clash
NACJAF
IRL + conformer
0.00
0.00
1.00
0.03
0.07
IRL + conformer (relaxed)
0.00
0.00
1.00
0.03
0.00
XAFPAY
IRL + conformer
0.00
0.00
0.00
0.00
0.37
IRL + conformer (relaxed)
0.00
0.00
0.00
0.00
0.00
XAFQIH
IRL + conformer
0.00
0.00
0.00
0.00
0.53
IRL + conformer (relaxed)
0.00
0.00
0.00
0.00
0.00
Appendix
Table 9: Blind-test metrics using generated conformer inputs. Each row has one K=30 block and therefore no error bars.
Target
Method
Distinct components C
Effective components eff
Dominant-mode fraction maxkpk
NACJAF
CG-OMatG
23.0±1.0
18.4±1.2
0.13±0.01
NACJAF
CG-OMatG-IRL
21.6±1.2
15.1±1.5
0.17±0.01
XAFPAY
CG-OMatG
27.5±0.6
25.2±1.2
0.08±0.01
XAFPAY
CG-OMatG-IRL
26.9±0.4
24.4±0.8
0.08±0.01
XAFQIH
CG-OMatG
27.1±0.8
23.9±1.9
0.10±0.02
XAFQIH
CG-OMatG-IRL
21.5±1.0
12.9±1.7
0.23±0.04
Appendix
Table 10: Diversity of generated blind-test structures under COMPACK component clustering.
Figure 11: CG-OMatG training and validation losses on the CSD database, decomposed by manifold component: lattice shape Sym3+ , cell and molecule orientations on SO(3) , and fractional coordinates on T3 . Weight norm is the global L2 norm of all trainable parameters. Note that the constant term is dropped from this loss, so it can be below zero, unlike the usual normalizing-flow presentation.
Figure 12: CG-OMatG training and validation losses on the OMC database, decomposed by manifold component: lattice shape Sym3+ , cell and molecule orientations on SO(3) , and fractional coordinates on T3 . Weight norm is the global L2 norm of all trainable parameters. Note that the constant term is dropped from this loss, so it can be below zero, unlike the usual normalizing-flow presentation. Occasional jumps in the loss curves reflect checkpoint restarts after improper GPU resource allocation on the cluster.
Department of Computer Science, University of North Carolina at Charlotte · Department of Mechanical Engineering and Engineering Science, University of North Carolina at Charlotte · North Carolina Battery Complexity, Autonomous Vehicle and Electrification (BATT CAVE) Research Center
Department of Computer Science and Engineering University of South Carolina Columbia, SC 29201 · Department of Computer Science and Cybersecurity University of North Georgia Columbia, SC 29201
State Key Laboratory of High Performance Ceramics, Shanghai Institute of Ceramics, Chinese Academy of Sciences, 1295 Dingxi Road, Shanghai, 200050, China. · Center of Materials Science and Optoelectronics Engineering, University of Chinese Academy of Sciences, Beijing, 100049, China. · State Key Laboratory of Multimodal Artificial Intelligence Systems, Institute of Automation, Chinese Academy of Sciences, 95 Zhongguancun East Road, Beijing, 100190, China. +5