126 lines
5.0 KiB
Python
126 lines
5.0 KiB
Python
"""The evaluation runner: Dataset x predictions -> EvalReport.
|
|
|
|
Pure orchestration, no I/O hidden inside: predictions arrive as a list
|
|
(loaded from a jsonl of model outputs, a Session store, or built inline),
|
|
results aggregate into an EvalReport that visualizers consume.
|
|
|
|
Judge wiring: pass judge=<callable(messages)->str> once a ModelAdapter
|
|
exists; llm_judge recipes work immediately after that, no recipe change.
|
|
"""
|
|
|
|
import traceback
|
|
from typing import Callable, Dict, Iterable, List, Optional, Sequence, Union
|
|
|
|
from ..data.dataset import Dataset
|
|
from ..data.sample import Sample
|
|
from .aggregator import mean as _mean_agg
|
|
from .recipe import EvalRecipe
|
|
from .record import EvalReport, SampleResult
|
|
from .scorer import ScoreContext
|
|
|
|
|
|
def evaluate(
|
|
dataset: Union[Dataset, List[Sample]],
|
|
predictions: Sequence[Union[str, Dict]],
|
|
recipe: Optional[EvalRecipe] = None,
|
|
*,
|
|
model: str = '',
|
|
judge: Optional[Callable] = None,
|
|
extra_metadata: Optional[Dict] = None,
|
|
) -> EvalReport:
|
|
"""Score a dataset against raw predictions.
|
|
|
|
dataset: a Dataset or a plain list of Samples (views/slices).
|
|
predictions: str per sample (raw model output) or dicts with
|
|
{'raw': str, 'group_key': ..., 'metadata': {...}} overrides.
|
|
"""
|
|
samples: List[Sample] = list(dataset)
|
|
spec = getattr(dataset, 'spec', None)
|
|
ds_name = spec.name if spec is not None else samples[0].metadata.get('dataset', 'adhoc') if samples else 'adhoc'
|
|
ds_subset = spec.subset if spec is not None else ''
|
|
if len(predictions) != len(samples):
|
|
raise ValueError(f'{len(predictions)} predictions for {len(samples)} samples')
|
|
|
|
if recipe is None:
|
|
from .recipe import get_eval
|
|
|
|
recipe = get_eval(ds_name)
|
|
extractor = recipe.resolve_extract()
|
|
scorers = recipe.resolve_scorers()
|
|
aggregators = recipe.resolve_aggregators()
|
|
ctx = ScoreContext(judge=judge, params={})
|
|
|
|
results: List[SampleResult] = []
|
|
for sample, pred in zip(samples, predictions):
|
|
raw = pred if isinstance(pred, str) else str(pred.get('raw', ''))
|
|
override = {} if isinstance(pred, str) else pred
|
|
result = SampleResult(
|
|
sample_id=sample.id,
|
|
dataset=ds_name,
|
|
subset=ds_subset,
|
|
task_type=sample.task_type,
|
|
raw_prediction=raw,
|
|
target=sample.target,
|
|
group_key=override.get('group_key')
|
|
or sample.metadata.get('group_key')
|
|
or (sample.metadata.get('task_id') or sample.metadata.get('id') or ''),
|
|
metadata={k: v for k, v in (sample.metadata or {}).items()
|
|
if k in ('category', 'subject', 'test_category', 'bin', 'difficulty')},
|
|
)
|
|
if isinstance(pred, dict) and pred.get('metadata'):
|
|
result.metadata.update(pred['metadata'])
|
|
try:
|
|
value, ok, note = extractor(raw, sample)
|
|
result.extracted_prediction = value
|
|
result.extraction_ok = ok
|
|
result.extraction_note = note
|
|
if not ok:
|
|
result.extraction_note = note or 'extractor returned not-ok'
|
|
for metric, scorer in scorers.items():
|
|
try:
|
|
scores, details = scorer(value if ok else '', sample.target, sample, ctx)
|
|
result.scores.update(scores)
|
|
result.score_details.update(details)
|
|
except Exception as e: # one metric failing must not kill the run
|
|
result.scores[metric] = 0.0
|
|
result.score_details[metric] = {'error': f'{type(e).__name__}: {e}'}
|
|
except Exception as e:
|
|
result.error = f'{type(e).__name__}: {e}\n{traceback.format_exc(limit=2)}'
|
|
results.append(result)
|
|
|
|
report = EvalReport(
|
|
dataset=ds_name,
|
|
recipe=recipe.name or dataset.spec.name,
|
|
model=model,
|
|
num_samples=len(results),
|
|
num_failed_extractions=sum(1 for r in results if not r.extraction_ok),
|
|
samples=results,
|
|
)
|
|
_aggregate_into(report, results, recipe, aggregators)
|
|
if extra_metadata:
|
|
report.metric_groups['run_info'] = {k: v for k, v in extra_metadata.items()
|
|
if isinstance(v, (int, float, str))}
|
|
return report
|
|
|
|
|
|
def _aggregate_into(report: EvalReport, results, recipe: EvalRecipe, aggregators) -> None:
|
|
for metric in recipe.scorers:
|
|
agg = aggregators.get(metric)
|
|
if agg is None:
|
|
agg = _mean_agg
|
|
try:
|
|
out = agg(results, metric)
|
|
except Exception as e:
|
|
report.metric_groups[f'agg_error_{metric}'] = {'error': str(e)[:200]}
|
|
continue
|
|
if isinstance(out, dict):
|
|
report.metric_groups[metric] = out
|
|
vals = [v for v in out.values() if isinstance(v, (int, float))]
|
|
if vals:
|
|
report.metrics[metric] = sum(vals) / len(vals)
|
|
else:
|
|
report.metrics[metric] = float(out)
|
|
report.metrics['extraction_failure_rate'] = (
|
|
report.num_failed_extractions / report.num_samples if report.num_samples else 0.0
|
|
)
|