CoolFace
Apppublic

DGXAI/driftcall

sourceHugging Faceapache-2.0updated 5mo agoView on Hugging Face
0likes
step_22_summary.py181 linesDownload Raw Back to cells
1"""Cell 22 — Markdown summary table (baseline → final → Δ).2 3Renders the markdown table that drives DESIGN.md §15 pitch 2:00–2:404"before/after" slide. Per evaluation.md §3.3, §3.4, §3.5:5 6- Per-reward baseline mean + 95% CI → final mean + 95% CI → paired Δ.7- Per-language breakdown table (n_episodes, reward_mean, R1..R5 means).8- Drift-detection latency before/after row.9 10Hard rules:11- No LLM-as-judge; static AST scan via ``_NO_LLM_JUDGE_FORBIDDEN_IMPORTS``.12- Every numeric cell rounds to 3 decimals.13"""14 15from __future__ import annotations16 17import math18from typing import TYPE_CHECKING19 20if TYPE_CHECKING:  # pragma: no cover - typing only21    from cells.step_18_eval_baseline import EvalReport, PerLanguageReport22 23 24__all__ = [25    "format_per_language_table",26    "format_per_reward_table",27    "print_summary_table",28]29 30 31_NO_LLM_JUDGE_FORBIDDEN_IMPORTS: frozenset[str] = frozenset(32    {"openai", "anthropic", "vertexai", "google.generativeai", "cohere"},33)34 35_REWARD_KEYS: tuple[str, ...] = ("reward", "r1", "r2", "r3", "r4", "r5")36 37 38def _fmt_ci(triple: tuple[float, float, float]) -> str:39    mean, lo, hi = triple40    if math.isnan(mean):41        return "NaN"42    return f"{mean:.3f} [{lo:.3f}, {hi:.3f}]"43 44 45def _fmt_paired(triple: tuple[float, float, float] | None) -> str:46    if triple is None:47        return "—"48    mean, lo, hi = triple49    if math.isnan(mean):50        return "NaN"51    sign = "+" if mean >= 0 else ""52    return f"{sign}{mean:.3f} [{lo:.3f}, {hi:.3f}]"53 54 55def format_per_reward_table(baseline: EvalReport, final: EvalReport) -> str:56    """Markdown table: per-reward baseline mean+CI → final mean+CI → Δ with CI."""57    paired_block = final.breakdown.get("paired_ci", {})58    if not isinstance(paired_block, dict):59        paired_block = {}60 61    lines: list[str] = []62    lines.append("| Reward | Baseline mean [95% CI] | Final mean [95% CI] | Δ paired [95% CI] |")63    lines.append("|--------|------------------------|---------------------|-------------------|")64    for key in _REWARD_KEYS:65        base_ci = getattr(baseline, f"{key}_mean_ci")66        final_ci = getattr(final, f"{key}_mean_ci")67        paired = paired_block.get(key)68        lines.append(69            f"| {key.upper():6s} | {_fmt_ci(base_ci):22s} | "70            f"{_fmt_ci(final_ci):19s} | {_fmt_paired(paired):17s} |",71        )72    return "\n".join(lines)73 74 75def _fmt_lang_cell(value: float) -> str:76    if math.isnan(value):77        return "NaN"78    return f"{value:.3f}"79 80 81def _per_lang_lookup(report: EvalReport) -> dict[str, PerLanguageReport]:82    return {pl.language: pl for pl in report.per_language}83 84 85def format_per_language_table(baseline: EvalReport, final: EvalReport) -> str:86    """Markdown table: per-language reward_mean baseline → final."""87    base_lookup = _per_lang_lookup(baseline)88    final_lookup = _per_lang_lookup(final)89    languages = sorted(set(base_lookup) | set(final_lookup))90 91    lines: list[str] = []92    lines.append(93        "| Language | n_episodes | Baseline reward_mean | Final reward_mean | Δ reward_mean |",94    )95    lines.append(96        "|----------|------------|----------------------|-------------------|---------------|",97    )98    for lang in languages:99        b = base_lookup.get(lang)100        f = final_lookup.get(lang)101        n = max(b.n_episodes if b else 0, f.n_episodes if f else 0)102        b_mean = b.reward_mean if b else float("nan")103        f_mean = f.reward_mean if f else float("nan")104        if math.isnan(b_mean) or math.isnan(f_mean):105            delta_str = "—"106        else:107            delta = f_mean - b_mean108            sign = "+" if delta >= 0 else ""109            delta_str = f"{sign}{delta:.3f}"110        lines.append(111            f"| {lang:8s} | {n:10d} | {_fmt_lang_cell(b_mean):20s} | "112            f"{_fmt_lang_cell(f_mean):17s} | {delta_str:13s} |",113        )114    return "\n".join(lines)115 116 117def _fmt_latency(value: float) -> str:118    if math.isnan(value):119        return "NaN"120    return f"{value:.2f}"121 122 123def format_drift_latency_table(baseline: EvalReport, final: EvalReport) -> str:124    """Markdown table: drift-detection latency p50/p95 baseline vs final."""125    bl = baseline.drift_detection_latency126    fl = final.drift_detection_latency127    lines: list[str] = []128    lines.append("| Stage | Baseline p50 | Baseline p95 | Final p50 | Final p95 | Undetected |")129    lines.append("|-------|--------------|--------------|-----------|-----------|------------|")130    lines.append(131        f"| Stage 2 | {_fmt_latency(bl.stage2_median):12s} | "132        f"{_fmt_latency(bl.stage2_p95):12s} | "133        f"{_fmt_latency(fl.stage2_median):9s} | "134        f"{_fmt_latency(fl.stage2_p95):9s} | "135        f"{fl.undetected_count:10d} |",136    )137    lines.append(138        f"| Stage 3 | {_fmt_latency(bl.stage3_median):12s} | "139        f"{_fmt_latency(bl.stage3_p95):12s} | "140        f"{_fmt_latency(fl.stage3_median):9s} | "141        f"{_fmt_latency(fl.stage3_p95):9s} | "142        f"{bl.undetected_count:10d} |",143    )144    return "\n".join(lines)145 146 147def print_summary_table(baseline: EvalReport, final: EvalReport) -> str:148    """Top-level entry point — emit the full multi-section markdown summary."""149    sections: list[str] = []150    sections.append("# DriftCall — Baseline → Final summary")151    sections.append("")152    sections.append(f"**Baseline model:** `{baseline.model_path}`")153    sections.append(f"**Final model:** `{final.model_path}`")154    sections.append(f"**Episodes:** baseline {baseline.n_episodes}, final {final.n_episodes}")155    sections.append("")156    sections.append("## Per-reward (mean + 95% CI)")157    sections.append("")158    sections.append(format_per_reward_table(baseline, final))159    sections.append("")160    sections.append("## Per-language breakdown")161    sections.append("")162    sections.append(format_per_language_table(baseline, final))163    sections.append("")164    sections.append("## Drift-detection latency")165    sections.append("")166    sections.append(format_drift_latency_table(baseline, final))167    sections.append("")168 169    # Reward-hacking offenses summary (DESIGN.md §15 pitch).170    sections.append("## Reward-hacking offenses (final vs baseline)")171    sections.append("")172    sections.append("| Class | Baseline | Final |")173    sections.append("|-------|----------|-------|")174    keys = sorted(set(baseline.reward_hacking_offenses) | set(final.reward_hacking_offenses))175    for key in keys:176        b_count = baseline.reward_hacking_offenses.get(key, 0)177        f_count = final.reward_hacking_offenses.get(key, 0)178        sections.append(f"| {key:22s} | {b_count:8d} | {f_count:5d} |")179    sections.append("")180    return "\n".join(sections)181