Skip to content

Predicting solubility with AttentiveFP (ESOL)

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 AttentiveFP (ESOL)

Predict how well molecules dissolve in water. AttentiveFP (Xiong et al., 2019) applies graph attention between neighboring atoms and then between each atom and the whole molecule, with GRU updates in between.

Same model as PyG's examples/attentive_fp.py.

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). AttentiveFPFeatures computes the atom and bond features from the paper (39 per atom, 10 per bond) from each molecule's SMILES string, using RDKit.

import keras
from k3_node.datasets import MoleculeNet
from k3_node.loader import DataLoader
from k3_node.models import AttentiveFP
from k3_node.transforms import AttentiveFPFeatures

dataset = MoleculeNet("data/AFP_Mol", "ESOL", pre_transform=AttentiveFPFeatures()).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=200, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=200)
test_loader = DataLoader(test_dataset, batch_size=200)
print(train_dataset[0])

Define the model

class Net(keras.Model):
    def __init__(self):
        super().__init__()
        self.afp = AttentiveFP(in_channels=39, hidden_channels=200, out_channels=1, edge_dim=10,
                               num_layers=2, num_timesteps=2, dropout=0.2)

    def call(self, data, training=False):
        return self.afp(data.x, data.edge_index, data.edge_attr, data.batch, batch_size=data.num_graphs,
                        training=training)


model = Net()

Train

The loss is the mean squared error; the root mean squared error (RMSE) is reported.

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=10 ** -2.5, weight_decay=10 ** -5),
    loss="mse",
    metrics=[keras.metrics.RootMeanSquaredError(name="rmse")],
)
model.fit(train_loader, validation_data=val_loader, epochs=200, verbose=2)

Evaluate

loss, rmse = model.evaluate(test_loader, verbose=0)
print(f"Test RMSE: {rmse:.4f}")