Skip to content

Using known labels as features (UniMP)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: PubMed (Planetoid)
Description: UniMP (Shi et al., 2020) feeds part of the known training labels into the model, added to the node features, and learns to predict the labels of the other nodes with graph transformer layers.

View in Colab   GitHub source


Using known labels as features (UniMP)

UniMP (Shi et al., 2020) feeds part of the known training labels into the model, added to the node features, and learns to predict the labels of the other nodes with graph transformer layers. In every training step a random 65% of the training labels are visible and the model is trained on the remaining 35%.

Same model as PyG's examples/unimp_arxiv.py; on PubMed instead of ogbn-arxiv, to keep the example small, trained for 200 epochs instead of PyG's 500 to keep the notebook quick.

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

PubMed (19,717 papers, 3 topics) stands in for the much larger graph of PyG's example. With split="full", every paper outside the validation and test sets is a training paper.

import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.layers import TransformerConv
from k3_node.models import MaskLabel
from k3_node.training import gradient_step

dataset = Planetoid("data/Planetoid", name="PubMed", split="full")
data = dataset[0]
print(data)

Define the model

MaskLabel adds a learned embedding of the label to the features of the nodes whose label is visible.

class UniMP(keras.Model):
    def __init__(self, in_channels, num_classes, hidden_channels, num_layers, heads, dropout=0.3):
        super().__init__()
        self.label_emb = MaskLabel(num_classes, in_channels)
        self.convs, self.norms = [], []
        for i in range(1, num_layers + 1):
            last = i == num_layers
            self.convs.append(TransformerConv(in_channels, num_classes if last else hidden_channels // heads, heads,
                                              concat=not last, beta=True, dropout=dropout))
            in_channels = hidden_channels
            if not last:
                self.norms.append(keras.layers.LayerNormalization())

    def call(self, x, y, edge_index, label_mask, training=False):
        x = self.label_emb(x, y, label_mask)
        for conv, norm in zip(self.convs, self.norms):
            x = ops.relu(norm(conv(x, edge_index, training=training)))
        return self.convs[-1](x, edge_index, training=training)


model = UniMP(dataset.num_features, dataset.num_classes, hidden_channels=64, num_layers=3, heads=2)

Train

Every step hides a different random part of the training labels.

optimizer = keras.optimizers.Adam(learning_rate=0.001, weight_decay=0.0005)
loss_fn = keras.losses.SparseCategoricalCrossentropy(from_logits=True)


def accuracy(out, mask):
    correct = ops.cast(ops.argmax(out, axis=-1) == ops.cast(data.y, "int32"), "float32")
    return float(ops.sum(correct * ops.cast(mask, "float32")) / ops.sum(ops.cast(mask, "float32")))


def train_step(label_rate=0.65):  # the share of training labels used as input
    propagation_mask = MaskLabel.ratio_mask(data.train_mask, ratio=label_rate)
    supervision_mask = ops.logical_xor(data.train_mask, propagation_mask)

    def loss():
        out = model(data.x, data.y, data.edge_index, propagation_mask, training=True)
        return loss_fn(data.y, out, sample_weight=ops.cast(supervision_mask, "float32"))

    return gradient_step(loss, model.trainable_variables, optimizer)


model(data.x, data.y, data.edge_index, data.train_mask)  # create the weights
epochs = 200
for epoch in range(1, epochs + 1):
    loss = train_step()
    if epoch % 10 == 0:
        val_accuracy = accuracy(model(data.x, data.y, data.edge_index, data.train_mask), data.val_mask)
        print(f"Epoch {epoch:03d}: loss {loss:.4f}, validation accuracy {val_accuracy:.4f}")

Evaluate

For testing, the labels of both training and validation nodes are visible.

visible = ops.logical_or(data.train_mask, data.val_mask)
print(f"Test accuracy: {accuracy(model(data.x, data.y, data.edge_index, visible), data.test_mask):.4f}")