Skip to content

Memory-based Graph Pooling (MemPool) on PROTEINS

Author: K3-Node Team
Backend: Multi-Backend
Dataset: PROTEINS_full (TUDataset)
Description: Classify proteins as enzymes or non-enzymes (PROTEINS_full, which includes extra node attributes).

View in Colab   GitHub source


Memory-based Graph Pooling (MemPool) on PROTEINS

Classify proteins as enzymes or non-enzymes (PROTEINS_full, which includes extra node attributes). Two residual GAT blocks compute node embeddings; memory-based pooling (Khasahmadi et al., 2020) then softly assigns nodes to 10 learned "memory" clusters and finally to a single vector per protein. A KL loss keeps the cluster assignments sharp.

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

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

Load the data

import keras
from keras import ops
from k3_node.datasets import TUDataset
from k3_node.layers import GATConv, MemPooling
from k3_node.loader import DataLoader
from k3_node.models import DeepGCNLayer

dataset = TUDataset("data/TU", name="PROTEINS_full", use_node_attr=True).shuffle()
n = (len(dataset) + 9) // 10  # 10% test, 10% validation, 80% training
test_loader = DataLoader(dataset[:n], batch_size=20)
val_loader = DataLoader(dataset[n:2 * n], batch_size=20)
train_dataset = dataset[2 * n:]
train_loader = DataLoader(train_dataset, batch_size=20, shuffle=True)
print(dataset)

Define the model

The KL losses of both pooling steps are added with add_loss, divided by the number of training graphs as in PyG.

class MemPoolNet(keras.Model):
    def __init__(self, hidden_channels, out_channels, dropout, kl_weight):
        super().__init__()
        self.kl_weight = kl_weight
        self.lin = keras.layers.Dense(hidden_channels)
        self.convs = [DeepGCNLayer(GATConv(hidden_channels, hidden_channels, dropout=dropout),
                                   keras.layers.BatchNormalization(), keras.layers.LeakyReLU(),
                                   block="res+", dropout=dropout) for _ in range(2)]
        self.mem1 = MemPooling(hidden_channels, 80, heads=5, num_clusters=10)
        self.mem2 = MemPooling(80, out_channels, heads=5, num_clusters=1)
        self.dropout = keras.layers.Dropout(dropout)

    def call(self, data, training=False):
        x = self.lin(data.x)
        for conv in self.convs:
            x = conv(x, data.edge_index, training=training)
        x, S1 = self.mem1(x, data.batch, batch_size=data.num_graphs)  # [graphs, 10 clusters, 80]
        x = self.dropout(ops.leaky_relu(x), training=training)
        x, S2 = self.mem2(x)  # [graphs, 1, classes]
        self.add_loss(self.kl_weight * (MemPooling.kl_loss(S1) + MemPooling.kl_loss(S2)))
        return ops.squeeze(x, axis=1)


model = MemPoolNet(hidden_channels=32, out_channels=dataset.num_classes, dropout=0.1,
                   kl_weight=1 / len(train_dataset))

Train

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=1e-3, weight_decay=4e-5),
    loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    metrics=["accuracy"],
)
model.fit(train_loader, validation_data=val_loader, epochs=200, verbose=2)

Evaluate

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