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]