Skip to main content

K3-Node: Multi-Backend Graph Neural Networks

K3-Node Logo

Documentation PyTorch tests TensorFlow tests JAX tests License Backends Code style: black


K3-Node is a next-generation graph neural network (GNN) library built natively on Keras 3. Write your GNN models once and execute seamlessly across TensorFlow, PyTorch, and JAX with full hardware acceleration (NVIDIA GPUs, Apple Silicon, Google Cloud TPUs).

K3-Node achieves 100% public API parity with PyTorch Geometric (PyG) and incorporates state-of-the-art foundation models and architectures from Spektral and StellarGraph.

📖 Documentation: https://anas-rz.github.io/k3-node/
📋 Porting Checklist & Parity Status: Checklist.md


Key Features

  • 🔄 True Multi-Backend Freedom: Switch between PyTorch, TensorFlow, and JAX with a single environment variable (KERAS_BACKEND=torch|tensorflow|jax).
  • 🧠 Pre-trained Foundation Models: Out-of-the-box architectures and checkpoint loaders for GraphMAE2, Graphormer (2D & 3D), GraphGPS, GROVER, and Mole-BERT.
  • ⚡ 65+ Convolution Layers: Full PyG parity (GCNConv, GATv2Conv, TransformerConv, GPSConv, PNAConv, SchNet, DimeNetPlusPlus, ViSNet, etc.).
  • 📊 26 Aggregation Operators: From elementary aggregations (sum, mean, max, softmax, powermean) to neural aggregations (SetTransformer, GraphMultisetTransformer, Set2Set, DeepSets, LSTMAggregation).
  • 🌐 31 Pooling Operators: Global readouts (global_add_pool, global_mean_pool), hierarchical coarsening (TopKPooling, SAGPooling, ASAPooling, EdgePooling, ClusterPooling), and 3D spatial pooling (voxel_grid, fps, knn, radius).
  • 🧱 Dense & Scalable GNNs: Dense matrix convolutions (DenseGCNConv, DenseGATConv), spectral pooling (DMoNPooling, dense_diff_pool, dense_mincut_pool), and linear-complexity graph transformers (SGFormer, LPFormer, Polynormer).
  • 🧭 Knowledge Graph Embeddings: Multi-relational link prediction with TransE, RotatE, DistMult, ComplEx, and framework-agnostic negative sampling loaders.
  • 📦 Data, Loaders & Transforms: Full suite of graph data structures (Data, HeteroData, Batch), mini-batch samplers (NeighborLoader, ClusterLoader, GraphSAINTSampler), and 62+ graph and 3D point cloud transforms.
  • ✅ Rigorous Verification: 700+ unit tests on every backend, training tests that check each layer's weights actually learn, compiled-vs-eager and cross-backend consistency tests, and numerical parity tests against PyTorch Geometric and reference checkpoints.

Installation

# git should be installed
pip install git+https://github.com/anas-rz/k3-node/

# with the extra packages the example notebooks use (scikit-learn, rdflib, matplotlib)
pip install "k3-node[examples] @ git+https://github.com/anas-rz/k3-node"

Selecting your Backend

Configure your preferred backend before importing k3_node:

export KERAS_BACKEND="torch"       # or "tensorflow" or "jax"

Or programmatically in Python:

import os
os.environ["KERAS_BACKEND"] = "torch"  # Must be set before importing k3_node / keras
import k3_node

Quickstart

Building a Graph Convolutional Network

import keras
from keras import ops
import k3_node.layers as gnn_layers
from k3_node.data import Data

class GCN(keras.Model):
    def __init__(self, in_channels, hidden_channels, out_channels):
        super().__init__()
        self.conv1 = gnn_layers.GCNConv(in_channels, hidden_channels)
        self.conv2 = gnn_layers.GCNConv(hidden_channels, out_channels)

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

# Instantiate model
model = GCN(in_channels=16, hidden_channels=32, out_channels=7)

# Forward pass on graph data
x = ops.ones((10, 16))
edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")

out = model(x, edge_index)
print("Output shape:", out.shape)  # (10, 7)

Training in a Few Lines

The task estimators in k3_node.tasks pick the loss, readout and metrics for you:

from k3_node.datasets import Planetoid
from k3_node.tasks import NodeClassifier

cora = Planetoid("data/Planetoid", name="Cora")[0]

classifier = NodeClassifier(backbone="gcn", hidden_channels=64, num_layers=2, dropout=0.5)
classifier.fit(cora, epochs=100, lr=0.01)
print(classifier.evaluate(cora, mask="test_mask"))

GraphClassifier, GraphRegressor, NodeRegressor and LinkPredictor work the same way.

Example Notebooks

The examples/ folder has 90+ notebooks that follow the architectures of PyG's examples, written with keras.Model.fit and K3-Node's loaders. They cover node, link and graph classification, knowledge graphs, molecules (including pre-trained DimeNet, DimeNet++ and SchNet on QM9), point clouds, temporal graphs and large-graph mini-batching. Each notebook opens in Colab and runs on any backend: change KERAS_BACKEND in its first cell. Browse them in the documentation.


Pre-trained Foundation Models

K3-Node provides ready-to-use architectures and automated checkpoint loading for state-of-the-art graph foundation models:

1. GraphMAE2 (Self-Supervised Masked Autoencoder)

from k3_node.models import GraphMAE2
from k3_node.models.graphmae2 import load_graphmae2_weights

model = GraphMAE2(
    in_dim=100,
    num_hidden=512,
    out_dim=100,
    num_layers=4,
    encoder_type="gat",
    decoder_type="gat"
)
# Load reference pre-trained weights
load_graphmae2_weights(model, "checkpoints/graphmae2_ogbn_arxiv.pt")

2. Graphormer (2D Molecular & 3D Structural Transformer)

from k3_node.models import Graphormer, Graphormer3D
from k3_node.models.graphormer import load_graphormer_weights

# 2D Graphormer (PCQM4Mv2)
model_2d = Graphormer(num_layers=12, num_heads=32, embed_dim=768)
load_graphormer_weights(model_2d, "checkpoints/graphormer_pcqm4mv2.pt")

# 3D Graphormer (OC20 Catalyst Adsorption & Molecular Conformations)
model_3d = Graphormer3D(num_layers=12, num_heads=32, embed_dim=768)

3. GraphGPS (Hybrid Local MPNN + Global Transformer)

from k3_node.models import GPSModel
from k3_node.models.gps_model import load_gps_model_weights

model = GPSModel(
    channels=64,
    num_layers=5,
    local_gnn_type="GINE",
    global_model_type="Transformer"
)
load_gps_model_weights(model, "checkpoints/graphgps_zinc.pt")

4. GROVER (Self-Supervised Message Passing Transformer)

from k3_node.models import GROVER, GROVEREmbedding
from k3_node.models.grover import load_grover_weights

model = GROVER(hidden_size=128, num_layers=3, num_heads=4)
load_grover_weights(model, "checkpoints/grover_base.pt")

5. Mole-BERT (Masked Chemical Graph Representation)

from k3_node.models import MoleBERT
from k3_node.models.mole_bert import load_mole_bert_weights

model = MoleBERT(num_layer=5, emb_dim=300, drop_ratio=0.5)
load_mole_bert_weights(model, "checkpoints/Mole-BERT.pth")

What's Included

Package Status Contents
k3_node.layers.conv ✅ 65/65 GCNConv, GATConv, GATv2Conv, SAGEConv, GINConv, GPSConv, TransformerConv, PNAConv, SchNet, DimeNetPlusPlus, ViSNet, etc.
k3_node.layers.pool ✅ 31/31 global_add_pool, global_mean_pool, TopKPooling, SAGPooling, ASAPooling, EdgePooling, ClusterPooling, voxel_grid, fps, graclus, etc.
k3_node.layers.aggr ✅ 26/26 SumAggregation, MeanAggregation, SoftmaxAggregation, PowerMeanAggregation, MultiAggregation, SetTransformerAggregation, Set2Set, etc.
k3_node.layers.norm ✅ 11/11 GraphNorm, PairNorm, DiffGroupNorm, MessageNorm, MeanSubtractionNorm, BatchNorm, LayerNorm, HeteroBatchNorm, etc.
k3_node.layers.dense ✅ 11/11 DenseGCNConv, DenseGATConv, DenseGINConv, DenseSAGEConv, DMoNPooling, dense_diff_pool, dense_mincut_pool, Linear, etc.
k3_node.layers.kge ✅ 5/5 KGEModel, TransE, RotatE, DistMult, ComplEx, KGTripletLoader.
k3_node.models ✅ 46/46 MLP, GAE, VGAE, DeepGraphInfomax, Node2Vec, LabelPropagation, LINKX, LightGCN, SGFormer, LPFormer, Polynormer, etc.
Foundation Models ✅ 5/5 GraphMAE2, Graphormer (2D/3D), GPSModel, GROVER, MoleBERT with pre-trained weight conversion.
k3_node.data ✅ 19/19 Data, HeteroData, Batch, TemporalData, HypergraphData, InMemoryDataset, FeatureStore, GraphStore, etc.
k3_node.loader ✅ 26/26 DataLoader, NeighborLoader, LinkNeighborLoader, ClusterLoader, GraphSAINTSampler, ShaDowKHopSampler, etc.
k3_node.transforms ✅ 62/62 Topology rewiring, positional encodings (LapPE, RWPE, GPSE), spectral diffusion (GDC), and 3D point cloud transforms.

Testing & Verification

Run the comprehensive test suite across backends:

# Run all unit tests
pytest k3_node/

# Run training tests (each layer's weights learn; slower, not run in CI)
pytest tests_training/

# Run reference parity check against PyTorch implementations
pytest tests_reference/

License

This project is licensed under the MIT License - see the LICENSE file for details.

Metadata

Release files for k3-node 1.0.0

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for k3-node 1.0.0
File Size Uploaded
k3_node-1.0.0.tar.gz 551.7 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for k3-node 1.0.0
File Interpreter ABI Platform
k3_node-1.0.0-py3-none-any.whl Python 3 none any Details

Total release size: 1.4 MB

Release files / k3_node-1.0.0.tar.gz

Download URL k3_node-1.0.0.tar.gz
Size 551.7 kB
Tags Source
SHA-256 checksum
How to use checksums
bedd53c14b96aa581b513c055dca759fc266549362bd4813f706ca43eac26d13
BLAKE2b-256 checksum
How to use checksums
bd51364c35fe1457454d1cc76bc1808644d5ba774c157c370e2dc31efbcbb57e
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Sep 30, 2026.

Transparency log

Release files / k3_node-1.0.0-py3-none-any.whl

Download URL k3_node-1.0.0-py3-none-any.whl
Size 814.1 kB
Tags Python 3
SHA-256 checksum
How to use checksums
7c2e2d17b0c9c66f86b0f3b7b6b9cd14669dae4f695061f39e42d66b1c183a9d
BLAKE2b-256 checksum
How to use checksums
e88c81e5823e7cd9848ce8c66054c4408916dcbf0098b9c7fb0105ed0c73dd85
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
Yes
Uploaded via twine/7.0.0 CPython/3.13.14

Provenance

Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.

PyPI Publish Attestation

PyPI verified that this artifact, at this checksum, originated from the publisher listed below.

Signed by GitHub Actions, verified by PyPI on Sep 30, 2026.

Transparency log

Release history Release notifications | RSS feed

This release

1.0.0 This release

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page