Graph Isomorphism Network (GIN) on MUTAG
Author: K3-Node Team
Backend: Multi-Backend
Dataset: MUTAG (TUDataset)
Description: Predict whether a molecule is mutagenic.
Graph Isomorphism Network (GIN) on MUTAG
Predict whether a molecule is mutagenic. Each of the 188 molecules in MUTAG is a small graph (atoms are nodes, bonds are edges), and the task is to classify whole graphs. A GIN (Xu et al., 2019) computes atom embeddings, sums them into one vector per molecule, and classifies that vector.
Same model and settings as PyG's examples/mutag_gin.py.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
Use 90% of the molecules for training and 10% for testing. A DataLoader groups graphs into mini-batches and can be passed straight to model.fit.
import keras
from keras import ops
from k3_node.datasets import TUDataset
from k3_node.layers import GINConv, global_add_pool
from k3_node.loader import DataLoader
from k3_node.models import MLP
dataset = TUDataset("data/TU", name="MUTAG").shuffle()
train_loader = DataLoader(dataset[:0.9], batch_size=128, shuffle=True)
test_loader = DataLoader(dataset[0.9:], batch_size=128)
print(dataset)
Define the model
Five GIN layers, then global_add_pool sums the atom embeddings of each molecule
(data.batch says which molecule each atom belongs to), and an MLP classifies the result.
class GIN(keras.Model):
def __init__(self, in_channels, hidden_channels, out_channels, num_layers):
super().__init__()
self.convs = []
for _ in range(num_layers):
mlp = MLP([in_channels, hidden_channels, hidden_channels])
self.convs.append(GINConv(mlp, train_eps=False))
in_channels = hidden_channels
self.mlp = MLP([hidden_channels, hidden_channels, out_channels], norm=None, dropout=0.5)
def call(self, data, training=False):
x = data.x
for conv in self.convs:
x = ops.relu(conv(x, data.edge_index, training=training))
x = global_add_pool(x, data.batch, data.num_graphs)
return self.mlp(x, training=training)
model = GIN(dataset.num_features, hidden_channels=32, out_channels=dataset.num_classes, num_layers=5)
Train
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.01),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
metrics=["accuracy"],
)
model.fit(train_loader, validation_data=test_loader, epochs=100, verbose=2)