52 lines
1.6 KiB
Python
52 lines
1.6 KiB
Python
"""evalharness.data -- the data layer.
|
|
|
|
Usage:
|
|
from evalharness.data import get_dataset, list_datasets
|
|
|
|
ds = get_dataset('gsm8k') # lazy handle, nothing downloaded
|
|
for s in ds: # first use triggers materialize (download->convert->cache)
|
|
...
|
|
"""
|
|
|
|
import importlib
|
|
import pkgutil
|
|
from pathlib import Path
|
|
from typing import List
|
|
|
|
from .dataset import Dataset
|
|
from .registry import DATASET_REGISTRY, DatasetProvider, get_dataset_provider, register_dataset
|
|
from .sample import ChatMessage, Sample, SandboxSpec, ToolInfo
|
|
from .spec import DatasetSpec, FieldSpec
|
|
|
|
__all__ = [
|
|
'Dataset', 'DatasetSpec', 'FieldSpec', 'Sample', 'ChatMessage', 'SandboxSpec', 'ToolInfo',
|
|
'register_dataset', 'get_dataset', 'list_datasets', 'get_dataset_provider',
|
|
]
|
|
|
|
|
|
def _discover_builtin_datasets() -> None:
|
|
"""Import every plugin module under ./datasets (import = register)."""
|
|
pkg_dir = Path(__file__).parent / 'datasets'
|
|
if not pkg_dir.exists():
|
|
return
|
|
for info in pkgutil.iter_modules([str(pkg_dir)]):
|
|
importlib.import_module(f'{__name__}.datasets.{info.name}')
|
|
|
|
|
|
_discover_builtin_datasets()
|
|
|
|
|
|
def get_dataset(name: str, **overrides) -> Dataset:
|
|
"""Return a lazy Dataset handle by name. No download happens here."""
|
|
provider = get_dataset_provider(name)
|
|
spec = provider.spec
|
|
if overrides:
|
|
import dataclasses
|
|
|
|
spec = dataclasses.replace(spec, **overrides)
|
|
return Dataset(spec, provider.resolve_record_fn())
|
|
|
|
|
|
def list_datasets() -> List[DatasetSpec]:
|
|
return [DATASET_REGISTRY.get(n).spec for n in DATASET_REGISTRY.names()]
|