- perf_stats aggregator lives in eval/, not model/: the import failed silently and EVERY perf column was empty (not just ttft). Now warns on stderr instead of swallowing. - repeats > 1 get their own checkpoint key (:rep2, :rep3, ...): repeat 2 previously restored repeat 1's predictions and finished instantly with identical scores. rep1 keeps the legacy key (existing checkpoints still resume). - repeats summary: report the MEAN score and aggregate time/tokens over ALL runs (was: last run only). - README: six-benchmark command as the primary example. Co-Authored-By: Claude <noreply@anthropic.com>
101 lines
3.5 KiB
Python
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)
|