CoolFace
Modelpublic

OneScience-Group/SurfDock

sourceHugging Facemitupdated 17d agoView on Hugging Face
0likes25downloads
torsion.py99 linesDownload Raw Back to utils
1import networkx as nx2import numpy as np3import torch, copy4from scipy.spatial.transform import Rotation as R5from torch_geometric.utils import to_networkx6from torch_geometric.data import Data7 8"""9    Preprocessing and computation for torsional updates to conformers10"""11 12 13def get_transformation_mask(pyg_data):14    15    G = to_networkx(pyg_data.to_homogeneous(), to_undirected=False)16    17    to_rotate = []18    edges = pyg_data['ligand', 'ligand'].edge_index.T.numpy()19    # undireted graph, so edges are duplicated, only keep one torsion per edge20    for i in range(0, edges.shape[0], 2):21        assert edges[i, 0] == edges[i+1, 1]22 23        G2 = G.to_undirected()24        G2.remove_edge(*edges[i])25        if not nx.is_connected(G2):26            l = list(sorted(nx.connected_components(G2), key=len)[0])27            if len(l) > 1:28                29                if edges[i, 0] in l:30                    to_rotate.append([])31                    to_rotate.append(l)32                else:33                    to_rotate.append(l)34                    to_rotate.append([])35                continue36        to_rotate.append([])37        to_rotate.append([])38 39    mask_edges = np.asarray([0 if len(l) == 0 else 1 for l in to_rotate], dtype=bool)40    mask_rotate = np.zeros((np.sum(mask_edges), len(G.nodes())), dtype=bool)41    idx = 042    for i in range(len(G.edges())):43        if mask_edges[i]:44            mask_rotate[idx][np.asarray(to_rotate[i], dtype=int)] = True45            idx += 146 47    return mask_edges, mask_rotate48 49 50def modify_conformer_torsion_angles(pos, edge_index, mask_rotate, torsion_updates, as_numpy=False):51    pos = copy.deepcopy(pos)52    if type(pos) != np.ndarray: pos = pos.cpu().numpy()53 54    for idx_edge, e in enumerate(edge_index.cpu().numpy()):55        if torsion_updates[idx_edge] == 0:56            continue57        u, v = e[0], e[1]58 59        # check if need to reverse the edge, v should be connected to the part that gets rotated60 61        if type(mask_rotate) is list:62            mask_rotate = mask_rotate[0]63        assert not mask_rotate[idx_edge, u]64        assert mask_rotate[idx_edge, v]65 66        rot_vec = pos[u] - pos[v]  # convention: positive rotation if pointing inwards67        rot_vec = rot_vec * torsion_updates[idx_edge] / np.linalg.norm(rot_vec) # idx_edge!68        rot_mat = R.from_rotvec(rot_vec).as_matrix()69        pos[mask_rotate[idx_edge]] = (pos[mask_rotate[idx_edge]] - pos[v]) @ rot_mat.T + pos[v]70    if not as_numpy: pos = torch.from_numpy(pos.astype(np.float32))71    return pos72 73 74def perturb_batch(data, torsion_updates, split=False, return_updates=False):75    if type(data) is Data:76        return modify_conformer_torsion_angles(data.pos,77                                               data.edge_index.T[data.edge_mask],78                                               data.mask_rotate, torsion_updates)79    pos_new = [] if split else copy.deepcopy(data.pos)80    edges_of_interest = data.edge_index.T[data.edge_mask]81    idx_node = 082    idx_edges = 083    torsion_update_list = []84    for i, mask_rotate in enumerate(data.mask_rotate):85        pos = data.pos[idx_node:idx_node + mask_rotate.shape[1]]86        edges = edges_of_interest[idx_edges:idx_edges + mask_rotate.shape[0]] - idx_node87        torsion_update = torsion_updates[idx_edges:idx_edges + mask_rotate.shape[0]]88        torsion_update_list.append(torsion_update)89        pos_new_ = modify_conformer_torsion_angles(pos, edges, mask_rotate, torsion_update)90        if split:91            pos_new.append(pos_new_)92        else:93            pos_new[idx_node:idx_node + mask_rotate.shape[1]] = pos_new_94 95        idx_node += mask_rotate.shape[1]96        idx_edges += mask_rotate.shape[0]97    if return_updates:98        return pos_new, torsion_update_list99    return pos_new