Skip to content

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.

View in Colab   GitHub source


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"

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

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

Evaluate

loss, accuracy = model.evaluate(FullGraphDataset(data, node_type="author", mask="test_mask"), verbose=0)
print(f"Test accuracy: {accuracy:.4f}")