Skip to main content

tabpfn-graph

tabpfn-graph turns a collection of graphs into a schema-stable pandas table and applies TabPFN or any scikit-learn-compatible estimator. It targets single-target, graph-level classification and regression.

from tabpfn_graph import GraphClassifier, GraphRegressor

clf = GraphClassifier().fit(graphs_train, y_train)
predictions = clf.predict(graphs_test)

reg = GraphRegressor(estimator=my_pipeline).fit(graphs_train, y_train)

The representation uses established graph descriptors: Local Degree/Topological Profile statistics, Weisfeiler–Lehman subtree hashes over categorical node/edge labels, attribute distribution aggregation, and optional centrality, paths, motifs, Laplacian, and NetLSD summaries. This follows the graph-level descriptor evidence from LTP and MOLTOP. The broader graph-to-table foundation-model pattern has also been explored for node tasks by G2T-FM and TabPFN-GN; those papers do not imply that this package implements their node-level methods.

Weighted, directed, bipartite, disconnected, and size-heterogeneous graphs each need a different configuration, and fit warns when the training data looks like a mismatch. docs/graph-types.md in the source repository states what each family needs.

Install

pip install tabpfn-graph

Python 3.11–3.14 is supported. The standard installation includes TabPFN, NetworkX (3.5 or newer, because WL subtree hashes changed in that release), NumPy, pandas, SciPy, scikit-learn, and joblib. Optional extras are:

pip install 'tabpfn-graph[fast]'    # Networkit for compatible primitives
pip install 'tabpfn-graph[client]'  # use hosted clients as user-supplied estimators
pip install 'tabpfn-graph[dev]'     # tests, lint, typing, packaging

The first local TabPFN fit may require accepting the checkpoint terms and downloading model weights. Package source code is Apache-2.0; TabPFN checkpoints have separate terms. More detail is available in docs/model-access.md in the source repository.

NetworkX quickstarts

Classification with the default local TabPFN:

import networkx as nx
from tabpfn_graph import GraphClassifier

graphs = [nx.path_graph(5), nx.cycle_graph(5), nx.star_graph(4), nx.complete_graph(5)]
y = [0, 1, 0, 1]

clf = GraphClassifier(random_state=0).fit(graphs, y)
labels = clf.predict([nx.path_graph(7), nx.cycle_graph(7)])
probabilities = clf.predict_proba([nx.path_graph(7), nx.cycle_graph(7)])

Regression uses the same extraction contract:

from tabpfn_graph import GraphRegressor

reg = GraphRegressor(random_state=0).fit(graphs, [1.2, 2.5, 0.8, 4.1])
values = reg.predict(graphs)

The default local estimators are created lazily during fit, with TabPFN's local text transformation enabled. Importing or constructing GraphClassifier() does not access model weights.

PyTorch Geometric datasets

PyG is intentionally optional. Pass a homogeneous iterable of torch_geometric.data.Data objects when it is installed:

from torch_geometric.datasets import TUDataset
from tabpfn_graph import GraphClassifier

dataset = TUDataset(root="data/TU", name="MUTAG")
graphs = [data for data in dataset]
y = [int(data.y.item()) for data in dataset]
model = GraphClassifier().fit(graphs[:150], y[:150])
prediction = model.predict(graphs[150:])

PyG's common doubled-edge representation is recognized: if every non-loop arc has a reciprocal arc with matching multiplicity, each pair is collapsed into one undirected logical edge.

Standalone feature extraction

from tabpfn_graph import GraphFeatureExtractor

extractor = GraphFeatureExtractor(features="balanced", n_jobs=-1, random_state=0)
X_train = extractor.fit_transform(graphs_train)
X_test = extractor.transform(graphs_test)

assert list(X_train.columns) == list(X_test.columns)

X_train and X_test are pandas DataFrames. Numeric values remain numeric, graph-level categoricals use pandas categorical dtype, and semantic documents use pandas string dtype.

Custom estimators

No text encoding is inserted for custom estimators. For numeric-only feature tables:

from sklearn.ensemble import RandomForestClassifier
from tabpfn_graph import GraphClassifier

model = GraphClassifier(
    features="balanced",
    estimator=RandomForestClassifier(n_estimators=500, random_state=0),
).fit(graphs_train, y_train)

XGBoost and CatBoost-style estimators work the same way:

from xgboost import XGBClassifier

model = GraphClassifier(
    features="balanced",
    estimator=XGBClassifier(n_estimators=500, random_state=0),
).fit(graphs_train, y_train)

CatBoost can consume categorical/text columns when configured with their column names or indices. Hosted TabPFN clients can be passed through estimator=; the package does not require a particular client API beyond sklearn-style fit and predict.

For a text-aware sklearn pipeline, explicitly select and transform the document column:

from sklearn.compose import ColumnTransformer
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.linear_model import LogisticRegression
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler
from tabpfn_graph import GraphClassifier, GraphFeatureExtractor

extractor = GraphFeatureExtractor(
    features=("basic", "text"),
    text_attributes=("node.description",),
)
preprocess = ColumnTransformer([
    ("text", TfidfVectorizer(), "node_attr__description__document"),
    ("numeric", StandardScaler(), ["basic__n_nodes", "basic__n_edges_native"]),
])
estimator = make_pipeline(preprocess, LogisticRegression(random_state=0))
model = GraphClassifier(feature_extractor=extractor, estimator=estimator)

Weighted and directed graphs

Structural descriptors ignore edge weights unless you name the attribute, because a weight has no universal meaning. Path and betweenness descriptors need a traversal cost, so the semantics are explicit rather than guessed:

extractor = GraphFeatureExtractor(
    edge_weight="weight",
    edge_weight_semantics="similarity",  # cost = 1 / w; use "distance" when w is a length
)

That switches strength profiles, weighted clustering, PageRank, betweenness, assortativity, shortest paths, and the Laplacian spectrum to their weighted definitions. When edge_weight is left unset but the training graphs carry numeric edge attributes, fit warns.

Directed input additionally gets reciprocity, in/out degree summaries and their correlation, strongly connected components, acyclicity, and PageRank on a direction-preserving projection of the native arcs; PageRank on the undirected projection is close to a rescaled degree. Mixed batches are supported and yield one schema, with undirected graphs described as their own symmetrization.

Diagnostics

extractor = GraphFeatureExtractor().fit(graphs_train)
extractor.diagnostics_        # sizes, directedness, bipartiteness, WL vocabulary, versions

fit warns about ignored edge weights, mixed directedness, extreme size heterogeneity, all-bipartite datasets, and aggressive WL pruning. fit_transform also records how many columns are constant or exactly duplicated on the training data, available standalone as tabpfn_graph.column_report(frame). Setting prune_uninformative=True turns that into a fit-learned schema decision and drops those columns.

Feature selection

Presets are fast, balanced (default), and comprehensive:

  • fast: native/basic topology and node, edge, graph metadata, and text aggregation.
  • balanced: fast plus clustering/core/LDP profiles and two-iteration hashed WL counts.
  • comprehensive: balanced plus centrality, paths, motifs, Laplacian, and NetLSD summaries.

Or select groups directly:

extractor = GraphFeatureExtractor(features=("basic", "wl", "spectral"))

Valid groups are basic, local_profile, attributes, text, wl, centrality, paths, motifs, and spectral. Expensive groups preflight against max_exact_nodes=2000. To make the change in semantics explicit, larger graphs require either a raised exact limit or allow_approximate=True, which switches to pivot-sampled betweenness, sampled-source distances, rescaled triangle counts, and a truncated-spectrum heat trace — estimators of the whole-graph quantity, recorded per graph in approx__active.

Scope

Single-target graph-level classification and regression on homogeneous NetworkX or PyG batches. Heterogeneous graphs (HeteroData), temporal graphs, and multilabel targets are out of scope.

This is an alpha release. The descriptors are established ones and the schema contract is tested, but the package is not backed by a broad benchmark study; treat it as a descriptor baseline rather than a method shown to beat tuned GBDT-on-descriptors or GNNs.

Detailed feature semantics, graph-family guidance, evaluation guidance, and reproducible benchmarks are kept in the source repository under docs/ and benchmarks/. Benchmark code is not part of the published wheel or source distribution.

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

tabpfn_graph-0.2.0.tar.gz (46.3 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

tabpfn_graph-0.2.0-py3-none-any.whl (33.0 kB view details)

Uploaded Python 3

File details

Details for the file tabpfn_graph-0.2.0.tar.gz.

File metadata

  • Download URL: tabpfn_graph-0.2.0.tar.gz
  • Upload date:
  • Size: 46.3 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for tabpfn_graph-0.2.0.tar.gz
Algorithm Hash digest
SHA256 faf20b502ee07f3487f2bd11c30fa19b6d9533d9e9e4e788913e923f48bbfa1f
MD5 29e83f0a59d837d7295a9ae68b46ccea
BLAKE2b-256 357dcb9492d21acfe41e9479dd6d15446dd062f2ab6f749ae0b1036da23b7630

See more details on using hashes here.

Provenance

The following attestation bundles were made for tabpfn_graph-0.2.0.tar.gz:

Publisher: publish.yml on m-herre/tabpfn-graph

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file tabpfn_graph-0.2.0-py3-none-any.whl.

File metadata

  • Download URL: tabpfn_graph-0.2.0-py3-none-any.whl
  • Upload date:
  • Size: 33.0 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for tabpfn_graph-0.2.0-py3-none-any.whl
Algorithm Hash digest
SHA256 0c46ea6251427764104c425dc17509b60b8a15262e5b7eed18db653a84f0bc2c
MD5 0fb8d35fa7133639e81fad5cb4603249
BLAKE2b-256 ecad4be60951bc35414fc7d30908a02d9cd6911ef5d2346a1359d5a2fc20501d

See more details on using hashes here.

Provenance

The following attestation bundles were made for tabpfn_graph-0.2.0-py3-none-any.whl:

Publisher: publish.yml on m-herre/tabpfn-graph

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Release history Release notifications | RSS feed

This release

0.2.0 This release

2 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