Skip to content

Benchmarks

Loader registry and contract

Look up loaders with get_class(name), list them with list_names(), or add one with register, inherited from ClassRegistry.

dataset_registry module-attribute

dataset_registry = DatasetRegistry()

DatasetLoaderProtocol

Bases: Protocol

Loader for one registered NL2Q benchmark.

Implementations provide access to benchmark tasks, database connectors, and the evaluation metrics appropriate for the dataset.

Attributes:

Name Type Description
name str

Unique registry key for the dataset (e.g. "bird-sql").

splits list[str]

Non-empty available data splits, with the default first.

default_metrics list[str]

Metric names evaluated by default for this dataset. Can be overridden via --metrics on the evaluate CLI.

name class-attribute

name: str

splits class-attribute

splits: list[str]

default_metrics class-attribute

default_metrics: list[str]

installation class-attribute

installation: BenchmarkInstallation

get_databases

get_databases(split: str) -> list[str]

Returns the list of database names available in the given split.

get_tasks_async async

get_tasks_async(
    split: str, databases: list[str] | None = None
) -> Sequence[NL2QTask]

Loads tasks for a split, optionally filtered to specific databases.

get_db_connectors_async async

get_db_connectors_async(
    split: str, databases: list[str] | None = None
) -> Mapping[str, DataConnector]

Creates database connectors keyed by database name.

get_split_async async

get_split_async(
    split: str,
    databases: list[str] | None = None,
    subsample_size: int | None = None,
    qids: list[str] | None = None,
) -> NL2QDataset

Load selected tasks and the connectors they require.

QID filtering precedes deterministic sampling. Unknown QIDs and invalid sample sizes raise ValueError.

DatasetRegistry

DatasetRegistry()

Bases: ClassRegistry[DatasetLoaderProtocol]

Lazily populated registry of benchmark loaders.

get_class

get_class(name: str) -> type[DatasetLoaderProtocol]

list_names

list_names() -> list[str]

register

register(cls: type[_T]) -> type[_T]

Implement a loader

For reusable splits, implement DatasetLoaderProtocol:

  • Declare name, available splits, default_metrics, and a BenchmarkInstallation describing the required local data.
  • Implement get_databases, get_tasks_async, and get_db_connectors_async.
  • Implement get_split_async to select tasks first, then open only the needed connectors and return an NL2QDataset. Reuse select_tasks for QID filtering and deterministic sampling, and selected_databases to find required databases.

Register the loader with dataset_registry.register(YourLoader). Registrations apply to the current Python process. For managed database services, supply a BenchmarkRuntime with start, stop, and readiness callbacks.

Task selection

QID filtering precedes deterministic sampling. Unknown QIDs and invalid sample sizes raise ValueError.

select_tasks

select_tasks(
    tasks: list[TaskT],
    qids: list[str] | None = None,
    subsample_size: int | None = None,
) -> list[TaskT]

Filter by QID, then take a deterministic random sample.

selected_databases

selected_databases(tasks: Sequence[NL2QTask]) -> list[str]

Return task database names in first-seen order.

Built-in loaders

Loaders are re-exported from tabulaflow.research.benchmarks. See Benchmarks for setup requirements.

BirdSQLDatasetLoader

BirdSQLDatasetLoader(
    directory: str | None = None,
    column_meaning_directory: str | None = None,
    max_concurrency: int = 16,
    connector_config: SQLConnectorConfig | None = None,
)

name class-attribute

name: str = 'bird-sql'

splits class-attribute

splits: list[str] = ['dev', 'dev_20251106', 'train']

installation class-attribute

installation: BenchmarkInstallation = BenchmarkInstallation(
    name=name,
    required_paths=(
        "dev_20240627/dev.json",
        "dev_20240627/dev_databases",
        "dev_20251106/dev.json",
        "train/train.json",
        "train/train_databases",
        "column_meaning/dev_column_meaning.json",
        "column_meaning/train_column_meaning.json",
    ),
    fetch=_fetch_bird_sql,
)

default_metrics class-attribute

default_metrics: list[str] = [
    "bird_sql_ex",
    "simple_ex",
    "bird_sql_ex_soft",
    "executable",
    "gold_executable",
    "gold_result_not_empty",
    "pred_success",
    "raw_pred_bird_sql_ex",
    "raw_pred_simple_ex",
    "schema_linking_stats",
]

directory instance-attribute

directory = str(
    default_directory if directory is None else directory
)

column_meaning_directory instance-attribute

column_meaning_directory = str(
    default_directory / "column_meaning"
    if column_meaning_directory is None
    else column_meaning_directory
)

max_concurrency instance-attribute

max_concurrency = max_concurrency

connector_config instance-attribute

connector_config = (
    SQLConnectorConfig(
        schema_cache_mode="read_write",
        query_timeout_seconds=300,
    )
    if connector_config is None
    else connector_config
)

get_databases

get_databases(split: str) -> list[str]

get_tasks_async async

get_tasks_async(
    split: str,
    databases: list[str] | None = None,
    difficulty: str
    | Literal["simple", "moderate", "challenging"]
    | None = None,
) -> list[SimpleNL2QTask]

get_db_connectors_async async

get_db_connectors_async(
    split: str, databases: list[str] | None = None
) -> dict[str, SQLConnector]

get_split_async async

get_split_async(
    split: str,
    databases: list[str] | None = None,
    subsample_size: int | None = None,
    qids: list[str] | None = None,
    difficulty: str
    | Literal["simple", "moderate", "challenging"]
    | None = None,
) -> NL2QDataset

Spider2SnowDatasetLoader

Spider2SnowDatasetLoader(
    directory: str | None = None,
    sf_user: Optional[str] = None,
    sf_password: Optional[str] = None,
    sf_account: Optional[str] = None,
    connector_config: SQLConnectorConfig | None = None,
)

Loader for Spider 2.0 Snowflake.

Initializes the Spider 2.0 Snowflake dataset loader.

Parameters:

Name Type Description Default
directory str | None

Path to the spider2-snow data directory.

None
sf_user Optional[str]

Snowflake username. Falls back to SF_USER env var.

None
sf_password Optional[str]

Snowflake password. Falls back to SF_PASSWORD env var.

None
sf_account Optional[str]

Snowflake account identifier. Falls back to SF_ACCOUNT env var.

None

name class-attribute

name: str = 'spider2-snow'

splits class-attribute

splits: list[str] = ['test']

installation class-attribute

installation: BenchmarkInstallation = BenchmarkInstallation(
    name=name,
    required_paths=(
        "spider2-snow.jsonl",
        "evaluation_suite/gold/exec_result",
        "resource/databases",
    ),
    fetch=_fetch_spider2_snow,
)

default_metrics class-attribute

default_metrics: list[str] = [
    "spider2_ex",
    "simple_ex",
    "executable",
    "gold_executable",
    "gold_result_not_empty",
    "pred_success",
    "schema_linking_stats",
]

directory instance-attribute

directory = str(
    self.installation.directory
    if directory is None
    else directory
)

sf_user instance-attribute

sf_user = sf_user

sf_password instance-attribute

sf_password = sf_password

sf_account instance-attribute

sf_account = sf_account

connector_config instance-attribute

connector_config = (
    SQLConnectorConfig(
        schema_cache_mode="read_write",
        query_timeout_seconds=300,
    )
    if connector_config is None
    else connector_config
)

get_databases

get_databases(split: str) -> list[str]

get_tasks_async async

get_tasks_async(
    split: str, databases: list[str] | None = None
) -> list[SimpleNL2QTask]

get_db_connectors_async async

get_db_connectors_async(
    split: str, databases: list[str] | None = None
) -> dict[str, SQLConnector]

get_split_async async

get_split_async(
    split: str,
    databases: list[str] | None = None,
    subsample_size: int | None = None,
    qids: list[str] | None = None,
) -> NL2QDataset

Spider2LiteDatasetLoader

Spider2LiteDatasetLoader(
    directory: str | None = None,
    sf_user: Optional[str] = None,
    sf_password: Optional[str] = None,
    sf_account: Optional[str] = None,
    google_cloud_project: Optional[str] = None,
    google_application_credentials: Optional[str] = None,
    connector_config: SQLConnectorConfig | None = None,
)

Loader for Spider 2.0-Lite (BigQuery, Snowflake, SQLite).

Initializes the Spider 2.0-Lite dataset loader.

Parameters:

Name Type Description Default
directory str | None

Path to the spider2-lite data directory.

None
sf_user Optional[str]

Snowflake username. Falls back to SF_USER env var.

None
sf_password Optional[str]

Snowflake password. Falls back to SF_PASSWORD env var.

None
sf_account Optional[str]

Snowflake account identifier. Falls back to SF_ACCOUNT env var.

None
google_cloud_project Optional[str]

GCP project used for BigQuery billing. Falls back to GOOGLE_CLOUD_PROJECT env var.

None
google_application_credentials Optional[str]

Path to a GCP service account JSON key file. Falls back to GOOGLE_APPLICATION_CREDENTIALS env var.

None

name class-attribute

name: str = 'spider2-lite'

splits class-attribute

splits: list[str] = ['test']

installation class-attribute

installation: BenchmarkInstallation = BenchmarkInstallation(
    name=name,
    required_paths=(
        "spider2-lite.jsonl",
        "evaluation_suite/gold/exec_result",
        "resource/databases/spider2-localdb",
    ),
    fetch=_fetch_spider2_lite,
)

default_metrics class-attribute

default_metrics: list[str] = [
    "spider2_ex",
    "simple_ex",
    "executable",
    "gold_executable",
    "gold_result_not_empty",
    "pred_success",
    "schema_linking_stats",
]

directory instance-attribute

directory = str(
    self.installation.directory
    if directory is None
    else directory
)

sf_user instance-attribute

sf_user = sf_user

sf_password instance-attribute

sf_password = sf_password

sf_account instance-attribute

sf_account = sf_account

google_cloud_project instance-attribute

google_cloud_project = google_cloud_project

google_application_credentials instance-attribute

google_application_credentials = (
    google_application_credentials
)

connector_config instance-attribute

connector_config = (
    SQLConnectorConfig(
        schema_cache_mode="read_write",
        query_timeout_seconds=300,
    )
    if connector_config is None
    else connector_config
)

get_databases

get_databases(split: str) -> list[str]

get_tasks_async async

get_tasks_async(
    split: str, databases: list[str] | None = None
) -> list[SimpleNL2QTask]

get_db_connectors_async async

get_db_connectors_async(
    split: str, databases: list[str] | None = None
) -> dict[str, SQLConnector]

Return DB connectors keyed by database name.

Dispatches to BigQuery, Snowflake, or SQLite based on the resource directory layout.

get_split_async async

get_split_async(
    split: str,
    databases: list[str] | None = None,
    subsample_size: int | None = None,
    qids: list[str] | None = None,
) -> NL2QDataset

Spider2DbtDatasetLoader

Spider2DbtDatasetLoader(
    directory: str | None = None,
    max_concurrency: int = 16,
    connector_config: SQLConnectorConfig | None = None,
    workspace_dir: str | Path | None = None,
)

Loader for Spider 2.0-DBT (DuckDB dbt transformation tasks).

Initializes the Spider 2.0-DBT dataset loader.

Parameters:

Name Type Description Default
directory str | None

Path to the spider2-dbt data directory.

None
max_concurrency int

Maximum concurrent DuckDB connections.

16
connector_config SQLConnectorConfig | None

Configuration for source database connectors.

None
workspace_dir str | Path | None

Root for isolated project copies. When provided, loaded splits are ready to run with DbtAgent.

None

name class-attribute

name: str = 'spider2-dbt'

splits class-attribute

splits: list[str] = ['test']

installation class-attribute

installation: BenchmarkInstallation = BenchmarkInstallation(
    name=name,
    required_paths=(
        "examples/spider2-dbt.jsonl",
        "examples",
        "evaluation_suite/gold",
    ),
    fetch=_fetch_spider2_dbt,
)

default_metrics class-attribute

default_metrics: list[str] = [
    "spider2_duckdb_match",
    "executable",
    "pred_success",
]

directory instance-attribute

directory = str(
    self.installation.directory
    if directory is None
    else directory
)

max_concurrency instance-attribute

max_concurrency = max_concurrency

workspace_dir instance-attribute

workspace_dir = (
    None if workspace_dir is None else Path(workspace_dir)
)

connector_config instance-attribute

connector_config = (
    SQLConnectorConfig(
        schema_cache_mode="read_write",
        query_timeout_seconds=300,
    )
    if connector_config is None
    else connector_config
)

get_databases

get_databases(split: str) -> list[str]

get_tasks_async async

get_tasks_async(
    split: str, databases: list[str] | None = None
) -> list[DbtTask]

get_db_connectors_async async

get_db_connectors_async(
    split: str, databases: list[str] | None = None
) -> dict[str, SQLConnector]

Return DuckDB connectors keyed by instance_id.

Each dbt project directory contains a .duckdb file that serves as the source database for that project.

get_split_async async

get_split_async(
    split: str,
    databases: list[str] | None = None,
    subsample_size: int | None = None,
    qids: list[str] | None = None,
) -> NL2QDataset

BeaverDatasetLoader

BeaverDatasetLoader(
    directory: str | None = None,
    dw_port: int = 3311,
    nw_port: int = 3312,
    connector_config: SQLConnectorConfig | None = None,
)

name class-attribute

name: str = 'beaver'

splits class-attribute

splits: list[str] = ['test']

installation class-attribute

installation: BenchmarkInstallation = BEAVER_INSTALLATION

runtime class-attribute

runtime: BenchmarkRuntime = BEAVER_RUNTIME

default_metrics class-attribute

default_metrics: list[str] = [
    "simple_ex",
    "executable",
    "gold_executable",
    "gold_result_not_empty",
    "pred_success",
]

directory instance-attribute

directory = str(
    self.installation.directory
    if directory is None
    else directory
)

dw_dbms_port instance-attribute

dw_dbms_port = dw_port

nw_dbms_port instance-attribute

nw_dbms_port = nw_port

connector_config instance-attribute

connector_config = (
    SQLConnectorConfig(
        schema_cache_mode="read_write",
        query_timeout_seconds=300,
    )
    if connector_config is None
    else connector_config
)

get_databases

get_databases(split: str) -> list[str]

get_tasks_async async

get_tasks_async(
    split: str, databases: list[str] | None = None
) -> list[SimpleNL2QTask]

get_db_connectors_async async

get_db_connectors_async(
    split: str, databases: list[str] | None = None
) -> dict[str, SQLConnector]

get_split_async async

get_split_async(
    split: str,
    databases: list[str] | None = None,
    subsample_size: int | None = None,
    qids: list[str] | None = None,
) -> NL2QDataset

ARCSDatasetLoader

ARCSDatasetLoader(
    directory: str | None = None,
    max_concurrency: int = 16,
    include_taxonomy: bool = False,
    connector_config: SQLConnectorConfig | None = None,
)

name class-attribute

name: str = 'arcs'

splits class-attribute

splits: list[str] = ['test', 'test_unsampled']

installation class-attribute

installation: BenchmarkInstallation = BenchmarkInstallation(
    name=name,
    required_paths=(
        "tasks/tasks_unsampled.json",
        "tasks/tasks_gold_intended_query_ids.json",
        "databases/sqlite",
        "databases/column_meanings.json",
    ),
)

default_metrics class-attribute

default_metrics: list[str] = [
    "simple_ex",
    "executable",
    "gold_executable",
    "gold_result_not_empty",
    "pred_success",
    "ambig_point_stats",
    "gold_ambig_point_stats",
    "found_one",
]

directory instance-attribute

directory = str(
    self.installation.directory
    if directory is None
    else directory
)

max_concurrency instance-attribute

max_concurrency = max_concurrency

include_taxonomy instance-attribute

include_taxonomy = include_taxonomy

connector_config instance-attribute

connector_config = (
    SQLConnectorConfig(
        schema_cache_mode="read_write",
        query_timeout_seconds=300,
    )
    if connector_config is None
    else connector_config
)

get_databases

get_databases(split: str) -> list[str]

get_tasks_async async

get_tasks_async(
    split: str, databases: list[str] | None = None
) -> list[AmbigNL2QTask]

get_db_connectors_async async

get_db_connectors_async(
    split: str, databases: list[str] | None = None
) -> dict[str, SQLConnector]

get_split_async async

get_split_async(
    split: str,
    databases: list[str] | None = None,
    subsample_size: int | None = None,
    qids: list[str] | None = None,
) -> NL2QDataset

AmbrosiaSDatasetLoader

AmbrosiaSDatasetLoader(
    directory: str | None = None,
    max_concurrency: int = 16,
    include_taxonomy: bool = False,
    connector_config: SQLConnectorConfig | None = None,
)

name class-attribute

name: str = 'ambrosia-s'

splits class-attribute

splits: list[str] = ['test', 'few_shot_examples']

installation class-attribute

installation: BenchmarkInstallation = BenchmarkInstallation(
    name=name,
    required_paths=(
        "ambrosia/ambrosia.csv",
        "db_list.txt",
        "ambrosia_test_processed.json",
        "ambrosia_few_shot_examples_processed.json",
    ),
    fetch=_fetch_ambrosia,
)

default_metrics class-attribute

default_metrics: list[str] = [
    "simple_ex",
    "executable",
    "gold_executable",
    "gold_result_not_empty",
    "pred_success",
    "ambig_point_stats",
    "gold_ambig_point_stats",
    "found_one",
]

directory instance-attribute

directory = str(
    self.installation.directory
    if directory is None
    else directory
)

max_concurrency instance-attribute

max_concurrency = max_concurrency

include_taxonomy instance-attribute

include_taxonomy = include_taxonomy

connector_config instance-attribute

connector_config = (
    SQLConnectorConfig(
        schema_cache_mode="read_write",
        query_timeout_seconds=300,
    )
    if connector_config is None
    else connector_config
)

db_list instance-attribute

db_list: list[str] = [
    (line.strip())
    for line in (f.readlines())
    if line.strip()
]

get_databases

get_databases(split: str) -> list[str]

get_tasks_async async

get_tasks_async(
    split: str, databases: list[str] | None = None
) -> list[AmbigNL2QTask]

Load tasks from JSON and transform to AmbigNL2QTask schema.

Parameters:

Name Type Description Default
split str

Dataset split to load

required
databases list[str] | None

Optional list of databases to filter

None

Returns:

Type Description
list[AmbigNL2QTask]

List of validated AmbigNL2QTask instances

get_db_connectors_async async

get_db_connectors_async(
    split: str, databases: list[str] | None = None
) -> dict[str, SQLConnector]

get_split_async async

get_split_async(
    split: str,
    databases: list[str] | None = None,
    subsample_size: int | None = None,
    qids: list[str] | None = None,
) -> NL2QDataset

CypherBenchDatasetLoader

CypherBenchDatasetLoader(
    directory: str | None = None,
    neo4j_host: str = "localhost",
    neo4j_user: str = "neo4j",
    neo4j_password: str = "cypherbench",
    graph_ports: Mapping[str, int] | None = None,
    connector_config: Neo4jConnectorConfig | None = None,
)

Loader for CypherBench (text-to-Cypher over Neo4j property graphs).

Initializes the CypherBench dataset loader.

Parameters:

Name Type Description Default
directory str | None

Path containing train.json and test.json.

None
neo4j_host str

Bolt host for deployed graphs.

'localhost'
neo4j_user str

Neo4j username.

'neo4j'
neo4j_password str

Neo4j password.

'cypherbench'
graph_ports Mapping[str, int] | None

Optional overrides for graph name -> host Bolt port.

None

name class-attribute

name: str = 'cypherbench'

splits class-attribute

splits: list[str] = ['test', 'train']

installation class-attribute

installation: BenchmarkInstallation = (
    CYPHERBENCH_INSTALLATION
)

runtime class-attribute

runtime: BenchmarkRuntime = CYPHERBENCH_RUNTIME

default_metrics class-attribute

default_metrics: list[str] = [
    "cypherbench_ex",
    "simple_ex",
    "psjs",
    "executable",
    "gold_executable",
    "gold_result_not_empty",
    "pred_success",
]

directory instance-attribute

directory = str(
    self.installation.directory
    if directory is None
    else directory
)

neo4j_host instance-attribute

neo4j_host = neo4j_host

neo4j_user instance-attribute

neo4j_user = neo4j_user

neo4j_password instance-attribute

neo4j_password = neo4j_password

connector_config instance-attribute

connector_config = (
    Neo4jConnectorConfig(
        schema_cache_mode="read_write",
        query_timeout_seconds=120,
        graph_schema_introspection_mode="full_scan",
    )
    if connector_config is None
    else connector_config
)

get_databases

get_databases(split: str) -> list[str]

get_tasks_async async

get_tasks_async(
    split: str, databases: list[str] | None = None
) -> list[SimpleNL2QTask]

get_db_connectors_async async

get_db_connectors_async(
    split: str, databases: list[str] | None = None
) -> dict[str, Neo4jConnector]

get_split_async async

get_split_async(
    split: str,
    databases: list[str] | None = None,
    subsample_size: int | None = None,
    qids: list[str] | None = None,
) -> NL2QDataset

Installation and runtime requirements

preflight_benchmark checks local installation and managed runtime readiness. It does not download data or start database services.

preflight_benchmark async

preflight_benchmark(
    name: str,
    split: str,
    databases: list[str] | None = None,
) -> None

Validate local data and managed runtime readiness for a benchmark.

BenchmarkInstallation dataclass

BenchmarkInstallation(
    name: str,
    required_paths: tuple[str, ...],
    fetch: FetchFunction | None = None,
)

Local installation requirements for one benchmark.

name instance-attribute

name: str

required_paths instance-attribute

required_paths: tuple[str, ...]

fetch class-attribute instance-attribute

fetch: FetchFunction | None = None

directory property

directory: Path

missing_paths property

missing_paths: tuple[str, ...]

is_installed property

is_installed: bool

require

require() -> None

Raise an actionable error when the benchmark is not installed.

install async

install(
    *,
    force: bool = False,
    progress: ProgressCallback | None = None,
) -> Path

Download, verify, and atomically install the benchmark.

BenchmarkInstallationError

Bases: RuntimeError

Raised when benchmark data cannot be installed.

BenchmarkRuntime dataclass

BenchmarkRuntime(
    start_action: RuntimeAction,
    stop_action: RuntimeAction,
    ready_action: RuntimeCheck,
    splits: tuple[str, ...] = (),
    default_split: str | None = None,
    supports_database_selection: bool = False,
    endpoint_resolver: RuntimeEndpointResolver
    | None = None,
    authentication: str | None = None,
)

Start and stop operations for a managed benchmark database runtime.

start_action instance-attribute

start_action: RuntimeAction

stop_action instance-attribute

stop_action: RuntimeAction

ready_action instance-attribute

ready_action: RuntimeCheck

splits class-attribute instance-attribute

splits: tuple[str, ...] = ()

default_split class-attribute instance-attribute

default_split: str | None = None

supports_database_selection class-attribute instance-attribute

supports_database_selection: bool = False

endpoint_resolver class-attribute instance-attribute

endpoint_resolver: RuntimeEndpointResolver | None = None

authentication class-attribute instance-attribute

authentication: str | None = None

resolve_split

resolve_split(split: str | None) -> str | None

Validate and resolve an optional runtime split.

endpoints

endpoints(
    split: str | None, databases: list[str] | None = None
) -> Mapping[str, str]

Return connection URLs for the selected managed databases.

start async

start(
    split: str | None,
    progress: ProgressCallback,
    databases: list[str] | None = None,
) -> None

stop async

stop(
    split: str | None,
    progress: ProgressCallback,
    databases: list[str] | None = None,
) -> None

require_ready async

require_ready(
    benchmark: str,
    split: str | None,
    databases: list[str] | None = None,
) -> None

Raise an actionable error when the requested runtime is not ready.

BenchmarkRuntimeError

Bases: RuntimeError

Raised when a managed benchmark runtime operation fails.

Installation and runtime callbacks

Custom installations supply a fetch callback. Custom runtimes supply start, stop, and readiness callbacks. Progress callbacks receive status messages.

ProgressCallback module-attribute

ProgressCallback = Callable[[str], None]

FetchFunction module-attribute

FetchFunction = Callable[
    [Path, ProgressCallback], Awaitable[None]
]

RuntimeAction module-attribute

RuntimeAction = Callable[
    [str | None, list[str] | None, ProgressCallback],
    Awaitable[None],
]

RuntimeCheck module-attribute

RuntimeCheck = Callable[
    [str | None, list[str] | None], Awaitable[bool]
]

RuntimeEndpointResolver module-attribute

RuntimeEndpointResolver = Callable[
    [str | None, list[str] | None], Mapping[str, str]
]

ReadinessCheck module-attribute

ReadinessCheck = Callable[[], Awaitable[bool]]