Skip to content

Node clustering with an adversarially regularized autoencoder (ARGVA)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Learn node embeddings of the Cora citation graph without labels, then cluster them with k-means.

View in Colab   GitHub source


Node clustering with an adversarially regularized autoencoder (ARGVA)

Learn node embeddings of the Cora citation graph without labels, then cluster them with k-means. The variational graph autoencoder (Pan et al., 2018) learns to reconstruct the graph's edges, while a discriminator pushes its embeddings to look like samples of a normal distribution.

Same model as PyG's examples/argva_node_clustering.py.

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

RandomLinkSplit hides 5% of the edges for validation and 10% for testing. Each split has the edges used for message passing (edge_index), the true edges to predict (pos_edge_label_index) and, for validation and testing, as many non-edges (neg_edge_label_index).

import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.layers import GCNConv
from k3_node.models import ARGVA
from k3_node.transforms import RandomLinkSplit

transform = RandomLinkSplit(num_val=0.05, num_test=0.1, is_undirected=True, split_labels=True,
                            add_negative_train_samples=False)
dataset = Planetoid("data/Planetoid", name="Cora", transform=transform)
train_data, val_data, test_data = dataset[0]
print(train_data)

Define the model

The encoder outputs the mean and log standard deviation of every node's embedding. The discriminator tries to tell the embeddings from random normal samples.

class Encoder(keras.Model):
    def __init__(self, in_channels, hidden_channels, out_channels):
        super().__init__()
        self.conv1 = GCNConv(in_channels, hidden_channels)
        self.conv_mu = GCNConv(hidden_channels, out_channels)
        self.conv_logstd = GCNConv(hidden_channels, out_channels)

    def call(self, x, edge_index):
        x = ops.relu(self.conv1(x, edge_index))
        return self.conv_mu(x, edge_index), self.conv_logstd(x, edge_index)


discriminator = keras.Sequential([
    keras.layers.Dense(64, activation="relu"),
    keras.layers.Dense(64, activation="relu"),
    keras.layers.Dense(32),
])
model = ARGVA(Encoder(dataset.num_features, hidden_channels=32, out_channels=32), discriminator)

Train

Each epoch updates the discriminator 5 times, then the encoder once. The validation scores measure how well the embeddings predict the hidden edges.

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.005),
    discriminator_optimizer=keras.optimizers.Adam(learning_rate=0.001),
)
history = model.fit(train_data, epochs=150, validation_data=val_data)

Evaluate

Cluster the embeddings into 7 groups with k-means and compare them with the true classes of the papers, which the model never saw.

from sklearn.cluster import KMeans
from sklearn.metrics import completeness_score, homogeneity_score, v_measure_score

z = ops.convert_to_numpy(model.embed(test_data))
clusters = KMeans(n_clusters=7, random_state=0, n_init="auto").fit_predict(z)
labels = ops.convert_to_numpy(test_data.y)

print(model.evaluate(test_data))  # link prediction on the test edges
print(f"Completeness: {completeness_score(labels, clusters):.3f}, "
      f"Homogeneity: {homogeneity_score(labels, clusters):.3f}, NMI: {v_measure_score(labels, clusters):.3f}")

Visualize

t-SNE projects the embeddings to 2D; colors are the true classes.

import matplotlib.pyplot as plt
from sklearn.manifold import TSNE

points = TSNE(n_components=2).fit_transform(z)
plt.figure(figsize=(8, 8))
plt.scatter(points[:, 0], points[:, 1], c=labels, s=10, cmap="tab10")
plt.axis("off")
plt.show()