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).
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"
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)