Skip to content

Pre-trained positional and structural encodings (GPSE)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: ESOL (MoleculeNet)
Description: GPSE (Cantürk et al., 2023) is a GNN pre-trained to predict many positional and structural encodings of nodes (random-walk statistics, eigenvectors, ...).

View in Colab   GitHub source


Pre-trained positional and structural encodings (GPSE)

GPSE (Cantürk et al., 2023) is a GNN pre-trained to predict many positional and structural encodings of nodes (random-walk statistics, eigenvectors, ...). Its node representations can be added to the input of any other GNN. Here a GCN predicts the solubility of molecules, with or without GPSE encodings.

Same model as PyG's examples/gpse.py; on ESOL instead of ZINC.

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; as in PyG's ZINC example, every atom is described by its type only (the atomic number).

import keras
import numpy as np
from keras import ops
from k3_node.datasets import MoleculeNet
from k3_node.layers import GCNConv, global_mean_pool
from k3_node.loader import DataLoader
from k3_node.models import GPSE, MLP, GPSENodeEncoder
from k3_node.models.gpse import precompute_gpse

use_gpse = "molpcba"  # which pre-trained GPSE to use, or None for the baseline without GPSE


def atom_types(data):
    data.x = data.x[:, :1]  # atomic number
    return data


dataset = MoleculeNet("data/MoleculeNet", "ESOL", transform=atom_types).shuffle()
n = len(dataset) // 10
train_dataset, test_dataset = dataset[n:], dataset[:n]

Compute the GPSE encodings

GPSE.from_pretrained downloads the pre-trained weights (about 250 MB) and precompute_gpse adds a 512-dimensional encoding per atom, pestat_GPSE, to every molecule.

if use_gpse:
    gpse = GPSE.from_pretrained(use_gpse, root="data/GPSE_pretrained")
    train_dataset = precompute_gpse(gpse, train_dataset)
    test_dataset = precompute_gpse(gpse, test_dataset)
train_loader = DataLoader(train_dataset, batch_size=256, shuffle=True)
test_loader = DataLoader(test_dataset, batch_size=256)
print(train_dataset[0])

Define the model

The atom type embedding (32 values) is concatenated with 32 values derived from the GPSE encoding (GPSENodeEncoder); without GPSE, a linear layer maps the embedding to 64 values. Eight GCN layers with residual connections follow, then a mean readout.

class GNNStackStage(keras.layers.Layer):  # GCN layers with residual connections and L2 normalization
    def __init__(self, channels, num_layers):
        super().__init__()
        self.convs = [GCNConv(channels, channels) for _ in range(num_layers)]

    def call(self, x, edge_index):
        for conv in self.convs:
            x = x + conv(x, edge_index)
        return x / ops.maximum(ops.norm(x, axis=-1, keepdims=True), 1e-12)


class GPSEPlusGNN(keras.Model):
    def __init__(self, dim_emb, dim_conv, num_layers, dim_pe_in, dim_pe_out, gpse):
        super().__init__()
        self.use_gpse = gpse
        self.encoder1 = keras.layers.Embedding(120, dim_emb - dim_pe_out)  # atom type embedding
        self.encoder2 = (GPSENodeEncoder(dim_emb, dim_pe_in, dim_pe_out, expand_x=False) if gpse
                         else keras.layers.Dense(dim_emb))
        self.premp = MLP([dim_emb, dim_emb, dim_conv])
        self.gnn = GNNStackStage(dim_conv, num_layers)
        self.dropout = keras.layers.Dropout(0.5)
        self.postmp = MLP([dim_conv, 1])

    def call(self, data, training=False):
        x = self.encoder1(ops.cast(data.x[:, 0], "int32"))
        x = self.encoder2(x, data.pestat_GPSE, training=training) if self.use_gpse else self.encoder2(x)
        x = self.gnn(self.premp(x, training=training), data.edge_index)
        x = self.dropout(global_mean_pool(x, data.batch, data.num_graphs), training=training)
        return self.postmp(x, training=training)


model = GPSEPlusGNN(dim_emb=64, dim_conv=128, num_layers=8, dim_pe_in=512, dim_pe_out=32, gpse=bool(use_gpse))

Train

model.compile(optimizer=keras.optimizers.Adam(learning_rate=0.001, weight_decay=5e-4), loss="mse")
model.fit(train_loader, validation_data=test_loader, epochs=100, verbose=2)

Evaluate

mse = model.evaluate(test_loader, verbose=0)
print(f"Test MSE: {mse:.4f}")