diff --git a/pyproject.toml b/pyproject.toml
index 379b60c..9505e1b 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -81,6 +81,8 @@ filterwarnings = [
"ignore:custom data:UserWarning",
"ignore:`torch.jit.script` is deprecated:DeprecationWarning",
"ignore:`torch.jit.load` is deprecated:DeprecationWarning",
+ "ignore:`torch.jit.save` is deprecated:DeprecationWarning",
+ "ignore:`torch.jit.save` is not supported:DeprecationWarning",
"ignore:`compute_requested_neighbors_from_options` is deprecated and will be removed in a future version:UserWarning",
"ignore:hf_xet\\.download_files\\(\\) is deprecated:DeprecationWarning",
"ignore:Due to '_pack_', the '.*' Structure will use memory layout compatible with MSVC:DeprecationWarning",
@@ -98,4 +100,6 @@ filterwarnings = [
"ignore:the '.*' quantity is deprecated:DeprecationWarning",
"ignore:ModelOutput.quantity is deprecated:DeprecationWarning",
"ignore:Setting the shape on a NumPy array has been deprecated in NumPy 2.5:DeprecationWarning",
+ "ignore:currentThread\\(\\) is deprecated:DeprecationWarning",
+ "ignore:unclosed file:ResourceWarning", # i-PI's initializer leaks its handle
]
diff --git a/src/flashmd/ipi.py b/src/flashmd/ipi.py
index 1cbdbaf..f429c7e 100644
--- a/src/flashmd/ipi.py
+++ b/src/flashmd/ipi.py
@@ -2,8 +2,10 @@
import ase.units
import numpy as np
import torch
+from ipi.engine.barostats import BaroBZP, BaroMTK
from ipi.engine.motion.dynamics import NPTIntegrator, NVEIntegrator, NVTIntegrator
from ipi.utils.depend import dstrip
+from ipi.utils.mathtools import matrix_exp
from ipi.utils.mathtools import random_rotation as random_rotation_matrix
from ipi.utils.messages import info, verbosity
from ipi.utils.units import Constants
@@ -206,35 +208,75 @@ def nvt_stepper(motion, *_, **__):
return nvt_stepper
-def _qbaro(baro):
+def _project_baro(baro):
+ """Zero the cell momentum components held fixed by `hfix`/`vol_constraint`.
+
+ The removed kinetic energy is credited to the barostat thermostat's `ethermo`, as
+ in i-PI's `BaroMTK.pstep`, so that the conserved quantity stays correct.
+ """
+
+ if not baro.vol_constraint and np.all(baro.hmask == 1.0):
+ return
+
+ baro.thermostat.ethermo += baro.kin
+ baro.p *= baro.hmask
+ if baro.vol_constraint:
+ baro.p -= np.eye(3) * np.trace(baro.p) / 3.0
+ baro.thermostat.ethermo -= baro.kin
+
+
+def _qbaro(baro, mode):
"""Propagation step for the cell volume (adjusting atomic positions and momenta)."""
- v = baro.p[0] / baro.m[0]
halfdt = (
baro.qdt
) # this is set to half the inner loop in all integrators that use a barostat
- expq, expp = (np.exp(v * halfdt), np.exp(-v * halfdt))
- baro.nm.qnm[0, :] *= expq
- baro.nm.pnm[0, :] *= expp
- baro.cell.h *= expq
+ if mode == "isotropic":
+ v = baro.p[0] / baro.m[0]
+ expq, expp = (np.exp(v * halfdt), np.exp(-v * halfdt))
+
+ baro.nm.qnm[0, :] *= expq
+ baro.nm.pnm[0, :] *= expp
+ baro.cell.h *= expq
+ elif mode == "flexible":
+ # the barostat thermostat kicks all six components, so `p` cannot be assumed
+ # to be projected here
+ _project_baro(baro)
+
+ v = baro.p / baro.m[0]
+ expq, expp = (matrix_exp(v * halfdt), matrix_exp(-v * halfdt))
+
+ baro.nm.qnm[0] = (dstrip(baro.nm.qnm)[0].reshape(-1, 3) @ expq.T).reshape(-1)
+ baro.nm.pnm[0] = (dstrip(baro.nm.pnm)[0].reshape(-1, 3) @ expp.T).reshape(-1)
+ baro.cell.h = expq @ dstrip(baro.cell.h)
+ else:
+ raise TypeError(f"Unknown barostat mode '{mode}'.")
-def _pbaro(baro):
+def _pbaro(baro, mode):
"""Propagation step for the cell momentum (adjusting atomic positions and momenta)."""
# we are assuming then that p the coupling between p^2 and dp/dt only involves the fast force
dt = baro.pdt[0]
-
- # computes the pressure associated with the forces at the outer level MTS level.
- press = np.trace(baro.stress_mts(0)) / 3.0
- # integerates the kinetic part of the pressure with the force at the inner-most level.
nbeads = baro.beads.nbeads
- baro.p += (
- 3.0
- * dt
- * (baro.cell.V * (press - nbeads * baro.pext) + Constants.kb * baro.temp)
- )
+
+ if mode == "isotropic":
+ # computes the pressure associated with the forces at the outer level MTS level.
+ press = np.trace(baro.stress_mts(0)) / 3.0
+ # integerates the kinetic part of the pressure with the force at the inner-most level.
+ baro.p += (
+ 3.0
+ * dt
+ * (baro.cell.V * (press - nbeads * baro.pext) + Constants.kb * baro.temp)
+ )
+ elif mode == "flexible":
+ stress = np.triu(dstrip(baro.stress_mts(0)) - nbeads * np.eye(3) * baro.pext)
+ baro.p += dt * (baro.cell.V * stress + Constants.kb * baro.temp * baro.L)
+
+ _project_baro(baro)
+ else:
+ raise TypeError(f"Unknown barostat mode '{mode}'.")
def get_npt_stepper(
@@ -251,6 +293,16 @@ def get_npt_stepper(
f"Base i-PI integrator is of type {motion.integrator.__class__.__name__}, use a NPT setup."
)
+ if type(motion.barostat) is BaroBZP:
+ baro_mode = "isotropic"
+ elif type(motion.barostat) is BaroMTK:
+ baro_mode = "flexible"
+ else:
+ raise TypeError(
+ f"Barostat is of type {motion.barostat.__class__.__name__}, use an "
+ "'isotropic' (BZP) or 'flexible' (MTK) barostat."
+ )
+
if use_standard_vv:
# use the standard velocity Verlet integrator
vv_step = get_standard_vv_step(
@@ -273,15 +325,15 @@ def npt_stepper(motion, *_, **__):
info("@flashmd: Barostat thermo", verbosity.debug)
motion.barostat.thermostat.step()
info("@flashmd: Barostat q", verbosity.debug)
- _qbaro(motion.barostat)
+ _qbaro(motion.barostat, baro_mode)
info("@flashmd: Barostat p", verbosity.debug)
- _pbaro(motion.barostat)
+ _pbaro(motion.barostat, baro_mode)
info("@flashmd: FlashVV", verbosity.debug)
vv_step(motion)
info("@flashmd: Barostat p", verbosity.debug)
- _pbaro(motion.barostat)
+ _pbaro(motion.barostat, baro_mode)
info("@flashmd: Barostat q", verbosity.debug)
- _qbaro(motion.barostat)
+ _qbaro(motion.barostat, baro_mode)
info("@flashmd: Barostat thermo", verbosity.debug)
motion.barostat.thermostat.step()
info("@flashmd: Particle thermo", verbosity.debug)
diff --git a/tests/test_ipi.py b/tests/test_ipi.py
new file mode 100644
index 0000000..8907818
--- /dev/null
+++ b/tests/test_ipi.py
@@ -0,0 +1,111 @@
+import ase.build
+import ase.io
+import numpy as np
+import pytest
+import torch
+from ipi.scripting import InteractiveSimulation
+from ipi.utils.depend import dstrip
+
+from flashmd import get_pretrained
+from flashmd.ipi import get_npt_stepper
+
+
+# same models as the other tests, so nothing extra is downloaded
+TIME_STEP = 64
+
+# input file for i-PI to run a simulation in the isotheral-isobaric ensemble with
+# flexible cell vectors. `run_npt` below replaces the {variables} with concrete values
+# for each run. `barostat_extra` is used to test different barostat options.
+INPUT = """
+ 1
+
+ 32123
+
+ metatomic
+ {{model: {model}, template: ./structure.xyz, device: {device}}}
+
+
+
+
+ ./structure.xyz
+ 300
+
+
+ 300
+ 0
+
+
+
+ {time_step}
+ 100
+
+ 2000
+ 100
+ {barostat_extra}
+
+
+
+
+
+"""
+
+
+@pytest.fixture(scope="module")
+def models(tmp_path_factory):
+ """The MLIP (as a file, for i-PI's metatomic driver) and the FlashMD model."""
+ device = "cuda" if torch.cuda.is_available() else "cpu"
+ mlip, flashmd_model = get_pretrained("pet-omatpes-v2", TIME_STEP)
+ mlip_path = tmp_path_factory.mktemp("models") / "mlip.pt"
+ mlip.save(str(mlip_path))
+ return mlip_path, flashmd_model, device
+
+
+@pytest.fixture(autouse=True)
+def in_tmp_path(monkeypatch, tmp_path):
+ monkeypatch.chdir(tmp_path)
+
+
+def run_npt(models, barostat_extra, n_steps=10):
+ """Run the NPT stepper and return the initial and the final cell."""
+ mlip_path, flashmd_model, device = models
+
+ # define a simple test system:
+ # fcc Al in a diagonal cell, so `h[0, 1]` starts at zero up to round-off
+ ase.io.write("structure.xyz", ase.build.bulk("Al", "fcc", cubic=True))
+
+ # create a simulation with a FlashMD NPT stepper
+ simulation = InteractiveSimulation(
+ # replace {variables} in INPUT with concrete values for this run
+ INPUT.format(
+ model=mlip_path,
+ device=device,
+ time_step=TIME_STEP,
+ barostat_extra=barostat_extra,
+ )
+ )
+ motion = simulation.syslist[0].motion
+ step = get_npt_stepper(simulation, flashmd_model, device)
+
+ # run a few steps of NPT and track the cell
+ cell = initial_cell = dstrip(motion.cell.h).copy() # type: ignore
+ for _ in range(n_steps):
+ step(motion)
+ cell = dstrip(motion.cell.h).copy() # type: ignore
+
+ # an unstable piston explodes the cell, and the neighbor list with it, so stop
+ # here instead of leaving the model to churn on a nonsensical structure
+ ratio = np.linalg.det(cell) / np.linalg.det(initial_cell)
+ assert 0.5 < ratio < 2.0, f"the barostat went unstable: V / V0 = {ratio}"
+ return initial_cell, cell
+
+
+def test_hfix_freezes_cell_component(models):
+ """`hfix` keeps a cell component that starts at zero at zero."""
+ initial_cell, cell = run_npt(models, " [ xy ] ")
+ assert cell[0, 1] == pytest.approx(initial_cell[0, 1], abs=1e-10)
+
+
+def test_vol_constraint_conserves_volume(models):
+ """`vol_constraint` keeps the cell volume constant."""
+ initial_cell, cell = run_npt(models, " True ")
+ assert np.linalg.det(cell) == pytest.approx(np.linalg.det(initial_cell), rel=1e-10)
diff --git a/tox.ini b/tox.ini
index 7883750..2970287 100644
--- a/tox.ini
+++ b/tox.ini
@@ -45,6 +45,7 @@ description = Run package tests with pytest
passenv = *
deps =
pytest
+ ipi
changedir = tests
commands =
pytest \