Dynamic Neighborhood Aggregation (DNA) on Cora
Author: K3-Node Team
Backend: Multi-Backend
Dataset: Cora (Planetoid)
Description: Classify papers in the Cora citation network by topic.
Dynamic Neighborhood Aggregation (DNA) on Cora
Classify papers in the Cora citation network by topic. DNA (Fey, 2019) lets every layer attend over the representations produced by all previous layers, so each node picks how far into the graph it looks. The nodes are split randomly into 20% training, 20% validation and 60% test.
Same model as PyG's examples/dna.py.
Install K3-Node, then choose a backend: "tensorflow", "torch" or "jax"
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 RandomNodeSplit
dataset = Planetoid("data/Planetoid", name="Cora", transform=RandomNodeSplit(split="train_rest", num_val=0.2, num_test=0.6))
data = dataset[0]
print(data)
Define the model
from k3_node.layers import DNAConv
class DNA(keras.Model):
def __init__(self, in_channels, hidden_channels, out_channels, num_layers, heads, groups):
super().__init__()
self.hidden_channels = hidden_channels
self.dropout = keras.layers.Dropout(0.5)
self.lin1 = keras.layers.Dense(hidden_channels, activation="relu")
self.convs = [DNAConv(hidden_channels, heads, groups, dropout=0.8) for _ in range(num_layers)]
self.lin2 = keras.layers.Dense(out_channels)
def call(self, data, training=False):
x = self.dropout(self.lin1(data.x), training=training)
x_all = ops.expand_dims(x, 1) # [num_nodes, layers so far, channels]
for conv in self.convs:
x = ops.relu(conv(x_all, data.edge_index, training=training))
x_all = ops.concatenate([x_all, ops.expand_dims(x, 1)], axis=1)
x = self.dropout(x_all[:, -1], training=training)
return self.lin2(x)
model = DNA(dataset.num_features, hidden_channels=128, out_channels=dataset.num_classes,
num_layers=5, heads=8, groups=16)
Train
FullGraphDataset feeds the whole graph to Keras; mask selects which nodes count in the loss and the accuracy.
model.compile(
optimizer=keras.optimizers.Adam(learning_rate=0.005, weight_decay=0.0005),
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,
verbose=2,
)