Multi-label protein function prediction with GAT (PPI)
Author: K3-Node Team
Backend: Multi-Backend
Dataset: PPI
Description: Predict the functions of proteins in protein-protein interaction (PPI) graphs.
Multi-label protein function prediction with GAT (PPI)
Predict the functions of proteins in protein-protein interaction (PPI) graphs. Every protein (node) can have several of 121 functions, so this is multi-label classification: the model outputs one score per function. The model trains on 20 graphs and is tested on 2 graphs it has never seen (inductive learning). Three graph attention layers (Veličković et al., 2018) with residual connections compute the scores.
Same model as PyG's examples/ppi.py.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
import keras
from keras import ops
from k3_node.datasets import PPI
from k3_node.layers import GATConv
from k3_node.loader import DataLoader
from k3_node.metrics import F1Score
train_dataset = PPI("data/PPI", split="train")
val_dataset = PPI("data/PPI", split="val")
test_dataset = PPI("data/PPI", split="test")
train_loader = DataLoader(train_dataset, batch_size=1, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=2)
test_loader = DataLoader(test_dataset, batch_size=2)
print(train_dataset[0])
Define the model
class GAT(keras.Model):
def __init__(self, in_channels, out_channels):
super().__init__()
self.conv1 = GATConv(in_channels, 256, heads=4, residual=True)
self.conv2 = GATConv(4 * 256, 256, heads=4, residual=True)
self.conv3 = GATConv(4 * 256, out_channels, heads=6, concat=False, residual=True)
def call(self, data, training=False):
x = ops.elu(self.conv1(data.x, data.edge_index, training=training))
x = ops.elu(self.conv2(x, data.edge_index, training=training))
return self.conv3(x, data.edge_index, training=training)
model = GAT(train_dataset.num_features, train_dataset.num_classes)
Train
Each of the 121 outputs is a yes/no prediction, so the loss is binary cross-entropy. Performance is measured with the micro-averaged F1 score.
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.005),
loss=keras.losses.BinaryCrossentropy(from_logits=True),
metrics=[F1Score(average="micro", from_logits=True)],
)
model.fit(train_loader, validation_data=val_loader, epochs=100, verbose=2)