diff --git a/src/or_audit/eval/cartesian.py b/src/or_audit/eval/cartesian.py index 7d86cdd..7b890ae 100644 --- a/src/or_audit/eval/cartesian.py +++ b/src/or_audit/eval/cartesian.py @@ -185,8 +185,14 @@ def _independent_case_count( ) task_group = stage.independent_case_groups.get(task.id, task.id) - split_items = manifest.items_for_split(target_split) - evaluated_items = set(split_items[:trials] if trials is not None else split_items) + all_inputs = load_items(root / task.environment.inputs_path) + allowed_items = set(manifest.items_for_split(target_split)) + filtered_inputs = [item for item in all_inputs if str(item["id"]) in allowed_items] + evaluated_items = ( + {str(item["id"]) for item in filtered_inputs[:trials]} + if trials is not None + else {str(item["id"]) for item in filtered_inputs} + ) for entry in matching: if any(item in evaluated_items for item in entry.item_ids): if unit == "patient" and entry.patient_id: @@ -308,8 +314,14 @@ def run_cartesian_job( raise TaskContractError( f"stage {stage.name} independent_case_groups keys must exactly match task ids" ) + for task_dir, task, _, _, _, trials in planned: - assert_trial_capacity(task, task_dir, trials or 0) + task_split = ( + stage.split + if stage.split + else (str(stage.name) if task.environment.splits_path else None) + ) + assert_trial_capacity(task, task_dir, trials or 0, split=task_split) observed_cases = _independent_case_count(planned, stage) if observed_cases != stage.independent_cases: raise TaskContractError( @@ -321,6 +333,11 @@ def run_cartesian_job( outcomes: list[str] = [] observed_units = 0 for task_dir, task, agent, agent_dir, dirname, pair_trials in planned: + task_split = ( + stage.split + if stage is not None and stage.split + else (str(stage.name) if stage is not None and task.environment.splits_path else None) + ) result: JobResult = run_job( task=task, task_dir=task_dir, @@ -329,6 +346,7 @@ def run_cartesian_job( out=out / dirname, n=pair_trials, gym_factory=gym_factory, + split=task_split, ) pairs.append( PairRecord( diff --git a/src/or_audit/eval/runner.py b/src/or_audit/eval/runner.py index 6508809..9c373b7 100644 --- a/src/or_audit/eval/runner.py +++ b/src/or_audit/eval/runner.py @@ -71,7 +71,14 @@ def assert_trial_capacity(task: TaskSpec, task_dir: Path, n: int, split: str | N input_ids = {str(item["id"]) for item in all_items} manifest.validate_against_input_items(input_ids, strict=False) target_split = split or "test" - available = len(manifest.items_for_split(target_split)) + available_items = manifest.items_for_split(target_split) + if not available_items: + available_splits = sorted({e.split for e in manifest.entries}) + raise TaskContractError( + f"task {task.id} requests split {target_split!r} but " + f"manifest only defines splits: {available_splits}" + ) + available = len(available_items) if n > available: raise TaskContractError( f"task {task.id} split {target_split!r} has {available} input items; " @@ -306,6 +313,10 @@ def run_job( episodes = n if n is not None else task.environment.n_eval_episodes if episodes < 1: raise TaskContractError(f"n must be >= 1, got {episodes}") + if split and not task.environment.splits_path: + raise TaskContractError( + f"task {task.id} has no declared splits_path; cannot execute explicit split {split!r}" + ) target_split = ( split if split @@ -692,11 +703,7 @@ def _run_predictions( f"task {task.id} requests split {target_split!r} but " f"manifest only defines splits: {available_splits}" ) - split_order = { - item_id: idx for idx, item_id in enumerate(manifest.items_for_split(target_split)) - } filtered = [item for item in inputs if str(item["id"]) in allowed_items] - filtered.sort(key=lambda item: split_order.get(str(item["id"]), 0)) if not filtered: raise TaskContractError( f"task {task.id} has no input items matching split {target_split!r}" @@ -887,11 +894,7 @@ def _run_interactive( f"task {task.id} requests split {target_split!r} but " f"manifest only defines splits: {available_splits}" ) - split_order = { - item_id: idx for idx, item_id in enumerate(manifest.items_for_split(target_split)) - } filtered = [item for item in inputs if str(item["id"]) in allowed_items] - filtered.sort(key=lambda item: split_order.get(str(item["id"]), 0)) if not filtered: raise TaskContractError( f"task {task.id} has no input items matching split {target_split!r}" @@ -1096,6 +1099,7 @@ def replay_job( out=out, n=int(config["n"]), gym_factory=gym_factory, + split=config.get("split") or previous.split or None, ) if rerun.head != previous.head: raise TaskContractError(f"replay head mismatch: stored {previous.head} reran {rerun.head}") diff --git a/tests/test_split_manifest.py b/tests/test_split_manifest.py index 996fea4..d240fe6 100644 --- a/tests/test_split_manifest.py +++ b/tests/test_split_manifest.py @@ -679,10 +679,8 @@ def test_runner_filters_interleaved_splits_without_leakage(tmp_path: Path) -> No ) assert result.n == 2 assert result.split == "test" - for trial in result.trials: - item_id = trial.trajectory[0].get("input", {}).get("id") - assert item_id in ("clip-001", "clip-002") - assert "train" not in str(item_id) + executed_ids = [trial.trajectory[0].get("input", {}).get("id") for trial in result.trials] + assert executed_ids == ["clip-001", "clip-002"] def test_cross_split_resume_refused(tmp_path: Path) -> None: @@ -757,3 +755,141 @@ def test_cross_split_resume_refused(tmp_path: Path) -> None: split="train", resume=True, ) + + +def test_replay_job_preserves_split(tmp_path: Path) -> None: + from or_audit.eval.loader import load_agent, load_task + from or_audit.eval.runner import replay_job, run_job + + root = Path(__file__).resolve().parents[1] + task_src = root / "docs/examples/tasks/video-nextstep" + agent_src = root / "docs/examples/agents/example-video-predictor" + + out = tmp_path / "job-replay-split" + result = run_job( + task=load_task(task_src), + task_dir=task_src, + agent=load_agent(agent_src), + agent_dir=agent_src, + out=out, + n=2, + split="test", + ) + assert result.split == "test" + replayed = replay_job(out, load_task=load_task, load_agent=load_agent) + assert replayed.split == "test" + assert replayed.head == result.head + + +def test_cartesian_independent_cases_matches_input_subset(tmp_path: Path) -> None: + import shutil + + from or_audit.eval.cartesian import run_cartesian_job + from or_audit.eval.job_config import resolve_job + + root = Path(__file__).resolve().parents[1] + task_src = root / "docs/examples/tasks/video-nextstep" + agent_src = root / "docs/examples/agents/example-video-predictor" + + task_dir = tmp_path / "task-input-subset" + shutil.copytree(task_src, task_dir) + + # Manifest defines 4 test items across 4 cases + splits_data = { + "format_version": "1", + "dataset_id": "subset-test", + "dataset_revision": "1.0", + "disjoint_by": ["case"], + "entries": [ + {"case_id": "c1", "episode_id": "e1", "split": "test", "item_ids": ["clip-001"]}, + {"case_id": "c2", "episode_id": "e2", "split": "test", "item_ids": ["clip-002"]}, + {"case_id": "c3", "episode_id": "e3", "split": "test", "item_ids": ["clip-003"]}, + {"case_id": "c4", "episode_id": "e4", "split": "test", "item_ids": ["clip-004"]}, + ], + } + (task_dir / "splits.json").write_text(json.dumps(splits_data), encoding="utf-8") + + # But inputs.json only has clip-001 and clip-002 (2 items) + inputs_data = { + "items": [ + {"id": "clip-001", "media": "public://clip-1"}, + {"id": "clip-002", "media": "public://clip-2"}, + ] + } + (task_dir / "inputs.json").write_text(json.dumps(inputs_data), encoding="utf-8") + labels_data = { + "items": [ + {"id": "clip-001", "next_step": "advance", "outcome": "continue", "unsafe": False}, + {"id": "clip-002", "next_step": "advance", "outcome": "continue", "unsafe": False}, + ] + } + (task_dir / "labels.json").write_text(json.dumps(labels_data), encoding="utf-8") + + job = tmp_path / "job-cartesian-subset" + job.mkdir() + body = f"""format_version = "1" +id = "subset-stage" +n = 2 +tasks = [{json.dumps(str(task_dir))}] +agents = [{json.dumps(str(agent_src))}] +[stage] +name = "qualification" +split = "test" +evaluation_unit = "scored clip" +target_units = 2 +independent_case_unit = "case" +independent_case_key = "id" +independent_cases = 2 +scenarios = ["video-nextstep"] +operator_contexts = ["offline"] +stop_conditions = ["stop on any hard gate failure"] +prerequisites = ["integration-smoke", "pilot"] +""" + (job / "job.toml").write_text(body, encoding="utf-8") + manifest = run_cartesian_job(resolve_job(job), out=tmp_path / "out-subset") + assert manifest.stage is not None + assert manifest.stage.independent_cases == 2 + from or_audit.eval.job import read_job_result + + assert len(manifest.pairs) == 1 + pair_result = read_job_result(tmp_path / "out-subset" / manifest.pairs[0].dir) + assert pair_result.split == "test" + pair_ids = [trial.trajectory[0].get("input", {}).get("id") for trial in pair_result.trials] + assert pair_ids == ["clip-001", "clip-002"] + + +def test_explicit_split_without_splits_path_refused(tmp_path: Path) -> None: + import shutil + + from or_audit.eval.loader import load_agent, load_task + from or_audit.eval.runner import run_job + + root = Path(__file__).resolve().parents[1] + task_src = root / "docs/examples/tasks/video-nextstep" + agent_src = root / "docs/examples/agents/example-video-predictor" + + task_dir = tmp_path / "task-no-split-decl" + shutil.copytree(task_src, task_dir) + + # Remove splits_path from task.toml + task_toml = task_dir / "task.toml" + old_text = task_toml.read_text(encoding="utf-8") + task_toml.write_text(old_text.replace('splits_path = "splits.json"', ""), encoding="utf-8") + + task = load_task(task_dir) + agent = load_agent(agent_src) + out = tmp_path / "out" + + with pytest.raises( + TaskContractError, + match="has no declared splits_path; cannot execute explicit split", + ): + run_job( + task=task, + task_dir=task_dir, + agent=agent, + agent_dir=agent_src, + out=out, + n=1, + split="test", + )