Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions examples/relational/amazon_parquet.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
# KumoRelational on rel-amazon from Parquet

This example predicts the RelBench `rel-amazon` `user-churn` task without loading the full database into pandas. Install SDM with its test dependencies and `relbench`, then run from the repository root:

```bash
python examples/relational/amazon_parquet.py
```

RelBench downloads the dataset to its cache on first use. Use `--context-size`, `--query-size`, and `--num-neighbors` to change the sampled workload. The default uses 512 training and validation task rows as context, 64 test task rows as queries, and two hops with eight neighbors each. The script removes target labels from queries before sampling.

`ParquetRelationalSampler` builds its graph from relationship keys and timestamps in the Parquet tables, using the same CPU neighbor sampler as `RelationalData.sampler()`. Polars reads the selected feature rows after sampling. The resulting `RelatedTables` go directly to `KumoRelational.fit()` and `predict()`; no full `RelationalData` is created from the database.

The keys, timestamps, graph index, and categorical dictionaries still occupy RAM. Feature rows are read from Parquet for each sample. This approach is useful when the full feature tables exceed memory but the graph metadata fits. The example reports accuracy for one small query batch and is not a full benchmark evaluation.

On a 15 GiB RAM machine with an NVIDIA L4, the default run completed against 23,218,245 source rows (6.82 GiB compressed Parquet). It sampled 5,018 related context rows and 1,019 related query rows, with 4.95 GiB peak process RAM and 0.7344 query accuracy. On the same machine, a direct attempt to load the full pandas database through `task.get_db()` was killed before sampling began.

The focused sampler test in `test/relational/test_parquet_sampler.py` compares sampled task rows, relationship links, and related table contents with the in-memory sampler, including categorical values and temporal sampling.
145 changes: 145 additions & 0 deletions examples/relational/amazon_parquet.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,145 @@
# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Run KumoRelational on rel-amazon with Parquet-backed sampling."""

import argparse
import resource
from typing import cast

import pandas as pd
import polars as pl
import pyarrow.parquet as pq
import relbench
import torch

import sdm


def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--context-size", type=int, default=512)
parser.add_argument("--query-size", type=int, default=64)
parser.add_argument("--num-neighbors", type=int, nargs="+", default=[8, 8])
parser.add_argument("--seed", type=int, default=0)
args = parser.parse_args()
torch.manual_seed(args.seed)

dataset = relbench.load_dataset("rel-amazon")
task = dataset.load_task("user-churn")
tables = {
name: dataset.db_dir / f"{name}.parquet"
for name in dataset.manifest.tables
}

stypes = {}
for name, spec in dataset.manifest.tables.items():
sample = (
pl.scan_parquet(tables[name])
.head(10_000)
.collect(engine="streaming")
.to_pandas()
)
keys = dict.fromkeys(spec.fkeys, "id")
if spec.pkey is not None:
keys[spec.pkey] = "id"
stypes[name] = sdm.infer_stypes(
sample,
overrides=keys,
text="drop",
unsupported="drop",
)

relationships = [
{
"left_table": name,
"left_column": column,
"right_table": other,
"right_column": cast(str, dataset.manifest.tables[other].pkey),
}
for name, spec in dataset.manifest.tables.items()
for column, other in spec.fkeys.items()
]
time_columns = {
name: spec.time_col
for name, spec in dataset.manifest.tables.items()
if spec.time_col is not None
}
sampler = sdm.ParquetRelationalSampler(
tables=tables,
stypes=stypes,
relationships=relationships,
time_columns=time_columns,
)

train = task.get_table("train", mask_input_cols=False).df
val = task.get_table("val", mask_input_cols=False).df
test = task.get_table("test", mask_input_cols=False).df
task_table = sdm.TableTensor.from_pandas(
df=pd.concat([train, val, test], ignore_index=True),
stypes={
task.entity_col: "id",
task.time_col: "datetime",
task.target_col: "categorical",
},
)
context, query = task_table.split([len(train) + len(val), len(test)])
context = context[torch.randperm(len(context))[: args.context_size]]
query = query[: args.query_size]

task_link = {
"task_column": task.entity_col,
"table": task.entity_table,
"table_column": cast(
str, dataset.manifest.tables[task.entity_table].pkey
),
}
sampled_context = sampler(
task_table=context,
task_link=task_link,
num_neighbors=args.num_neighbors,
task_time_column=task.time_col,
)
sampled_query = sampler(
task_table=query.drop_columns(task.target_col),
task_link=task_link,
num_neighbors=args.num_neighbors,
task_time_column=task.time_col,
)

source_rows = sum(
pq.ParquetFile(path).metadata.num_rows for path in tables.values()
)
context_rows = sum(
len(table) for table in sampled_context.related_tables.tables.values()
)
query_rows = sum(
len(table) for table in sampled_query.related_tables.tables.values()
)
print(f"Source rows on disk: {source_rows:,}")
print(f"Context related rows: {context_rows:,}")
print(f"Query related rows: {query_rows:,}")

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = sdm.models.KumoRelational(task="classification", device=device)
sampled_context = sampled_context.to(device)
sampled_query = sampled_query.to(device)
with torch.amp.autocast(device.type, enabled=device.type == "cuda"):
model.fit(
x=sampled_context.task_table.drop_columns(task.target_col),
y=sampled_context.task_table[task.target_col],
related_tables=sampled_context.related_tables,
)
prediction = model.predict(*sampled_query)

scores, target = sdm.evaluation.to_class_indices(
prediction, cast(sdm.TableTensor, query[task.target_col].to(device))
)
accuracy = (scores.argmax(-1) == target).float().mean()
print(f"Query accuracy: {accuracy:.4f}")
peak_gib = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 2**20
print(f"Peak process RAM: {peak_gib:.2f} GiB")


if __name__ == "__main__":
main()
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,7 @@ test = [
"pytest",
"pytest-cov",
"pandas",
"polars>=1.30",
"cudf-cu13>=26.8; sys_platform == 'linux'",
"pyg-lib; sys_platform == 'linux'",
"pylibcugraph-cu13>=26.8; sys_platform == 'linux'",
Expand Down
2 changes: 2 additions & 0 deletions sdm/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
RelationalData,
TaskLink,
RelatedTables,
ParquetRelationalSampler,
)
from sdm.processing import Recipe
from sdm import models, evaluation, explain
Expand Down Expand Up @@ -50,6 +51,7 @@
"RelationalData",
"TaskLink",
"RelatedTables",
"ParquetRelationalSampler",
"Recipe",
"models",
"evaluation",
Expand Down
5 changes: 4 additions & 1 deletion sdm/relational/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,12 +5,15 @@

from sdm.relational.data import Relationship, RelationalData
from sdm.relational.task import TaskLink, RelatedTables
from sdm.relational.sampler import RelationalSampler
from sdm.relational.sampler import RelationalSampler, RelationalSamplerOutput
from sdm.relational.parquet import ParquetRelationalSampler

__all__ = [
"Relationship",
"RelationalData",
"TaskLink",
"RelatedTables",
"RelationalSampler",
"RelationalSamplerOutput",
"ParquetRelationalSampler",
]
191 changes: 191 additions & 0 deletions sdm/relational/parquet.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,191 @@
# SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

from collections.abc import Collection, Mapping, Sequence
from pathlib import Path
from typing import Literal, cast

import torch

from sdm import Stype, StypeLike, TableTensor
from sdm.relational import (
RelatedTables,
RelationalData,
RelationalSamplerOutput,
Relationship,
TaskLink,
)

_ROW = "__sdm_disk_row__"
_EXAMPLE = "__example__"


class ParquetRelationalSampler:
"""Sample relational Parquet tables while keeping feature columns on disk.

The relationship keys and time columns are held in CPU memory and sampled
by the same PyG backend as :class:`RelationalSampler`. Polars reads only
sampled feature rows from Parquet. The task table must be two-dimensional.
Source row order must match the order used to create any in-memory
:class:`RelationalData` being compared.

Args:
tables: Parquet file path for each table.
stypes: Semantic types to load from each table. Include every key used
by ``relationships`` and the task link as ``"id"``.
relationships: Joins between the source tables.
time_columns: Datetime column used to constrain sampling by time for
each time-aware table.
"""

def __init__(
self,
tables: Mapping[str, str | Path],
stypes: Mapping[str, Mapping[str, StypeLike]],
relationships: Collection[
Relationship | Mapping[str, str | Sequence[str]]
],
time_columns: Mapping[str, str] | None = None,
) -> None:
try:
import polars as pl # noqa: PLC0415
except ImportError as error:
raise ImportError(
"ParquetRelationalSampler requires 'polars'"
) from error

self._pl = pl
self.paths = {name: Path(path) for name, path in tables.items()}
self.stypes = {
name: {column: Stype(stype) for column, stype in schema.items()}
for name, schema in stypes.items()
}
self.time_columns = dict(time_columns or {})
self._categories: dict[str, dict[str, list[object]]] = {}

graph_tables = {}
for name, path in self.paths.items():
schema = self.stypes[name]
self._categories[name] = {}
for column, stype in schema.items():
if stype != Stype.categorical:
continue
categories = (
pl.scan_parquet(path)
.select(column)
.unique(maintain_order=True)
.filter(pl.col(column).is_not_null())
.collect(engine="streaming")
)
self._categories[name][column] = categories[column].to_list()
columns = [
column
for column, stype in schema.items()
if stype == Stype.id or column == self.time_columns.get(name)
]
if (
_ROW in schema
or _ROW in pl.scan_parquet(path).collect_schema()
):
raise ValueError(f"Column {_ROW!r} is reserved")
frame = pl.scan_parquet(path, row_index_name=_ROW).select(
pl.col(_ROW).cast(pl.Int64), *columns
)
graph_tables[name] = TableTensor.from_pandas(
df=frame.collect(engine="streaming").to_pandas(),
stypes={
_ROW: Stype.id,
**{column: schema[column] for column in columns},
},
)

self.data = RelationalData(
tables=graph_tables,
relationships=relationships,
)
self._sampler = self.data.sampler(time_columns=self.time_columns)

def __call__(
self,
task_table: TableTensor,
task_link: TaskLink | Mapping[str, str | Sequence[str]],
num_neighbors: Sequence[int],
task_time_column: str | None = None,
temporal_strategy: Literal["last", "uniform"] = "last",
) -> RelationalSamplerOutput:
"""Alias of :meth:`sample`."""
return self.sample(
task_table=task_table,
task_link=task_link,
num_neighbors=num_neighbors,
task_time_column=task_time_column,
temporal_strategy=temporal_strategy,
)

def sample(
self,
task_table: TableTensor,
task_link: TaskLink | Mapping[str, str | Sequence[str]],
num_neighbors: Sequence[int],
task_time_column: str | None = None,
temporal_strategy: Literal["last", "uniform"] = "last",
) -> RelationalSamplerOutput:
"""Return sampled task and related tables for model input.

Args:
task_table: Task rows with entity IDs and optional timestamps.
task_link: Link from task rows to a source table.
num_neighbors: Number of neighbors per relationship at each hop.
task_time_column: Datetime column in ``task_table``.
temporal_strategy: ``"last"`` or ``"uniform"`` sampling.

Returns:
The same output type as :meth:`RelationalSampler.sample`.
"""
if task_table.dim() != 2:
raise ValueError("Task table needs to be two-dimensional")

sampled = self._sampler.sample(
task_table=task_table,
task_link=task_link,
num_neighbors=num_neighbors,
task_time_column=task_time_column,
temporal_strategy=temporal_strategy,
)
import pandas as pd

tables: dict[str, TableTensor] = {}
for name, graph_table in sampled.related_tables.tables.items():
assert isinstance(graph_table, TableTensor)
index = graph_table[_ROW].id[..., 0].numpy()
columns = list(self.stypes[name])
frame = (
self._pl.scan_parquet(self.paths[name], row_index_name=_ROW)
.filter(self._pl.col(_ROW).is_in(index.tolist()))
.select(_ROW, *columns)
.collect(engine="streaming")
.to_pandas()
.set_index(_ROW)
.loc[index]
.reset_index(drop=True)
)
for column, categories in self._categories[name].items():
frame[column] = pd.Categorical(
frame[column], categories=categories
)
hydrated = TableTensor.from_pandas(
df=frame,
stypes=self.stypes[name],
)
tables[name] = cast(
TableTensor,
torch.cat([hydrated, graph_table[_EXAMPLE]], dim=-1),
)

return RelationalSamplerOutput(
task_table=sampled.task_table,
related_tables=cast(
RelatedTables[TableTensor],
sampled.related_tables.replace_tables(tables),
),
)
Loading