Predicting solubility with Principal Neighbourhood Aggregation
Author: K3-Node Team
Backend: Multi-Backend
Dataset: ESOL (MoleculeNet)
Description: Predict how well molecules dissolve in water.
Predicting solubility with Principal Neighbourhood Aggregation
Predict how well molecules dissolve in water. PNA (Corso et al., 2020) combines several aggregators (mean, min, max, standard deviation) with scalers that depend on each node's degree, so that it can tell apart neighborhoods that simple sums or means confuse.
Same model as PyG's examples/pna.py; on ESOL instead of ZINC, trained for 200 epochs instead of PyG's 300 to keep the notebook quick.
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.
import keras
from keras import ops
from k3_node.datasets import MoleculeNet
from k3_node.loader import DataLoader
from k3_node.layers import BatchNorm, PNAConv, global_add_pool
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/MoleculeNet", "ESOL", transform=atom_and_bond_types).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=128, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=128)
test_loader = DataLoader(test_dataset, batch_size=128)
print(train_dataset[0])
Define the model
The degree scalers need the histogram of node degrees in the training set.
deg = PNAConv.get_degree_histogram(train_dataset)
class PNA(keras.Model):
def __init__(self, deg):
super().__init__()
self.node_emb = keras.layers.Embedding(120, 75) # one embedding per chemical element
self.edge_emb = keras.layers.Embedding(22, 50) # one embedding per bond type
self.convs = [PNAConv(75, 75, aggregators=["mean", "min", "max", "std"],
scalers=["identity", "amplification", "attenuation"], deg=deg, edge_dim=50,
towers=5, pre_layers=1, post_layers=1, divide_input=False) for _ in range(4)]
self.batch_norms = [BatchNorm(75) for _ in range(4)]
self.mlp = keras.Sequential([keras.layers.Dense(50, activation="relu"),
keras.layers.Dense(25, activation="relu"), keras.layers.Dense(1)])
def call(self, data, training=False):
x = self.node_emb(data.x)
edge_attr = self.edge_emb(data.edge_attr)
for conv, batch_norm in zip(self.convs, self.batch_norms):
x = ops.relu(batch_norm(conv(x, data.edge_index, edge_attr), training=training))
return self.mlp(global_add_pool(x, data.batch, data.num_graphs))
model = PNA(deg)
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=200, callbacks=[reduce_lr], verbose=2)