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