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