fix(gsm8k): patch experience flywheel merge, provenance, and layering

- merge_compacted_runs merges key-by-key without re-expanding prior counts
- load default train_sample cases for live operation_class resolution
- hash full scout row evidence in source_report_hash and source_run_id
- replace scripts.gsm8k_frontier_report import with local _extract_category
This commit is contained in:
Shay 2026-06-17 21:25:36 -07:00
parent 0ca48cc9a3
commit 9e7432748d
2 changed files with 345 additions and 38 deletions

View file

@ -12,7 +12,7 @@ Trust boundary:
from __future__ import annotations from __future__ import annotations
import json import re
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from typing import Any, Literal from typing import Any, Literal
@ -20,13 +20,16 @@ from typing import Any, Literal
from formation.hashing import canonical_json, sha256_of from formation.hashing import canonical_json, sha256_of
from evals.gsm8k_math.practice.v1.runner import classify_operation from evals.gsm8k_math.practice.v1.runner import classify_operation
from evals.gsm8k_math.train_sample.v1.runner import _CASES_PATH, _load_cases
from evals.gsm8k_math.train_sample.v1.scout import ( from evals.gsm8k_math.train_sample.v1.scout import (
SealedAttemptScoutRow, SealedAttemptScoutRow,
build_scout_rows,
build_scout_summary, build_scout_summary,
classify_delta_kind, classify_delta_kind,
) )
from scripts.gsm8k_frontier_report import _classify_reason, _extract_category
_RECOGNIZED_NO_INJ = "candidate_graph: recognizer matched but produced no injection"
_CATEGORY_RE = re.compile(r"category=([a-zA-Z0-9_]+)")
_UNKNOWN_OPERATION_CLASS = "unknown"
SCHEMA_VERSION = 1 SCHEMA_VERSION = 1
ADR = "experience_flywheel_pr1" ADR = "experience_flywheel_pr1"
@ -230,7 +233,34 @@ def compute_dedupe_key(record: ExperienceRecord) -> str:
return sha256_of(payload) return sha256_of(payload)
def _extract_category(reason: str) -> str | None:
"""Pull (category=...) from recognized-no-injection refusal reasons."""
if _RECOGNIZED_NO_INJ not in reason:
return None
match = _CATEGORY_RE.search(reason)
return match.group(1) if match else None
def _scout_row_evidence(rows: list[dict[str, Any]] | None) -> list[dict[str, Any]]:
"""Compact per-case scout evidence for provenance hashing (not raw traces)."""
evidence: list[dict[str, Any]] = []
for row in rows or []:
evidence.append(
{
"case_id": row["case_id"],
"served_status": row["served_status"],
"aggressive_status": row["aggressive_status"],
"failure_family": row["failure_family"],
"trace_key": row["trace_key"],
"candidate_lift_family": row.get("candidate_lift_family"),
"first_failed_step": row.get("first_failed_step"),
}
)
return sorted(evidence, key=lambda item: item["case_id"])
def compute_run_id(scout_summary: dict[str, Any]) -> str: def compute_run_id(scout_summary: dict[str, Any]) -> str:
"""Identity of one scout run — includes per-case row evidence, not aggregates alone."""
payload = { payload = {
"schema_version": scout_summary.get("schema_version"), "schema_version": scout_summary.get("schema_version"),
"adr": scout_summary.get("adr"), "adr": scout_summary.get("adr"),
@ -239,15 +269,26 @@ def compute_run_id(scout_summary: dict[str, Any]) -> str:
"serving_counts": scout_summary.get("serving_counts"), "serving_counts": scout_summary.get("serving_counts"),
"sealed_counts": scout_summary.get("sealed_counts"), "sealed_counts": scout_summary.get("sealed_counts"),
"delta_counts": scout_summary.get("delta_counts"), "delta_counts": scout_summary.get("delta_counts"),
"rows": _scout_row_evidence(scout_summary.get("rows")),
} }
return sha256_of(payload) return sha256_of(payload)
def compute_report_hash(scout_summary: dict[str, Any]) -> str: def compute_report_hash(scout_summary: dict[str, Any]) -> str:
payload = {k: v for k, v in scout_summary.items() if k != "rows"} """Hash full load-bearing scout summary evidence, including compact row payloads."""
payload = dict(scout_summary)
if "rows" in payload:
payload["rows"] = _scout_row_evidence(payload["rows"])
return sha256_of(payload) return sha256_of(payload)
def _resolve_operation_class(raw_case: dict[str, Any]) -> str:
expr = raw_case.get("answer_expression", "")
if not expr:
return _UNKNOWN_OPERATION_CLASS
return classify_operation(expr)
def _arithmetic_chain_signature( def _arithmetic_chain_signature(
*, *,
delta_kind: str, delta_kind: str,
@ -572,7 +613,7 @@ def records_from_scout_rows(
out: list[ExperienceRecord] = [] out: list[ExperienceRecord] = []
for row in rows: for row in rows:
raw_case = (cases_by_id or {}).get(row.case_id, {}) raw_case = (cases_by_id or {}).get(row.case_id, {})
op_class = classify_operation(raw_case.get("answer_expression", "")) op_class = _resolve_operation_class(raw_case)
category = ( category = (
_extract_category(row.refusal_reason or "") _extract_category(row.refusal_reason or "")
if row.refusal_reason if row.refusal_reason
@ -638,37 +679,79 @@ def compact_records(
return tuple(sorted(compacted, key=lambda c: (c.case_id, c.dedupe_key))) return tuple(sorted(compacted, key=lambda c: (c.case_id, c.dedupe_key)))
def _merge_evidence_refs(
prior: tuple[str, ...],
new: tuple[str, ...],
) -> tuple[str, ...]:
return tuple(sorted(set(prior) | set(new)))
def _append_status_transitions(
prior: tuple[str, ...],
new: tuple[str, ...],
) -> tuple[str, ...]:
merged = list(prior)
for transition in new:
if not merged or merged[-1] != transition:
merged.append(transition)
return tuple(merged)
def _merge_compacted_pair(
prior: CompactedExperienceRecord,
new: CompactedExperienceRecord,
) -> CompactedExperienceRecord:
return CompactedExperienceRecord(
dedupe_key=prior.dedupe_key,
record_id=new.record_id,
case_id=new.case_id,
serving_status=new.serving_status,
sealed_status=new.sealed_status,
gold_answer=new.gold_answer,
sealed_answer=new.sealed_answer,
serving_refusal_family=new.serving_refusal_family,
sealed_failure_family=new.sealed_failure_family,
candidate_family=new.candidate_family,
first_missing_primitive=new.first_missing_primitive,
arithmetic_chain_signature=new.arithmetic_chain_signature,
positive_evidence_refs=_merge_evidence_refs(
prior.positive_evidence_refs, new.positive_evidence_refs
),
negative_evidence_refs=_merge_evidence_refs(
prior.negative_evidence_refs, new.negative_evidence_refs
),
hazard_tags=new.hazard_tags,
recommended_action=new.recommended_action,
promotion_status=new.promotion_status,
count=prior.count + new.count,
first_seen_run_id=prior.first_seen_run_id,
last_seen_run_id=new.last_seen_run_id,
status_transitions=_append_status_transitions(
prior.status_transitions, new.status_transitions
),
source_report_hash=new.source_report_hash,
)
def merge_compacted_runs( def merge_compacted_runs(
prior: tuple[CompactedExperienceRecord, ...], prior: tuple[CompactedExperienceRecord, ...],
new_records: tuple[ExperienceRecord, ...], new_records: tuple[ExperienceRecord, ...],
) -> tuple[CompactedExperienceRecord, ...]: ) -> tuple[CompactedExperienceRecord, ...]:
"""Merge prior compacted state with records from a new scout run.""" """Merge prior compacted state with records from a new scout run.
revived = [
ExperienceRecord( O(number of compacted records) never re-expands prior counts.
record_id=c.record_id, """
case_id=c.case_id, new_compacted = compact_records(new_records)
serving_status=c.serving_status, merged: dict[str, CompactedExperienceRecord] = {
sealed_status=c.sealed_status, record.dedupe_key: record for record in prior
gold_answer=c.gold_answer, }
sealed_answer=c.sealed_answer, for new_record in new_compacted:
serving_refusal_family=c.serving_refusal_family, existing = merged.get(new_record.dedupe_key)
sealed_failure_family=c.sealed_failure_family, if existing is None:
candidate_family=c.candidate_family, merged[new_record.dedupe_key] = new_record
first_missing_primitive=c.first_missing_primitive, else:
arithmetic_chain_signature=c.arithmetic_chain_signature, merged[new_record.dedupe_key] = _merge_compacted_pair(existing, new_record)
positive_evidence_refs=c.positive_evidence_refs, return tuple(sorted(merged.values(), key=lambda c: (c.case_id, c.dedupe_key)))
negative_evidence_refs=c.negative_evidence_refs,
hazard_tags=c.hazard_tags,
recommended_action=c.recommended_action,
promotion_status=c.promotion_status,
source_run_id=c.last_seen_run_id,
source_report_hash=c.source_report_hash,
)
for c in prior
for _ in range(c.count)
]
combined = tuple(revived) + new_records
return compact_records(combined)
def build_family_summaries( def build_family_summaries(
@ -788,15 +871,15 @@ def build_experience_report(
prior_compacted: tuple[CompactedExperienceRecord, ...] | None = None, prior_compacted: tuple[CompactedExperienceRecord, ...] | None = None,
include_raw_records: bool = False, include_raw_records: bool = False,
) -> dict[str, Any]: ) -> dict[str, Any]:
loaded_cases = cases
if scout_summary is None: if scout_summary is None:
scout_summary = build_scout_summary(cases, include_rows=True) if loaded_cases is None:
loaded_cases = _load_cases(_CASES_PATH)
scout_summary = build_scout_summary(loaded_cases, include_rows=True)
elif "rows" not in scout_summary: elif "rows" not in scout_summary:
raise ValueError("scout_summary must include rows") raise ValueError("scout_summary must include rows")
cases_by_id = {c["case_id"]: c for c in (cases or [])} cases_by_id = {c["case_id"]: c for c in (loaded_cases or [])}
if not cases_by_id and scout_summary.get("rows"):
for row in scout_summary["rows"]:
cases_by_id.setdefault(row["case_id"], {})
records = records_from_scout_summary(scout_summary, cases_by_id) records = records_from_scout_summary(scout_summary, cases_by_id)
if prior_compacted: if prior_compacted:
@ -896,6 +979,7 @@ __all__ = [
"compute_record_id", "compute_record_id",
"compute_report_hash", "compute_report_hash",
"compute_run_id", "compute_run_id",
"_extract_category",
"load_compacted_from_report", "load_compacted_from_report",
"merge_compacted_runs", "merge_compacted_runs",
"records_from_scout_rows", "records_from_scout_rows",

View file

@ -9,6 +9,9 @@ import pytest
from evals.gsm8k_math.runner import CaseOutcome from evals.gsm8k_math.runner import CaseOutcome
from evals.gsm8k_math.train_sample.v1.experience import ( from evals.gsm8k_math.train_sample.v1.experience import (
CompactedExperienceRecord,
ExperienceRecord,
_extract_category,
build_experience_report, build_experience_report,
compact_records, compact_records,
compute_dedupe_key, compute_dedupe_key,
@ -230,6 +233,176 @@ def test_merge_compacted_runs_increments_count():
assert merged[0].count == 2 assert merged[0].count == 2
def _experience_record(
*,
case_id: str = "gsm8k-train-sample-v1-0003",
serving_status: str = "refused",
sealed_status: str = "correct",
promotion_status: str = "candidate",
signature: str = "lift_refused_to_correct|additive|recognizer_injection|abc123",
source_run_id: str = "run-a",
source_report_hash: str = "hash-a",
positive_refs: tuple[str, ...] = (),
negative_refs: tuple[str, ...] = (),
) -> ExperienceRecord:
record = ExperienceRecord(
record_id="",
case_id=case_id,
serving_status=serving_status, # type: ignore[arg-type]
sealed_status=sealed_status, # type: ignore[arg-type]
gold_answer="864",
sealed_answer="864",
serving_refusal_family="lift_family",
sealed_failure_family="lift_family",
candidate_family="relation_hypothesis:discrete_count_statement",
first_missing_primitive="relation_hypothesis",
arithmetic_chain_signature=signature,
positive_evidence_refs=positive_refs,
negative_evidence_refs=negative_refs,
hazard_tags=(),
recommended_action="action",
promotion_status=promotion_status, # type: ignore[arg-type]
source_run_id=source_run_id,
source_report_hash=source_report_hash,
)
return ExperienceRecord(
record_id=compute_record_id(record),
case_id=record.case_id,
serving_status=record.serving_status,
sealed_status=record.sealed_status,
gold_answer=record.gold_answer,
sealed_answer=record.sealed_answer,
serving_refusal_family=record.serving_refusal_family,
sealed_failure_family=record.sealed_failure_family,
candidate_family=record.candidate_family,
first_missing_primitive=record.first_missing_primitive,
arithmetic_chain_signature=record.arithmetic_chain_signature,
positive_evidence_refs=record.positive_evidence_refs,
negative_evidence_refs=record.negative_evidence_refs,
hazard_tags=record.hazard_tags,
recommended_action=record.recommended_action,
promotion_status=record.promotion_status,
source_run_id=record.source_run_id,
source_report_hash=record.source_report_hash,
)
def test_merge_preserves_prior_transition_history():
prior_rec = _experience_record(
promotion_status="candidate",
serving_status="refused",
source_run_id="run-prior",
source_report_hash="hash-prior",
positive_refs=("scout:prior=1",),
)
prior = compact_records((prior_rec,))
promoted = _experience_record(
promotion_status="promoted_in_pr",
serving_status="correct",
source_run_id="run-new",
source_report_hash="hash-new",
positive_refs=("scout:new=2",),
)
merged = merge_compacted_runs(prior, (promoted,))
assert merged[0].status_transitions == (
"refused/correct:candidate",
"correct/correct:promoted_in_pr",
)
assert merged[0].first_seen_run_id == "run-prior"
assert merged[0].last_seen_run_id == "run-new"
def test_merge_accumulates_evidence_refs():
row = _lift_row()
scout = _scout_summary_from_rows((row,))
recs = records_from_scout_rows((row,), scout_summary=scout)
prior = compact_records(recs)
prior_record = CompactedExperienceRecord(
dedupe_key=prior[0].dedupe_key,
record_id=prior[0].record_id,
case_id=prior[0].case_id,
serving_status=prior[0].serving_status,
sealed_status=prior[0].sealed_status,
gold_answer=prior[0].gold_answer,
sealed_answer=prior[0].sealed_answer,
serving_refusal_family=prior[0].serving_refusal_family,
sealed_failure_family=prior[0].sealed_failure_family,
candidate_family=prior[0].candidate_family,
first_missing_primitive=prior[0].first_missing_primitive,
arithmetic_chain_signature=prior[0].arithmetic_chain_signature,
positive_evidence_refs=("scout:alpha=1", "scout:beta=2"),
negative_evidence_refs=("scout:neg=1",),
hazard_tags=prior[0].hazard_tags,
recommended_action=prior[0].recommended_action,
promotion_status=prior[0].promotion_status,
count=1,
first_seen_run_id="run-prior",
last_seen_run_id="run-prior",
status_transitions=prior[0].status_transitions,
source_report_hash="hash-prior",
)
new_rec = ExperienceRecord(
record_id=prior[0].record_id,
case_id=prior[0].case_id,
serving_status=prior[0].serving_status,
sealed_status=prior[0].sealed_status,
gold_answer=prior[0].gold_answer,
sealed_answer=prior[0].sealed_answer,
serving_refusal_family=prior[0].serving_refusal_family,
sealed_failure_family=prior[0].sealed_failure_family,
candidate_family=prior[0].candidate_family,
first_missing_primitive=prior[0].first_missing_primitive,
arithmetic_chain_signature=prior[0].arithmetic_chain_signature,
positive_evidence_refs=("scout:beta=2", "scout:gamma=3"),
negative_evidence_refs=("scout:neg=2",),
hazard_tags=prior[0].hazard_tags,
recommended_action=prior[0].recommended_action,
promotion_status=prior[0].promotion_status,
source_run_id="run-new",
source_report_hash="hash-new",
)
merged = merge_compacted_runs((prior_record,), (new_rec,))
assert merged[0].positive_evidence_refs == (
"scout:alpha=1",
"scout:beta=2",
"scout:gamma=3",
)
assert merged[0].negative_evidence_refs == ("scout:neg=1", "scout:neg=2")
def test_merge_scales_with_compacted_records_not_prior_count():
row = _lift_row()
scout = _scout_summary_from_rows((row,))
recs = records_from_scout_rows((row,), scout_summary=scout)
prior = compact_records(recs)
heavy = CompactedExperienceRecord(
dedupe_key=prior[0].dedupe_key,
record_id=prior[0].record_id,
case_id=prior[0].case_id,
serving_status=prior[0].serving_status,
sealed_status=prior[0].sealed_status,
gold_answer=prior[0].gold_answer,
sealed_answer=prior[0].sealed_answer,
serving_refusal_family=prior[0].serving_refusal_family,
sealed_failure_family=prior[0].sealed_failure_family,
candidate_family=prior[0].candidate_family,
first_missing_primitive=prior[0].first_missing_primitive,
arithmetic_chain_signature=prior[0].arithmetic_chain_signature,
positive_evidence_refs=prior[0].positive_evidence_refs,
negative_evidence_refs=prior[0].negative_evidence_refs,
hazard_tags=prior[0].hazard_tags,
recommended_action=prior[0].recommended_action,
promotion_status=prior[0].promotion_status,
count=10_000,
first_seen_run_id="run-heavy",
last_seen_run_id="run-heavy",
status_transitions=prior[0].status_transitions,
source_report_hash="hash-heavy",
)
merged = merge_compacted_runs((heavy,), recs)
assert merged[0].count == 10_001
def test_blocked_family_cannot_be_candidate_in_summary(): def test_blocked_family_cannot_be_candidate_in_summary():
rows = (_lift_row("gsm8k-train-sample-v1-0003"), _sealed_wrong_row()) rows = (_lift_row("gsm8k-train-sample-v1-0003"), _sealed_wrong_row())
scout = _scout_summary_from_rows(rows) scout = _scout_summary_from_rows(rows)
@ -391,4 +564,54 @@ def test_injected_scout_adapter_produces_retained_records(injected_scout_summary
assert report["retained_record_count"] >= 2 assert report["retained_record_count"] >= 2
statuses = {r["promotion_status"] for r in report["case_records"]} statuses = {r["promotion_status"] for r in report["case_records"]}
assert "candidate" in statuses assert "candidate" in statuses
assert "blocked_by_wrong_risk" in statuses assert "blocked_by_wrong_risk" in statuses
def test_extract_category_canonical_no_injection_reason():
reason = (
"candidate_graph: recognizer matched but produced no injection "
"(category=discrete_count_statement)"
)
assert _extract_category(reason) == "discrete_count_statement"
def test_extract_category_returns_none_for_unrelated_reason():
assert _extract_category("candidate_graph: no admissible candidate for statement") is None
def test_report_hash_differs_when_row_evidence_differs():
base = _scout_summary_from_rows((_lift_row(),))
other = dict(base)
other["rows"] = [
{
**_lift_row().as_dict(),
"case_id": "gsm8k-train-sample-v1-9999",
}
]
assert compute_report_hash(base) != compute_report_hash(other)
def test_report_hash_stable_for_identical_input():
scout = _scout_summary_from_rows((_lift_row(), _sealed_wrong_row()))
assert compute_report_hash(scout) == compute_report_hash(scout)
def test_live_default_report_uses_real_operation_classes():
report = build_experience_report()
op_classes = {
rec["arithmetic_chain_signature"].split("|")[1]
for rec in report["case_records"]
}
assert op_classes - {"unknown", "additive"}
assert any(
cls in {"multiplicative", "divisive"}
for cls in op_classes
)
def test_scout_summary_without_cases_uses_unknown_operation_class():
row = _lift_row()
scout = _scout_summary_from_rows((row,))
recs = records_from_scout_rows((row,), scout_summary=scout, cases_by_id={})
assert len(recs) == 1
assert recs[0].arithmetic_chain_signature.split("|")[1] == "unknown"