diff --git a/skillopt/optimizer/slow_update.py b/skillopt/optimizer/slow_update.py index a2264ec0..d51af3d9 100644 --- a/skillopt/optimizer/slow_update.py +++ b/skillopt/optimizer/slow_update.py @@ -21,6 +21,7 @@ import json import os import traceback +from typing import Any from skillopt.model import chat_optimizer from skillopt.prompts import load_prompt @@ -156,6 +157,37 @@ def _read_trajectory(rollout_dir: str, task_id: str) -> str: # ── Structured comparison pairs ───────────────────────────────────────────── +def _is_result_success(res: dict | Any) -> bool: + """Determine whether a rollout result indicates success across metrics. + + Handles explicit 'hard' flags, general 'score' metrics, 'exact_match', + and soft threshold metrics to prevent miscategorizing benchmark outcomes. + """ + if not res or not isinstance(res, dict): + return False + if "hard" in res: + try: + return float(res["hard"]) >= 0.5 + except (ValueError, TypeError): + return bool(res["hard"]) + if "score" in res: + try: + return float(res["score"]) >= 0.5 + except (ValueError, TypeError): + return bool(res["score"]) + if "exact_match" in res: + try: + return float(res["exact_match"]) >= 0.5 + except (ValueError, TypeError): + return bool(res["exact_match"]) + if "soft" in res: + try: + return float(res["soft"]) >= 1.0 - 1e-6 + except (ValueError, TypeError): + return bool(res["soft"]) + return False + + def build_comparison_pairs( results_prev: list[dict], results_curr: list[dict], @@ -183,8 +215,8 @@ def build_comparison_pairs( tid = str(item.get("id", "")) prev = prev_by_id.get(tid, {}) curr = curr_by_id.get(tid, {}) - prev_ok = bool(prev.get("hard", 0)) - curr_ok = bool(curr.get("hard", 0)) + prev_ok = _is_result_success(prev) + curr_ok = _is_result_success(curr) if not prev_ok and curr_ok: category = "improved" @@ -201,13 +233,13 @@ def build_comparison_pairs( "category": category, "prev": { "hard": int(prev_ok), - "soft": float(prev.get("soft", 0.0)), + "soft": float(prev.get("soft", 1.0 if prev_ok else 0.0)), "predicted_answer": prev.get("predicted_answer", prev.get("answer", "N/A")), "fail_reason": prev.get("fail_reason", ""), }, "curr": { "hard": int(curr_ok), - "soft": float(curr.get("soft", 0.0)), + "soft": float(curr.get("soft", 1.0 if curr_ok else 0.0)), "predicted_answer": curr.get("predicted_answer", curr.get("answer", "N/A")), "fail_reason": curr.get("fail_reason", ""), }, diff --git a/tests/test_slow_update_robustness.py b/tests/test_slow_update_robustness.py new file mode 100644 index 00000000..22fd4f63 --- /dev/null +++ b/tests/test_slow_update_robustness.py @@ -0,0 +1,133 @@ +"""Tests for slow update field manipulation and longitudinal comparison robustness.""" + +from __future__ import annotations + +import json +import os +import tempfile + +from skillopt.optimizer.slow_update import ( + _is_result_success, + _strip_all_slow_update_fields, + build_comparison_pairs, + extract_slow_update_field, + has_slow_update_field, + inject_empty_slow_update_field, + replace_slow_update_field, + save_comparison_pairs, +) + + +def test_is_result_success_handles_diverse_metrics() -> None: + # Explicit hard boolean / float + assert _is_result_success({"hard": 1}) is True + assert _is_result_success({"hard": 0}) is False + assert _is_result_success({"hard": 1.0}) is True + assert _is_result_success({"hard": "true"}) is True + + # General score metric + assert _is_result_success({"score": 1.0}) is True + assert _is_result_success({"score": 0.0}) is False + assert _is_result_success({"score": 0.8}) is True + assert _is_result_success({"score": 0.2}) is False + + # Exact match metric + assert _is_result_success({"exact_match": 1}) is True + assert _is_result_success({"exact_match": 0}) is False + + # Soft metric threshold + assert _is_result_success({"soft": 1.0}) is True + assert _is_result_success({"soft": 0.5}) is False + + # Empty or invalid inputs + assert _is_result_success({}) is False + assert _is_result_success(None) is False + assert _is_result_success("not-a-dict") is False + + +def test_build_comparison_pairs_categorization() -> None: + items = [ + {"id": "task-1", "question": "Task 1 description"}, + {"id": "task-2", "question": "Task 2 description"}, + {"id": "task-3", "question": "Task 3 description"}, + {"id": "task-4", "question": "Task 4 description"}, + ] + + # task-1: improved (fail -> pass) + # task-2: regressed (pass -> fail) + # task-3: persistent_fail (fail -> fail) + # task-4: stable_success (pass -> pass) + results_prev = [ + {"id": "task-1", "score": 0.0, "predicted_answer": "wrong1"}, + {"id": "task-2", "exact_match": 1.0, "predicted_answer": "correct2"}, + {"id": "task-3", "hard": 0, "fail_reason": "timeout"}, + {"id": "task-4", "hard": 1, "predicted_answer": "correct4"}, + ] + results_curr = [ + {"id": "task-1", "score": 1.0, "predicted_answer": "correct1"}, + {"id": "task-2", "exact_match": 0.0, "predicted_answer": "wrong2"}, + {"id": "task-3", "hard": 0, "fail_reason": "wrong_syntax"}, + {"id": "task-4", "hard": 1, "predicted_answer": "correct4"}, + ] + + pairs = build_comparison_pairs(results_prev, results_curr, items) + assert len(pairs) == 4 + + by_id = {p["id"]: p for p in pairs} + assert by_id["task-1"]["category"] == "improved" + assert by_id["task-1"]["prev"]["hard"] == 0 + assert by_id["task-1"]["curr"]["hard"] == 1 + + assert by_id["task-2"]["category"] == "regressed" + assert by_id["task-2"]["prev"]["hard"] == 1 + assert by_id["task-2"]["curr"]["hard"] == 0 + + assert by_id["task-3"]["category"] == "persistent_fail" + assert by_id["task-3"]["prev"]["hard"] == 0 + assert by_id["task-3"]["curr"]["hard"] == 0 + + assert by_id["task-4"]["category"] == "stable_success" + assert by_id["task-4"]["prev"]["hard"] == 1 + assert by_id["task-4"]["curr"]["hard"] == 1 + + +def test_slow_update_field_lifecycle() -> None: + skill = "# Main Skill\n\nRule 1: Always verify assumptions." + + assert not has_slow_update_field(skill) + injected = inject_empty_slow_update_field(skill) + assert has_slow_update_field(injected) + assert extract_slow_update_field(injected) == "" + + # Idempotent inject + assert inject_empty_slow_update_field(injected) == injected + + # Replace field with guidance + guidance = "Avoid premature tool exit on partial output." + updated = replace_slow_update_field(injected, guidance) + assert has_slow_update_field(updated) + assert extract_slow_update_field(updated) == guidance + + # Stripping all fields + stripped = _strip_all_slow_update_fields(updated) + assert not has_slow_update_field(stripped) + assert stripped == skill.rstrip() + + +def test_save_comparison_pairs_writes_valid_json() -> None: + pairs = [ + { + "id": "item-1", + "task": "Test task", + "category": "improved", + "prev": {"hard": 0}, + "curr": {"hard": 1}, + } + ] + with tempfile.TemporaryDirectory() as tmpdir: + out_file = os.path.join(tmpdir, "comparison.json") + save_comparison_pairs(pairs, out_file) + assert os.path.exists(out_file) + with open(out_file, encoding="utf-8") as f: + loaded = json.load(f) + assert loaded == pairs