Skip to content

MinCut Pooling on PROTEINS

Author: K3-Node Team
Backend: Multi-Backend
Dataset: PROTEINS (TUDataset)
Description: Classify proteins as enzymes or non-enzymes.

View in Colab   GitHub source


MinCut Pooling on PROTEINS

Classify proteins as enzymes or non-enzymes. Each of the 1,113 proteins in the PROTEINS dataset is a graph whose nodes are secondary-structure elements, connected when they are close in the 3D structure. After one GCN layer, each protein is converted to a dense (padded) representation with to_dense_batch and to_dense_adj, and pooled twice into fewer clusters. MinCut pooling (Bianchi et al., 2020) learns cluster assignments that cut as few edges as possible, using two auxiliary losses (mincut and orthogonality).

Same model as PyG's examples/proteins_mincut_pool.py; trained for 200 epochs instead of PyG's 15,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 math
import keras
from keras import ops
from k3_node.datasets import TUDataset
from k3_node.layers import DenseGraphConv, GCNConv, dense_mincut_pool, to_dense_adj, to_dense_batch
from k3_node.loader import DataLoader

dataset = TUDataset("data/TU", name="PROTEINS").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_loader = DataLoader(dataset[2 * n:], batch_size=20, shuffle=True)
print(dataset)
avg_num_nodes = sum(graph.num_nodes for graph in dataset) / len(dataset)

Define the model

The mincut and orthogonality losses are added with add_loss.

class MinCutNet(keras.Model):
    def __init__(self, in_channels, out_channels, avg_num_nodes, hidden_channels=32):
        super().__init__()
        self.conv1 = GCNConv(in_channels, hidden_channels)
        num_clusters = math.ceil(0.5 * avg_num_nodes)
        self.pool1 = keras.layers.Dense(num_clusters)  # cluster assignment scores
        self.conv2 = DenseGraphConv(hidden_channels, hidden_channels)
        num_clusters = math.ceil(0.5 * num_clusters)
        self.pool2 = keras.layers.Dense(num_clusters)
        self.conv3 = DenseGraphConv(hidden_channels, hidden_channels)
        self.lin1 = keras.layers.Dense(hidden_channels, activation="relu")
        self.lin2 = keras.layers.Dense(out_channels)

    def call(self, data):
        x = ops.relu(self.conv1(data.x, data.edge_index))
        x, mask = to_dense_batch(x, data.batch, dim_size=data.num_graphs)
        adj = to_dense_adj(data.edge_index, data.batch, batch_size=data.num_graphs)
        x, adj, mincut_loss1, ortho_loss1 = dense_mincut_pool(x, adj, self.pool1(x), mask)
        x = ops.relu(self.conv2(x, adj))
        x, adj, mincut_loss2, ortho_loss2 = dense_mincut_pool(x, adj, self.pool2(x))
        x = self.conv3(x, adj)
        self.add_loss(mincut_loss1 + mincut_loss2 + ortho_loss1 + ortho_loss2)
        return self.lin2(self.lin1(ops.mean(x, axis=1)))


model = MinCutNet(dataset.num_features, dataset.num_classes, avg_num_nodes)

Train

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=5e-4, weight_decay=1e-4),
    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}")