OneScience-Group/DiffDock
031
1from e3nn import o32import torch3from torch import nn4from torch.nn import functional as F5from torch_cluster import radius, radius_graph6from torch_geometric.utils import subgraph7from torch_scatter import scatter_mean8import numpy as np9 10from onescience.datapipes.diffdock.process_mols import lig_feature_dims, rec_residue_feature_dims, rec_atom_feature_dims11from onescience.utils.diffdock import so3, torus12 13from .layers import GaussianSmearing, AtomEncoder14from .tensor_layers import get_irrep_seq, TensorProductConvLayer15 16try:17 from esm.pretrained import load_model_and_alphabet18except ImportError:19 load_model_and_alphabet = None20 21AGGREGATORS = {"mean": lambda x: torch.mean(x, dim=1),22 "max": lambda x: torch.max(x, dim=1)[0],23 "min": lambda x: torch.min(x, dim=1)[0],24 "std": lambda x: torch.std(x, dim=1)}25 26 27class AAModel(torch.nn.Module):28 def __init__(self, t_to_sigma, device, timestep_emb_func, in_lig_edge_features=4, sigma_embed_dim=32, sh_lmax=2,29 ns=16, nv=4, num_conv_layers=2, lig_max_radius=5, rec_max_radius=30, cross_max_distance=250,30 center_max_distance=30, distance_embed_dim=32, cross_distance_embed_dim=32, no_torsion=False,31 scale_by_sigma=True, norm_by_sigma=True, use_second_order_repr=False, batch_norm=True,32 dynamic_max_cross=False, dropout=0.0, smooth_edges=False, odd_parity=False,33 separate_noise_schedule=False, lm_embedding_type=None, confidence_mode=False,34 confidence_dropout=0, confidence_no_batchnorm = False,35 asyncronous_noise_schedule=False, affinity_prediction=False, parallel=1,36 parallel_aggregators="mean max min std", num_confidence_outputs=1, atom_num_confidence_outputs=1, fixed_center_conv=False,37 no_aminoacid_identities=False, include_miscellaneous_atoms=False,38 differentiate_convolutions=True, tp_weights_layers=2, num_prot_emb_layers=0,39 reduce_pseudoscalars=False, embed_also_ligand=False, atom_confidence=False, sidechain_pred=False,40 depthwise_convolution=False, crop_beyond=None):41 super(AAModel, self).__init__()42 assert (not no_aminoacid_identities) or (lm_embedding_type is None), "no language model emb without identities"43 assert not sidechain_pred, "sidechain prediction not implemented/makes sense for all atom model"44 assert not depthwise_convolution, "depthwise convolution not implemented for all atom model"45 if parallel > 1: assert affinity_prediction46 47 self.t_to_sigma = t_to_sigma48 self.in_lig_edge_features = in_lig_edge_features49 sigma_embed_dim *= (3 if separate_noise_schedule else 1)50 self.sigma_embed_dim = sigma_embed_dim51 self.lig_max_radius = lig_max_radius52 self.rec_max_radius = rec_max_radius53 self.cross_max_distance = cross_max_distance54 self.dynamic_max_cross = dynamic_max_cross55 self.center_max_distance = center_max_distance56 self.distance_embed_dim = distance_embed_dim57 self.cross_distance_embed_dim = cross_distance_embed_dim58 self.sh_irreps = o3.Irreps.spherical_harmonics(lmax=sh_lmax)59 self.ns, self.nv = ns, nv60 self.scale_by_sigma = scale_by_sigma61 self.norm_by_sigma = norm_by_sigma62 self.device = device63 self.no_torsion = no_torsion64 self.smooth_edges = smooth_edges65 self.odd_parity = odd_parity66 self.num_conv_layers = num_conv_layers67 self.timestep_emb_func = timestep_emb_func68 self.separate_noise_schedule = separate_noise_schedule69 self.confidence_mode = confidence_mode70 self.num_conv_layers = num_conv_layers71 self.num_prot_emb_layers = num_prot_emb_layers72 self.asyncronous_noise_schedule = asyncronous_noise_schedule73 self.affinity_prediction = affinity_prediction74 self.parallel, self.parallel_aggregators = parallel, parallel_aggregators.split(' ')75 self.fixed_center_conv = fixed_center_conv76 self.no_aminoacid_identities = no_aminoacid_identities77 self.differentiate_convolutions = differentiate_convolutions78 self.reduce_pseudoscalars = reduce_pseudoscalars79 self.atom_confidence = atom_confidence80 self.atom_num_confidence_outputs = atom_num_confidence_outputs81 self.crop_beyond = crop_beyond82 83 self.lm_embedding_type = lm_embedding_type84 if lm_embedding_type is None:85 lm_embedding_dim = 086 elif lm_embedding_type == "precomputed":87 lm_embedding_dim=128088 else:89 if load_model_and_alphabet is None:90 raise ImportError(91 "esm is required when lm_embedding_type is not None and not 'precomputed'."92 )93 lm, alphabet = load_model_and_alphabet(lm_embedding_type)94 self.batch_converter = alphabet.get_batch_converter()95 lm.lm_head = torch.nn.Identity()96 lm.contact_head = torch.nn.Identity()97 lm_embedding_dim = lm.embed_dim98 self.lm = lm99 100 # embedding layers101 atom_encoder_class = AtomEncoder102 self.lig_node_embedding = atom_encoder_class(emb_dim=ns, feature_dims=lig_feature_dims, sigma_embed_dim=sigma_embed_dim)103 self.lig_edge_embedding = nn.Sequential(nn.Linear(in_lig_edge_features + sigma_embed_dim + distance_embed_dim, ns),nn.ReLU(),nn.Dropout(dropout),nn.Linear(ns, ns))104 105 self.rec_sigma_embedding = nn.Sequential(nn.Linear(sigma_embed_dim, ns), nn.ReLU(), nn.Dropout(dropout), nn.Linear(ns, ns))106 self.rec_node_embedding = atom_encoder_class(emb_dim=ns, feature_dims=rec_residue_feature_dims, sigma_embed_dim=0, lm_embedding_dim=lm_embedding_dim)107 self.rec_edge_embedding = nn.Sequential(nn.Linear(distance_embed_dim, ns), nn.ReLU(), nn.Dropout(dropout), nn.Linear(ns, ns))108 self.atom_node_embedding = atom_encoder_class(emb_dim=ns, feature_dims=rec_atom_feature_dims, sigma_embed_dim=0)109 self.atom_edge_embedding = nn.Sequential(nn.Linear(distance_embed_dim, ns), nn.ReLU(), nn.Dropout(dropout), nn.Linear(ns, ns))110 111 self.lr_edge_embedding = nn.Sequential(nn.Linear(sigma_embed_dim + cross_distance_embed_dim, ns), nn.ReLU(), nn.Dropout(dropout),nn.Linear(ns, ns))112 self.ar_edge_embedding = nn.Sequential(nn.Linear(distance_embed_dim, ns), nn.ReLU(), nn.Dropout(dropout),nn.Linear(ns, ns))113 self.la_edge_embedding = nn.Sequential(nn.Linear(sigma_embed_dim + cross_distance_embed_dim, ns), nn.ReLU(), nn.Dropout(dropout),nn.Linear(ns, ns))114 115 self.lig_distance_expansion = GaussianSmearing(0.0, lig_max_radius, distance_embed_dim)116 self.rec_distance_expansion = GaussianSmearing(0.0, rec_max_radius, distance_embed_dim)117 self.cross_distance_expansion = GaussianSmearing(0.0, cross_max_distance, cross_distance_embed_dim)118 119 irrep_seq = get_irrep_seq(ns, nv, use_second_order_repr, reduce_pseudoscalars)120 assert not include_miscellaneous_atoms, "currently not supported"121 122 rec_emb_layers = []123 for i in range(num_prot_emb_layers):124 in_irreps = irrep_seq[min(i, len(irrep_seq) - 1)]125 out_irreps = irrep_seq[min(i + 1, len(irrep_seq) - 1)]126 layer = TensorProductConvLayer(127 in_irreps=in_irreps,128 sh_irreps=self.sh_irreps,129 out_irreps=out_irreps,130 n_edge_features=3 * ns,131 hidden_features=3 * ns,132 residual=True,133 batch_norm=batch_norm,134 dropout=dropout,135 faster=sh_lmax == 1 and not use_second_order_repr,136 tp_weights_layers=tp_weights_layers,137 edge_groups=1 if not differentiate_convolutions else 4,138 )139 rec_emb_layers.append(layer)140 self.rec_emb_layers = nn.ModuleList(rec_emb_layers)141 142 self.embed_also_ligand = embed_also_ligand143 if embed_also_ligand:144 lig_emb_layers = []145 for i in range(num_prot_emb_layers):146 in_irreps = irrep_seq[min(i, len(irrep_seq) - 1)]147 out_irreps = irrep_seq[min(i + 1, len(irrep_seq) - 1)]148 layer = TensorProductConvLayer(149 in_irreps=in_irreps,150 sh_irreps=self.sh_irreps,151 out_irreps=out_irreps,152 n_edge_features=3 * ns,153 hidden_features=3 * ns,154 residual=True,155 batch_norm=batch_norm,156 dropout=dropout,157 faster=sh_lmax == 1 and not use_second_order_repr,158 tp_weights_layers=tp_weights_layers,159 edge_groups=1,160 )161 lig_emb_layers.append(layer)162 self.lig_emb_layers = nn.ModuleList(lig_emb_layers)163 164 # convolutional layers165 conv_layers = []166 for i in range(num_prot_emb_layers, num_prot_emb_layers + num_conv_layers):167 in_irreps = irrep_seq[min(i, len(irrep_seq) - 1)]168 out_irreps = irrep_seq[min(i + 1, len(irrep_seq) - 1)]169 layer = TensorProductConvLayer(170 in_irreps=in_irreps,171 sh_irreps=self.sh_irreps,172 out_irreps=out_irreps,173 n_edge_features=3 * ns,174 hidden_features=3 * ns,175 residual=True,176 batch_norm=batch_norm,177 dropout=dropout,178 faster=sh_lmax == 1 and not use_second_order_repr,179 tp_weights_layers=tp_weights_layers,180 edge_groups=1 if not differentiate_convolutions else (3 if i == num_prot_emb_layers + num_conv_layers - 1 else 9),181 )182 conv_layers.append(layer)183 self.conv_layers = nn.ModuleList(conv_layers)184 185 # confidence and affinity prediction layers186 if self.confidence_mode:187 if self.affinity_prediction:188 if self.parallel > 1:189 output_confidence_dim = 1 + ns190 else:191 output_confidence_dim = num_confidence_outputs + 1192 else:193 output_confidence_dim = num_confidence_outputs194 195 input_size = ns + (nv if reduce_pseudoscalars else ns) if num_conv_layers + num_prot_emb_layers >= 3 else ns196 197 if self.atom_confidence:198 self.atom_confidence_predictor = nn.Sequential(199 nn.Linear(input_size, ns),200 nn.BatchNorm1d(ns) if not confidence_no_batchnorm else nn.Identity(),201 nn.ReLU(),202 nn.Dropout(confidence_dropout),203 nn.Linear(ns, ns),204 nn.BatchNorm1d(ns) if not confidence_no_batchnorm else nn.Identity(),205 nn.ReLU(),206 nn.Dropout(confidence_dropout),207 nn.Linear(ns, atom_num_confidence_outputs + ns)208 )209 input_size = ns210 211 self.confidence_predictor = nn.Sequential(212 nn.Linear(input_size, ns),213 nn.BatchNorm1d(ns) if not confidence_no_batchnorm else nn.Identity(),214 nn.ReLU(),215 nn.Dropout(confidence_dropout),216 nn.Linear(ns, ns),217 nn.BatchNorm1d(ns) if not confidence_no_batchnorm else nn.Identity(),218 nn.ReLU(),219 nn.Dropout(confidence_dropout),220 nn.Linear(ns, output_confidence_dim)221 )222 223 if self.parallel > 1:224 self.affinity_predictor = nn.Sequential(225 nn.Linear(len(self.parallel_aggregators) * ns, ns),226 nn.BatchNorm1d(ns) if not confidence_no_batchnorm else nn.Identity(),227 nn.ReLU(),228 nn.Dropout(confidence_dropout),229 nn.Linear(ns, ns),230 nn.BatchNorm1d(ns) if not confidence_no_batchnorm else nn.Identity(),231 nn.ReLU(),232 nn.Dropout(confidence_dropout),233 nn.Linear(ns, 1)234 )235 236 else:237 # convolution for translational and rotational scores238 self.center_distance_expansion = GaussianSmearing(0.0, center_max_distance, distance_embed_dim)239 self.center_edge_embedding = nn.Sequential(240 nn.Linear(distance_embed_dim + sigma_embed_dim, ns),241 nn.ReLU(),242 nn.Dropout(dropout),243 nn.Linear(ns, ns)244 )245 246 self.final_conv = TensorProductConvLayer(247 in_irreps=self.conv_layers[-1].out_irreps,248 sh_irreps=self.sh_irreps,249 out_irreps=f'2x1o + 2x1e' if not self.odd_parity else '1x1o + 1x1e',250 n_edge_features=2 * ns,251 residual=False,252 dropout=dropout,253 batch_norm=batch_norm254 )255 256 self.tr_final_layer = nn.Sequential(nn.Linear(1 + sigma_embed_dim, ns),nn.Dropout(dropout), nn.ReLU(), nn.Linear(ns, 1))257 self.rot_final_layer = nn.Sequential(nn.Linear(1 + sigma_embed_dim, ns),nn.Dropout(dropout), nn.ReLU(), nn.Linear(ns, 1))258 259 if not no_torsion:260 # convolution for torsional score261 self.final_edge_embedding = nn.Sequential(262 nn.Linear(distance_embed_dim, ns),263 nn.ReLU(),264 nn.Dropout(dropout),265 nn.Linear(ns, ns)266 )267 self.final_tp_tor = o3.FullTensorProduct(self.sh_irreps, "2e")268 self.tor_bond_conv = TensorProductConvLayer(269 in_irreps=self.conv_layers[-1].out_irreps,270 sh_irreps=self.final_tp_tor.irreps_out,271 out_irreps=f'{ns}x0o + {ns}x0e' if not self.odd_parity else f'{ns}x0o',272 n_edge_features=3 * ns,273 residual=False,274 dropout=dropout,275 batch_norm=batch_norm276 )277 self.tor_final_layer = nn.Sequential(278 nn.Linear(2 * ns if not self.odd_parity else ns, ns, bias=False),279 nn.Tanh(),280 nn.Dropout(dropout),281 nn.Linear(ns, 1, bias=False)282 )283 284 @staticmethod285 def _resolve_edge_store(data, primary_key, fallback_key):286 edge_types = getattr(data, "edge_types", ())287 if primary_key in edge_types:288 return data[primary_key]289 if fallback_key in edge_types:290 return data[fallback_key]291 try:292 return data[primary_key]293 except Exception:294 return data[fallback_key]295 296 def _ligand_edge_store(self, data):297 return self._resolve_edge_store(298 data,299 ("ligand", "ligand"),300 ("ligand", "lig_bond", "ligand"),301 )302 303 def _receptor_edge_store(self, data):304 return self._resolve_edge_store(305 data,306 ("receptor", "receptor"),307 ("receptor", "rec_contact", "receptor"),308 )309 310 def _atom_edge_store(self, data):311 return self._resolve_edge_store(312 data,313 ("atom", "atom"),314 ("atom", "atom_contact", "atom"),315 )316 317 def _atom_receptor_edge_store(self, data):318 return self._resolve_edge_store(319 data,320 ("atom", "receptor"),321 ("atom", "atom_rec_contact", "receptor"),322 )323 324 def embedding(self, data):325 receptor_edge_store = self._receptor_edge_store(data)326 atom_edge_store = self._atom_edge_store(data)327 atom_receptor_edge_store = self._atom_receptor_edge_store(data)328 if not hasattr(data['receptor'], "rec_node_attr"):329 if self.lm_embedding_type not in [None, 'precomputed']:330 sequences = [s for l in data['receptor'].sequence for s in l]331 if isinstance(sequences[0], list):332 sequences = [s for l in sequences for s in l]333 sequences = [(i, s) for i, s in enumerate(sequences)]334 batch_labels, batch_strs, batch_tokens = self.batch_converter(sequences)335 out = self.lm(batch_tokens.to(data['receptor'].x.device), repr_layers=[self.lm.num_layers], return_contacts=False)336 rec_lm_emb = torch.cat([t[:len(sequences[i][1])] for i, t in enumerate(out['representations'][self.lm.num_layers])], dim=0)337 data['receptor'].x = torch.cat([data['receptor'].x, rec_lm_emb], dim=-1)338 339 rec_node_attr, rec_edge_attr, rec_edge_sh, rec_edge_weight = self.build_rec_conv_graph(data)340 rec_node_attr = self.rec_node_embedding(rec_node_attr)341 rec_edge_attr = self.rec_edge_embedding(rec_edge_attr)342 343 atom_node_attr, atom_edge_attr, atom_edge_sh, atom_edge_weight = self.build_atom_conv_graph(data)344 atom_node_attr = self.atom_node_embedding(atom_node_attr)345 atom_edge_attr = self.atom_edge_embedding(atom_edge_attr)346 347 ar_edge_attr, ar_edge_sh, ar_edge_weight = self.build_cross_rec_conv_graph(data)348 ar_edge_attr = self.ar_edge_embedding(ar_edge_attr)349 350 rec_edge_index = receptor_edge_store.edge_index.clone()351 atom_edge_index = atom_edge_store.edge_index.clone()352 ar_edge_index = atom_receptor_edge_store.edge_index.clone()353 354 node_attr = torch.cat([rec_node_attr, atom_node_attr], dim=0)355 ar_edge_index[0] = ar_edge_index[0] + len(rec_node_attr)356 edge_index = torch.cat([rec_edge_index, ar_edge_index, atom_edge_index + len(rec_node_attr), torch.flip(ar_edge_index, dims=[0])], dim=1)357 edge_attr = torch.cat([rec_edge_attr, ar_edge_attr, atom_edge_attr, ar_edge_attr], dim=0)358 edge_sh = torch.cat([rec_edge_sh, ar_edge_sh, atom_edge_sh, ar_edge_sh], dim=0)359 edge_weight = torch.cat([rec_edge_weight, ar_edge_weight, atom_edge_weight, ar_edge_weight], dim=0) \360 if torch.is_tensor(rec_edge_weight) else torch.ones((len(edge_index[0]), 1), device=edge_index.device)361 s1, s2, s3 = len(rec_edge_index[0]), len(rec_edge_index[0]) + len(ar_edge_index[0]), len(rec_edge_index[0]) + len(ar_edge_index[0]) + len(atom_edge_index[0])362 363 for l in range(len(self.rec_emb_layers)):364 edge_attr_ = torch.cat(365 [edge_attr, node_attr[edge_index[0], :self.ns], node_attr[edge_index[1], :self.ns]], -1)366 if self.differentiate_convolutions: edge_attr_ = [edge_attr_[:s1], edge_attr_[s1:s2], edge_attr_[s2:s3], edge_attr_[s3:]]367 node_attr = self.rec_emb_layers[l](node_attr, edge_index, edge_attr_, edge_sh, edge_weight=edge_weight)368 369 370 data['receptor'].rec_node_attr = node_attr[:len(rec_node_attr)]371 receptor_edge_store.rec_edge_attr = rec_edge_attr372 receptor_edge_store.edge_sh = rec_edge_sh373 receptor_edge_store.edge_weight = rec_edge_weight374 375 data['atom'].atom_node_attr = node_attr[len(rec_node_attr):]376 atom_edge_store.atom_edge_attr = atom_edge_attr377 atom_edge_store.edge_sh = atom_edge_sh378 atom_edge_store.edge_weight = atom_edge_weight379 380 atom_receptor_edge_store.edge_attr = ar_edge_attr381 atom_receptor_edge_store.edge_sh = ar_edge_sh382 atom_receptor_edge_store.edge_weight = ar_edge_weight383 384 # receptor embedding385 rec_sigma_emb = self.rec_sigma_embedding(self.timestep_emb_func(data.complex_t['tr']))386 rec_node_attr = data['receptor'].rec_node_attr + 0387 rec_node_attr[:, :self.ns] = rec_node_attr[:, :self.ns] + rec_sigma_emb[data['receptor'].batch]388 rec_edge_attr = receptor_edge_store.rec_edge_attr + rec_sigma_emb[data['receptor'].batch[receptor_edge_store.edge_index[0]]]389 390 # atom embedding391 atom_node_attr = data['atom'].atom_node_attr + 0392 atom_node_attr[:, :self.ns] = atom_node_attr[:, :self.ns] + rec_sigma_emb[data['atom'].batch]393 atom_edge_attr = atom_edge_store.atom_edge_attr + rec_sigma_emb[data['atom'].batch[atom_edge_store.edge_index[0]]]394 395 # atom-receptor embedding396 ar_edge_attr = atom_receptor_edge_store.edge_attr + rec_sigma_emb[data['atom'].batch[atom_receptor_edge_store.edge_index[0]]]397 398 # ligand embedding399 lig_node_attr, lig_edge_index, lig_edge_attr, lig_edge_sh, lig_edge_weight = self.build_lig_conv_graph(data)400 lig_node_attr = self.lig_node_embedding(lig_node_attr)401 lig_edge_attr = self.lig_edge_embedding(lig_edge_attr)402 403 if self.embed_also_ligand:404 for l in range(len(self.lig_emb_layers)):405 edge_attr_ = torch.cat([lig_edge_attr, lig_node_attr[lig_edge_index[0], :self.ns], lig_node_attr[lig_edge_index[1], :self.ns]], -1)406 lig_node_attr = self.lig_emb_layers[l](lig_node_attr, lig_edge_index, edge_attr_, lig_edge_sh, edge_weight=lig_edge_weight)407 408 else:409 lig_node_attr = F.pad(lig_node_attr, (0, rec_node_attr.shape[-1] - lig_node_attr.shape[-1]))410 411 return lig_node_attr, lig_edge_index, lig_edge_attr, lig_edge_sh, lig_edge_weight, \412 rec_node_attr, receptor_edge_store.edge_index, rec_edge_attr, receptor_edge_store.edge_sh, receptor_edge_store.edge_weight, \413 atom_node_attr, atom_edge_store.edge_index, atom_edge_attr, atom_edge_store.edge_sh, atom_edge_store.edge_weight, \414 atom_receptor_edge_store.edge_index, ar_edge_attr, atom_receptor_edge_store.edge_sh, atom_receptor_edge_store.edge_weight415 416 def forward(self, data):417 if self.crop_beyond is not None:418 # TODO missing filtering atoms419 raise NotImplementedError420 ligand_pos = data['ligand'].pos421 receptor_pos = data['receptor'].pos422 residues_to_keep = torch.any(torch.sum((ligand_pos.unsqueeze(0) - receptor_pos.unsqueeze(1)) ** 2, -1) < self.crop_beyond ** 2, dim=1)423 424 data['receptor'].pos = data['receptor'].pos[residues_to_keep]425 data['receptor'].x = data['receptor'].x[residues_to_keep]426 data['receptor'].side_chain_vecs = data['receptor'].side_chain_vecs[residues_to_keep]427 data['receptor', 'rec_contact', 'receptor'].edge_index = subgraph(residues_to_keep, data['receptor', 'rec_contact', 'receptor'].edge_index, relabel_nodes=True)[0]428 429 if self.no_aminoacid_identities:430 data['receptor'].x = data['receptor'].x * 0431 432 if not self.confidence_mode:433 tr_sigma, rot_sigma, tor_sigma = self.t_to_sigma(*[data.complex_t[noise_type] for noise_type in ['tr', 'rot', 'tor']])434 else:435 tr_sigma, rot_sigma, tor_sigma = [data.complex_t[noise_type] for noise_type in ['tr', 'rot', 'tor']]436 437 lig_node_attr, lig_edge_index, lig_edge_attr, lig_edge_sh, lig_edge_weight, rec_node_attr, \438 rec_edge_index, rec_edge_attr, rec_edge_sh, rec_edge_weight,\439 atom_node_attr, atom_edge_index, atom_edge_attr, atom_edge_sh, atom_edge_weight, \440 ar_edge_index, ar_edge_attr, ar_edge_sh, ar_edge_weight = self.embedding(data)441 442 # build lig cross graph443 cross_cutoff = (tr_sigma * 3 + 20).unsqueeze(1) if self.dynamic_max_cross else self.cross_max_distance444 lr_edge_index, lr_edge_attr, lr_edge_sh, lr_edge_weight, la_edge_index, la_edge_attr, \445 la_edge_sh, la_edge_weight = self.build_cross_lig_conv_graph(data, cross_cutoff)446 lr_edge_attr= self.lr_edge_embedding(lr_edge_attr)447 la_edge_attr = self.la_edge_embedding(la_edge_attr)448 449 n_lig, n_rec = len(lig_node_attr), len(rec_node_attr)450 451 node_attr = torch.cat([lig_node_attr, rec_node_attr, atom_node_attr], dim=0)452 rec_edge_index, atom_edge_index, lr_edge_index, la_edge_index, ar_edge_index = rec_edge_index.clone(), atom_edge_index.clone(), lr_edge_index.clone(), la_edge_index.clone(), ar_edge_index.clone()453 rec_edge_index[0], rec_edge_index[1] = rec_edge_index[0] + n_lig, rec_edge_index[1] + n_lig454 atom_edge_index[0], atom_edge_index[1] = atom_edge_index[0] + n_lig + n_rec, atom_edge_index[1] + n_lig + n_rec455 lr_edge_index[1] = lr_edge_index[1] + n_lig456 la_edge_index[1] = la_edge_index[1] + n_lig + n_rec457 ar_edge_index[0], ar_edge_index[1] = ar_edge_index[0] + n_lig + n_rec, ar_edge_index[1] + n_lig458 459 edge_index = torch.cat([lig_edge_index, lr_edge_index, la_edge_index, rec_edge_index,460 torch.flip(lr_edge_index, dims=[0]), torch.flip(ar_edge_index, dims=[0]),461 atom_edge_index, torch.flip(la_edge_index, dims=[0]), ar_edge_index], dim=1)462 edge_attr = torch.cat([lig_edge_attr, lr_edge_attr, la_edge_attr, rec_edge_attr, lr_edge_attr,463 ar_edge_attr, atom_edge_attr, la_edge_attr, ar_edge_attr], dim=0)464 edge_sh = torch.cat([lig_edge_sh, lr_edge_sh, la_edge_sh, rec_edge_sh, lr_edge_sh, ar_edge_sh,465 atom_edge_sh, la_edge_sh, ar_edge_sh], dim=0)466 edge_weight = torch.cat([lig_edge_weight, lr_edge_weight, la_edge_weight, rec_edge_weight, lr_edge_weight,467 ar_edge_weight, atom_edge_weight, la_edge_weight, ar_edge_weight],468 dim=0) if torch.is_tensor(lig_edge_weight) else torch.ones((len(edge_index[0]), 1),469 device=edge_index.device)470 s1, s2, s3, s4, s5, s6, s7, s8, _ = tuple(np.cumsum(list(map(len, [lig_edge_attr, lr_edge_attr, la_edge_attr,471 rec_edge_attr, lr_edge_attr, ar_edge_attr, atom_edge_attr, la_edge_attr, ar_edge_attr]))).tolist())472 473 for l in range(len(self.conv_layers)):474 if l < len(self.conv_layers) - 1:475 edge_attr_ = torch.cat([edge_attr, node_attr[edge_index[0], :self.ns], node_attr[edge_index[1], :self.ns]], -1)476 if self.differentiate_convolutions: edge_attr_ = [edge_attr_[:s1], edge_attr_[s1:s2], edge_attr_[s2:s3], edge_attr_[s3:s4],477 edge_attr_[s4:s5], edge_attr_[s5:s6], edge_attr_[s6:s7], edge_attr_[s7:s8], edge_attr_[s8:]]478 node_attr = self.conv_layers[l](node_attr, edge_index, edge_attr_, edge_sh, edge_weight=edge_weight)479 else:480 edge_attr_ = torch.cat([edge_attr[:s3], node_attr[edge_index[0, :s3], :self.ns], node_attr[edge_index[1, :s3], :self.ns]], -1)481 if self.differentiate_convolutions: edge_attr_ = [edge_attr_[:s1], edge_attr_[s1:s2], edge_attr_[s2:s3]]482 node_attr = self.conv_layers[l](node_attr, edge_index[:, :s3], edge_attr_, edge_sh[:s3], edge_weight=edge_weight[:s3])483 484 lig_node_attr = node_attr[:len(lig_node_attr)]485 486 # confidence and affinity prediction487 if self.confidence_mode:488 scalar_lig_attr = torch.cat([lig_node_attr[:,:self.ns], lig_node_attr[:,-(self.nv if self.reduce_pseudoscalars else self.ns):] ], dim=1) \489 if self.num_conv_layers + self.num_prot_emb_layers >= 3 else lig_node_attr[:,:self.ns]490 491 if self.atom_confidence:492 scalar_lig_attr = self.atom_confidence_predictor(scalar_lig_attr)493 atom_confidence = scalar_lig_attr[:, :self.atom_num_confidence_outputs]494 scalar_lig_attr = scalar_lig_attr[:, self.atom_num_confidence_outputs:]495 else:496 atom_confidence = torch.zeros((len(lig_node_attr),), device=lig_node_attr.device)497 498 confidence = self.confidence_predictor(scatter_mean(scalar_lig_attr, data['ligand'].batch, dim=0)).squeeze(dim=-1)499 500 if self.parallel > 1:501 confidence, affinity = confidence[:, 0], confidence[:, 1:]502 confidence = confidence.reshape(data.num_graphs, self.parallel)503 affinity = affinity.reshape(data.num_graphs, self.parallel, -1)504 affinity = torch.cat([AGGREGATORS[agg](affinity) for agg in self.parallel_aggregators], dim=-1)505 affinity = self.affinity_predictor(affinity).squeeze(dim=-1)506 confidence = confidence, affinity507 return confidence, atom_confidence508 assert self.parallel == 1509 510 # compute translational and rotational score vectors511 center_edge_index, center_edge_attr, center_edge_sh = self.build_center_conv_graph(data)512 center_edge_attr = self.center_edge_embedding(center_edge_attr)513 if self.fixed_center_conv:514 center_edge_attr = torch.cat([center_edge_attr, lig_node_attr[center_edge_index[1], :self.ns]], -1)515 else:516 center_edge_attr = torch.cat([center_edge_attr, lig_node_attr[center_edge_index[0], :self.ns]], -1)517 global_pred = self.final_conv(lig_node_attr, center_edge_index, center_edge_attr, center_edge_sh, out_nodes=data.num_graphs)518 519 tr_pred = global_pred[:, :3] + (global_pred[:, 6:9] if not self.odd_parity else 0)520 rot_pred = global_pred[:, 3:6] + (global_pred[:, 9:] if not self.odd_parity else 0)521 522 if self.separate_noise_schedule:523 data.graph_sigma_emb = torch.cat([self.timestep_emb_func(data.complex_t[noise_type]) for noise_type in ['tr', 'rot', 'tor']], dim=1)524 elif self.asyncronous_noise_schedule:525 data.graph_sigma_emb = self.timestep_emb_func(data.complex_t['t'])526 else: # tr rot and tor noise is all the same in this case527 data.graph_sigma_emb = self.timestep_emb_func(data.complex_t['tr'])528 529 # adjust the magniture of the score vectors530 tr_norm = torch.linalg.vector_norm(tr_pred, dim=1).unsqueeze(1)531 tr_pred = tr_pred / tr_norm * self.tr_final_layer(torch.cat([tr_norm, data.graph_sigma_emb], dim=1))532 533 rot_norm = torch.linalg.vector_norm(rot_pred, dim=1).unsqueeze(1)534 rot_pred = rot_pred / rot_norm * self.rot_final_layer(torch.cat([rot_norm, data.graph_sigma_emb], dim=1))535 536 if self.scale_by_sigma:537 tr_pred = tr_pred / tr_sigma.unsqueeze(1)538 rot_pred = rot_pred * so3.score_norm(rot_sigma.cpu()).unsqueeze(1).to(data['ligand'].x.device)539 540 if self.no_torsion or data['ligand'].edge_mask.sum() == 0: return tr_pred, rot_pred, torch.empty(0,device=self.device), None541 542 # torsional components543 tor_bonds, tor_edge_index, tor_edge_attr, tor_edge_sh, tor_edge_weight = self.build_bond_conv_graph(data)544 tor_bond_vec = data['ligand'].pos[tor_bonds[1]] - data['ligand'].pos[tor_bonds[0]]545 tor_bond_attr = lig_node_attr[tor_bonds[0]] + lig_node_attr[tor_bonds[1]]546 547 tor_bonds_sh = o3.spherical_harmonics("2e", tor_bond_vec, normalize=True, normalization='component')548 tor_edge_sh = self.final_tp_tor(tor_edge_sh, tor_bonds_sh[tor_edge_index[0]])549 550 tor_edge_attr = torch.cat([tor_edge_attr, lig_node_attr[tor_edge_index[1], :self.ns],551 tor_bond_attr[tor_edge_index[0], :self.ns]], -1)552 tor_pred = self.tor_bond_conv(lig_node_attr, tor_edge_index, tor_edge_attr, tor_edge_sh,553 out_nodes=data['ligand'].edge_mask.sum(), reduce='mean', edge_weight=tor_edge_weight)554 tor_pred = self.tor_final_layer(tor_pred).squeeze(1)555 ligand_edge_store = self._ligand_edge_store(data)556 edge_sigma = tor_sigma[data['ligand'].batch][ligand_edge_store.edge_index[0]][data['ligand'].edge_mask]557 558 if self.scale_by_sigma:559 tor_pred = tor_pred * torch.sqrt(torch.tensor(torus.score_norm(edge_sigma.cpu().numpy())).float()560 .to(data['ligand'].x.device))561 return tr_pred, rot_pred, tor_pred, None562 563 def get_edge_weight(self, edge_vec, max_norm):564 if self.smooth_edges:565 normalised_norm = torch.clip(edge_vec.norm(dim=-1) * np.pi / max_norm, max=np.pi)566 return 0.5 * (torch.cos(normalised_norm) + 1.0).unsqueeze(-1)567 return 1.0568 569 def build_lig_conv_graph(self, data):570 # build the graph between ligand atoms571 if self.separate_noise_schedule:572 data['ligand'].node_sigma_emb = torch.cat(573 [self.timestep_emb_func(data['ligand'].node_t[noise_type]) for noise_type in ['tr', 'rot', 'tor']],574 dim=1)575 elif self.asyncronous_noise_schedule:576 data['ligand'].node_sigma_emb = self.timestep_emb_func(data['ligand'].node_t['t'])577 else:578 data['ligand'].node_sigma_emb = self.timestep_emb_func(579 data['ligand'].node_t['tr']) # tr rot and tor noise is all the same580 581 if self.parallel == 1:582 radius_edges = radius_graph(data['ligand'].pos, self.lig_max_radius, data['ligand'].batch)583 else:584 batches = torch.zeros(data.num_graphs, device=data['ligand'].x.device).long()585 batches = batches.index_add(0, data['ligand'].batch, torch.ones(len(data['ligand'].batch), device=data['ligand'].x.device).long())586 outer_batches = data.num_graphs587 b = [torch.ones(batches[i].item()//self.parallel, device=data['ligand'].x.device).long() * (self.parallel * i + j)588 for i in range(outer_batches) for j in range(self.parallel)]589 data['ligand'].batch_parallel = torch.cat(b)590 radius_edges = radius_graph(data['ligand'].pos, self.lig_max_radius, data['ligand'].batch_parallel)591 ligand_edge_store = self._ligand_edge_store(data)592 edge_index = torch.cat([ligand_edge_store.edge_index, radius_edges], 1).long()593 edge_attr = torch.cat([594 ligand_edge_store.edge_attr,595 torch.zeros(radius_edges.shape[-1], self.in_lig_edge_features, device=data['ligand'].x.device)596 ], 0)597 598 edge_sigma_emb = data['ligand'].node_sigma_emb[edge_index[0].long()]599 edge_attr = torch.cat([edge_attr, edge_sigma_emb], 1)600 node_attr = torch.cat([data['ligand'].x, data['ligand'].node_sigma_emb], 1)601 602 src, dst = edge_index603 edge_vec = data['ligand'].pos[dst.long()] - data['ligand'].pos[src.long()]604 edge_length_emb = self.lig_distance_expansion(edge_vec.norm(dim=-1))605 606 edge_attr = torch.cat([edge_attr, edge_length_emb], 1)607 edge_sh = o3.spherical_harmonics(self.sh_irreps, edge_vec, normalize=True, normalization='component')608 edge_weight = self.get_edge_weight(edge_vec, self.lig_max_radius)609 610 return node_attr, edge_index, edge_attr, edge_sh, edge_weight611 612 def build_rec_conv_graph(self, data):613 # build the graph between receptor residues614 node_attr = data['receptor'].x615 616 # this assumes the edges were already created in preprocessing since protein's structure is fixed617 edge_index = self._receptor_edge_store(data).edge_index618 src, dst = edge_index619 edge_vec = data['receptor'].pos[dst.long()] - data['receptor'].pos[src.long()]620 621 edge_attr = self.rec_distance_expansion(edge_vec.norm(dim=-1))622 edge_sh = o3.spherical_harmonics(self.sh_irreps, edge_vec, normalize=True, normalization='component')623 edge_weight = self.get_edge_weight(edge_vec, self.rec_max_radius)624 625 return node_attr, edge_attr, edge_sh, edge_weight626 627 def build_atom_conv_graph(self, data):628 # build the graph between receptor atoms629 node_attr = data['atom'].x630 631 # this assumes the edges were already created in preprocessing since protein's structure is fixed632 edge_index = self._atom_edge_store(data).edge_index633 src, dst = edge_index634 edge_vec = data['atom'].pos[dst.long()] - data['atom'].pos[src.long()]635 636 edge_attr = self.lig_distance_expansion(edge_vec.norm(dim=-1))637 edge_sh = o3.spherical_harmonics(self.sh_irreps, edge_vec, normalize=True, normalization='component')638 edge_weight = self.get_edge_weight(edge_vec, self.lig_max_radius)639 640 return node_attr, edge_attr, edge_sh, edge_weight641 642 def build_cross_lig_conv_graph(self, data, lr_cross_distance_cutoff):643 # build the cross edges between ligand atoms and receptor residues + atoms644 645 # LIGAND to RECEPTOR646 if torch.is_tensor(lr_cross_distance_cutoff):647 # different cutoff for every graph648 lr_edge_index = radius(data['receptor'].pos / lr_cross_distance_cutoff[data['receptor'].batch],649 data['ligand'].pos / lr_cross_distance_cutoff[data['ligand'].batch], 1,650 data['receptor'].batch, data['ligand'].batch, max_num_neighbors=10000)651 else:652 lr_edge_index = radius(data['receptor'].pos, data['ligand'].pos, lr_cross_distance_cutoff,653 data['receptor'].batch, data['ligand'].batch, max_num_neighbors=10000)654 655 lr_edge_vec = data['receptor'].pos[lr_edge_index[1].long()] - data['ligand'].pos[lr_edge_index[0].long()]656 lr_edge_length_emb = self.cross_distance_expansion(lr_edge_vec.norm(dim=-1))657 lr_edge_sigma_emb = data['ligand'].node_sigma_emb[lr_edge_index[0].long()]658 lr_edge_attr = torch.cat([lr_edge_sigma_emb, lr_edge_length_emb], 1)659 lr_edge_sh = o3.spherical_harmonics(self.sh_irreps, lr_edge_vec, normalize=True, normalization='component')660 661 cutoff_d = lr_cross_distance_cutoff[data['ligand'].batch[lr_edge_index[0]]].squeeze() \662 if torch.is_tensor(lr_cross_distance_cutoff) else lr_cross_distance_cutoff663 lr_edge_weight = self.get_edge_weight(lr_edge_vec, cutoff_d)664 665 # LIGAND to ATOM666 la_edge_index = radius(data['atom'].pos, data['ligand'].pos, self.lig_max_radius,667 data['atom'].batch, data['ligand'].batch, max_num_neighbors=10000)668 669 la_edge_vec = data['atom'].pos[la_edge_index[1].long()] - data['ligand'].pos[la_edge_index[0].long()]670 la_edge_length_emb = self.lig_distance_expansion(la_edge_vec.norm(dim=-1))671 la_edge_sigma_emb = data['ligand'].node_sigma_emb[la_edge_index[0].long()]672 la_edge_attr = torch.cat([la_edge_sigma_emb, la_edge_length_emb], 1)673 la_edge_sh = o3.spherical_harmonics(self.sh_irreps, la_edge_vec, normalize=True, normalization='component')674 la_edge_weight = self.get_edge_weight(la_edge_vec, self.lig_max_radius)675 676 return lr_edge_index, lr_edge_attr, lr_edge_sh, lr_edge_weight, la_edge_index, la_edge_attr, \677 la_edge_sh, la_edge_weight678 679 def build_cross_rec_conv_graph(self, data):680 # build the cross edges between ligan atoms, receptor residues and receptor atoms681 682 # ATOM to RECEPTOR683 ar_edge_index = self._atom_receptor_edge_store(data).edge_index684 ar_edge_vec = data['receptor'].pos[ar_edge_index[1].long()] - data['atom'].pos[ar_edge_index[0].long()]685 ar_edge_attr = self.rec_distance_expansion(ar_edge_vec.norm(dim=-1))686 ar_edge_sh = o3.spherical_harmonics(self.sh_irreps, ar_edge_vec, normalize=True, normalization='component')687 ar_edge_weight = 1688 689 return ar_edge_attr, ar_edge_sh, ar_edge_weight690 691 def build_center_conv_graph(self, data):692 # build the filter for the convolution of the center with the ligand atoms693 # for translational and rotational score694 edge_index = torch.cat([data['ligand'].batch.unsqueeze(0), torch.arange(len(data['ligand'].batch)).to(data['ligand'].x.device).unsqueeze(0)], dim=0)695 696 center_pos, count = torch.zeros((data.num_graphs, 3)).to(data['ligand'].x.device), torch.zeros((data.num_graphs, 3)).to(data['ligand'].x.device)697 center_pos.index_add_(0, index=data['ligand'].batch, source=data['ligand'].pos)698 center_pos = center_pos / torch.bincount(data['ligand'].batch).unsqueeze(1)699 700 edge_vec = data['ligand'].pos[edge_index[1]] - center_pos[edge_index[0]]701 edge_attr = self.center_distance_expansion(edge_vec.norm(dim=-1))702 edge_sigma_emb = data['ligand'].node_sigma_emb[edge_index[1].long()]703 edge_attr = torch.cat([edge_attr, edge_sigma_emb], 1)704 edge_sh = o3.spherical_harmonics(self.sh_irreps, edge_vec, normalize=True, normalization='component')705 return edge_index, edge_attr, edge_sh706 707 def build_bond_conv_graph(self, data):708 # build graph for the pseudotorque layer709 bonds = self._ligand_edge_store(data).edge_index[:, data['ligand'].edge_mask].long()710 bond_pos = (data['ligand'].pos[bonds[0]] + data['ligand'].pos[bonds[1]]) / 2711 bond_batch = data['ligand'].batch[bonds[0]]712 edge_index = radius(data['ligand'].pos, bond_pos, self.lig_max_radius, batch_x=data['ligand'].batch, batch_y=bond_batch)713 714 edge_vec = data['ligand'].pos[edge_index[1]] - bond_pos[edge_index[0]]715 edge_attr = self.lig_distance_expansion(edge_vec.norm(dim=-1))716 717 edge_attr = self.final_edge_embedding(edge_attr)718 edge_sh = o3.spherical_harmonics(self.sh_irreps, edge_vec, normalize=True, normalization='component')719 edge_weight = self.get_edge_weight(edge_vec, self.lig_max_radius)720 721 return bonds, edge_index, edge_attr, edge_sh, edge_weight722 