Team Ai
Modelpublic

OneScience-Group/DiffDock

sourceHugging Facemitupdated 2mo agoView on Hugging Face
0likes52downloads
old_aa_model.py583 linesDownload Raw Back to models
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