GPSite / scripts /model.py
anzhi2710gmailcom's picture
Upload folder using huggingface_hub
40e5504 verified
Raw
History Blame Contribute Delete
11.8 kB
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.data as data
from torch_scatter import scatter_mean
import torch_geometric
from torch_geometric.nn import radius_graph, TransformerConv
############## Model ##############
class GNNLayer(nn.Module):
def __init__(self, num_hidden, dropout=0.2, num_heads=4):
super(GNNLayer, self).__init__()
self.dropout = nn.Dropout(dropout)
self.norm = nn.ModuleList([nn.LayerNorm(num_hidden) for _ in range(2)])
self.attention = TransformerConv(in_channels=num_hidden, out_channels=int(num_hidden / num_heads), heads=num_heads, dropout = dropout, edge_dim = num_hidden, root_weight=False)
self.PositionWiseFeedForward = nn.Sequential(
nn.Linear(num_hidden, num_hidden*4),
nn.ReLU(),
nn.Linear(num_hidden*4, num_hidden)
)
self.edge_update = EdgeMLP(num_hidden, dropout)
self.context = Context(num_hidden)
def forward(self, h_V, edge_index, h_E, batch_id):
dh = self.attention(h_V, edge_index, h_E)
h_V = self.norm[0](h_V + self.dropout(dh))
# Position-wise feedforward
dh = self.PositionWiseFeedForward(h_V)
h_V = self.norm[1](h_V + self.dropout(dh))
# update edge
h_E = self.edge_update(h_V, edge_index, h_E)
# context node update
h_V = self.context(h_V, batch_id)
return h_V, h_E
class EdgeMLP(nn.Module):
def __init__(self, num_hidden, dropout=0.2):
super(EdgeMLP, self).__init__()
self.dropout = nn.Dropout(dropout)
self.norm = nn.BatchNorm1d(num_hidden)
self.W11 = nn.Linear(3*num_hidden, num_hidden, bias=True)
self.W12 = nn.Linear(num_hidden, num_hidden, bias=True)
self.act = torch.nn.GELU()
def forward(self, h_V, edge_index, h_E):
src_idx = edge_index[0]
dst_idx = edge_index[1]
h_EV = torch.cat([h_V[src_idx], h_E, h_V[dst_idx]], dim=-1)
h_message = self.W12(self.act(self.W11(h_EV)))
h_E = self.norm(h_E + self.dropout(h_message))
return h_E
class Context(nn.Module):
def __init__(self, num_hidden):
super(Context, self).__init__()
self.V_MLP_g = nn.Sequential(
nn.Linear(num_hidden,num_hidden),
nn.ReLU(),
nn.Linear(num_hidden,num_hidden),
nn.Sigmoid()
)
def forward(self, h_V, batch_id):
c_V = scatter_mean(h_V, batch_id, dim=0)
h_V = h_V * self.V_MLP_g(c_V[batch_id])
return h_V
class Graph_encoder(nn.Module):
def __init__(self, node_in_dim, edge_in_dim, hidden_dim, num_layers=4, drop_rate=0.2):
super(Graph_encoder, self).__init__()
self.node_embedding = nn.Linear(node_in_dim, hidden_dim, bias=True)
self.edge_embedding = nn.Linear(edge_in_dim, hidden_dim, bias=True)
self.norm_nodes = nn.BatchNorm1d(hidden_dim)
self.norm_edges = nn.BatchNorm1d(hidden_dim)
self.W_v = nn.Linear(hidden_dim, hidden_dim, bias=True)
self.W_e = nn.Linear(hidden_dim, hidden_dim, bias=True)
self.layers = nn.ModuleList(
GNNLayer(num_hidden=hidden_dim, dropout=drop_rate, num_heads=4)
for _ in range(num_layers))
def forward(self, h_V, edge_index, h_E, batch_id):
h_V = self.W_v(self.norm_nodes(self.node_embedding(h_V)))
h_E = self.W_e(self.norm_edges(self.edge_embedding(h_E)))
for layer in self.layers:
h_V, h_E = layer(h_V, edge_index, h_E, batch_id)
return h_V
class GPSite(nn.Module):
def __init__(self, node_input_dim, edge_input_dim, hidden_dim, num_layers, augment_eps, dropout, task_list):
super(GPSite, self).__init__()
self.augment_eps = augment_eps
self.Graph_encoder = Graph_encoder(node_in_dim=node_input_dim, edge_in_dim=edge_input_dim, hidden_dim=hidden_dim, num_layers=num_layers, drop_rate=dropout)
self.task_list = task_list
for task in self.task_list:
self.add_module("FC_{}1".format(task), nn.Linear(hidden_dim, hidden_dim, bias=True))
self.add_module("FC_{}2".format(task), nn.Linear(hidden_dim, 1, bias=True))
# Initialization
for p in self.parameters():
if p.dim() > 1:
nn.init.xavier_uniform_(p)
def forward(self, X, h_V, edge_index, batch_id):
# Data augmentation
if self.training and self.augment_eps > 0:
X = X + self.augment_eps * torch.randn_like(X)
h_V = h_V + self.augment_eps * torch.randn_like(h_V)
h_V_geo, h_E = get_geo_feat(X, edge_index)
h_V = torch.cat([h_V, h_V_geo], dim=-1)
h_V = self.Graph_encoder(h_V, edge_index, h_E, batch_id) # [num_residue, hidden_dim]
output = []
for task in self.task_list:
emb = F.elu(self._modules["FC_{}1".format(task)](h_V))
emb = self._modules["FC_{}2".format(task)](emb)
output.append(emb)
output = torch.cat(output, dim=1)
return output
############## DataLoader ##############
class ProteinGraphDataset(data.Dataset):
def __init__(self, ID_list, outpath, radius=15):
super(ProteinGraphDataset, self).__init__()
self.IDs = ID_list
self.path = outpath
self.radius = radius
def __len__(self): return len(self.IDs)
def __getitem__(self, idx): return self._featurize_graph(idx)
def _featurize_graph(self, idx):
name = self.IDs[idx]
with torch.no_grad():
X = torch.load(self.path + "pdb/" + name + ".tensor")
prottrans_feat = torch.load(self.path + "ProtTrans/" + name + ".tensor")
dssp_feat = torch.load(self.path + 'DSSP/' + name + ".tensor")
pre_computed_node_feat = torch.cat([prottrans_feat, dssp_feat], dim=-1)
X_ca = X[:, 1]
edge_index = radius_graph(X_ca, r=self.radius, loop=True, max_num_neighbors = 1000, num_workers = 8)
graph_data = torch_geometric.data.Data(name=name, X=X, node_feat=pre_computed_node_feat, edge_index=edge_index)
return graph_data
############## Geometric Featurizer ##############
def get_geo_feat(X, edge_index):
pos_embeddings = _positional_embeddings(edge_index)
node_angles = _get_angle(X)
node_dist, edge_dist = _get_distance(X, edge_index)
node_direction, edge_direction, edge_orientation = _get_direction_orientation(X, edge_index)
geo_node_feat = torch.cat([node_angles, node_dist, node_direction], dim=-1)
geo_edge_feat = torch.cat([pos_embeddings, edge_orientation, edge_dist, edge_direction], dim=-1)
return geo_node_feat, geo_edge_feat
def _positional_embeddings(edge_index, num_embeddings=16):
d = edge_index[0] - edge_index[1]
frequency = torch.exp(
torch.arange(0, num_embeddings, 2, dtype=torch.float32, device=edge_index.device)
* -(np.log(10000.0) / num_embeddings)
)
angles = d.unsqueeze(-1) * frequency
PE = torch.cat((torch.cos(angles), torch.sin(angles)), -1)
return PE
def _get_angle(X, eps=1e-7):
# psi, omega, phi
X = torch.reshape(X[:, :3], [3*X.shape[0], 3])
dX = X[1:] - X[:-1]
U = F.normalize(dX, dim=-1)
u_2 = U[:-2]
u_1 = U[1:-1]
u_0 = U[2:]
# Backbone normals
n_2 = F.normalize(torch.cross(u_2, u_1), dim=-1)
n_1 = F.normalize(torch.cross(u_1, u_0), dim=-1)
# Angle between normals
cosD = torch.sum(n_2 * n_1, -1)
cosD = torch.clamp(cosD, -1 + eps, 1 - eps)
D = torch.sign(torch.sum(u_2 * n_1, -1)) * torch.acos(cosD)
D = F.pad(D, [1, 2]) # This scheme will remove phi[0], psi[-1], omega[-1]
D = torch.reshape(D, [-1, 3])
dihedral = torch.cat([torch.cos(D), torch.sin(D)], 1)
# alpha, beta, gamma
cosD = (u_2 * u_1).sum(-1) # alpha_{i}, gamma_{i}, beta_{i+1}
cosD = torch.clamp(cosD, -1 + eps, 1 - eps)
D = torch.acos(cosD)
D = F.pad(D, [1, 2])
D = torch.reshape(D, [-1, 3])
bond_angles = torch.cat((torch.cos(D), torch.sin(D)), 1)
node_angles = torch.cat((dihedral, bond_angles), 1)
return node_angles # dim = 12
def _rbf(D, D_min=0., D_max=20., D_count=16):
'''
Returns an RBF embedding of `torch.Tensor` `D` along a new axis=-1.
That is, if `D` has shape [...dims], then the returned tensor will have shape [...dims, D_count].
'''
D_mu = torch.linspace(D_min, D_max, D_count, device=D.device)
D_mu = D_mu.view([1, -1])
D_sigma = (D_max - D_min) / D_count
D_expand = torch.unsqueeze(D, -1)
RBF = torch.exp(-((D_expand - D_mu) / D_sigma) ** 2)
return RBF
def _get_distance(X, edge_index):
atom_N = X[:,0] # [L, 3]
atom_Ca = X[:,1]
atom_C = X[:,2]
atom_O = X[:,3]
atom_R = X[:,4]
node_list = ['Ca-N', 'Ca-C', 'Ca-O', 'N-C', 'N-O', 'O-C', 'R-N', 'R-Ca', "R-C", 'R-O']
node_dist = []
for pair in node_list:
atom1, atom2 = pair.split('-')
E_vectors = vars()['atom_' + atom1] - vars()['atom_' + atom2]
rbf = _rbf(E_vectors.norm(dim=-1))
node_dist.append(rbf)
node_dist = torch.cat(node_dist, dim=-1) # dim = [N, 10 * 16]
atom_list = ["N", "Ca", "C", "O", "R"]
edge_dist = []
for atom1 in atom_list:
for atom2 in atom_list:
E_vectors = vars()['atom_' + atom1][edge_index[0]] - vars()['atom_' + atom2][edge_index[1]]
rbf = _rbf(E_vectors.norm(dim=-1))
edge_dist.append(rbf)
edge_dist = torch.cat(edge_dist, dim=-1) # dim = [E, 25 * 16]
return node_dist, edge_dist
def _get_direction_orientation(X, edge_index): # N, CA, C, O, R
X_N = X[:,0] # [L, 3]
X_Ca = X[:,1]
X_C = X[:,2]
u = F.normalize(X_Ca - X_N, dim=-1)
v = F.normalize(X_C - X_Ca, dim=-1)
b = F.normalize(u - v, dim=-1)
n = F.normalize(torch.cross(u, v), dim=-1)
local_frame = torch.stack([b, n, torch.cross(b, n)], dim=-1) # [L, 3, 3] (3 column vectors)
node_j, node_i = edge_index
t = F.normalize(X[:, [0,2,3,4]] - X_Ca.unsqueeze(1), dim=-1) # [L, 4, 3]
node_direction = torch.matmul(t, local_frame).reshape(t.shape[0], -1) # [L, 4 * 3]
t = F.normalize(X[node_j] - X_Ca[node_i].unsqueeze(1), dim=-1) # [E, 5, 3]
edge_direction_ji = torch.matmul(t, local_frame[node_i]).reshape(t.shape[0], -1) # [E, 5 * 3]
t = F.normalize(X[node_i] - X_Ca[node_j].unsqueeze(1), dim=-1) # [E, 5, 3]
edge_direction_ij = torch.matmul(t, local_frame[node_j]).reshape(t.shape[0], -1) # [E, 5 * 3] # slightly improve performance
edge_direction = torch.cat([edge_direction_ji, edge_direction_ij], dim = -1) # [E, 2 * 5 * 3]
r = torch.matmul(local_frame[node_i].transpose(-1,-2), local_frame[node_j]) # [E, 3, 3]
edge_orientation = _quaternions(r) # [E, 4]
return node_direction, edge_direction, edge_orientation
def _quaternions(R):
""" Convert a batch of 3D rotations [R] to quaternions [Q]
R [E,3,3]
Q [E,4]
"""
diag = torch.diagonal(R, dim1=-2, dim2=-1)
Rxx, Ryy, Rzz = diag.unbind(-1)
magnitudes = 0.5 * torch.sqrt(torch.abs(1 + torch.stack([
Rxx - Ryy - Rzz,
- Rxx + Ryy - Rzz,
- Rxx - Ryy + Rzz
], -1)))
_R = lambda i,j: R[:,i,j]
signs = torch.sign(torch.stack([
_R(2,1) - _R(1,2),
_R(0,2) - _R(2,0),
_R(1,0) - _R(0,1)
], -1))
xyz = signs * magnitudes
# The relu enforces a non-negative trace
w = torch.sqrt(F.relu(1 + diag.sum(-1, keepdim=True))) / 2.
Q = torch.cat((xyz, w), -1)
Q = F.normalize(Q, dim=-1)
return Q