Relational deep learning with a heterogeneous GraphSAGE
Author: K3-Node Team
Backend: Multi-Backend
Dataset: DBLP
Description: Relational deep learning (Fey et al., 2024) treats a relational database as a heterogeneous graph: every table row is a node and every foreign-key link an edge.
Relational deep learning with a heterogeneous GraphSAGE
Relational deep learning (Fey et al., 2024) treats a relational database as a heterogeneous graph: every table row is a node and every foreign-key link an edge. PyG's example learns on RelBench databases (with PyTorch Frame table encoders); RelBench is not available in K3-Node, so this notebook applies the same model to the DBLP graph (authors, papers, terms and conferences) and predicts the research area of every author: node features are encoded per node type, a heterogeneous GraphSAGE passes messages along every kind of link, and an MLP head predicts.
Same model as PyG's examples/rdl.py; on the DBLP graph instead of RelBench databases (without the table and time encoders it needs).
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
Load the data
Conferences have no features, so they get a constant one.
import keras
import numpy as np
from keras import ops
from k3_node.datasets import DBLP
from k3_node.layers import HeteroConv, SAGEConv
from k3_node.loader import FullGraphDataset
dataset = DBLP("data/DBLP")
data = dataset[0]
data["conference"].x = np.ones((data["conference"].num_nodes, 1), dtype="float32")
print(data)
Define the model
HeteroConv applies one SAGEConv per edge type and sums the results per node type; each layer is followed by a layer normalization per node type and a ReLU.
class HeteroGraphSAGE(keras.layers.Layer):
def __init__(self, node_types, edge_types, channels, aggr="sum", num_layers=2):
super().__init__()
self.convs = [HeteroConv({et: SAGEConv((channels, channels), channels, aggr=aggr) for et in edge_types},
aggr="sum") for _ in range(num_layers)]
self.norms = [{nt: keras.layers.LayerNormalization() for nt in node_types} for _ in range(num_layers)]
def call(self, x_dict, edge_index_dict):
for conv, norms in zip(self.convs, self.norms):
x_dict = conv(x_dict, edge_index_dict)
x_dict = {nt: ops.relu(norms[nt](x)) for nt, x in x_dict.items()}
return x_dict
class Model(keras.Model):
def __init__(self, node_types, edge_types, channels, out_channels, entity):
super().__init__()
self.entity = entity
self.encoder = {nt: keras.layers.Dense(channels) for nt in node_types} # one feature encoder per type
self.gnn = HeteroGraphSAGE(node_types, edge_types, channels)
self.head = keras.layers.Dense(out_channels)
def call(self, data):
x_dict = {nt: self.encoder[nt](x) for nt, x in data.x_dict.items()}
x_dict = self.gnn(x_dict, data.edge_index_dict)
return self.head(x_dict[self.entity])
model = Model(data.node_types, data.edge_types, channels=128, out_channels=4, entity="author")
Train
FullGraphDataset with node_type feeds the whole heterogeneous graph; the model receives data.x_dict and data.edge_index_dict (one entry per node and edge type), and the labels and mask come from the given node type.
Keras warns that some weights of the last layer get no gradients: they update node types other than authors, which the loss never looks at. PyG's model has the same unused weights.
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.005),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
weighted_metrics=["accuracy"],
)
model.fit(
FullGraphDataset(data, node_type="author", mask="train_mask"),
validation_data=FullGraphDataset(data, node_type="author", mask="val_mask"),
epochs=10,
verbose=2,
)