94 lines
3.1 KiB
Python
94 lines
3.1 KiB
Python
"""Dataset registry: decorator registration + name lookup with suggestions.
|
|
|
|
Registration happens at import time ("import = register"). The registry maps
|
|
name -> DatasetProvider; a Dataset is only *materialized* on first use.
|
|
"""
|
|
|
|
import difflib
|
|
from typing import Callable, Dict, List, Optional, Union
|
|
|
|
from .spec import DatasetSpec, FieldSpec
|
|
|
|
# A provider factory: called once, returns how records become Samples.
|
|
ProviderFactory = Callable[[], Union[FieldSpec, Callable, None]]
|
|
|
|
|
|
class DatasetProvider:
|
|
"""A registered dataset: metadata + the recipe for converting records."""
|
|
|
|
def __init__(self, spec: DatasetSpec, factory: Optional[ProviderFactory] = None):
|
|
self.spec = spec
|
|
self._factory = factory
|
|
self._record_fn = None
|
|
self._resolved = False
|
|
|
|
def resolve_record_fn(self) -> Callable:
|
|
"""Resolve the record->Sample converter, lazily and exactly once."""
|
|
if not self._resolved:
|
|
result = self._factory() if self._factory else None
|
|
if isinstance(result, FieldSpec):
|
|
from .loader import field_spec_to_record_fn
|
|
|
|
self._record_fn = field_spec_to_record_fn(result)
|
|
elif callable(result):
|
|
self._record_fn = result
|
|
elif result is None:
|
|
self._record_fn = field_spec_to_record_fn(FieldSpec())
|
|
else:
|
|
raise TypeError(f'{self.spec.name}: factory must return FieldSpec or callable, got {type(result)}')
|
|
self._resolved = True
|
|
return self._record_fn
|
|
|
|
|
|
class Registry:
|
|
"""Minimal dict-like registry with duplicate protection and suggestions."""
|
|
|
|
def __init__(self, kind: str):
|
|
self.kind = kind
|
|
self._items: Dict[str, DatasetProvider] = {}
|
|
|
|
def register(self, name: str, item: DatasetProvider) -> DatasetProvider:
|
|
if name in self._items:
|
|
raise ValueError(f'{self.kind} {name!r} is already registered')
|
|
self._items[name] = item
|
|
return item
|
|
|
|
def get(self, name: str) -> DatasetProvider:
|
|
if name not in self._items:
|
|
suggestions = difflib.get_close_matches(name, self._items.keys(), n=3)
|
|
hint = f" Did you mean: {', '.join(suggestions)}?" if suggestions else ''
|
|
raise KeyError(f'unknown {self.kind} {name!r}.{hint}')
|
|
return self._items[name]
|
|
|
|
def names(self) -> List[str]:
|
|
return sorted(self._items)
|
|
|
|
def __contains__(self, name: str) -> bool:
|
|
return name in self._items
|
|
|
|
def __len__(self) -> int:
|
|
return len(self._items)
|
|
|
|
|
|
DATASET_REGISTRY = Registry('dataset')
|
|
|
|
|
|
def register_dataset(spec: DatasetSpec):
|
|
"""Decorator: register a dataset plugin.
|
|
|
|
Usage:
|
|
@register_dataset(DatasetSpec(name='gsm8k', source=...))
|
|
def gsm8k():
|
|
return lambda record: Sample(...) # or FieldSpec(...), or None
|
|
"""
|
|
|
|
def decorator(factory: ProviderFactory) -> ProviderFactory:
|
|
DATASET_REGISTRY.register(spec.name, DatasetProvider(spec, factory))
|
|
return factory
|
|
|
|
return decorator
|
|
|
|
|
|
def get_dataset_provider(name: str) -> DatasetProvider:
|
|
return DATASET_REGISTRY.get(name)
|