Feature-wise Linear Modulation (FiLM) on PPI
Author: K3-Node Team
Backend: Multi-Backend
Dataset: PPI
Description: Predict the functions of proteins in protein-protein interaction networks.
Feature-wise Linear Modulation (FiLM) on PPI
Predict the functions of proteins in protein-protein interaction networks. GNN-FiLM (Brockschmidt, 2020) lets each node's own features decide how the messages from its neighbors are scaled and shifted.
Same model as PyG's examples/film.py; trained for 200 epochs instead of PyG's 500 to keep the notebook quick.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
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 FiLMConv
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=2, 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 FiLM(keras.Model):
def __init__(self, in_channels, hidden_channels, out_channels, num_layers, dropout):
super().__init__()
self.convs = [FiLMConv(in_channels, hidden_channels)]
self.convs += [FiLMConv(hidden_channels, hidden_channels) for _ in range(num_layers - 2)]
self.convs += [FiLMConv(hidden_channels, out_channels, act=None)]
self.norms = [keras.layers.BatchNormalization() for _ in range(num_layers - 1)]
self.dropout = keras.layers.Dropout(dropout)
def call(self, data, training=False):
x = data.x
for conv, norm in zip(self.convs[:-1], self.norms):
x = norm(conv(x, data.edge_index), training=training)
x = self.dropout(x, training=training)
return self.convs[-1](x, data.edge_index)
model = FiLM(train_dataset.num_features, hidden_channels=320, out_channels=train_dataset.num_classes,
num_layers=4, dropout=0.1)
Train
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.01),
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=200, verbose=2)