Skip to content

Multi-label protein function prediction with Cluster-GCN (PPI)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: PPI
Description: Predict the functions of proteins in protein-protein interaction (PPI) graphs.

View in Colab   GitHub source


Multi-label protein function prediction with Cluster-GCN (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). Cluster-GCN (Chiang et al., 2019) splits the training graphs into 50 clusters of densely connected nodes and trains on one cluster at a time, which keeps memory use low on large graphs. The model is a 6-layer GraphSAGE network with batch normalization.

Same model as PyG's examples/cluster_gcn_ppi.py.

Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"

!pip install k3-node[examples]
import os
os.environ["KERAS_BACKEND"] = "tensorflow"

Load the data

All training graphs are combined into one big graph and partitioned into 50 clusters; ClusterLoader yields one cluster (with the edges inside it) per step.

import keras
from keras import ops
from k3_node.data import Batch
from k3_node.datasets import PPI
from k3_node.layers import BatchNorm, SAGEConv
from k3_node.loader import ClusterData, ClusterLoader, 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_data = Batch.from_data_list(list(train_dataset))
cluster_data = ClusterData(train_data, num_parts=50)
train_loader = ClusterLoader(cluster_data, batch_size=1, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=2)
test_loader = DataLoader(test_dataset, batch_size=2)
print(cluster_data)

Define the model

class GraphSAGE(keras.Model):
    def __init__(self, in_channels, hidden_channels, out_channels, num_layers):
        super().__init__()
        sizes = [in_channels] + [hidden_channels] * (num_layers - 1)
        self.convs = [SAGEConv(c, hidden_channels) for c in sizes[:-1]] + [SAGEConv(hidden_channels, out_channels)]
        self.batch_norms = [BatchNorm(hidden_channels) for _ in range(num_layers - 1)]
        self.dropout = keras.layers.Dropout(0.2)

    def call(self, data, training=False):
        x = data.x
        for conv, batch_norm in zip(self.convs[:-1], self.batch_norms):
            x = ops.relu(batch_norm(conv(x, data.edge_index), training=training))
            x = self.dropout(x, training=training)
        return self.convs[-1](x, data.edge_index)


model = GraphSAGE(train_dataset.num_features, 1024, train_dataset.num_classes, num_layers=6)

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.01),
    loss=keras.losses.BinaryCrossentropy(from_logits=True),
    metrics=[F1Score(average="micro", from_logits=True)],
)
model.fit(train_loader, validation_data=val_loader, epochs=200, verbose=2)

Evaluate

loss, f1 = model.evaluate(test_loader, verbose=0)
print(f"Test F1: {f1:.4f}")