diff --git a/.github/workflows/eval.yml b/.github/workflows/eval.yml index f1c031ca6..a2444340a 100644 --- a/.github/workflows/eval.yml +++ b/.github/workflows/eval.yml @@ -144,6 +144,26 @@ jobs: print(f'::warning title=Fidelity recall::{result.failed_cases} case(s) fell below 0.9 recall: {result.errors[:3]}') " + # Real-dataset recall on the prose path (HotpotQA): does the ground-truth + # answer survive compressing the supporting context? Uses the production + # routing path, so prose flows through Kompress (ModernBERT) — allowed here + # because the weekly job installs [all]. Non-blocking and defensive: a + # dataset download or model failure warns rather than fails the job. + - name: Dataset recall report — HotpotQA (model-allowed, non-blocking) + run: | + python -c " + try: + from headroom.evals.datasets import load_hotpotqa + from headroom.evals.runners.compression_only import CompressionOnlyRunner + suite = load_hotpotqa(n=50) + result = CompressionOnlyRunner().evaluate_dataset_recall(suite) + print(f'HotpotQA answer recall: {result.passed_cases}/{result.total_cases} probeable cases >=0.9, avg compression {result.avg_compression_ratio:.1%}') + if result.failed_cases: + print(f'::warning title=Dataset recall::{result.failed_cases} HotpotQA case(s) lost the answer under compression') + except Exception as e: + print(f'::warning title=Dataset recall::skipped (dataset/model unavailable): {e}') + " || true + - name: Upload results if: always() uses: actions/upload-artifact@v7 diff --git a/headroom/evals/runners/compression_only.py b/headroom/evals/runners/compression_only.py index 4ef71dbc0..7bf3e2fe2 100644 --- a/headroom/evals/runners/compression_only.py +++ b/headroom/evals/runners/compression_only.py @@ -14,7 +14,10 @@ import json import logging import time from dataclasses import dataclass, field -from typing import Any +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from headroom.evals.core import EvalSuite logger = logging.getLogger(__name__) @@ -242,6 +245,91 @@ class CompressionOnlyRunner: errors=errors, ) + def evaluate_dataset_recall( + self, + suite: EvalSuite, + recall_threshold: float = 0.9, + min_answer_chars: int = 4, + ) -> CompressionOnlyResult: + """Compress each QA case's context and check its answer survives. + + The probe is the case's ``ground_truth`` answer. A case only counts when + the answer literally appears in the original context (otherwise survival + is not measurable); trivial answers (too short, or yes/no) are skipped. + Compression uses the production routing path, so prose flows through + Kompress (ModernBERT) — intended for the model-allowed weekly job, not + the hermetic per-PR gate. + + Args: + suite: An EvalSuite of QA cases (e.g. from ``load_hotpotqa``). + recall_threshold: Minimum answer recall for a case to pass. + min_answer_chars: Answers shorter than this are skipped as un-probeable. + """ + from headroom.evals.metrics import compute_information_recall + from headroom.transforms.content_router import ContentRouter + + trivial = {"yes", "no", "true", "false"} + start_time = time.time() + router = ContentRouter() + passed = 0 + failed = 0 + total_original = 0 + total_compressed = 0 + details: list[dict[str, Any]] = [] + errors: list[str] = [] + + for case in suite.cases: + answer = (case.ground_truth or "").strip() + if len(answer) < min_answer_chars or answer.lower() in trivial: + continue # un-probeable: survival of this answer carries no signal + if answer.lower() not in case.context.lower(): + continue # answer not literally in context; nothing to measure + + original_tokens = self._estimate_tokens(case.context) + try: + compressed = router.compress(case.context).compressed + compressed_tokens = self._estimate_tokens(compressed) + recall = compute_information_recall(case.context, compressed, [answer])["recall"] + is_pass = recall >= recall_threshold + + total_original += original_tokens + total_compressed += compressed_tokens + passed += is_pass + failed += not is_pass + details.append( + { + "id": case.id, + "passed": is_pass, + "recall": recall, + "answer": answer, + "compression_ratio": 1 - (compressed_tokens / original_tokens) + if original_tokens > 0 + else 0, + } + ) + except Exception as e: + failed += 1 + errors.append(f"Dataset recall error for {case.id}: {e}") + details.append({"id": case.id, "passed": False, "error": str(e)}) + + total_cases = passed + failed + ratios = [d["compression_ratio"] for d in details if "compression_ratio" in d] + + return CompressionOnlyResult( + benchmark=f"dataset_recall:{suite.name}", + total_cases=total_cases, + passed_cases=passed, + failed_cases=failed, + accuracy_rate=passed / total_cases if total_cases > 0 else 0.0, + avg_compression_ratio=sum(ratios) / len(ratios) if ratios else 0.0, + total_original_tokens=total_original, + total_compressed_tokens=total_compressed, + total_tokens_saved=total_original - total_compressed, + duration_seconds=time.time() - start_time, + details=details, + errors=errors, + ) + def generate_ccr_test_cases(self, n: int = 50) -> list[dict[str, Any]]: """Generate synthetic test cases for CCR needle-retention testing. diff --git a/tests/test_dataset_recall_runner.py b/tests/test_dataset_recall_runner.py new file mode 100644 index 000000000..effa39ca4 --- /dev/null +++ b/tests/test_dataset_recall_runner.py @@ -0,0 +1,67 @@ +"""Hermetic unit tests for CompressionOnlyRunner.evaluate_dataset_recall. + +Exercises the dataset-recall plumbing with synthetic JSON-array contexts (which +route through SmartCrusher / Rust — no model, no network) so it runs in the +standard [dev] shard. The weekly job drives the same method with real prose +datasets (HotpotQA), which is intentionally not exercised here. +""" + +from __future__ import annotations + +import json + +from headroom.evals.core import EvalCase, EvalSuite +from headroom.evals.runners.compression_only import CompressionOnlyRunner + + +def _array_context_with(answer: str) -> str: + """A JSON-array tool output whose error row embeds ``answer`` (a kept row).""" + rows = [{"seq": i, "level": "INFO", "status": "ok", "msg": f"heartbeat {i}"} for i in range(30)] + rows[14] = {"seq": 14, "level": "ERROR", "status": "failed", "msg": answer} + return json.dumps(rows) + + +def _suite() -> EvalSuite: + answer = "PaymentService NullPointerException at charge line 88" + return EvalSuite( + name="synthetic", + cases=[ + # Probeable: answer is in an error row -> retained -> recall 1.0. + EvalCase( + id="probeable", + context=_array_context_with(answer), + query="what failed?", + ground_truth=answer, + ), + # Skipped: trivial yes/no answer. + EvalCase( + id="trivial", + context=_array_context_with(answer), + query="did it fail?", + ground_truth="yes", + ), + # Skipped: answer not present in the context at all. + EvalCase( + id="absent", + context=_array_context_with(answer), + query="?", + ground_truth="totally-absent-token-xyz", + ), + ], + ) + + +def test_dataset_recall_counts_only_probeable_cases() -> None: + result = CompressionOnlyRunner().evaluate_dataset_recall(_suite()) + # Only the "probeable" case is measurable; trivial + absent are skipped. + assert result.total_cases == 1 + assert result.passed_cases == 1 + assert result.accuracy_rate == 1.0 + assert result.benchmark == "dataset_recall:synthetic" + + +def test_dataset_recall_empty_suite_is_safe() -> None: + result = CompressionOnlyRunner().evaluate_dataset_recall(EvalSuite(name="empty", cases=[])) + assert result.total_cases == 0 + assert result.accuracy_rate == 0.0 + assert result.errors == []