CoolFace
Modelpublic

OneScience-Group/DiffDock

sourceHugging Facemitupdated 1mo agoView on Hugging Face
0likes31downloads
aa_model.py722 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_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