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
20 changes: 16 additions & 4 deletions tmol/score/genbonded/_genbonded_energy_term.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,8 @@ def __init__(self, param_db: ParameterDatabase, device: torch.device):
# regardless of which block types are later loaded.
self._type_to_idx = self.gen_database.make_type_to_idx()
self._all_type_names = self.gen_database.all_type_names()
self._torsion_params_cache = {}
self._improper_params_cache = {}

@classmethod
def class_name(cls):
Expand Down Expand Up @@ -191,6 +193,7 @@ def resolve_torsion_params(self, block_type: RefinedResidueType, torsions):
"""
kept = []
rows = []
cache = self._torsion_params_cache

for i, j, k, l in torsions:
t1 = self.get_atom_chem_type(block_type, i)
Expand All @@ -206,9 +209,12 @@ def resolve_torsion_params(self, block_type: RefinedResidueType, torsions):

# find_torsion_params tries both forward and reversed directions
# internally and returns the most specific match.
entry = self.gen_database.find_torsion_params(
t1, t2, t3, t4, bond_type_int, is_ring
)
key = (t1, t2, t3, t4, bond_type_int, is_ring)
try:
entry = cache[key]
except KeyError:
entry = self.gen_database.find_torsion_params(*key)
cache[key] = entry
if entry is not None:
kept.append((i, j, k, l))
# Rosetta's calculate_offset() zeros the minimum when the
Expand Down Expand Up @@ -239,6 +245,7 @@ def resolve_improper_params(self, block_type: RefinedResidueType, impropers):
"""
kept = []
rows = []
cache = self._improper_params_cache

for quad in impropers:
center, n1, n2, n3 = quad
Expand All @@ -247,7 +254,12 @@ def resolve_improper_params(self, block_type: RefinedResidueType, impropers):
t2 = self.get_atom_chem_type(block_type, n2)
t3 = self.get_atom_chem_type(block_type, n3)

entry = self.gen_database.find_improper_params(tc, t1, t2, t3)
key = (tc, t1, t2, t3)
try:
entry = cache[key]
except KeyError:
entry = self.gen_database.find_improper_params(*key)
cache[key] = entry
if entry is not None:
kept.append(quad)
rows.append([entry.k, entry.delta])
Expand Down
60 changes: 60 additions & 0 deletions tmol/tests/score/genbonded/test_genbonded_energy_term.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,60 @@
import numpy
import torch

from tmol.database.scoring._genbonded import GenBondedDatabase
from tmol.score.genbonded import GenBondedEnergyTerm


def test_genbonded_parameter_lookups_are_reused(
fresh_default_packed_block_types, default_database, monkeypatch
):
calls = {"torsion": 0, "improper": 0}
original_torsion = GenBondedDatabase.find_torsion_params
original_improper = GenBondedDatabase.find_improper_params

def counted_torsion(database, *args):
calls["torsion"] += 1
return original_torsion(database, *args)

def counted_improper(database, *args):
calls["improper"] += 1
return original_improper(database, *args)

monkeypatch.setattr(GenBondedDatabase, "find_torsion_params", counted_torsion)
monkeypatch.setattr(GenBondedDatabase, "find_improper_params", counted_improper)

term = GenBondedEnergyTerm(default_database, torch.device("cpu"))
blocks = fresh_default_packed_block_types.active_block_types
torsion_block, torsions = next(
(block, subgraphs)
for block in blocks
if (subgraphs := term.find_torsion_subgraphs(block.bond_indices))
)
improper_block, impropers = next(
(block, subgraphs)
for block in blocks
if (subgraphs := term.find_improper_subgraphs(block.bond_indices))
)

first_torsions, first_torsion_params = term.resolve_torsion_params(
torsion_block, torsions
)
first_impropers, first_improper_params = term.resolve_improper_params(
improper_block, impropers
)
first_call_counts = calls.copy()

second_torsions, second_torsion_params = term.resolve_torsion_params(
torsion_block, torsions
)
second_impropers, second_improper_params = term.resolve_improper_params(
improper_block, impropers
)

assert first_call_counts["torsion"] > 0
assert first_call_counts["improper"] > 0
assert calls == first_call_counts
assert second_torsions == first_torsions
assert second_impropers == first_impropers
numpy.testing.assert_array_equal(second_torsion_params, first_torsion_params)
numpy.testing.assert_array_equal(second_improper_params, first_improper_params)
Loading