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
24 changes: 21 additions & 3 deletions src/or_audit/eval/cartesian.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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(
Expand All @@ -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,
Expand All @@ -329,6 +346,7 @@ def run_cartesian_job(
out=out / dirname,
n=pair_trials,
gym_factory=gym_factory,
split=task_split,
)
pairs.append(
PairRecord(
Expand Down
22 changes: 13 additions & 9 deletions src/or_audit/eval/runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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; "
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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}"
Expand Down Expand Up @@ -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}"
Expand Down Expand Up @@ -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}")
Expand Down
144 changes: 140 additions & 4 deletions tests/test_split_manifest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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",
)
Loading