Skip to content

Predicting solubility with Principal Neighbourhood Aggregation

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

!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.

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)

Evaluate

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