"""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 `_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)