Counting triangles with supervised SAGPool (TRIANGLES)
Author: K3-Node Team
Backend: Multi-Backend
Dataset: TRIANGLES (TUDataset)
Description: Count the triangles in a graph.
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"
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)