Skip to content

Graph Isomorphism Network (GIN) on MUTAG

Author: K3-Node Team
Backend: Multi-Backend
Dataset: MUTAG (TUDataset)
Description: Predict whether a molecule is mutagenic.

View in Colab   GitHub source


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"

!pip install k3-node[examples]
import os
os.environ["KERAS_BACKEND"] = "tensorflow"

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)

Evaluate

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