MinCut Pooling on PROTEINS
Author: K3-Node Team
Backend: Multi-Backend
Dataset: PROTEINS (TUDataset)
Description: Classify proteins as enzymes or non-enzymes.
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"
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)