Skip to content

Predicting solubility with a graph transformer (GraphGPS)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: ESOL (MoleculeNet)
Description: Predict how well molecules dissolve in water.

View in Colab   GitHub source


Predicting solubility with a graph transformer (GraphGPS)

Predict how well molecules dissolve in water. GraphGPS (Rampášek et al., 2022) combines, in every layer, local message passing (here GINE) with global attention between all atoms of a molecule. Random-walk positional encodings tell the attention where each atom sits in the graph.

Same model as PyG's examples/graph_gps.py; on ESOL instead of ZINC, with standard multi-head attention (PyG also offers Performer attention).

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

ESOL contains 1,128 molecules and their water solubility (log mol/L), which the model predicts. PyG's example uses ZINC, where every atom and bond has a single type; atom_and_bond_types keeps the same information here: the atomic number of each atom and the type of each bond. AddRandomWalkPE adds 20 random-walk positional encodings per atom as data.pe.

import keras
from keras import ops
from k3_node.datasets import MoleculeNet
from k3_node.loader import DataLoader
from k3_node.layers import GINEConv, GPSConv, global_add_pool
from k3_node.transforms import AddRandomWalkPE, Compose


def atom_and_bond_types(data):
    data.x = data.x[:, 0]  # atomic number
    data.edge_attr = data.edge_attr[:, 0]  # bond type
    return data


dataset = MoleculeNet("data/ESOL-PE", "ESOL", pre_transform=Compose([atom_and_bond_types, AddRandomWalkPE(walk_length=20, attr_name="pe")])).shuffle()
n = len(dataset) // 10  # 10% validation, 10% test, 80% training
val_dataset, test_dataset, train_dataset = dataset[:n], dataset[n:2 * n], dataset[2 * n:]
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=32)
test_loader = DataLoader(test_dataset, batch_size=32)
print(train_dataset[0])

Define the model

Each atom embedding is the concatenation of its element embedding and its normalized positional encoding. Each GPSConv layer has 4 attention heads.

class GPS(keras.Model):
    def __init__(self, channels, pe_dim, num_layers):
        super().__init__()
        self.node_emb = keras.layers.Embedding(120, channels - pe_dim)
        self.pe_lin = keras.layers.Dense(pe_dim)
        self.pe_norm = keras.layers.BatchNormalization()
        self.edge_emb = keras.layers.Embedding(22, channels)
        self.convs = []
        for _ in range(num_layers):
            mlp = keras.Sequential([keras.layers.Dense(channels, activation="relu"), keras.layers.Dense(channels)])
            self.convs.append(GPSConv(channels, GINEConv(mlp), heads=4))
        self.mlp = keras.Sequential([keras.layers.Dense(channels // 2, activation="relu"),
                                     keras.layers.Dense(channels // 4, activation="relu"), keras.layers.Dense(1)])

    def call(self, data, training=False):
        pe = self.pe_norm(data.pe, training=training)
        x = ops.concatenate([self.node_emb(data.x), self.pe_lin(pe)], axis=1)
        edge_attr = self.edge_emb(data.edge_attr)
        for conv in self.convs:
            x = conv(x, data.edge_index, data.batch, batch_size=data.num_graphs, edge_attr=edge_attr,
                     training=training)
        return self.mlp(global_add_pool(x, data.batch, data.num_graphs))


model = GPS(channels=64, pe_dim=8, num_layers=10)

Train

The loss is the mean absolute error (MAE). As in PyG, the learning rate is halved whenever the validation error stops improving.

model.compile(optimizer=keras.optimizers.Adam(learning_rate=0.001), loss="mae")
reduce_lr = keras.callbacks.ReduceLROnPlateau(factor=0.5, patience=20, min_lr=0.00001)
model.fit(train_loader, validation_data=val_loader, epochs=100, callbacks=[reduce_lr], verbose=2)

Evaluate

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