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,
]