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