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