Skip to content

Predicting quantum properties of molecules with an MPNN (QM9)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: QM9
Description: Predict a quantum-mechanical property (here the dipole moment, target 0) of small molecules from their structure.

View in Colab   GitHub source


Predicting quantum properties of molecules with an MPNN (QM9)

Predict a quantum-mechanical property (here the dipole moment, target 0) of small molecules from their structure. The model is the message passing neural network of Gilmer et al., 2017: an NNConv layer whose weights depend on the bond type and the distance between atoms, applied three times with a GRU in between, and a Set2Set readout.

Same model as PyG's examples/qm9_nn_conv.py; on 6,000 of the QM9 molecules, trained for 50 epochs instead of 300 (every epoch takes a few minutes on a CPU: each edge of the complete graphs gets its own 64x64 weight matrix).

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

Every molecule becomes a complete graph (every atom connected to every other one); the edge features are the bond type (zero where there is no bond) plus the distance between the atoms. The target is standardized. To keep the example small, it uses 6,000 of the 130,831 molecules.

import keras
from keras import ops
from k3_node.datasets import QM9
from k3_node.layers import NNConv, Set2Set
from k3_node.loader import DataLoader
from k3_node.transforms import CompleteGraph, Compose, Distance

target = 0  # the dipole moment
dataset = QM9("data/QM9").shuffle()[:6000]
mean, std = dataset.mean(target), dataset.std(target)


def select_target(data):  # standardized value of the chosen target
    data.y = (data.y[:, target] - mean) / std
    return data


dataset.transform = Compose([select_target, CompleteGraph(), Distance(norm=False)])
test_dataset, val_dataset, train_dataset = dataset[:500], dataset[500:1000], dataset[1000:]
test_loader = DataLoader(test_dataset, batch_size=128)
val_loader = DataLoader(val_dataset, batch_size=128)
train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True)
print(train_dataset[0])

Define the model

class MPNN(keras.Model):
    def __init__(self, in_channels, dim=64):
        super().__init__()
        self.lin0 = keras.layers.Dense(dim, activation="relu")
        edge_network = keras.Sequential([keras.layers.Dense(128, activation="relu"), keras.layers.Dense(dim * dim)])
        self.conv = NNConv(dim, dim, edge_network, aggr="mean")
        self.gru = keras.layers.GRUCell(dim)
        self.set2set = Set2Set(dim, processing_steps=3)
        self.lin1 = keras.layers.Dense(dim, activation="relu")
        self.lin2 = keras.layers.Dense(1)

    def call(self, data):
        out = h = self.lin0(data.x)
        for _ in range(3):
            m = ops.relu(self.conv(out, data.edge_index, data.edge_attr))
            out, [h] = self.gru(m, [h])
        out = self.set2set(out, data.batch, dim_size=data.num_graphs)
        return ops.squeeze(self.lin2(self.lin1(out)), axis=-1)


model = MPNN(dataset.num_features)

Train

The loss is the mean squared error of the standardized target; the learning rate is reduced when the validation error stops improving.

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

Evaluate

The error is converted back to the target's unit (Debye).

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