Skip to content

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.

View in Colab   GitHub source


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"

!pip install k3-node[examples]
import os
os.environ["KERAS_BACKEND"] = "tensorflow"

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