Skip to content

DMoN 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


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"

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

Evaluate

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