Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 10 additions & 2 deletions eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -384,11 +384,19 @@ async def run_eval(
logger.info("submit_mode_enabled", n_all=len(questions), n_resolved=len(iteration_resolved))
else:
seen_ids: set[str] = set()
questions = []
question_by_id: dict[str, Question] = {}
for q in iteration_resolved:
if q.id not in seen_ids:
seen_ids.add(q.id)
questions.append(_build_question(q))
built = _build_question(q)
question_by_id[q.id] = built
elif isinstance(q.resolution_dates, list):
existing = question_by_id[q.id]
if isinstance(existing.resolution_dates, list):
merged = list(dict.fromkeys(existing.resolution_dates + q.resolution_dates))
if len(merged) > len(existing.resolution_dates):
question_by_id[q.id] = existing.model_copy(update={"resolution_dates": merged})
questions = list(question_by_id.values())
logger.info("forecasting_questions", n_base=len(questions), n_resolved=len(iteration_resolved))
model_slug = _model_slug(agent_name, run_label=run_label, prompt_variant=prompt_variant)

Expand Down
156 changes: 156 additions & 0 deletions tests/test_eval.py
Original file line number Diff line number Diff line change
Expand Up @@ -433,6 +433,162 @@ def _counting_forecaster(question, **kwargs):
assert "mq1" in eval_result.forecasts


class TestMultiRoundForecastMatching:
"""Regression tests for multi-round forecast matching.

When a multi-horizon dataset question appears in multiple rounds,
each round has different resolution_dates (derived from forecast_due_date
+ FORECAST_HORIZONS_IN_DAYS). The dedup must merge resolution_dates across
rounds so the forecaster produces composite keys for ALL horizons.
"""

def test_multi_round_merges_resolution_dates(self, tmp_path: Path, monkeypatch: object) -> None:
"""Question in 2 rounds should have resolution_dates merged."""
import eval as eval_mod

resolved = [
ResolvedQuestion(
id="dq1", source="fred", question="Dataset Q",
outcome=1, forecast_due_date="2024-01-01",
resolution_dates=["2024-06-29", "2024-12-31"],
resolution_date="2024-06-29",
),
ResolvedQuestion(
id="dq1", source="fred", question="Dataset Q",
outcome=0, forecast_due_date="2024-01-01",
resolution_dates=["2024-06-29", "2024-12-31"],
resolution_date="2024-12-31",
),
ResolvedQuestion(
id="dq1", source="fred", question="Dataset Q",
outcome=1, forecast_due_date="2024-02-01",
resolution_dates=["2024-07-30", "2025-01-31"],
resolution_date="2024-07-30",
),
ResolvedQuestion(
id="dq1", source="fred", question="Dataset Q",
outcome=0, forecast_due_date="2024-02-01",
resolution_dates=["2024-07-30", "2025-01-31"],
resolution_date="2025-01-31",
),
]
question_sets = [
QuestionSet(
forecast_due_date="2024-01-01", question_set="set_0",
questions=[
Question(id="dq1", source="fred", question="Dataset Q",
resolution_dates=["2024-06-29", "2024-12-31"]),
],
),
QuestionSet(
forecast_due_date="2024-02-01", question_set="set_1",
questions=[
Question(id="dq1", source="fred", question="Dataset Q",
resolution_dates=["2024-07-30", "2025-01-31"]),
],
),
QuestionSet(forecast_due_date="2024-03-01", question_set="set_2", questions=[]),
QuestionSet(forecast_due_date="2024-04-01", question_set="set_3", questions=[]),
]

results_dir = tmp_path / "results"
results_dir.mkdir()

call_resolution_dates: list[list[str]] = []

def _tracking_multi_forecaster(
question: Question, resolution_dates: list[str], **kwargs: object,
) -> list[float]:
call_resolution_dates.append(resolution_dates)
return [0.6] * len(resolution_dates)

monkeypatch.setattr(eval_mod, "RESULTS_DIR", results_dir)
monkeypatch.setattr(eval_mod, "load_data", lambda: (question_sets, resolved))
monkeypatch.setattr(eval_mod, "CACHE_DIR", tmp_path / "cache")

asyncio.run(run_eval(
_dummy_forecaster, n_held_out=2, raw=True,
multi_forecaster=_tracking_multi_forecaster,
))

assert len(call_resolution_dates) == 1, "Should call multi_forecaster once per base question"
assert set(call_resolution_dates[0]) == {
"2024-06-29", "2024-12-31", "2024-07-30", "2025-01-31",
}, "Should merge resolution_dates across rounds"

def test_multi_round_no_missing_forecasts(self, tmp_path: Path, monkeypatch: object) -> None:
"""All resolved entries across rounds should match a forecast (n_missing=0)."""
import eval as eval_mod

resolved = [
ResolvedQuestion(
id="dq1", source="fred", question="Dataset Q",
outcome=1, forecast_due_date="2024-01-01",
resolution_dates=["2024-06-29", "2024-12-31"],
resolution_date="2024-06-29",
),
ResolvedQuestion(
id="dq1", source="fred", question="Dataset Q",
outcome=0, forecast_due_date="2024-01-01",
resolution_dates=["2024-06-29", "2024-12-31"],
resolution_date="2024-12-31",
),
ResolvedQuestion(
id="dq1", source="fred", question="Dataset Q",
outcome=1, forecast_due_date="2024-02-01",
resolution_dates=["2024-07-30", "2025-01-31"],
resolution_date="2024-07-30",
),
ResolvedQuestion(
id="dq1", source="fred", question="Dataset Q",
outcome=0, forecast_due_date="2024-02-01",
resolution_dates=["2024-07-30", "2025-01-31"],
resolution_date="2025-01-31",
),
]
question_sets = [
QuestionSet(
forecast_due_date="2024-01-01", question_set="set_0",
questions=[
Question(id="dq1", source="fred", question="Dataset Q",
resolution_dates=["2024-06-29", "2024-12-31"]),
],
),
QuestionSet(
forecast_due_date="2024-02-01", question_set="set_1",
questions=[
Question(id="dq1", source="fred", question="Dataset Q",
resolution_dates=["2024-07-30", "2025-01-31"]),
],
),
QuestionSet(forecast_due_date="2024-03-01", question_set="set_2", questions=[]),
QuestionSet(forecast_due_date="2024-04-01", question_set="set_3", questions=[]),
]

results_dir = tmp_path / "results"
results_dir.mkdir()

def _multi_forecaster(
question: Question, resolution_dates: list[str], **kwargs: object,
) -> list[float]:
return [0.6] * len(resolution_dates)

monkeypatch.setattr(eval_mod, "RESULTS_DIR", results_dir)
monkeypatch.setattr(eval_mod, "load_data", lambda: (question_sets, resolved))
monkeypatch.setattr(eval_mod, "CACHE_DIR", tmp_path / "cache")

eval_result = asyncio.run(run_eval(
_dummy_forecaster, n_held_out=2, raw=True,
multi_forecaster=_multi_forecaster,
))

assert eval_result.scoring.n_missing == 0, (
f"Expected 0 missing forecasts, got {eval_result.scoring.n_missing}. "
f"Forecasts: {sorted(eval_result.forecasts.keys())}"
)
assert eval_result.scoring.n_dataset == 4


class TestDifficultyAdjustmentLogging:
def test_skip_message_includes_reason(self, tmp_path: Path, monkeypatch: object, caplog: object) -> None:
"""difficulty_adjustment_skipped log should include reason field."""
Expand Down
Loading