Skip to content

Efficient Graph Convolution (EGC) on BACE molecules

Author: K3-Node Team
Backend: Multi-Backend
Dataset: BACE (MoleculeNet)
Description: Predict whether a molecule inhibits the BACE-1 enzyme (a binary label per molecule).

View in Colab   GitHub source


Efficient Graph Convolution (EGC) on BACE molecules

Predict whether a molecule inhibits the BACE-1 enzyme (a binary label per molecule). EGC (Tailor et al., 2022) is a memory-efficient convolution: every node combines a few shared "basis" filters with node-specific weights. Molecules are graphs of atoms, so this is graph classification: atom embeddings are averaged into one vector per molecule.

Same model as PyG's examples/egc.py; on MoleculeNet BACE instead of ogbg-molhiv to keep the dataset small.

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

The molecules are shuffled and split 80/10/10 into training, validation and test sets.

import keras
from keras import ops
from k3_node.datasets import MoleculeNet
from k3_node.layers import EGConv, global_mean_pool
from k3_node.loader import DataLoader
from k3_node.models import AtomEncoder

dataset = MoleculeNet("data/MoleculeNet", name="BACE").shuffle()
train_loader = DataLoader(dataset[:0.8], batch_size=32, shuffle=True)
val_loader = DataLoader(dataset[0.8:0.9], batch_size=256)
test_loader = DataLoader(dataset[0.9:], batch_size=256)
print(dataset)

Define the model

AtomEncoder turns the integer atom features into vectors; four residual EGC layers follow.

class EGC(keras.Model):
    def __init__(self, hidden_channels, num_layers, num_heads, num_bases):
        super().__init__()
        self.encoder = AtomEncoder(hidden_channels)
        self.convs = [EGConv(hidden_channels, hidden_channels, ["symnorm"], num_heads, num_bases)
                      for _ in range(num_layers)]
        self.norms = [keras.layers.BatchNormalization() for _ in range(num_layers)]
        self.mlp = keras.Sequential([
            keras.layers.Dense(hidden_channels // 2, use_bias=False),
            keras.layers.BatchNormalization(),
            keras.layers.ReLU(),
            keras.layers.Dense(hidden_channels // 4, use_bias=False),
            keras.layers.BatchNormalization(),
            keras.layers.ReLU(),
            keras.layers.Dense(1),
        ])

    def call(self, data, training=False):
        x = self.encoder(data.x)
        for conv, norm in zip(self.convs, self.norms):
            x = x + ops.relu(norm(conv(x, data.edge_index), training=training))
        x = global_mean_pool(x, data.batch, data.num_graphs)
        return self.mlp(x, training=training)


model = EGC(hidden_channels=236, num_layers=4, num_heads=4, num_bases=4)

Train

The learning rate is halved when the validation ROC-AUC stops improving.

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=1e-4),
    loss=keras.losses.BinaryCrossentropy(from_logits=True),
    metrics=[keras.metrics.AUC(from_logits=True, name="auc")],
)
reduce_lr = keras.callbacks.ReduceLROnPlateau(monitor="val_auc", mode="max", factor=0.5, patience=20, min_lr=1e-5)
model.fit(train_loader, validation_data=val_loader, epochs=30, callbacks=[reduce_lr], verbose=2)

Evaluate

loss, auc = model.evaluate(test_loader, verbose=0)
print(f"Test ROC-AUC: {auc:.4f}")