Multi-label protein function prediction with a deep GCNII (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 a deep GCNII (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). The model is a 9-layer GCNII (Chen et al., 2020): every layer mixes in the initial node representation and adds an identity mapping, which lets deep GCNs train well.
Same model as PyG's examples/gcn2_ppi.py; trained for 200 epochs instead of PyG's 2,000 to keep the notebook quick.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
GCNNorm computes the normalized GCN edge weights once, when the data is first processed.
import keras
from keras import ops
from k3_node.datasets import PPI
from k3_node.layers import GCN2Conv
from k3_node.loader import DataLoader
from k3_node.metrics import F1Score
from k3_node.transforms import GCNNorm
train_dataset = PPI("data/GCN2_PPI", split="train", pre_transform=GCNNorm())
val_dataset = PPI("data/GCN2_PPI", split="val", pre_transform=GCNNorm())
test_dataset = PPI("data/GCN2_PPI", split="test", pre_transform=GCNNorm())
train_loader = DataLoader(train_dataset, batch_size=2, 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 GCNII(keras.Model):
def __init__(self, hidden_channels, out_channels, num_layers, alpha, theta, shared_weights=True, dropout=0.0):
super().__init__()
self.lin_in = keras.layers.Dense(hidden_channels, activation="relu")
self.convs = [GCN2Conv(hidden_channels, alpha, theta, layer + 1, shared_weights, normalize=False)
for layer in range(num_layers)]
self.lin_out = keras.layers.Dense(out_channels)
self.dropout = keras.layers.Dropout(dropout)
def call(self, data, training=False):
x = x_0 = self.lin_in(self.dropout(data.x, training=training))
for conv in self.convs:
h = conv(self.dropout(x, training=training), x_0, data.edge_index, data.edge_weight)
x = ops.relu(h + x)
return self.lin_out(self.dropout(x, training=training))
model = GCNII(hidden_channels=2048, out_channels=train_dataset.num_classes, num_layers=9, alpha=0.5, theta=1.0,
shared_weights=False, dropout=0.2)
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.001),
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)