Skip to content

Graph classification with the Weisfeiler-Lehman kernel (ENZYMES)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: ENZYMES (TUDataset)
Description: Classify enzymes into 6 classes without training a neural network.

View in Colab   GitHub source


Graph classification with the Weisfeiler-Lehman kernel (ENZYMES)

Classify enzymes into 6 classes without training a neural network. WLConv runs the Weisfeiler-Lehman algorithm: each round, every node gets a new "color" that summarizes its own color and the colors of its neighbors. The histogram of colors in a graph is a fixed-length description of it, and a linear support vector machine (from scikit-learn) classifies these histograms.

Same model as PyG's examples/wl_kernel.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

Batch.from_data_list puts all graphs into one big batch.

import numpy as np
from sklearn.metrics import accuracy_score
from sklearn.svm import LinearSVC
from k3_node.data import Batch
from k3_node.datasets import TUDataset
from k3_node.layers import WLConv

dataset = TUDataset("data/TU", name="ENZYMES")
data = Batch.from_data_list(list(dataset))
print(data)

Compute the color histograms

Five rounds of WLConv give five histograms per graph, each describing a larger neighborhood around the nodes.

convs = [WLConv() for _ in range(5)]

x, histograms = data.x, []
for conv in convs:
    x = conv(x, data.edge_index)
    histograms.append(np.asarray(conv.histogram(x, data.batch, norm=True)))
print([h.shape for h in histograms])

Train and evaluate

For 10 random splits (80% training, 10% validation, 10% test), pick the histogram and SVM regularization C that work best on the validation graphs, then report the test accuracy.

import warnings
from sklearn.exceptions import ConvergenceWarning

warnings.filterwarnings("ignore", category=ConvergenceWarning)
y = np.asarray(data.y)
n = len(y)
test_accuracies = []
for run in range(10):
    perm = np.random.permutation(n)
    val, test, train = perm[:n // 10], perm[n // 10:n // 5], perm[n // 5:]
    best_val = 0
    for hist in histograms:
        for C in [1e3, 1e2, 1e1, 1e0, 1e-1, 1e-2, 1e-3]:
            svm = LinearSVC(C=C, tol=0.01).fit(hist[train], y[train])
            val_accuracy = accuracy_score(y[val], svm.predict(hist[val]))
            if val_accuracy > best_val:
                best_val = val_accuracy
                test_accuracy = accuracy_score(y[test], svm.predict(hist[test]))
    test_accuracies.append(test_accuracy)
    print(f"Run {run + 1:02d}: validation {best_val:.4f}, test {test_accuracy:.4f}")
print(f"Test accuracy: {np.mean(test_accuracies):.4f} ± {np.std(test_accuracies):.4f}")