Skip to content

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.

View in Colab   GitHub source


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"

!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 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)

Evaluate

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