OneScience-Group/DiffDock
052
1from e3nn import o32import torch3from torch import nn4from torch.nn import functional as F5from torch_cluster import radius, radius_graph6from torch_scatter import scatter, scatter_mean7import numpy as np8from e3nn.nn import BatchNorm9 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, OldAtomEncoder, AtomEncoder14from .tensor_layers import OldTensorProductConvLayer15 16AGGREGATORS = {"mean": lambda x: torch.mean(x, dim=1),17 "max": lambda x: torch.max(x, dim=1)[0],18 "min": lambda x: torch.min(x, dim=1)[0],19 "std": lambda x: torch.std(x, dim=1)}20 21 22class AAOldModel(torch.nn.Module):23 def __init__(self, t_to_sigma, device, timestep_emb_func, in_lig_edge_features=4, sigma_embed_dim=32, sh_lmax=2,24 ns=16, nv=4, num_conv_layers=2, lig_max_radius=5, rec_max_radius=30, cross_max_distance=250,25 center_max_distance=30, distance_embed_dim=32, cross_distance_embed_dim=32, no_torsion=False,26 scale_by_sigma=True, norm_by_sigma=True, use_second_order_repr=False, batch_norm=True,27 dynamic_max_cross=False, dropout=0.0, smooth_edges=False, odd_parity=False,28 separate_noise_schedule=False, lm_embedding_type=None, confidence_mode=False,29 confidence_dropout=0, confidence_no_batchnorm = False,30 asyncronous_noise_schedule=False, affinity_prediction=False, parallel=1,31 parallel_aggregators="mean max min std", num_confidence_outputs=1, fixed_center_conv=False,32 no_aminoacid_identities=False, include_miscellaneous_atoms=False, use_old_atom_encoder=False):33 super(AAOldModel, self).__init__()34 assert (not no_aminoacid_identities) or (lm_embedding_type is None), "no language model emb without identities"35 if parallel > 1: assert affinity_prediction36 self.t_to_sigma = t_to_sigma37 self.in_lig_edge_features = in_lig_edge_features38 sigma_embed_dim *= (3 if separate_noise_schedule else 1)39 self.sigma_embed_dim = sigma_embed_dim40 self.lig_max_radius = lig_max_radius41 self.rec_max_radius = rec_max_radius42 self.cross_max_distance = cross_max_distance43 self.dynamic_max_cross = dynamic_max_cross44 self.center_max_distance = center_max_distance45 self.distance_embed_dim = distance_embed_dim46 self.cross_distance_embed_dim = cross_distance_embed_dim47 self.sh_irreps = o3.Irreps.spherical_harmonics(lmax=sh_lmax)48 self.ns, self.nv = ns, nv49 self.scale_by_sigma = scale_by_sigma50 self.norm_by_sigma = norm_by_sigma51 self.device = device52 self.no_torsion = no_torsion53 self.smooth_edges = smooth_edges54 self.odd_parity = odd_parity55 self.num_conv_layers = num_conv_layers56 self.timestep_emb_func = timestep_emb_func57 self.separate_noise_schedule = separate_noise_schedule58 self.confidence_mode = confidence_mode59 self.num_conv_layers = num_conv_layers60 self.asyncronous_noise_schedule = asyncronous_noise_schedule61 self.affinity_prediction = affinity_prediction62 self.parallel, self.parallel_aggregators = parallel, parallel_aggregators.split(' ')63 self.fixed_center_conv = fixed_center_conv64 self.no_aminoacid_identities = no_aminoacid_identities65 66 # embedding layers67 atom_encoder_class = OldAtomEncoder if use_old_atom_encoder else AtomEncoder68 lm_embedding_dim = 0 if lm_embedding_type is None else 128069 self.lig_node_embedding = atom_encoder_class(emb_dim=ns, feature_dims=lig_feature_dims, sigma_embed_dim=sigma_embed_dim)70 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))71 72 self.rec_node_embedding = (73 atom_encoder_class(74 emb_dim=ns,75 feature_dims=rec_residue_feature_dims,76 sigma_embed_dim=sigma_embed_dim,77 lm_embedding_type=lm_embedding_type,78 )79 if use_old_atom_encoder80 else atom_encoder_class(81 emb_dim=ns,82 feature_dims=rec_residue_feature_dims,83 sigma_embed_dim=sigma_embed_dim,84 lm_embedding_dim=lm_embedding_dim,85 )86 )87 self.rec_edge_embedding = nn.Sequential(nn.Linear(sigma_embed_dim + distance_embed_dim, ns), nn.ReLU(), nn.Dropout(dropout),nn.Linear(ns, ns))88 89 self.atom_node_embedding = atom_encoder_class(emb_dim=ns, feature_dims=rec_atom_feature_dims, sigma_embed_dim=sigma_embed_dim)90 self.atom_edge_embedding = nn.Sequential(nn.Linear(sigma_embed_dim + distance_embed_dim, ns), nn.ReLU(), nn.Dropout(dropout),nn.Linear(ns, ns))91 92 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))93 self.ar_edge_embedding = nn.Sequential(nn.Linear(sigma_embed_dim + distance_embed_dim, ns), nn.ReLU(), nn.Dropout(dropout),nn.Linear(ns, ns))94 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))95 96 self.lig_distance_expansion = GaussianSmearing(0.0, lig_max_radius, distance_embed_dim)97 self.rec_distance_expansion = GaussianSmearing(0.0, rec_max_radius, distance_embed_dim)98 self.cross_distance_expansion = GaussianSmearing(0.0, cross_max_distance, cross_distance_embed_dim)99 100 if use_second_order_repr:101 irrep_seq = [102 f'{ns}x0e',103 f'{ns}x0e + {nv}x1o + {nv}x2e',104 f'{ns}x0e + {nv}x1o + {nv}x2e + {nv}x1e + {nv}x2o',105 f'{ns}x0e + {nv}x1o + {nv}x2e + {nv}x1e + {nv}x2o + {ns}x0o'106 ]107 else:108 irrep_seq = [109 f'{ns}x0e',110 f'{ns}x0e + {nv}x1o',111 f'{ns}x0e + {nv}x1o + {nv}x1e',112 f'{ns}x0e + {nv}x1o + {nv}x1e + {ns}x0o'113 ]114 115 # convolutional layers116 conv_layers = []117 for i in range(num_conv_layers):118 in_irreps = irrep_seq[min(i, len(irrep_seq) - 1)]119 out_irreps = irrep_seq[min(i + 1, len(irrep_seq) - 1)]120 parameters = {121 'in_irreps': in_irreps,122 'sh_irreps': self.sh_irreps,123 'out_irreps': out_irreps,124 'n_edge_features': 3 * ns,125 'residual': False,126 'batch_norm': batch_norm,127 'dropout': dropout128 }129 130 for _ in range(9): # 3 intra & 6 inter per each layer131 conv_layers.append(OldTensorProductConvLayer(**parameters))132 133 self.conv_layers = nn.ModuleList(conv_layers)134 135 # confidence and affinity prediction layers136 if self.confidence_mode:137 if self.affinity_prediction:138 if self.parallel > 1:139 output_confidence_dim = 1 + ns140 else:141 output_confidence_dim = num_confidence_outputs +1142 else:143 output_confidence_dim = num_confidence_outputs144 145 self.confidence_predictor = nn.Sequential(146 nn.Linear(2 * self.ns if num_conv_layers >= 3 else self.ns, ns),147 nn.BatchNorm1d(ns) if not confidence_no_batchnorm else nn.Identity(),148 nn.ReLU(),149 nn.Dropout(confidence_dropout),150 nn.Linear(ns, ns),151 nn.BatchNorm1d(ns) if not confidence_no_batchnorm else nn.Identity(),152 nn.ReLU(),153 nn.Dropout(confidence_dropout),154 nn.Linear(ns, output_confidence_dim)155 )156 157 if self.parallel > 1:158 self.affinity_predictor = nn.Sequential(159 nn.Linear(len(self.parallel_aggregators) * ns, ns),160 nn.BatchNorm1d(ns) if not confidence_no_batchnorm else nn.Identity(),161 nn.ReLU(),162 nn.Dropout(confidence_dropout),163 nn.Linear(ns, ns),164 nn.BatchNorm1d(ns) if not confidence_no_batchnorm else nn.Identity(),165 nn.ReLU(),166 nn.Dropout(confidence_dropout),167 nn.Linear(ns, 1)168 )169 170 else:171 # convolution for translational and rotational scores172 self.center_distance_expansion = GaussianSmearing(0.0, center_max_distance, distance_embed_dim)173 self.center_edge_embedding = nn.Sequential(174 nn.Linear(distance_embed_dim + sigma_embed_dim, ns),175 nn.ReLU(),176 nn.Dropout(dropout),177 nn.Linear(ns, ns)178 )179 180 self.final_conv = OldTensorProductConvLayer(181 in_irreps=self.conv_layers[-1].out_irreps,182 sh_irreps=self.sh_irreps,183 out_irreps=f'2x1o + 2x1e' if not self.odd_parity else '1x1o + 1x1e',184 n_edge_features=2 * ns,185 residual=False,186 dropout=dropout,187 batch_norm=batch_norm188 )189 190 self.tr_final_layer = nn.Sequential(nn.Linear(1 + sigma_embed_dim, ns),nn.Dropout(dropout), nn.ReLU(), nn.Linear(ns, 1))191 self.rot_final_layer = nn.Sequential(nn.Linear(1 + sigma_embed_dim, ns),nn.Dropout(dropout), nn.ReLU(), nn.Linear(ns, 1))192 193 if not no_torsion:194 # convolution for torsional score195 self.final_edge_embedding = nn.Sequential(196 nn.Linear(distance_embed_dim, ns),197 nn.ReLU(),198 nn.Dropout(dropout),199 nn.Linear(ns, ns)200 )201 self.final_tp_tor = o3.FullTensorProduct(self.sh_irreps, "2e")202 self.tor_bond_conv = OldTensorProductConvLayer(203 in_irreps=self.conv_layers[-1].out_irreps,204 sh_irreps=self.final_tp_tor.irreps_out,205 out_irreps=f'{ns}x0o + {ns}x0e' if not self.odd_parity else f'{ns}x0o',206 n_edge_features=3 * ns,207 residual=False,208 dropout=dropout,209 batch_norm=batch_norm210 )211 self.tor_final_layer = nn.Sequential(212 nn.Linear(2 * ns if not self.odd_parity else ns, ns, bias=False),213 nn.Tanh(),214 nn.Dropout(dropout),215 nn.Linear(ns, 1, bias=False)216 )217 218 @staticmethod219 def _resolve_edge_store(data, primary_key, fallback_key):220 edge_types = getattr(data, "edge_types", ())221 if primary_key in edge_types:222 return data[primary_key]223 if fallback_key in edge_types:224 return data[fallback_key]225 try:226 return data[primary_key]227 except Exception:228 return data[fallback_key]229 230 def _ligand_edge_store(self, data):231 return self._resolve_edge_store(232 data,233 ("ligand", "ligand"),234 ("ligand", "lig_bond", "ligand"),235 )236 237 def _receptor_edge_store(self, data):238 return self._resolve_edge_store(239 data,240 ("receptor", "receptor"),241 ("receptor", "rec_contact", "receptor"),242 )243 244 def _atom_edge_store(self, data):245 return self._resolve_edge_store(246 data,247 ("atom", "atom"),248 ("atom", "atom_contact", "atom"),249 )250 251 def _atom_receptor_edge_store(self, data):252 return self._resolve_edge_store(253 data,254 ("atom", "receptor"),255 ("atom", "atom_rec_contact", "receptor"),256 )257 258 def forward(self, data):259 if self.no_aminoacid_identities:260 data['receptor'].x = data['receptor'].x * 0261 262 if not self.confidence_mode:263 tr_sigma, rot_sigma, tor_sigma = self.t_to_sigma(*[data.complex_t[noise_type] for noise_type in ['tr', 'rot', 'tor']])264 else:265 tr_sigma, rot_sigma, tor_sigma = [data.complex_t[noise_type] for noise_type in ['tr', 'rot', 'tor']]266 267 # build ligand graph268 lig_node_attr, lig_edge_index, lig_edge_attr, lig_edge_sh, lig_edge_weight = self.build_lig_conv_graph(data)269 lig_node_attr = self.lig_node_embedding(lig_node_attr)270 lig_edge_attr = self.lig_edge_embedding(lig_edge_attr)271 272 # build receptor graph273 rec_node_attr, rec_edge_index, rec_edge_attr, rec_edge_sh, rec_edge_weight = self.build_rec_conv_graph(data)274 rec_node_attr = self.rec_node_embedding(rec_node_attr)275 rec_edge_attr = self.rec_edge_embedding(rec_edge_attr)276 277 # build atom graph278 atom_node_attr, atom_edge_index, atom_edge_attr, atom_edge_sh, atom_edge_weight = self.build_atom_conv_graph(data)279 atom_node_attr = self.atom_node_embedding(atom_node_attr)280 atom_edge_attr = self.atom_edge_embedding(atom_edge_attr)281 282 # build cross graph283 cross_cutoff = (tr_sigma * 3 + 20).unsqueeze(1) if self.dynamic_max_cross else self.cross_max_distance284 lr_edge_index, lr_edge_attr, lr_edge_sh, lr_edge_weight, la_edge_index, la_edge_attr, \285 la_edge_sh, la_edge_weight, ar_edge_index, ar_edge_attr, ar_edge_sh, ar_edge_weight = \286 self.build_cross_conv_graph(data, cross_cutoff)287 lr_edge_attr= self.lr_edge_embedding(lr_edge_attr)288 la_edge_attr = self.la_edge_embedding(la_edge_attr)289 ar_edge_attr = self.ar_edge_embedding(ar_edge_attr)290 291 for l in range(self.num_conv_layers):292 # LIGAND updates293 lig_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)294 lig_update = self.conv_layers[9*l](lig_node_attr, lig_edge_index, lig_edge_attr_, lig_edge_sh, edge_weight=lig_edge_weight)295 296 lr_edge_attr_ = torch.cat([lr_edge_attr, lig_node_attr[lr_edge_index[0], :self.ns], rec_node_attr[lr_edge_index[1], :self.ns]], -1)297 lr_update = self.conv_layers[9*l+1](rec_node_attr, lr_edge_index, lr_edge_attr_, lr_edge_sh,298 out_nodes=lig_node_attr.shape[0], edge_weight=lr_edge_weight)299 300 la_edge_attr_ = torch.cat([la_edge_attr, lig_node_attr[la_edge_index[0], :self.ns], atom_node_attr[la_edge_index[1], :self.ns]], -1)301 la_update = self.conv_layers[9*l+2](atom_node_attr, la_edge_index, la_edge_attr_, la_edge_sh,302 out_nodes=lig_node_attr.shape[0], edge_weight=la_edge_weight)303 304 if l != self.num_conv_layers-1: # last layer optimisation305 306 # ATOM UPDATES307 atom_edge_attr_ = torch.cat([atom_edge_attr, atom_node_attr[atom_edge_index[0], :self.ns], atom_node_attr[atom_edge_index[1], :self.ns]], -1)308 atom_update = self.conv_layers[9*l+3](atom_node_attr, atom_edge_index, atom_edge_attr_, atom_edge_sh, edge_weight=atom_edge_weight)309 310 al_edge_attr_ = torch.cat([la_edge_attr, atom_node_attr[la_edge_index[1], :self.ns], lig_node_attr[la_edge_index[0], :self.ns]], -1)311 al_update = self.conv_layers[9*l+4](lig_node_attr, torch.flip(la_edge_index, dims=[0]), al_edge_attr_,312 la_edge_sh, out_nodes=atom_node_attr.shape[0], edge_weight=la_edge_weight)313 314 ar_edge_attr_ = torch.cat([ar_edge_attr, atom_node_attr[ar_edge_index[0], :self.ns], rec_node_attr[ar_edge_index[1], :self.ns]],-1)315 ar_update = self.conv_layers[9*l+5](rec_node_attr, ar_edge_index, ar_edge_attr_, ar_edge_sh, out_nodes=atom_node_attr.shape[0], edge_weight=ar_edge_weight)316 317 # RECEPTOR updates318 rec_edge_attr_ = torch.cat([rec_edge_attr, rec_node_attr[rec_edge_index[0], :self.ns], rec_node_attr[rec_edge_index[1], :self.ns]], -1)319 rec_update = self.conv_layers[9*l+6](rec_node_attr, rec_edge_index, rec_edge_attr_, rec_edge_sh, edge_weight=rec_edge_weight)320 321 rl_edge_attr_ = torch.cat([lr_edge_attr, rec_node_attr[lr_edge_index[1], :self.ns], lig_node_attr[lr_edge_index[0], :self.ns]], -1)322 rl_update = self.conv_layers[9*l+7](lig_node_attr, torch.flip(lr_edge_index, dims=[0]), rl_edge_attr_,323 lr_edge_sh, out_nodes=rec_node_attr.shape[0], edge_weight=lr_edge_weight)324 325 ra_edge_attr_ = torch.cat([ar_edge_attr, rec_node_attr[ar_edge_index[1], :self.ns], atom_node_attr[ar_edge_index[0], :self.ns]], -1)326 ra_update = self.conv_layers[9*l+8](atom_node_attr, torch.flip(ar_edge_index, dims=[0]), ra_edge_attr_,327 ar_edge_sh, out_nodes=rec_node_attr.shape[0], edge_weight=ar_edge_weight)328 329 # padding original features and update features with residual updates330 lig_node_attr = F.pad(lig_node_attr, (0, lig_update.shape[-1] - lig_node_attr.shape[-1]))331 lig_node_attr = lig_node_attr + lig_update + la_update + lr_update332 333 if l != self.num_conv_layers - 1: # last layer optimisation334 atom_node_attr = F.pad(atom_node_attr, (0, atom_update.shape[-1] - atom_node_attr.shape[-1]))335 atom_node_attr = atom_node_attr + atom_update + al_update + ar_update336 rec_node_attr = F.pad(rec_node_attr, (0, rec_update.shape[-1] - rec_node_attr.shape[-1]))337 rec_node_attr = rec_node_attr + rec_update + ra_update + rl_update338 339 # confidence and affinity prediction340 if self.confidence_mode:341 scalar_lig_attr = torch.cat([lig_node_attr[:,:self.ns],lig_node_attr[:,-self.ns:]], dim=1) if self.num_conv_layers >= 3 else lig_node_attr[:,:self.ns]342 confidence = self.confidence_predictor(scatter_mean(scalar_lig_attr, data['ligand'].batch if self.parallel == 1 else data['ligand'].batch_parallel, dim=0)).squeeze(dim=-1)343 344 if self.parallel > 1:345 confidence, affinity = confidence[:, 0], confidence[:, 1:]346 confidence = confidence.reshape(data.num_graphs, self.parallel)347 affinity = affinity.reshape(data.num_graphs, self.parallel, -1)348 affinity = torch.cat([AGGREGATORS[agg](affinity) for agg in self.parallel_aggregators], dim=-1)349 affinity = self.affinity_predictor(affinity).squeeze(dim=-1)350 confidence = confidence, affinity351 return confidence352 assert self.parallel == 1353 354 # compute translational and rotational score vectors355 center_edge_index, center_edge_attr, center_edge_sh = self.build_center_conv_graph(data)356 center_edge_attr = self.center_edge_embedding(center_edge_attr)357 if self.fixed_center_conv:358 center_edge_attr = torch.cat([center_edge_attr, lig_node_attr[center_edge_index[1], :self.ns]], -1)359 else:360 center_edge_attr = torch.cat([center_edge_attr, lig_node_attr[center_edge_index[0], :self.ns]], -1)361 global_pred = self.final_conv(lig_node_attr, center_edge_index, center_edge_attr, center_edge_sh, out_nodes=data.num_graphs)362 363 tr_pred = global_pred[:, :3] + (global_pred[:, 6:9] if not self.odd_parity else 0)364 rot_pred = global_pred[:, 3:6] + (global_pred[:, 9:] if not self.odd_parity else 0)365 366 if self.separate_noise_schedule:367 data.graph_sigma_emb = torch.cat([self.timestep_emb_func(data.complex_t[noise_type]) for noise_type in ['tr', 'rot', 'tor']], dim=1)368 elif self.asyncronous_noise_schedule:369 data.graph_sigma_emb = self.timestep_emb_func(data.complex_t['t'])370 else: # tr rot and tor noise is all the same in this case371 data.graph_sigma_emb = self.timestep_emb_func(data.complex_t['tr'])372 373 # adjust the magniture of the score vectors374 tr_norm = torch.linalg.vector_norm(tr_pred, dim=1).unsqueeze(1)375 tr_pred = tr_pred / tr_norm * self.tr_final_layer(torch.cat([tr_norm, data.graph_sigma_emb], dim=1))376 377 rot_norm = torch.linalg.vector_norm(rot_pred, dim=1).unsqueeze(1)378 rot_pred = rot_pred / rot_norm * self.rot_final_layer(torch.cat([rot_norm, data.graph_sigma_emb], dim=1))379 380 if self.scale_by_sigma:381 tr_pred = tr_pred / tr_sigma.unsqueeze(1)382 rot_pred = rot_pred * so3.score_norm(rot_sigma.cpu()).unsqueeze(1).to(data['ligand'].x.device)383 384 if self.no_torsion or data['ligand'].edge_mask.sum() == 0: return tr_pred, rot_pred, torch.empty(0,device=self.device)385 386 # torsional components387 tor_bonds, tor_edge_index, tor_edge_attr, tor_edge_sh, tor_edge_weight = self.build_bond_conv_graph(data)388 tor_bond_vec = data['ligand'].pos[tor_bonds[1]] - data['ligand'].pos[tor_bonds[0]]389 tor_bond_attr = lig_node_attr[tor_bonds[0]] + lig_node_attr[tor_bonds[1]]390 391 tor_bonds_sh = o3.spherical_harmonics("2e", tor_bond_vec, normalize=True, normalization='component')392 tor_edge_sh = self.final_tp_tor(tor_edge_sh, tor_bonds_sh[tor_edge_index[0]])393 394 tor_edge_attr = torch.cat([tor_edge_attr, lig_node_attr[tor_edge_index[1], :self.ns],395 tor_bond_attr[tor_edge_index[0], :self.ns]], -1)396 tor_pred = self.tor_bond_conv(lig_node_attr, tor_edge_index, tor_edge_attr, tor_edge_sh,397 out_nodes=data['ligand'].edge_mask.sum(), reduce='mean', edge_weight=tor_edge_weight)398 tor_pred = self.tor_final_layer(tor_pred).squeeze(1)399 ligand_edge_store = self._ligand_edge_store(data)400 edge_sigma = tor_sigma[data['ligand'].batch][ligand_edge_store.edge_index[0]][data['ligand'].edge_mask]401 402 if self.scale_by_sigma:403 tor_pred = tor_pred * torch.sqrt(torch.tensor(torus.score_norm(edge_sigma.cpu().numpy())).float()404 .to(data['ligand'].x.device))405 return tr_pred, rot_pred, tor_pred406 407 def get_edge_weight(self, edge_vec, max_norm):408 if self.smooth_edges:409 normalised_norm = torch.clip(edge_vec.norm(dim=-1) * np.pi / max_norm, max=np.pi)410 return 0.5 * (torch.cos(normalised_norm) + 1.0).unsqueeze(-1)411 return 1.0412 413 def build_lig_conv_graph(self, data):414 # build the graph between ligand atoms415 if self.separate_noise_schedule:416 data['ligand'].node_sigma_emb = torch.cat(417 [self.timestep_emb_func(data['ligand'].node_t[noise_type]) for noise_type in ['tr', 'rot', 'tor']],418 dim=1)419 elif self.asyncronous_noise_schedule:420 data['ligand'].node_sigma_emb = self.timestep_emb_func(data['ligand'].node_t['t'])421 else:422 data['ligand'].node_sigma_emb = self.timestep_emb_func(423 data['ligand'].node_t['tr']) # tr rot and tor noise is all the same424 425 if self.parallel == 1:426 radius_edges = radius_graph(data['ligand'].pos, self.lig_max_radius, data['ligand'].batch)427 else:428 batches = torch.zeros(data.num_graphs, device=data['ligand'].x.device).long()429 batches = batches.index_add(0, data['ligand'].batch, torch.ones(len(data['ligand'].batch), device=data['ligand'].x.device).long())430 outer_batches = data.num_graphs431 b = [torch.ones(batches[i].item()//self.parallel, device=data['ligand'].x.device).long() * (self.parallel * i + j)432 for i in range(outer_batches) for j in range(self.parallel)]433 data['ligand'].batch_parallel = torch.cat(b)434 radius_edges = radius_graph(data['ligand'].pos, self.lig_max_radius, data['ligand'].batch_parallel)435 ligand_edge_store = self._ligand_edge_store(data)436 edge_index = torch.cat([ligand_edge_store.edge_index, radius_edges], 1).long()437 edge_attr = torch.cat([438 ligand_edge_store.edge_attr,439 torch.zeros(radius_edges.shape[-1], self.in_lig_edge_features, device=data['ligand'].x.device)440 ], 0)441 442 edge_sigma_emb = data['ligand'].node_sigma_emb[edge_index[0].long()]443 edge_attr = torch.cat([edge_attr, edge_sigma_emb], 1)444 node_attr = torch.cat([data['ligand'].x, data['ligand'].node_sigma_emb], 1)445 446 src, dst = edge_index447 edge_vec = data['ligand'].pos[dst.long()] - data['ligand'].pos[src.long()]448 edge_length_emb = self.lig_distance_expansion(edge_vec.norm(dim=-1))449 450 edge_attr = torch.cat([edge_attr, edge_length_emb], 1)451 edge_sh = o3.spherical_harmonics(self.sh_irreps, edge_vec, normalize=True, normalization='component')452 edge_weight = self.get_edge_weight(edge_vec, self.lig_max_radius)453 454 return node_attr, edge_index, edge_attr, edge_sh, edge_weight455 456 def build_rec_conv_graph(self, data):457 # build the graph between receptor residues458 if self.separate_noise_schedule:459 data['receptor'].node_sigma_emb = torch.cat(460 [self.timestep_emb_func(data['receptor'].node_t[noise_type]) for noise_type in ['tr', 'rot', 'tor']],461 dim=1)462 elif self.asyncronous_noise_schedule:463 data['receptor'].node_sigma_emb = self.timestep_emb_func(data['receptor'].node_t['t'])464 else:465 data['receptor'].node_sigma_emb = self.timestep_emb_func(466 data['receptor'].node_t['tr']) # tr rot and tor noise is all the same467 node_attr = torch.cat([data['receptor'].x, data['receptor'].node_sigma_emb], 1)468 469 # this assumes the edges were already created in preprocessing since protein's structure is fixed470 edge_index = self._receptor_edge_store(data).edge_index471 src, dst = edge_index472 edge_vec = data['receptor'].pos[dst.long()] - data['receptor'].pos[src.long()]473 #assert torch.all(edge_vec.norm(dim=-1) < self.rec_max_radius)474 475 edge_length_emb = self.rec_distance_expansion(edge_vec.norm(dim=-1))476 edge_sigma_emb = data['receptor'].node_sigma_emb[edge_index[0].long()]477 edge_attr = torch.cat([edge_sigma_emb, edge_length_emb], 1)478 edge_sh = o3.spherical_harmonics(self.sh_irreps, edge_vec, normalize=True, normalization='component')479 edge_weight = self.get_edge_weight(edge_vec, self.rec_max_radius)480 481 return node_attr, edge_index, edge_attr, edge_sh, edge_weight482 483 def build_atom_conv_graph(self, data):484 # build the graph between receptor atoms485 if self.separate_noise_schedule:486 data['atom'].node_sigma_emb = torch.cat([self.timestep_emb_func(data['atom'].node_t[noise_type]) for noise_type in ['tr', 'rot', 'tor']],dim=1)487 elif self.asyncronous_noise_schedule:488 data['atom'].node_sigma_emb = self.timestep_emb_func(data['atom'].node_t['t'])489 else:490 data['atom'].node_sigma_emb = self.timestep_emb_func(data['atom'].node_t['tr']) # tr rot and tor noise is all the same491 node_attr = torch.cat([data['atom'].x, data['atom'].node_sigma_emb], 1)492 493 # this assumes the edges were already created in preprocessing since protein's structure is fixed494 edge_index = self._atom_edge_store(data).edge_index495 src, dst = edge_index496 edge_vec = data['atom'].pos[dst.long()] - data['atom'].pos[src.long()]497 498 edge_length_emb = self.lig_distance_expansion(edge_vec.norm(dim=-1))499 edge_sigma_emb = data['atom'].node_sigma_emb[edge_index[0].long()]500 edge_attr = torch.cat([edge_sigma_emb, edge_length_emb], 1)501 edge_sh = o3.spherical_harmonics(self.sh_irreps, edge_vec, normalize=True, normalization='component')502 edge_weight = self.get_edge_weight(edge_vec, self.lig_max_radius)503 504 return node_attr, edge_index, edge_attr, edge_sh, edge_weight505 506 def build_cross_conv_graph(self, data, lr_cross_distance_cutoff):507 # build the cross edges between ligan atoms, receptor residues and receptor atoms508 509 # LIGAND to RECEPTOR510 if torch.is_tensor(lr_cross_distance_cutoff):511 # different cutoff for every graph512 lr_edge_index = radius(data['receptor'].pos / lr_cross_distance_cutoff[data['receptor'].batch],513 data['ligand'].pos / lr_cross_distance_cutoff[data['ligand'].batch], 1,514 data['receptor'].batch, data['ligand'].batch, max_num_neighbors=10000)515 else:516 lr_edge_index = radius(data['receptor'].pos, data['ligand'].pos, lr_cross_distance_cutoff,517 data['receptor'].batch, data['ligand'].batch, max_num_neighbors=10000)518 519 lr_edge_vec = data['receptor'].pos[lr_edge_index[1].long()] - data['ligand'].pos[lr_edge_index[0].long()]520 lr_edge_length_emb = self.cross_distance_expansion(lr_edge_vec.norm(dim=-1))521 lr_edge_sigma_emb = data['ligand'].node_sigma_emb[lr_edge_index[0].long()]522 lr_edge_attr = torch.cat([lr_edge_sigma_emb, lr_edge_length_emb], 1)523 lr_edge_sh = o3.spherical_harmonics(self.sh_irreps, lr_edge_vec, normalize=True, normalization='component')524 525 cutoff_d = lr_cross_distance_cutoff[data['ligand'].batch[lr_edge_index[0]]].squeeze() \526 if torch.is_tensor(lr_cross_distance_cutoff) else lr_cross_distance_cutoff527 lr_edge_weight = self.get_edge_weight(lr_edge_vec, cutoff_d)528 529 # LIGAND to ATOM530 la_edge_index = radius(data['atom'].pos, data['ligand'].pos, self.lig_max_radius,531 data['atom'].batch, data['ligand'].batch, max_num_neighbors=10000)532 533 la_edge_vec = data['atom'].pos[la_edge_index[1].long()] - data['ligand'].pos[la_edge_index[0].long()]534 la_edge_length_emb = self.cross_distance_expansion(la_edge_vec.norm(dim=-1))535 la_edge_sigma_emb = data['ligand'].node_sigma_emb[la_edge_index[0].long()]536 la_edge_attr = torch.cat([la_edge_sigma_emb, la_edge_length_emb], 1)537 la_edge_sh = o3.spherical_harmonics(self.sh_irreps, la_edge_vec, normalize=True, normalization='component')538 la_edge_weight = self.get_edge_weight(la_edge_vec, self.lig_max_radius)539 540 # ATOM to RECEPTOR541 ar_edge_index = self._atom_receptor_edge_store(data).edge_index542 ar_edge_vec = data['receptor'].pos[ar_edge_index[1].long()] - data['atom'].pos[ar_edge_index[0].long()]543 ar_edge_length_emb = self.rec_distance_expansion(ar_edge_vec.norm(dim=-1))544 ar_edge_sigma_emb = data['atom'].node_sigma_emb[ar_edge_index[0].long()]545 ar_edge_attr = torch.cat([ar_edge_sigma_emb, ar_edge_length_emb], 1)546 ar_edge_sh = o3.spherical_harmonics(self.sh_irreps, ar_edge_vec, normalize=True, normalization='component')547 ar_edge_weight = 1548 549 return lr_edge_index, lr_edge_attr, lr_edge_sh, lr_edge_weight, la_edge_index, la_edge_attr, \550 la_edge_sh, la_edge_weight, ar_edge_index, ar_edge_attr, ar_edge_sh, ar_edge_weight551 552 def build_center_conv_graph(self, data):553 # build the filter for the convolution of the center with the ligand atoms554 # for translational and rotational score555 edge_index = torch.cat([data['ligand'].batch.unsqueeze(0), torch.arange(len(data['ligand'].batch)).to(data['ligand'].x.device).unsqueeze(0)], dim=0)556 557 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)558 center_pos.index_add_(0, index=data['ligand'].batch, source=data['ligand'].pos)559 center_pos = center_pos / torch.bincount(data['ligand'].batch).unsqueeze(1)560 561 edge_vec = data['ligand'].pos[edge_index[1]] - center_pos[edge_index[0]]562 edge_attr = self.center_distance_expansion(edge_vec.norm(dim=-1))563 edge_sigma_emb = data['ligand'].node_sigma_emb[edge_index[1].long()]564 edge_attr = torch.cat([edge_attr, edge_sigma_emb], 1)565 edge_sh = o3.spherical_harmonics(self.sh_irreps, edge_vec, normalize=True, normalization='component')566 return edge_index, edge_attr, edge_sh567 568 def build_bond_conv_graph(self, data):569 # build graph for the pseudotorque layer570 bonds = self._ligand_edge_store(data).edge_index[:, data['ligand'].edge_mask].long()571 bond_pos = (data['ligand'].pos[bonds[0]] + data['ligand'].pos[bonds[1]]) / 2572 bond_batch = data['ligand'].batch[bonds[0]]573 edge_index = radius(data['ligand'].pos, bond_pos, self.lig_max_radius, batch_x=data['ligand'].batch, batch_y=bond_batch)574 575 edge_vec = data['ligand'].pos[edge_index[1]] - bond_pos[edge_index[0]]576 edge_attr = self.lig_distance_expansion(edge_vec.norm(dim=-1))577 578 edge_attr = self.final_edge_embedding(edge_attr)579 edge_sh = o3.spherical_harmonics(self.sh_irreps, edge_vec, normalize=True, normalization='component')580 edge_weight = self.get_edge_weight(edge_vec, self.lig_max_radius)581 582 return bonds, edge_index, edge_attr, edge_sh, edge_weight583 