Skip to content

Logging training to TensorBoard

Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Train the GCN of the Cora example and log the loss and accuracy of every epoch for TensorBoard.

View in Colab   GitHub source


Logging training to TensorBoard

Train the GCN of the Cora example and log the loss and accuracy of every epoch for TensorBoard. With Keras this only needs the TensorBoard callback.

Same model as PyG's examples/tensorboard_logging.py.

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

data holds one graph: node features x, edges edge_index, labels y, and masks marking the training, validation and test nodes.

import keras
from keras import ops
from k3_node.datasets import Planetoid
from k3_node.loader import FullGraphDataset
from k3_node.transforms import NormalizeFeatures
from k3_node.layers import GCNConv

dataset = Planetoid("data/Planetoid", name="Cora", transform=NormalizeFeatures())
data = dataset[0]
print(data)

Define the model

class GCN(keras.Model):
    def __init__(self, in_channels, out_channels):
        super().__init__()
        self.conv1 = GCNConv(in_channels, 16)
        self.conv2 = GCNConv(16, out_channels)
        self.dropout = keras.layers.Dropout(0.5)

    def call(self, data, training=False):
        x = self.dropout(ops.relu(self.conv1(data.x, data.edge_index)), training=training)
        return self.conv2(x, data.edge_index)


model = GCN(dataset.num_features, dataset.num_classes)

Train

The logs are written to logs/. Run tensorboard --logdir logs (or %tensorboard --logdir logs in a notebook, after %load_ext tensorboard) to see them.

model.compile(
    optimizer=keras.optimizers.Adam(learning_rate=0.01, weight_decay=5e-4),
    loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    weighted_metrics=["accuracy"],
)
model.fit(
    FullGraphDataset(data, mask="train_mask"),
    validation_data=FullGraphDataset(data, mask="val_mask"),
    epochs=200,
    callbacks=[keras.callbacks.TensorBoard(log_dir="logs")],
    verbose=2,
)

Evaluate

loss, accuracy = model.evaluate(FullGraphDataset(data, mask="test_mask"), verbose=0)
print(f"Test accuracy: {accuracy:.4f}")