Skip to content

Top-k 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


Top-k 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 each graph convolution, TopKPooling (Gao & Ji, 2019; Cangea et al., 2018) keeps only the 80% highest-scoring nodes, so the graph gets smaller layer by layer. A readout after each level is summed and classified.

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

10% of the proteins are used for testing, the rest for training.

import keras
from keras import ops
from k3_node.datasets import TUDataset
from k3_node.layers import GraphConv, TopKPooling, global_max_pool, global_mean_pool
from k3_node.loader import DataLoader

dataset = TUDataset("data/TU", name="PROTEINS").shuffle()
n = len(dataset) // 10
test_loader = DataLoader(dataset[:n], batch_size=60)
train_loader = DataLoader(dataset[n:], batch_size=60, shuffle=True)
print(dataset)

Define the model

Three rounds of convolution and pooling. data.batch tells which protein each node belongs to.

class TopKNet(keras.Model):
    def __init__(self, in_channels, out_channels):
        super().__init__()
        self.convs = [GraphConv(in_channels, 128), GraphConv(128, 128), GraphConv(128, 128)]
        self.pools = [TopKPooling(128, ratio=0.8) for _ in range(3)]
        self.lin1 = keras.layers.Dense(128, activation="relu")
        self.lin2 = keras.layers.Dense(64, activation="relu")
        self.lin3 = keras.layers.Dense(out_channels)
        self.dropout = keras.layers.Dropout(0.5)

    def call(self, data, training=False):
        x, edge_index, batch = data.x, data.edge_index, data.batch
        readout = 0
        for conv, pool in zip(self.convs, self.pools):
            x = ops.relu(conv(x, edge_index))
            x, edge_index, _, batch, _, _ = pool(x, edge_index, batch=batch)
            readout += ops.concatenate([global_max_pool(x, batch, data.num_graphs),
                                        global_mean_pool(x, batch, data.num_graphs)], axis=-1)
        x = self.dropout(self.lin1(readout), training=training)
        return self.lin3(self.lin2(x))


model = TopKNet(dataset.num_features, dataset.num_classes)

Train

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.0005),
    loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    metrics=["accuracy"],
)
model.fit(train_loader, validation_data=test_loader, epochs=200, verbose=2)

Evaluate

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