Pre-trained SchNet models on QM9
Author: K3-Node Team
Backend: Multi-Backend
Dataset: QM9
Description: Evaluate the official pre-trained SchNet models (Schütt et al., 2017) on QM9: one model per target property.
Pre-trained SchNet models on QM9
Evaluate the official pre-trained SchNet models (Schütt et al., 2017)
on QM9: one model per target property. SchNet learns continuous filters over the distances between
atoms. SchNet.from_qm9_pretrained downloads the SchNetPack weights, converts them to Keras (reading the
files needs PyTorch) and returns the data split the model was trained with.
Same model as PyG's examples/qm9_pretrained_schnet.py.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
import numpy as np
from keras import ops
from k3_node.datasets import QM9
from k3_node.loader import DataLoader
from k3_node.models import SchNet
from k3_node.training import no_grad
dataset = QM9("data/QM9")
Evaluate every target
The mean absolute error on each model's test molecules; energies in meV. To keep the example fast on a CPU, only the first 1,000 test molecules are evaluated.
num_test = 1000 # test molecules per target; None evaluates all of them (best on a GPU)
for target in range(12):
model, (train_dataset, val_dataset, test_dataset) = SchNet.from_qm9_pretrained("data/QM9/pretrained", dataset, target)
errors = []
for data in DataLoader(test_dataset[:num_test], batch_size=64):
with no_grad(): # evaluation only: nothing to remember for gradients
prediction = model(data.z, data.pos, data.batch, batch_size=data.num_graphs)[:, 0]
errors.append(np.abs(ops.convert_to_numpy(prediction) - ops.convert_to_numpy(data.y)[:, target]))
errors = np.concatenate(errors)
if target in [2, 3, 4, 6, 7, 8, 9, 10]: # energies: report meV instead of eV
errors = 1000 * errors
print(f"Target {target:02d}: MAE {errors.mean():.5f} ± {errors.std():.5f}")