Sirius16's picture
Upload 22 files
7e254aa verified
Raw
History Blame Contribute Delete
3.9 kB
import os
import pickle
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from flask import Flask, render_template, request, jsonify
from torch_geometric.nn import SAGEConv
class RenameUnpickler(pickle.Unpickler):
def find_class(self, module, name):
if module == 'numpy._core.multiarray':
module = 'numpy.core.multiarray'
elif module == 'numpy._core':
module = 'numpy.core'
return super().find_class(module, name)
class GraphSAGE(nn.Module):
def __init__(self, in_dim, hidden_dim, dropout=0.3):
super().__init__()
self.conv1 = SAGEConv(in_dim, hidden_dim, aggr='max')
self.conv2 = SAGEConv(hidden_dim, hidden_dim, aggr='max')
self.conv3 = SAGEConv(hidden_dim, hidden_dim, aggr='max')
self.dropout = nn.Dropout(dropout)
def forward(self, x, edge_index):
x = F.relu(self.conv1(x, edge_index))
x = self.dropout(x)
x = F.relu(self.conv2(x, edge_index))
x = self.dropout(x)
x = F.relu(self.conv3(x, edge_index))
x = self.dropout(x)
return x
class LinkPredictor(nn.Module):
def __init__(self, hidden_dim, dropout=0.3):
super().__init__()
self.mlp = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(hidden_dim, 1)
)
def forward(self, h, edge):
src = h[edge[:, 0]]
dst = h[edge[:, 1]]
return self.mlp(src * dst).squeeze()
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
with open("node2idx.pkl", "rb") as f:
node2idx = RenameUnpickler(f).load()
valid_proteins = list(node2idx.keys())
with open("protein_embeddings_protT5_1056.pkl", "rb") as f:
prot_embeddings = RenameUnpickler(f).load()
embedding_dim = len(next(iter(prot_embeddings.values())))
x_array = np.zeros((len(node2idx), embedding_dim))
for node, idx in node2idx.items():
x_array[idx] = prot_embeddings[node]
x = torch.tensor(x_array, dtype=torch.float).to(device)
edge_index = torch.load("edge_index.pt", map_location=device)
checkpoint = torch.load("graphsage_asd_model.pth", map_location=device)
sage = GraphSAGE(checkpoint['embedding_dim'], checkpoint['hidden_dim']).to(device)
link = LinkPredictor(checkpoint['hidden_dim']).to(device)
sage.load_state_dict(checkpoint['sage_state_dict'])
link.load_state_dict(checkpoint['link_state_dict'])
sage.eval()
link.eval()
with torch.no_grad():
h_global = sage(x, edge_index)
app = Flask(__name__)
@app.route('/')
def home():
return render_template('index.html')
@app.route('/api/proteins', methods=['GET'])
def get_proteins():
return jsonify(valid_proteins)
@app.route('/api/predict', methods=['POST'])
def predict():
try:
p1 = request.form.get('protein1', '').strip()
p2 = request.form.get('protein2', '').strip()
if not p1 or not p2:
return jsonify({"status": "error", "message": "ID Protein tidak valid"})
if p1 not in node2idx or p2 not in node2idx:
return jsonify({"status": "error", "message": "Protein tidak ditemukan dalam database graf"})
idx1 = node2idx[p1]
idx2 = node2idx[p2]
edge = torch.tensor([[idx1, idx2]], dtype=torch.long, device=device)
with torch.no_grad():
logits = link(h_global, edge)
prob = logits.sigmoid().item()
return jsonify({
"status": "success",
"protein1": p1,
"protein2": p2,
"probability": prob
})
except Exception as e:
return jsonify({"status": "error", "message": str(e)})
if __name__ == '__main__':
app.run(host='0.0.0.0', port=7860)