Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
44 changes: 41 additions & 3 deletions src/somd2/config/_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -671,7 +671,6 @@ def __init__(
self.num_lambda = num_lambda
self.lambda_values = lambda_values
self.lambda_energy = lambda_energy
self.lambda_schedule = lambda_schedule
self.charge_scale_factor = charge_scale_factor
self.swap_end_states = swap_end_states
self.shift_coulomb = shift_coulomb
Expand Down Expand Up @@ -732,6 +731,7 @@ def __init__(
self.taylor_power = taylor_power
self.beutler_alpha = beutler_alpha
self.beutler_fix_epsilon = beutler_fix_epsilon
self.lambda_schedule = lambda_schedule
self.somd1_compatibility = somd1_compatibility
self.pert_file = pert_file
self.auto_fix_minimise = auto_fix_minimise
Expand Down Expand Up @@ -1179,6 +1179,8 @@ def lambda_energy(self, lambda_energy):

@property
def lambda_schedule(self):
if self._lambda_schedule is None:
return self._build_deferred_schedule()
return self._lambda_schedule

@lambda_schedule.setter
Expand All @@ -1200,7 +1202,9 @@ def lambda_schedule(self, lambda_schedule):
self._lambda_schedule = _LambdaSchedule.standard_morph()
self._lambda_schedule_name = "standard_morph"
elif keyword == "charge_scaled_morph":
self._lambda_schedule = _LambdaSchedule.charge_scaled_morph(0.2)
self._lambda_schedule = _LambdaSchedule.charge_scaled_morph(
self._charge_scale_factor
)
self._lambda_schedule_name = "charge_scaled_morph"
elif keyword == "ring_break_morph":
from .._utils._schedules import (
Expand Down Expand Up @@ -1243,6 +1247,38 @@ def lambda_schedule(self, lambda_schedule):
self._lambda_schedule = _LambdaSchedule.standard_morph()
self._lambda_schedule_name = "standard_morph"

def _build_deferred_schedule(self):
"""
Build the ABFE schedules, which depend on the soft-core settings and
the lever of any Boresch restraint.
"""
fix_epsilon = self._softcore_form == "beutler" and self._beutler_fix_epsilon
restraint_lever = self._boresch_restraint_lever() or "split"
if self._lambda_schedule_name == "annihilate":
from .._utils._schedules import annihilate as _annihilate

return _annihilate(fix_epsilon=fix_epsilon, restraint_lever=restraint_lever)
elif self._lambda_schedule_name == "decouple":
from .._utils._schedules import decouple as _decouple

return _decouple(fix_epsilon=fix_epsilon, restraint_lever=restraint_lever)

def _boresch_restraint_lever(self):
"""
Return the lever of the Boresch restraints, or None if there are none.
"""
levers = {
restraint.restraint_lever()
for restraint in self._restraints or []
if isinstance(restraint, _sr.mm.BoreschRestraints)
}
if len(levers) > 1:
raise ValueError(
"All Boresch restraints must use the same 'restraint_lever', "
f"got {', '.join(sorted(levers))}."
)
return levers.pop() if levers else None

@property
def charge_scale_factor(self):
return self._charge_scale_factor
Expand All @@ -1256,7 +1292,9 @@ def charge_scale_factor(self, charge_scale_factor):
raise ValueError("'charge_scale_factor' must be a float")
self._charge_scale_factor = charge_scale_factor
# Update the lambda schedule if it is charge scaled morph.
if self._lambda_schedule == "charge_scaled_morph":
if getattr(self, "_lambda_schedule_name", None) == "charge_scaled_morph":
from sire.cas import LambdaSchedule as _LambdaSchedule

self._lambda_schedule = _LambdaSchedule.charge_scaled_morph(
self._charge_scale_factor
)
Expand Down
16 changes: 5 additions & 11 deletions src/somd2/runner/_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -442,17 +442,11 @@ def __init__(self, system, config):
self._config._extra_args["beutler_alpha"] = self._config.beutler_alpha

# Build deferred schedules now that the softcore form is known.
fix_epsilon = (
self._config.softcore_form == "beutler" and self._config.beutler_fix_epsilon
)
if self._config._lambda_schedule_name == "annihilate":
from .._utils._schedules import annihilate as _annihilate

self._config._lambda_schedule = _annihilate(fix_epsilon=fix_epsilon)
elif self._config._lambda_schedule_name == "decouple":
from .._utils._schedules import decouple as _decouple

self._config._lambda_schedule = _decouple(fix_epsilon=fix_epsilon)
if self._config._lambda_schedule_name in ("annihilate", "decouple"):
self._config._lambda_schedule = self._config._build_deferred_schedule()
if self._is_abfe_bound:
restraint_lever = self._config._boresch_restraint_lever() or "split"
_logger.info(f"Using the '{restraint_lever}' Boresch restraint lever.")

# Alchemical ions are real (non-ghost) atoms mutating identity (e.g. a
# water oxygen turning into Na+), not ghost-atom decoupling/annihilation
Expand Down
30 changes: 30 additions & 0 deletions tests/runner/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -158,6 +158,36 @@ def test_lambda_schedule_input_forms():
with pytest.raises(ValueError, match="Unable to interpret"):
Config(lambda_schedule="not_a_schedule")


def test_abfe_schedule_restraint_lever():
"""Validate that the ABFE schedules match the Boresch restraint lever."""
import pytest

split_levers = {"restraint_dihedral", "restraint_distance_angle"}

# No restraint, so the lever used by the restraint search.
config = Config(lambda_schedule="annihilate")
assert split_levers.issubset(config.lambda_schedule.get_levers())

combined = sr.mm.BoreschRestraints()
combined.set_restraint_lever("combined")
split = sr.mm.BoreschRestraints()
split.set_restraint_lever("split")

for name in ["annihilate", "decouple"]:
config = Config(lambda_schedule=name, restraints=combined)
levers = config.lambda_schedule.get_levers()
assert "restraint" in levers
assert split_levers.isdisjoint(levers)

config = Config(lambda_schedule=name, restraints=split)
assert split_levers.issubset(config.lambda_schedule.get_levers())

with pytest.raises(ValueError, match="restraint_lever"):
Config(
lambda_schedule="annihilate", restraints=[combined, split]
).lambda_schedule

# A stream file holding the wrong type of object.
wrong_path = os.path.join(tmpdir, "wrong.s3")
sr.stream.save(sr.cas.Symbol("x"), wrong_path)
Expand Down