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