Skip to content

DiffPool 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


DiffPool 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. DiffPool (Ying et al., 2018) learns to softly group nodes into clusters: one GNN computes node embeddings, another computes a cluster assignment, and the graph is coarsened twice. It works on dense (padded) graphs, so proteins are converted with ToDense and batched with DenseDataLoader.

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

Proteins with more than 150 nodes are skipped, and the rest are padded to 150 nodes.

import math
import keras
from keras import ops
from k3_node.datasets import TUDataset
from k3_node.layers import DenseSAGEConv, dense_diff_pool
from k3_node.loader import DenseDataLoader
from k3_node.transforms import ToDense

max_nodes = 150
# A separate folder, since the filtered dataset differs from the full one stored in "data/TU"
dataset = TUDataset("data/TU_dense", name="PROTEINS", transform=ToDense(max_nodes),
                    pre_filter=lambda data: data.num_nodes <= max_nodes).shuffle()
n = (len(dataset) + 9) // 10  # 10% test, 10% validation, 80% training
test_loader = DenseDataLoader(dataset[:n], batch_size=20)
val_loader = DenseDataLoader(dataset[n:2 * n], batch_size=20)
train_loader = DenseDataLoader(dataset[2 * n:], batch_size=20, shuffle=True)
print(dataset)

Define the model

The auxiliary link-prediction and entropy losses of DiffPool are added with add_loss, so fit minimizes them together with the classification loss.

class GNN(keras.layers.Layer):
    def __init__(self, in_channels, hidden_channels, out_channels, lin=True):
        super().__init__()
        self.convs = [DenseSAGEConv(in_channels, hidden_channels), DenseSAGEConv(hidden_channels, hidden_channels),
                      DenseSAGEConv(hidden_channels, out_channels)]
        self.norms = [keras.layers.BatchNormalization() for _ in range(3)]
        self.lin = keras.layers.Dense(out_channels, activation="relu") if lin else None

    def call(self, x, adj, mask=None, training=False):
        xs = []
        for conv, norm in zip(self.convs, self.norms):
            x = norm(ops.relu(conv(x, adj, mask)), training=training)
            xs.append(x)
        x = ops.concatenate(xs, axis=-1)
        return self.lin(x) if self.lin is not None else x


class DiffPool(keras.Model):
    def __init__(self, in_channels, out_channels, max_nodes):
        super().__init__()
        num_clusters = math.ceil(0.25 * max_nodes)
        self.gnn1_pool = GNN(in_channels, 64, num_clusters)
        self.gnn1_embed = GNN(in_channels, 64, 64, lin=False)
        num_clusters = math.ceil(0.25 * num_clusters)
        self.gnn2_pool = GNN(3 * 64, 64, num_clusters)
        self.gnn2_embed = GNN(3 * 64, 64, 64, lin=False)
        self.gnn3_embed = GNN(3 * 64, 64, 64, lin=False)
        self.lin1 = keras.layers.Dense(64, activation="relu")
        self.lin2 = keras.layers.Dense(out_channels)

    def call(self, data, training=False):
        x, adj, mask = data.x, data.adj, data.mask
        s = self.gnn1_pool(x, adj, mask, training=training)
        x = self.gnn1_embed(x, adj, mask, training=training)
        x, adj, link_loss1, entropy_loss1 = dense_diff_pool(x, adj, s, mask)
        s = self.gnn2_pool(x, adj, training=training)
        x = self.gnn2_embed(x, adj, training=training)
        x, adj, link_loss2, entropy_loss2 = dense_diff_pool(x, adj, s)
        x = self.gnn3_embed(x, adj, training=training)
        self.add_loss(link_loss1 + link_loss2 + entropy_loss1 + entropy_loss2)
        return self.lin2(self.lin1(ops.mean(x, axis=1)))


model = DiffPool(dataset.num_features, dataset.num_classes, max_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=150, verbose=2)

Evaluate

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