Graph Multiset Transformer (GMT) on PROTEINS
Author: K3-Node Team
Backend: Multi-Backend
Dataset: PROTEINS (TUDataset)
Description: Classify proteins as enzymes or non-enzymes.
Graph Multiset Transformer (GMT) 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. Three GCN layers compute node embeddings; the Graph Multiset Transformer (Baek et al., 2021) then pools each graph's nodes with attention, learning which nodes matter for the graph-level prediction.
Same model as PyG's examples/proteins_gmt.py.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
import keras
from keras import ops
from k3_node.datasets import TUDataset
from k3_node.layers import GCNConv, GraphMultisetTransformer
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=128)
val_loader = DataLoader(dataset[n:2 * n], batch_size=128)
train_loader = DataLoader(dataset[2 * n:], batch_size=128, shuffle=True)
print(dataset)
Define the model
class GMTNet(keras.Model):
def __init__(self, in_channels, out_channels):
super().__init__()
self.conv1 = GCNConv(in_channels, 32)
self.conv2 = GCNConv(32, 32)
self.conv3 = GCNConv(32, 32)
self.pool = GraphMultisetTransformer(96, k=10, heads=4)
self.lin1 = keras.layers.Dense(16, activation="relu")
self.dropout = keras.layers.Dropout(0.5)
self.lin2 = keras.layers.Dense(out_channels)
def call(self, data, training=False):
x1 = ops.relu(self.conv1(data.x, data.edge_index))
x2 = ops.relu(self.conv2(x1, data.edge_index))
x3 = ops.relu(self.conv3(x2, data.edge_index))
x = ops.concatenate([x1, x2, x3], axis=-1)
x = self.pool(x, index=data.batch, dim_size=data.num_graphs) # one vector per protein
x = self.dropout(self.lin1(x), training=training)
return self.lin2(x)
model = GMTNet(dataset.num_features, dataset.num_classes)
Train
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.001, weight_decay=1e-4),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
metrics=["accuracy"],
)
model.fit(train_loader, validation_data=val_loader, epochs=200, verbose=2)