diff --git a/src/somd2/config/_config.py b/src/somd2/config/_config.py index d712e40..467373d 100644 --- a/src/somd2/config/_config.py +++ b/src/somd2/config/_config.py @@ -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 @@ -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 @@ -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 @@ -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 ( @@ -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 @@ -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 ) diff --git a/src/somd2/runner/_base.py b/src/somd2/runner/_base.py index 91383d3..baa0701 100644 --- a/src/somd2/runner/_base.py +++ b/src/somd2/runner/_base.py @@ -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 diff --git a/tests/runner/test_config.py b/tests/runner/test_config.py index ce0168e..5abdcf2 100644 --- a/tests/runner/test_config.py +++ b/tests/runner/test_config.py @@ -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)