Skip to content

A GNN made heterogeneous, on a movie graph

Author: K3-Node Team
Backend: Multi-Backend
Dataset: IMDB
Description: PyG's RelBench example builds a two-layer GraphSAGE and converts it with to_hetero into a model that has separate weights for every node and edge type of a relational database.

View in Colab   GitHub source


A GNN made heterogeneous, on a movie graph

PyG's RelBench example builds a two-layer GraphSAGE and converts it with to_hetero into a model that has separate weights for every node and edge type of a relational database. This notebook builds the same heterogeneous model directly with HeteroConv and trains it on the IMDB graph (movies, directors and actors) to predict the genre of every movie. RelBench is not available in K3-Node.

Same model as PyG's examples/relbench_example.py; on the IMDB graph instead of a RelBench database, and for classification instead of regression.

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 keras
from keras import ops
from k3_node.datasets import IMDB
from k3_node.layers import HeteroConv, SAGEConv
from k3_node.loader import FullGraphDataset

dataset = IMDB("data/IMDB")
data = dataset[0]
print(data)

Define the model

What to_hetero produces from SAGEConv -> ReLU -> SAGEConv -> ReLU -> Linear: every SAGEConv becomes one convolution per edge type (summed per node type), and the final linear layer becomes one per node type. Only the parts that lead to the movie predictions are built.

class HeteroGNN(keras.Model):
    def __init__(self, node_types, edge_types, hidden_channels, out_channels, entity):
        super().__init__()
        self.entity = entity
        self.conv1 = HeteroConv({et: SAGEConv((-1, -1), hidden_channels) for et in edge_types}, aggr="sum")
        # Only the outputs for the predicted node type are needed after the last layer
        self.conv2 = HeteroConv({et: SAGEConv((-1, -1), hidden_channels) for et in edge_types if et[-1] == entity},
                                aggr="sum")
        self.lin = keras.layers.Dense(out_channels)

    def call(self, data):
        x_dict = {nt: ops.relu(x) for nt, x in self.conv1(data.x_dict, data.edge_index_dict).items()}
        x_dict = {nt: ops.relu(x) for nt, x in self.conv2(x_dict, data.edge_index_dict).items()}
        return self.lin(x_dict[self.entity])


model = HeteroGNN(data.node_types, data.edge_types, hidden_channels=64, out_channels=3, entity="movie")

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.

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="movie", mask="train_mask"),
    validation_data=FullGraphDataset(data, node_type="movie", mask="val_mask"),
    epochs=30,
    verbose=2,
)

Evaluate

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