Source code for tmol.score.hbond._hbond_energy_term

import torch

from ._hbond_dependent_term import HBondDependentTerm
from ._params import CompactedHBondDatabase
from .._atom_type_dependent_term import AtomTypeDependentTerm

from tmol.database import ParameterDatabase

from tmol.chemical import RefinedResidueType
from tmol.pose import (
    PackedBlockTypes,
    PoseStack,
)

_PACK_HBOND_ROTAMER_CANDIDATE_WINDOW = 32 * 1024


[docs] class HBondEnergyTerm(AtomTypeDependentTerm, HBondDependentTerm): tile_size: int = 32 hb_param_db: CompactedHBondDatabase def __init__(self, param_db: ParameterDatabase, device: torch.device): super(HBondEnergyTerm, self).__init__(param_db=param_db, device=device) self.tile_size = HBondEnergyTerm.tile_size self.hb_param_db = CompactedHBondDatabase.from_database( param_db.chemical, param_db.scoring.hbond, device ) # Cache of parameter tables converted to whichever coord dtype is used # in scoring. Avoids a per-forward .to(coords_dtype) on three tensors. # Keyed by dtype; stored value is a (pair_param, pair_poly, global_param) tuple. # pair_poly_table is always kept at float64 — the C++ kernel templates # HBondPolynomials on double independently of the Real coord dtype. self._param_tables_by_dtype: dict = {} def _param_tables(self, coords_dtype): tables = self._param_tables_by_dtype.get(coords_dtype) if tables is None: tables = ( self.hb_param_db.pair_param_table.to(coords_dtype), self.hb_param_db.pair_poly_table.to(torch.float64), self.hb_param_db.global_param_table.to(coords_dtype), ) self._param_tables_by_dtype[coords_dtype] = tables return tables @classmethod def class_name(cls): return "HBond" @classmethod def score_types(cls): import tmol.score.terms._hbond_creator return tmol.score.terms._hbond_creator.HBondTermCreator.score_types() def n_bodies(self): return 2 def setup_block_type(self, block_type: RefinedResidueType): super(HBondEnergyTerm, self).setup_block_type(block_type) def setup_packed_block_types(self, packed_block_types: PackedBlockTypes): super(HBondEnergyTerm, self).setup_packed_block_types(packed_block_types) def setup_poses(self, poses: PoseStack): super(HBondEnergyTerm, self).setup_poses(poses) # Count actual residues once, before scoring or CUDA Graph capture. # Padded stack size alone can overestimate work in jagged batches. poses._hbond_allow_split_pairs = ( poses.device.type == "cuda" and poses.block_type_ind.numel() >= 4096 and int((poses.block_type_ind >= 0).sum()) >= 2048 ) def pose_score_hbond(self, *args): from tmol.score.hbond.potentials import hbond_pose_scores common_args = args[:-4] pose_stack, hbond_params, block_pair_scoring, shared_block_neighbors = args[-4:] return hbond_pose_scores( *self._hbond_score_args( common_args, pose_stack, hbond_params, block_pair_scoring ), shared_block_neighbors, getattr(pose_stack, "_hbond_allow_split_pairs", False), ) def _hbond_score_args( self, common_args, pose_stack, hbond_params, block_pair_scoring ): from tmol.score.hbond.potentials import ( gen_hbond_bases, ) coords_dtype = common_args[0].dtype pair_param_table, pair_poly_table, global_param_table = self._param_tables( coords_dtype ) # Pairwise kernels route gradients through derived_atom_inds directly. with torch.no_grad(): derived_coords, derived_atom_inds = gen_hbond_bases( common_args[0], common_args[1], common_args[3], common_args[4], common_args[5], common_args[6], common_args[7], pose_stack.inter_residue_connections, pose_stack.packed_block_types.n_atoms, pose_stack.packed_block_types.n_conn, pose_stack.packed_block_types.conn_atom, pose_stack.packed_block_types.n_all_bonds, pose_stack.packed_block_types.all_bonds, pose_stack.packed_block_types.atom_all_bond_ranges, hbond_params.tile_n_donH, hbond_params.tile_n_acc, hbond_params.tile_donH_inds, hbond_params.tile_acc_inds, hbond_params.tile_acceptor_hybridization, hbond_params.is_hydrogen, ) score_args = ( *common_args, pose_stack.inter_residue_connections, pose_stack.inter_block_bondsep.near_blocks, pose_stack.inter_block_bondsep.bondsep, pose_stack.packed_block_types.n_atoms, pose_stack.packed_block_types.n_conn, pose_stack.packed_block_types.conn_atom, pose_stack.packed_block_types.n_all_bonds, pose_stack.packed_block_types.all_bonds, pose_stack.packed_block_types.atom_all_bond_ranges, pose_stack.packed_block_types.bond_separation, hbond_params.tile_n_donH, hbond_params.tile_n_acc, hbond_params.tile_donH_inds, hbond_params.tile_acc_inds, hbond_params.tile_donorH_type, hbond_params.tile_acceptor_type, hbond_params.tile_acceptor_hybridization, hbond_params.is_hydrogen, pair_param_table, pair_poly_table, global_param_table, derived_coords, derived_atom_inds, block_pair_scoring, ) return score_args def rotamer_score_hbond(self, *args): from tmol.score.hbond.potentials import ( hbond_rotamer_scores, hbond_rotamer_scores_shared, ) common_args = args[:-4] pose_stack, hbond_params, block_pair_scoring, shared_dispatch_indices = args[ -4: ] score_args = self._hbond_score_args( common_args, pose_stack, hbond_params, block_pair_scoring ) score_op = ( hbond_rotamer_scores_shared if shared_dispatch_indices.numel() != 0 else hbond_rotamer_scores ) if shared_dispatch_indices.numel() != 0: score_args = (*score_args, shared_dispatch_indices) return score_op(*score_args) def iter_packing_rotamer_scores(self, *args): """Yield canonical H-bond rotamer scores in bounded candidate pages.""" from tmol.score.hbond.potentials import ( hbond_rotamer_dispatch_page, hbond_rotamer_scores_shared, hbond_rotamer_spheres, ) *score_args, topology_only = args common_args = score_args[:-3] pose_stack, hbond_params, block_pair_scoring = score_args[-3:] first_rot_block_type = common_args[4] n_rots_for_block = common_args[10] rot_offset_for_block = common_args[11] lockstep_group_for_block = common_args[12] rot_spheres, block_spheres = hbond_rotamer_spheres( common_args[0], common_args[1], first_rot_block_type, common_args[7], n_rots_for_block, rot_offset_for_block, pose_stack.packed_block_types.n_atoms, ) n_poses, max_n_blocks = first_rot_block_type.shape candidates_per_pose = max_n_blocks * (max_n_blocks + 1) // 2 n_candidates = n_poses * candidates_per_pose prepared_score_args = None if not topology_only: prepared_score_args = self._hbond_score_args( common_args, pose_stack, hbond_params, block_pair_scoring ) for candidate_begin in range( 0, n_candidates, _PACK_HBOND_ROTAMER_CANDIDATE_WINDOW ): candidate_end = min( candidate_begin + _PACK_HBOND_ROTAMER_CANDIDATE_WINDOW, n_candidates, ) indices = hbond_rotamer_dispatch_page( first_rot_block_type, block_spheres, n_rots_for_block, rot_offset_for_block, rot_spheres, lockstep_group_for_block, self.get_block_neighbor_cutoff(), candidate_begin, candidate_end, ) if topology_only: yield None, indices del indices continue scores, indices = hbond_rotamer_scores_shared(*prepared_score_args, indices) yield scores, indices del scores, indices @property def score_only_in_no_grad(self): # CPU compilers can round score-only and derivative paths differently # (including float64 on ARM). Preserve tracked-input values there. return self.device.type == "cuda" def get_pose_score_term_function(self): return self.pose_score_hbond def get_rotamer_score_term_function(self): return self.rotamer_score_hbond def get_packing_rotamer_score_term_function(self): return self.iter_packing_rotamer_scores def get_block_neighbor_cutoff(self): return 5.5 def accepts_shared_rotamer_dispatch(self): return True def rotamer_dispatch_key(self): return "sphere_overlap" def get_score_term_attributes(self, pose_stack: PoseStack): self.setup_packed_block_types(pose_stack.packed_block_types) annotation = HBondDependentTerm.setup_packed_block_types( self, pose_stack.packed_block_types ) return [pose_stack, annotation]