Skip to content

Commit 9ff3631

Browse files
committed
feat: merge_these_atom_params
1 parent 9d1a137 commit 9ff3631

3 files changed

Lines changed: 40 additions & 14 deletions

File tree

meeko/molsetup.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1176,7 +1176,7 @@ def get_neighbors(self, atom_index: int):
11761176

11771177
# endregion
11781178

1179-
def merge_terminal_atoms(self, indices, merge_rmin_half=False) -> None:
1179+
def merge_terminal_atoms(self, indices, merge_rmin_half=False, merge_these_atom_params=()) -> None:
11801180
"""
11811181
Primarily for merging hydrogens, but will merge the data for any atom or pseudoatom that is bonded to only one
11821182
other atom.
@@ -1194,6 +1194,8 @@ def merge_terminal_atoms(self, indices, merge_rmin_half=False) -> None:
11941194
"""
11951195
if merge_rmin_half and "rmin_half" not in self.atom_params:
11961196
raise ValueError("can't merge rmin_half because it's not in atom_params")
1197+
if merge_rmin_half and "rmin_half" in merge_these_atom_params:
1198+
raise ValueError("merge_rmin_half=True wants to merge by volume, rmin_half in merge_these_atom_params wants to merge by radius")
11971199
for index in indices:
11981200
if len(self.get_neighbors(index)) != 1:
11991201
msg = "Atempted to merge atom %d with %d neighbors. "
@@ -1204,6 +1206,10 @@ def merge_terminal_atoms(self, indices, merge_rmin_half=False) -> None:
12041206
self.atoms[neighbor_index].charge += self.get_charge(index)
12051207
self.atoms[index].charge = 0.0
12061208
self.atoms[index].is_ignore = True
1209+
for param_key in merge_these_atom_params:
1210+
value = self.atom_params[param_key][index]
1211+
self.atom_params[param_key][neighbor_index] += value
1212+
self.atom_params[param_key][index] = 0.0
12071213
if not merge_rmin_half:
12081214
continue
12091215
r_neigh = self.atom_params["rmin_half"][neighbor_index]

meeko/preparation.py

Lines changed: 15 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -84,6 +84,7 @@ class MoleculePreparation:
8484
def __init__(
8585
self,
8686
merge_these_atom_types=("H",),
87+
merge_these_atom_params=(),
8788
merge_rmin_half=False,
8889
hydrate=False,
8990
flexible_amides=False,
@@ -120,6 +121,7 @@ def __init__(
120121
Parameters
121122
----------
122123
merge_these_atom_types
124+
merge_these_atom_params
123125
merge_rmin_half
124126
hydrate
125127
flexible_amides
@@ -162,6 +164,7 @@ def __init__(
162164

163165
self.deprecated_setup_access = None
164166
self.merge_these_atom_types = merge_these_atom_types
167+
self.merge_these_atom_params = merge_these_atom_params
165168
self.merge_rmin_half = merge_rmin_half
166169
self.hydrate = hydrate
167170
self.flexible_amides = flexible_amides
@@ -605,6 +608,17 @@ def prepare(
605608
self.dihedral_params,
606609
)
607610

611+
if self.crippen or self.crippen_as_solpar:
612+
add_crippen_to_molsetup(setup)
613+
if not self.crippen:
614+
setup.atom_params["ad4_sol_par"] = setup.atom_params.pop("crippen")
615+
elif self.crippen_as_solpar:
616+
setup.atom_params["ad4_sol_par"] = [value for value in setup.atom_params["crippen"]]
617+
618+
if self.override_ad4sol_par_including_q:
619+
qasp = self.override_ad4sol_par_including_q_qasp
620+
set_ad4sol_par_including_q(setup, qasp)
621+
608622
# Convert molecule to graph and apply trained Espaloma model
609623
# skip if charges are read from template
610624
if self.dihedral_model == "espaloma" or (self.charge_model == "espaloma" and self.compute_charges):
@@ -634,7 +648,7 @@ def prepare(
634648
for atom in setup.atoms:
635649
if atom.atom_type == atype_to_merge:
636650
indices.add(atom.index)
637-
setup.merge_terminal_atoms(indices, self.merge_rmin_half)
651+
setup.merge_terminal_atoms(indices, self.merge_rmin_half, self.merge_these_atom_params)
638652

639653
# 3. assign bond types
640654
# - all single bonds rotatable except some amides and SMARTS rigidification
@@ -672,17 +686,6 @@ def prepare(
672686
new_atom_info = orig_pdbinfo._replace(name=new_name)
673687
atom.pdbinfo = new_atom_info
674688

675-
if self.crippen or self.crippen_as_solpar:
676-
add_crippen_to_molsetup(setup)
677-
if not self.crippen:
678-
setup.atom_params["ad4_sol_par"] = setup.atom_params.pop("crippen")
679-
elif self.crippen_as_solpar:
680-
setup.atom_params["ad4_sol_par"] = [value for value in setup.atom_params["crippen"]]
681-
682-
if self.override_ad4sol_par_including_q:
683-
qasp = self.override_ad4sol_par_including_q_qasp
684-
set_ad4sol_par_including_q(setup, qasp)
685-
686689
if self.reactive_smarts is None:
687690
setups = [setup]
688691
else:

test/parameterization_test.py

Lines changed: 18 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -114,13 +114,30 @@ def test_mbondi():
114114
rdDistGeom.EmbedMolecule(mol, etkdg_v3)
115115
mk_prep = MoleculePreparation(load_atom_params="gb_mbondi")
116116
molsetup = mk_prep(mol)[0]
117-
print(molsetup.atom_params)
118117
assert 1.3 in molsetup.atom_params["mbondi_radius"] # H bound to N or C
119118
assert 0.8 in molsetup.atom_params["mbondi_radius"] # H bound to O or S
120119
assert 1.2 not in molsetup.atom_params["mbondi_radius"] # other H
121120
assert molsetup.atom_params["gb_screen"].count(0.72) == 4 # C
122121
assert molsetup.atom_params["gb_screen"].count(0.79) == 2 # N
123122

123+
def test_merge_crippen():
124+
mol = Chem.AddHs(Chem.MolFromSmiles("OC(F)F"))
125+
etkdg_v3 = rdDistGeom.ETKDGv3()
126+
rdDistGeom.EmbedMolecule(mol, etkdg_v3)
127+
mk_prep = MoleculePreparation(crippen=True, load_atom_params=["ad4_types"])
128+
molsetup = mk_prep(mol)[0]
129+
mk_prep = MoleculePreparation(crippen=True, load_atom_params=["ad4_types"], merge_these_atom_params=["crippen"])
130+
molsetup_merge = mk_prep(mol)[0]
131+
h_mask = np.array([atom.atom_type == "H" for atom in molsetup.atoms])
132+
crippen_orig = np.array(molsetup.atom_params["crippen"])
133+
crippen_merge = np.array(molsetup_merge.atom_params["crippen"])
134+
abs_diff = np.abs(crippen_merge - crippen_orig)
135+
assert abs(sum(crippen_orig) - sum(crippen_merge)) < 1e-8
136+
assert np.max(crippen_orig[h_mask]) > 1e-2
137+
assert np.max(crippen_merge[h_mask]) < 1e-8
138+
assert np.min(abs_diff[h_mask]) > 1e-2
139+
assert np.max(abs_diff[~h_mask]) > 1e-2
140+
124141
def test_override_ad4sol_par_including_q():
125142
qasp = 0.01097
126143
mol = Chem.AddHs(Chem.MolFromSmiles("c1nc[nH]c1CO"))

0 commit comments

Comments
 (0)