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