diff --git a/tmol/score/genbonded/_genbonded_energy_term.py b/tmol/score/genbonded/_genbonded_energy_term.py index 56a765fcc..b9eaeb96a 100644 --- a/tmol/score/genbonded/_genbonded_energy_term.py +++ b/tmol/score/genbonded/_genbonded_energy_term.py @@ -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): @@ -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) @@ -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 @@ -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 @@ -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]) diff --git a/tmol/tests/score/genbonded/test_genbonded_energy_term.py b/tmol/tests/score/genbonded/test_genbonded_energy_term.py new file mode 100644 index 000000000..f4f9872db --- /dev/null +++ b/tmol/tests/score/genbonded/test_genbonded_energy_term.py @@ -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)