101 lines
3.5 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:
provider = DatasetProvider(spec, factory)
# few-shot hook convention: a module-level `<name>_few_shot(split,
# subset, n) -> Optional[str]` next to the plugin is picked up here,
# so runner can inject official hand-written exemplars (e.g. bbh CoT).
hook = factory.__globals__.get(f'{spec.name}_few_shot')
if callable(hook):
provider.few_shot_hook = hook
DATASET_REGISTRY.register(spec.name, provider)
return factory
return decorator
def get_dataset_provider(name: str) -> DatasetProvider:
return DATASET_REGISTRY.get(name)