Source code for tmol.score.ljlk._ljlk_energy_term

import numpy
import torch

from .._annotation_cache import AnnotationKey, cached_annotation, store_annotation
from .._atom_type_dependent_term import AtomTypeDependentTerm
from .._bond_dependent_term import BondDependentTerm
from ._params import LJLKTypeParams, LJLKGlobalParams

from tmol.database import ParameterDatabase
from tmol.score.common import tile_subset_indices
from tmol.score.ljlk import LJLKParamResolver

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


[docs] class LJLKEnergyTerm(AtomTypeDependentTerm, BondDependentTerm): type_params: LJLKTypeParams global_params: LJLKGlobalParams tile_size: int = 32 def __init__(self, param_db: ParameterDatabase, device: torch.device): ljlk_param_resolver = LJLKParamResolver.from_database( param_db.chemical, param_db.scoring.ljlk, device=device ) super(LJLKEnergyTerm, self).__init__(param_db=param_db, device=device) self.type_params = ljlk_param_resolver.type_params self._ljlk_param_resolver = ljlk_param_resolver self.global_params = ljlk_param_resolver.global_params self._max_dis = float(param_db.scoring.ljlk.global_parameters.max_dis) self.tile_size = LJLKEnergyTerm.tile_size self.soft_repulsive = False self.rosetta_typed = frozenset(param_db.scoring.genbonded.rosetta_typed) self._ljlk_block_key = AnnotationKey.from_sources( param_db.chemical, settings=(self.tile_size,) ) self._ljlk_packed_key = AnnotationKey.from_sources( param_db.chemical, param_db.scoring.ljlk, settings=( self.type_params.lj_radius.device, self.tile_size, self.rosetta_typed, ), ) @classmethod def class_name(cls): return "LJLK" @classmethod def score_types(cls): import tmol.score.terms._ljlk_creator return tmol.score.terms._ljlk_creator.LJLKTermCreator.score_types() def n_bodies(self): return 2 def set_options(self, options: dict): self.soft_repulsive = options.get("soft_rep", False) def setup_block_type(self, block_type: RefinedResidueType): self._ljlk_param_resolver.validate_block_type(block_type) atom_params = super(LJLKEnergyTerm, self).setup_block_type(block_type) cached = cached_annotation(block_type, "_ljlk_annotation", self._ljlk_block_key) if cached is not None: return cached heavy_atoms_in_tile, n_in_tile = tile_subset_indices( atom_params[1], self.tile_size ) setattr(block_type, "ljlk_heavy_atoms_in_tile", heavy_atoms_in_tile) setattr(block_type, "ljlk_n_heavy_atoms_in_tile", n_in_tile) return store_annotation( block_type, "_ljlk_annotation", self._ljlk_block_key, (heavy_atoms_in_tile, n_in_tile), fields=("ljlk_heavy_atoms_in_tile", "ljlk_n_heavy_atoms_in_tile"), ) def setup_packed_block_types(self, packed_block_types: PackedBlockTypes): atom_params = super(LJLKEnergyTerm, self).setup_packed_block_types( packed_block_types ) cached = cached_annotation( packed_block_types, "_ljlk_annotation", self._ljlk_packed_key ) if cached is not None: return cached blocks = [ self.setup_block_type(bt) for bt in packed_block_types.active_block_types ] max_n_tiles = (packed_block_types.max_n_atoms - 1) // self.tile_size + 1 heavy_atoms_in_tile = numpy.full( (packed_block_types.n_types, max_n_tiles * self.tile_size), -1, dtype=numpy.int32, ) n_heavy_ats_in_tile = numpy.full( (packed_block_types.n_types, max_n_tiles), 0, dtype=numpy.int32, ) for i, (heavy, counts) in enumerate(blocks): i_n_tiles = counts.shape[0] i_n_tile_ats = i_n_tiles * self.tile_size heavy_atoms_in_tile[i, :i_n_tile_ats] = heavy n_heavy_ats_in_tile[i, :i_n_tiles] = counts setattr( packed_block_types, "ljlk_heavy_atoms_in_tile", torch.as_tensor(heavy_atoms_in_tile, device=self.device), ) setattr( packed_block_types, "ljlk_n_heavy_atoms_in_tile", torch.as_tensor(n_heavy_ats_in_tile, device=self.device), ) # A pair of ligand-typed atoms uses CP_CROSSOVER_3FULL: 1-4 pairs get # full weight (1.0), encoded by rewriting path_dist 3 and 4 to 5 so # connectivity_weight returns 1.0. The convention follows the atoms, # not the residue: a noncanonical whose sidechain is ligand-typed is # still a polymer. Keep bond_separation unchanged for hbond, which # excludes on a binary rule. rosetta_typed = self.rosetta_typed ljlk_bond_separation = packed_block_types.bond_separation.clone() ligand = numpy.zeros(ljlk_bond_separation.shape[:2], dtype=bool) all_ligand = numpy.zeros(len(ligand), dtype=numpy.int32) for i, bt in enumerate(packed_block_types.active_block_types): typed = [a.atom_type not in rosetta_typed for a in bt.atoms] ligand[i, : len(typed)] = typed all_ligand[i] = all(typed) ligand = torch.from_numpy(ligand).to(ljlk_bond_separation.device) both = ligand.unsqueeze(2) & ligand.unsqueeze(1) ljlk_bond_separation[ both & ((ljlk_bond_separation == 3) | (ljlk_bond_separation == 4)) ] = 5 setattr(packed_block_types, "ljlk_bond_separation", ljlk_bond_separation) setattr( packed_block_types, "ljlk_all_atoms_ligand_typed", torch.from_numpy(all_ligand).to(self.device), ) return store_annotation( packed_block_types, "_ljlk_annotation", self._ljlk_packed_key, ( packed_block_types.ljlk_n_heavy_atoms_in_tile, packed_block_types.ljlk_heavy_atoms_in_tile, atom_params[0], ljlk_bond_separation, packed_block_types.ljlk_all_atoms_ligand_typed, ), fields=( "ljlk_n_heavy_atoms_in_tile", "ljlk_heavy_atoms_in_tile", "ljlk_bond_separation", "ljlk_all_atoms_ligand_typed", ), bindings=tuple( (block_type, field) for block_type in packed_block_types.active_block_types for field in ( "ljlk_heavy_atoms_in_tile", "ljlk_n_heavy_atoms_in_tile", ) ), ) def setup_poses(self, poses: PoseStack): super(LJLKEnergyTerm, self).setup_poses(poses) def get_pose_score_term_function(self): from tmol.score.ljlk.potentials import ljlk_pose_scores return ljlk_pose_scores def get_rotamer_score_term_function(self): from tmol.score.ljlk.potentials import ljlk_rotamer_scores return ljlk_rotamer_scores def get_block_neighbor_cutoff(self): return self._max_dis def rotamer_dispatch_key(self): return "sphere_overlap" def get_score_term_attributes(self, pose_stack): annotation = self.setup_packed_block_types(pose_stack.packed_block_types) def _t(ts): return tuple(map(lambda t: t.to(torch.float), ts)) type_params = torch.stack( _t( [ self.type_params.lj_radius, torch.sqrt(self.type_params.lj_wdepth), self.type_params.lk_dgfree, self.type_params.lk_lambda, self.type_params.lk_volume, self.type_params.is_donor, self.type_params.is_hydroxyl, self.type_params.is_polarh, self.type_params.is_acceptor, self.type_params.is_carbon_lk, self.type_params.is_hydrogen, -self.type_params.lk_dgfree / (2 * 5.56832799683 * self.type_params.lk_lambda), 1 / (self.type_params.lk_lambda * self.type_params.lk_lambda), ] ), dim=1, ) global_params = torch.stack( _t( [ self.global_params.max_dis, ( self.global_params.lj_dlin_sigma_factor_soft if self.soft_repulsive else self.global_params.lj_dlin_sigma_factor ), self.global_params.lj_hbond_dis, self.global_params.lj_hbond_OH_donor_dis, self.global_params.lj_hbond_hdis, ] ), dim=1, ) return [ pose_stack.inter_block_bondsep.near_blocks, pose_stack.inter_block_bondsep.bondsep, pose_stack.packed_block_types.n_atoms, annotation[0], annotation[1], annotation[2], pose_stack.packed_block_types.n_conn, pose_stack.packed_block_types.conn_atom, annotation[3], annotation[4], type_params, global_params, # max_dis as host scalar for detect-neighbors call self._max_dis, ]