CoolFace
Apppublic

DGXAI/driftcall

sourceHugging Faceapache-2.0updated 5mo agoView on Hugging Face
0likes
step_19_eval_final.py233 linesDownload Raw Back to cells
1"""Cell 19 — Final evaluation harness (post-training LoRA).2 3Implements ``docs/modules/evaluation.md`` §2.1, §3.1, §3.3 (paired-difference),4§3.5 (drift-detection latency aggregation), §3.8, §5 ``EpisodeSetLeakError``.5 6Hard rules (evaluation.md §3.1, §6.1, §6.3):7- Same 50 episodes as baseline (paired); ``EpisodeSetLeakError`` raised on8  mismatch.9- Bootstrap CI seed for paired-difference is ``20260428`` (evaluation.md §2.4).10- Wall-clock budget 20 minutes — same ceiling as baseline.11- No LLM-as-judge; static AST scan via ``_NO_LLM_JUDGE_FORBIDDEN_IMPORTS``.12 13Heavy imports (``torch``) are deferred so this module imports cleanly on14CPU-only CI. The training-eval delegate is injected (see step_18).15"""16 17from __future__ import annotations18 19import time20from dataclasses import replace21from pathlib import Path22from typing import TYPE_CHECKING, Any23 24from cells.step_18_eval_baseline import (25    BUDGET_RUN_EVAL_SECONDS,26    DEFAULT_N_BOOT,27    DEFAULT_PAIRED_BOOTSTRAP_SEED,28    DriftDetectionLatency,29    EvalBudgetExceededError,30    EvalReport,31    EvaluationError,32    PerLanguageReport,33    TrainingEvalCallable,34    _check_catalogue_hashes,35    _episode_ids_from_breakdown,36    _validate_briefs_first_50,37    run_eval,38)39 40if TYPE_CHECKING:  # pragma: no cover - typing only41    from collections.abc import Callable, Sequence42 43 44__all__ = [45    "BUDGET_RUN_EVAL_SECONDS",46    "DEFAULT_PAIRED_BOOTSTRAP_SEED",47    "DriftDetectionLatency",48    "EpisodeSetLeakError",49    "EvalBudgetExceededError",50    "EvalReport",51    "PerLanguageReport",52    "assert_paired_episode_sets",53    "eval_final",54    "paired_difference_ci",55]56 57 58# ---------------------------------------------------------------------------59# Errors — evaluation.md §560# ---------------------------------------------------------------------------61 62 63class EpisodeSetLeakError(EvaluationError):64    """Baseline ``episode_ids`` ≠ final ``episode_ids`` — paired-comparison invariant violated."""65 66 67# ---------------------------------------------------------------------------68# Paired-difference CI — evaluation.md §2.469# ---------------------------------------------------------------------------70 71 72def paired_difference_ci(73    baseline_samples: tuple[float, ...],74    final_samples: tuple[float, ...],75    n_boot: int = DEFAULT_N_BOOT,76    rng_seed: int = DEFAULT_PAIRED_BOOTSTRAP_SEED,77) -> tuple[float, float, float]:78    """Bootstrap 95% CI on ``mean(final - baseline)`` — index-paired.79 80    evaluation.md §2.4: lengths must match (raises ``EpisodeSetLeakError``).81    Edge cases mirror :func:`bootstrap_ci`: empty → all-NaN; single → triple.82    """83    if len(baseline_samples) != len(final_samples):84        raise EpisodeSetLeakError(85            f"paired-comparison invariant: len(baseline)={len(baseline_samples)} "86            f"!= len(final)={len(final_samples)}",87        )88    n = len(baseline_samples)89    if n == 0:90        nan = float("nan")91        return nan, nan, nan92    diffs = tuple(f - b for b, f in zip(baseline_samples, final_samples, strict=True))93    mean = sum(diffs) / n94    if n == 1:95        return mean, mean, mean96    if all(d == diffs[0] for d in diffs):97        return mean, mean, mean98 99    import numpy as np100 101    rng = np.random.default_rng(rng_seed)102    arr = np.asarray(diffs, dtype=np.float64)103    idx = rng.integers(0, n, size=(n_boot, n))104    means = arr[idx].mean(axis=1)105    lo = float(np.percentile(means, 2.5))106    hi = float(np.percentile(means, 97.5))107    return float(mean), lo, hi108 109 110# ---------------------------------------------------------------------------111# Episode-set leak guard — evaluation.md §3.1112# ---------------------------------------------------------------------------113 114 115def assert_paired_episode_sets(baseline: EvalReport, final: EvalReport) -> None:116    """Raise ``EpisodeSetLeakError`` iff ``episode_ids`` tuples differ."""117    base_ids = _episode_ids_from_breakdown(baseline)118    final_ids = _episode_ids_from_breakdown(final)119    if base_ids != final_ids:120        raise EpisodeSetLeakError(121            "paired-comparison invariant violated — baseline.episode_ids != final.episode_ids; "122            "operator must re-run baseline against the current val split.",123        )124 125 126# ---------------------------------------------------------------------------127# Drift-detection-latency point extraction — evaluation.md §3.5128# ---------------------------------------------------------------------------129 130 131def _final_latency_point(report: EvalReport) -> tuple[float, float]:132    """Return ``(p50, p95)`` from the report's drift-detection latency."""133    lat = report.drift_detection_latency134    # Stage-3 takes precedence (final stage); falls back to stage-2 if Stage-3 NaN.135    p50 = lat.stage3_median136    p95 = lat.stage3_p95137    return float(p50), float(p95)138 139 140# ---------------------------------------------------------------------------141# Final-eval entry point — evaluation.md §2.2 ``eval_final.py``142# ---------------------------------------------------------------------------143 144 145def eval_final(146    checkpoint: Path,147    episodes: int = 50,148    *,149    baseline: EvalReport,150    training_eval: TrainingEvalCallable,151    briefs: Sequence[Any],152    catalogue_hashes: dict[str, str] | None = None,153    budget_seconds: int = BUDGET_RUN_EVAL_SECONDS,154    monotonic: Callable[[], float] | None = None,155) -> EvalReport:156    """Run the trained LoRA against the SAME 50 paired episodes used by baseline.157 158    evaluation.md §2.1, §3.1: rejects mismatched checkpoints; verifies catalogue159    hashes; computes paired-difference CIs and stores them under160    ``EvalReport.breakdown['paired_ci']``.161    """162    if not isinstance(checkpoint, Path):163        raise EvaluationError(164            f"checkpoint must be pathlib.Path; got {type(checkpoint).__name__}",165        )166    if episodes != 50:167        raise EvaluationError(168            f"eval_final expects episodes=50 (paired contract); got {episodes}",169        )170 171    selected = _validate_briefs_first_50(briefs)172    if catalogue_hashes is not None:173        _check_catalogue_hashes(selected, catalogue_hashes)174 175    # Pre-flight: episode_ids match baseline before launching rollout.176    expected_ids = tuple(row.episode_id for row in selected)177    base_ids = _episode_ids_from_breakdown(baseline)178    if base_ids and base_ids != expected_ids:179        raise EpisodeSetLeakError(180            "paired-comparison invariant violated at entry — baseline.episode_ids "181            "do not match val/briefs.jsonl[0:50]; re-run baseline first.",182        )183 184    clock = monotonic if monotonic is not None else time.monotonic185    started = clock()186 187    final_report = run_eval(188        checkpoint,189        episodes,190        training_eval=training_eval,191        briefs=briefs,192        catalogue_hashes=catalogue_hashes,193        budget_seconds=budget_seconds,194        monotonic=clock,195    )196    elapsed = clock() - started197    if elapsed > budget_seconds:198        raise EvalBudgetExceededError(199            f"eval_final wall-clock {elapsed:.1f}s exceeded {budget_seconds}s",200        )201 202    assert_paired_episode_sets(baseline, final_report)203 204    # Compute paired-difference CIs (evaluation.md §3.3).205    paired_ci = _build_paired_ci_block(baseline, final_report)206    breakdown = dict(final_report.breakdown)207    breakdown["paired_ci"] = paired_ci208    return replace(final_report, breakdown=breakdown)209 210 211def _build_paired_ci_block(212    baseline: EvalReport,213    final: EvalReport,214) -> dict[str, tuple[float, float, float]]:215    """Construct the ``breakdown['paired_ci']`` block for the blog narrative."""216    out: dict[str, tuple[float, float, float]] = {}217    base_samples: dict[str, tuple[float, ...]] = baseline.breakdown.get("samples", {})218    final_samples: dict[str, tuple[float, ...]] = final.breakdown.get("samples", {})219    for key in ("reward", "r1", "r2", "r3", "r4", "r5"):220        if key in base_samples and key in final_samples:221            out[key] = paired_difference_ci(222                tuple(base_samples[key]),223                tuple(final_samples[key]),224            )225 226    # Drift-latency delta — final p50 minus baseline p50 (lower is better).227    base_p50, _ = _final_latency_point(baseline)228    final_p50, _ = _final_latency_point(final)229    if not (base_p50 != base_p50 or final_p50 != final_p50):  # neither NaN230        delta = final_p50 - base_p50231        out["drift_latency_p50"] = (delta, delta, delta)232    return out233