Skip to content

GeniePath on PPI

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

View in Colab   GitHub source


GeniePath on PPI

Predict the functions of proteins in protein-protein interaction networks. GeniePath (Liu et al., 2019) explores the graph adaptively in two directions: attention decides which neighbors matter (breadth), and an LSTM decides how much information from farther away to keep (depth). This is the "lazy" variant used by PyG's example: all breadth steps are computed first, then the LSTM runs over them.

Same model as PyG's examples/geniepath.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

PPI has 24 protein-protein interaction graphs (one per human tissue). Each protein (node) has 121 binary labels, so a protein can have several functions at once. Training, validation and test use different graphs.

import keras
from keras import ops
from k3_node.datasets import PPI
from k3_node.loader import DataLoader
from k3_node.metrics import F1Score
from k3_node.layers import GATConv

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)

Define the model

class GeniePathLazy(keras.Model):
    def __init__(self, out_channels, dim=256, lstm_hidden=256, num_layers=4):
        super().__init__()
        self.lstm_hidden = lstm_hidden
        self.lin1 = keras.layers.Dense(dim)
        self.breadths = [GATConv(dim, dim, heads=1) for _ in range(num_layers)]
        self.depths = [keras.layers.LSTMCell(lstm_hidden, use_bias=False) for _ in range(num_layers)]
        self.lin2 = keras.layers.Dense(out_channels)

    def call(self, data):
        x = self.lin1(data.x)
        h = c = ops.zeros((ops.shape(x)[0], self.lstm_hidden))  # LSTM state, one per node
        breadth = [ops.tanh(gat(x, data.edge_index)) for gat in self.breadths]
        for b, lstm in zip(breadth, self.depths):
            x, (h, c) = lstm(ops.concatenate([b, x], axis=-1), [h, c])
        return self.lin2(x)


model = GeniePathLazy(train_dataset.num_classes)

Train

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.005),
    loss=keras.losses.BinaryCrossentropy(from_logits=True),  # one yes/no decision per label
    metrics=[F1Score(average="micro", from_logits=True, name="f1")],
)
model.fit(train_loader, validation_data=val_loader, epochs=100, verbose=2)

Evaluate

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