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.
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"
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.