Skip to content

Counting triangles with supervised SAGPool (TRIANGLES)

Author: K3-Node Team
Backend: Multi-Backend
Dataset: TRIANGLES (TUDataset)
Description: Count the triangles in a graph.

View in Colab   GitHub source


Counting triangles with supervised SAGPool (TRIANGLES)

Count the triangles in a graph. Following Knyazev et al., 2019, a GIN is combined with two SAGPooling layers (Lee et al., 2019), whose scores are supervised to highlight the nodes that belong to triangles. Nodes have no features, so their one-hot degree is used instead.

Same model as PyG's examples/triangles_sag_pool.py; on a subset of TRIANGLES to keep the dataset small, trained for 200 epochs instead of 300.

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

The given node importances go to data.attn; the node features become one-hot degrees. A subset of the 45,000 graphs keeps the example quick; the test graphs are larger than the training ones.

import keras
from keras import ops
from k3_node.datasets import TUDataset
from k3_node.layers import GCNConv, GINConv, SAGPooling, global_max_pool, global_mean_pool
from k3_node.loader import DataLoader
from k3_node.transforms import OneHotDegree


def handle_node_attention(data):
    data.attn = ops.reshape(ops.softmax(data.x, axis=0), (-1,))
    data.x = None
    return OneHotDegree(max_degree=14)(data)


dataset = TUDataset("data/TU", name="TRIANGLES", use_node_attr=True, transform=handle_node_attention)
train_loader = DataLoader(dataset[:5000], batch_size=60, shuffle=True)
val_loader = DataLoader(dataset[30000:31000], batch_size=60)
test_loader = DataLoader(dataset[35000:37000], 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 TriangleCounter(keras.Model):
    def __init__(self):
        super().__init__()
        self.convs = [GINConv(mlp()) for _ in range(3)]
        self.pools = [SAGPooling(64, min_score=0.001, GNN=GCNConv) for _ in range(2)]
        self.lin = keras.layers.Dense(1)

    def call(self, data):
        x, edge_index, batch = data.x, data.edge_index, data.batch
        for conv, pool in zip(self.convs[:2], self.pools):
            x = ops.relu(conv(x, edge_index))
            x, edge_index, _, batch, perm, score = pool(x, edge_index, batch=batch)
        x = ops.relu(self.convs[2](x, edge_index))
        out = self.lin(global_max_pool(x, batch, data.num_graphs))
        # Supervise the second pooling step (as in PyG); `perm` indexes the nodes kept by both steps
        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 = TriangleCounter()

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