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
3 changes: 2 additions & 1 deletion src/sire/restraints/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
"distance",
"boresch",
"boresch_search",
"check_boresch_search",
"inverse_bond",
"inverse_distance",
"rmsd",
Expand All @@ -26,4 +27,4 @@
rmsd,
)
from ._standard_state_correction import get_standard_state_correction
from ._boresch_search import boresch_search
from ._boresch_search import boresch_search, check_boresch_search
181 changes: 116 additions & 65 deletions src/sire/restraints/_boresch_search.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
Automatic Boresch restraint generation for ABFE simulations.
"""

__all__ = ["boresch_search"]
__all__ = ["boresch_search", "check_boresch_search"]

from collections import deque as _deque

Expand Down Expand Up @@ -51,6 +51,75 @@ def _build_triplets(connectivity, anchor_idx, mol, is_lig, ghost_elem, h_elem):
return triplets


def _find_ligand(system):
"""
Return the perturbable molecule and its non-ghost, non-H atom indices
(lambda=0 state), validating that there is enough to build a restraint.
"""
from ..legacy import Mol as _SireMol

pert_mols = system.molecules("property is_perturbable")
if pert_mols.num_molecules() != 1:
raise ValueError(
"System must contain exactly one perturbable molecule for Boresch "
f"restraint generation; found {pert_mols.num_molecules()}."
)
pert_mol = pert_mols.molecule(0)

ghost_elem = _SireMol.Element(0)
h_elem = _SireMol.Element("H")

heavy_idxs = []
for atom in pert_mol.atoms():
elem0 = atom.property("element0")
if elem0 != ghost_elem and elem0 != h_elem:
heavy_idxs.append(atom.index())

if len(heavy_idxs) < 3:
raise ValueError(
f"Ligand has only {len(heavy_idxs)} non-ghost, non-hydrogen atom(s); "
"need at least 3 for Boresch restraints."
)

return pert_mol, heavy_idxs


def _find_rxrx_ligand_atoms(pert_mol, heavy_idxs):
"""
Return the RXRX candidate anchor atoms (non-terminal heavy atoms) and the
ligand N/O atoms available as hydrogen-bond partners.
"""
from ..legacy import Mol as _SireMol

ghost_elem = _SireMol.Element(0)
h_elem = _SireMol.Element("H")
hbond_elems = (_SireMol.Element("N"), _SireMol.Element("O"))
connectivity = pert_mol.connectivity()

candidate_set = {
idx
for idx in heavy_idxs
if len(_nonH_bonded(connectivity, idx, pert_mol, True, ghost_elem, h_elem)) >= 2
}
if not candidate_set:
raise ValueError(
"No non-terminal ligand heavy atoms (bonded to >= 2 heavy atoms) "
"were found for RXRX restraint search."
)

hbond_idxs = [
idx
for idx in heavy_idxs
if pert_mol.atom(idx).property("element0") in hbond_elems
]
if not hbond_idxs:
raise ValueError(
"Ligand has no N/O heavy atoms to use as hydrogen-bond partners."
)

return candidate_set, hbond_idxs


def _circular_mean_std(values):
"""
Circular mean/std (radians) for periodic dihedral samples.
Expand Down Expand Up @@ -470,8 +539,6 @@ def _boresch_search_rxrx(

ghost_elem = _SireMol.Element(0)
h_elem = _SireMol.Element("H")
n_elem = _SireMol.Element("N")
o_elem = _SireMol.Element("O")

# -------------------------------------------------------------------------
# 1. Locate the perturbable molecule (ligand) and its candidate atoms.
Expand All @@ -480,48 +547,17 @@ def _boresch_search_rxrx(
# pool; the candidate anchor set is restricted to non-terminal atoms
# (bonded to >= 2 non-H heavy atoms), matching the RXRX paper's Figure 2B.
# -------------------------------------------------------------------------
pert_mols = system.molecules("property is_perturbable")
if pert_mols.num_molecules() != 1:
raise ValueError(
"System must contain exactly one perturbable molecule for Boresch "
f"restraint generation; found {pert_mols.num_molecules()}."
)
pert_mol = pert_mols.molecule(0)
pert_mol, lig_heavy_idxs = _find_ligand(system)
pert_mol_num = pert_mol.number()
lig_connectivity = pert_mol.connectivity()

lig_heavy_idxs = []
for atom in pert_mol.atoms():
elem0 = atom.property("element0")
if elem0 != ghost_elem and elem0 != h_elem:
lig_heavy_idxs.append(atom.index())

if len(lig_heavy_idxs) < 3:
raise ValueError(
f"Ligand has only {len(lig_heavy_idxs)} non-ghost, non-hydrogen atom(s); "
"need at least 3 for Boresch restraints."
)
candidate_set, lig_hbond_idxs = _find_rxrx_ligand_atoms(pert_mol, lig_heavy_idxs)

# Lambda=0 atomic masses, used for the mass-weighted centre of mass below
# (topology doesn't change between frames, so this is computed once).
lig_masses = _np.array(
[float(pert_mol.atom(idx).property("mass0").value()) for idx in lig_heavy_idxs]
)

candidate_set = set()
for idx in lig_heavy_idxs:
n_heavy = len(
_nonH_bonded(lig_connectivity, idx, pert_mol, True, ghost_elem, h_elem)
)
if n_heavy >= 2:
candidate_set.add(idx)

if not candidate_set:
raise ValueError(
"No non-terminal ligand heavy atoms (bonded to >= 2 heavy atoms) "
"were found for RXRX restraint search."
)

def _attached_H(connectivity, at_idx, mol, is_lig):
bonded = connectivity.connections_to(at_idx)
result = []
Expand All @@ -531,15 +567,6 @@ def _attached_H(connectivity, at_idx, mol, is_lig):
result.append(b_idx)
return result

lig_hbond_idxs = [
idx
for idx in lig_heavy_idxs
if pert_mol.atom(idx).property("element0") in (n_elem, o_elem)
]
if not lig_hbond_idxs:
raise ValueError(
"Ligand has no N/O heavy atoms to use as hydrogen-bond partners."
)
lig_donor_H = {
idx: _attached_H(lig_connectivity, idx, pert_mol, True)
for idx in lig_hbond_idxs
Expand Down Expand Up @@ -1105,31 +1132,12 @@ def _boresch_search_aldeghi(
# -------------------------------------------------------------------------
# 1. Locate the perturbable molecule (ligand).
# -------------------------------------------------------------------------
pert_mols = system.molecules("property is_perturbable")
if pert_mols.num_molecules() != 1:
raise ValueError(
"System must contain exactly one perturbable molecule for Boresch "
f"restraint generation; found {pert_mols.num_molecules()}."
)
pert_mol = pert_mols.molecule(0)
pert_mol, lig_atom_idxs = _find_ligand(system)
pert_mol_num = pert_mol.number()

ghost_elem = _SireMol.Element(0)
h_elem = _SireMol.Element("H")

# Collect non-ghost, non-H ligand AtomIdx values (lambda=0 state).
lig_atom_idxs = []
for atom in pert_mol.atoms():
elem0 = atom.property("element0")
if elem0 != ghost_elem and elem0 != h_elem:
lig_atom_idxs.append(atom.index())

if len(lig_atom_idxs) < 3:
raise ValueError(
f"Ligand has only {len(lig_atom_idxs)} non-ghost, non-hydrogen atom(s); "
"need at least 3 for Boresch restraints."
)

lig_connectivity = pert_mol.connectivity()

# -------------------------------------------------------------------------
Expand Down Expand Up @@ -1587,3 +1595,46 @@ def boresch_search(
raise ValueError(
f"Unknown 'protocol'={protocol!r}; must be 'rxrx' or 'aldeghi'."
)


def check_boresch_search(system, protocol="rxrx"):
"""
Check that the ligand in a system is suitable for ``boresch_search``
with the given protocol, without requiring any trajectory frames. This
allows topology-level failures, such as a ligand with no N/O atoms to
act as hydrogen-bond partners for the "rxrx" protocol, to be caught
before the restraint-search trajectory is generated.

Parameters
----------

system : sire.system.System
A Sire system containing exactly one perturbable molecule.

protocol : str
The restraint search protocol to check: ``"rxrx"`` (default) or
``"aldeghi"``.

Raises
------

ValueError
If ``boresch_search`` is guaranteed to fail for this system and
protocol. The message matches the one ``boresch_search`` would raise.
"""
from ..system import System as _System

if not isinstance(system, _System):
raise TypeError(
f"'system' must be of type 'sire.system.System', got {type(system)}"
)

if protocol not in ("rxrx", "aldeghi"):
raise ValueError(
f"Unknown 'protocol'={protocol!r}; must be 'rxrx' or 'aldeghi'."
)

pert_mol, heavy_idxs = _find_ligand(system)

if protocol == "rxrx":
_find_rxrx_ligand_atoms(pert_mol, heavy_idxs)
47 changes: 47 additions & 0 deletions tests/restraints/test_boresch_search.py
Original file line number Diff line number Diff line change
Expand Up @@ -241,3 +241,50 @@ def test_tight_cutoff_raises(self, abfe_system):

with pytest.raises(ValueError, match="cutoff"):
boresch_search(abfe_system, protocol="aldeghi", cutoff="0.1 A")


@pytest.fixture(scope="module")
def no_hbond_system(abfe_system):
"""The test system with the ligand's N/O atoms changed to carbon."""
mols = abfe_system.clone()
cursor = mols.molecules("property is_perturbable")[0].cursor()
for atom in cursor.atoms():
if atom["element0"].symbol() in ("N", "O"):
atom["element0"] = sr.mol.Element("C")
mols.update(cursor.commit())
return mols


class TestCheckBoreschSearch:
@pytest.mark.parametrize("protocol", PROTOCOLS)
def test_valid_system_passes(self, abfe_system, protocol):
from sire.restraints import check_boresch_search

check_boresch_search(abfe_system, protocol=protocol)

def test_no_frames_required(self, abfe_system):
from sire.restraints import check_boresch_search

mols = abfe_system.clone()
mols.delete_all_frames()
check_boresch_search(mols)

def test_rxrx_no_hbond_partners_raises(self, no_hbond_system):
from sire.restraints import boresch_search, check_boresch_search

with pytest.raises(ValueError, match="N/O"):
check_boresch_search(no_hbond_system, protocol="rxrx")

with pytest.raises(ValueError, match="N/O"):
boresch_search(no_hbond_system, protocol="rxrx")

def test_aldeghi_no_hbond_partners_passes(self, no_hbond_system):
from sire.restraints import check_boresch_search

check_boresch_search(no_hbond_system, protocol="aldeghi")

def test_invalid_protocol_raises(self, abfe_system):
from sire.restraints import check_boresch_search

with pytest.raises(ValueError, match="protocol"):
check_boresch_search(abfe_system, protocol="nonsense")