DMoN Pooling on PROTEINS
Author: K3-Node Team
Backend: Multi-Backend
Dataset: PROTEINS (TUDataset)
Description: Classify proteins as enzymes or non-enzymes.
DMoN 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.
DMoN (Tsitsulin et al., 2023) finds clusters by maximizing graph
modularity; its spectral and cluster-size losses are added to the classification loss.
Same model as PyG's examples/proteins_dmon_pool.py.
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, DMoNPooling, 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 spectral and cluster losses are added with add_loss.
class DMoNNet(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 = DMoNPooling([hidden_channels, hidden_channels], num_clusters)
self.conv2 = DenseGraphConv(hidden_channels, hidden_channels)
num_clusters = math.ceil(0.5 * num_clusters)
self.pool2 = DMoNPooling([hidden_channels, hidden_channels], 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, spectral_loss1, _, cluster_loss1 = self.pool1(x, adj, mask)
x = ops.relu(self.conv2(x, adj))
_, x, adj, spectral_loss2, _, cluster_loss2 = self.pool2(x, adj)
x = self.conv3(x, adj)
self.add_loss(spectral_loss1 + spectral_loss2 + cluster_loss1 + cluster_loss2)
return self.lin2(self.lin1(ops.mean(x, axis=1)))
model = DMoNNet(dataset.num_features, dataset.num_classes, avg_num_nodes)
Train
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.001),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
metrics=["accuracy"],
)
model.fit(train_loader, validation_data=val_loader, epochs=100, verbose=2)