Skip to content

Counting colors with supervised Top-k Pooling (COLORS-3)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: COLORS-3 (TUDataset)
Description: Count how many nodes of a given color a graph contains.

View in Colab   GitHub source


Counting colors with supervised Top-k Pooling (COLORS-3)

Count how many nodes of a given color a graph contains. In COLORS-3 each node has a one-hot color (plus a first feature marking the nodes that matter). Following Knyazev et al., 2019, a GIN is combined with TopKPooling whose node scores are supervised: they should highlight exactly the nodes that should be counted. The model trains on 500 graphs and is tested on larger ones.

Same model as PyG's examples/colors_topk_pool.py; trained for 200 epochs instead of PyG's 300 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

A small transform moves the first node feature into data.attn (the target node importances, normalized per graph) and keeps the colors as data.x.

import keras
from keras import ops
from k3_node.datasets import TUDataset
from k3_node.layers import GINConv, TopKPooling, global_add_pool, global_mean_pool
from k3_node.loader import DataLoader


def handle_node_attention(data):
    data.attn = ops.softmax(data.x[:, 0], axis=0)
    data.x = data.x[:, 1:]
    return data


dataset = TUDataset("data/TU", name="COLORS-3", use_node_attr=True, transform=handle_node_attention)
train_loader = DataLoader(dataset[:500], batch_size=60, shuffle=True)
val_loader = DataLoader(dataset[500:3000], batch_size=60)
test_loader = DataLoader(dataset[3000:], batch_size=60)
print(dataset)

Define the model

The pooling scores are also trained to match the given node importances (data.attn) with a KL-divergence loss, added with add_loss (weight 100, as in PyG).

def count_accuracy(y_true, y_pred):  # a prediction is correct if it rounds to the true count
    return ops.cast(ops.equal(ops.round(ops.squeeze(y_pred, -1)), ops.cast(y_true, y_pred.dtype)), "float32")


def mlp():
    return keras.Sequential([keras.layers.Dense(64, activation="relu"), keras.layers.Dense(64)])


class ColorCounter(keras.Model):
    def __init__(self, in_channels):
        super().__init__()
        self.conv1 = GINConv(mlp())
        self.pool1 = TopKPooling(in_channels, min_score=0.05)
        self.conv2 = GINConv(mlp())
        self.lin = keras.layers.Dense(1)

    def call(self, data):
        out = ops.relu(self.conv1(data.x, data.edge_index))
        out, edge_index, _, batch, perm, score = self.pool1(out, data.edge_index, batch=data.batch, attn=data.x)
        out = ops.relu(self.conv2(out, edge_index))
        out = self.lin(global_add_pool(out, batch, data.num_graphs))
        # KL divergence between the pooling scores and the target importances of the kept nodes
        target = ops.take(data.attn, perm, axis=0)
        attn_loss = keras.losses.kl_divergence(target[:, None], score[:, None])
        self.add_loss(100 * ops.mean(global_mean_pool(attn_loss[:, None], batch, data.num_graphs)))
        return out


model = ColorCounter(dataset.num_features)

Train

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.001),
    loss="mse",
    metrics=[count_accuracy],
)
model.fit(train_loader, validation_data=val_loader, epochs=200, verbose=2)

Evaluate

loss, accuracy = model.evaluate(test_loader, verbose=0)
print(f"Test accuracy (exact count): {accuracy:.4f}")