agent-harness / scripts /audit_study2_design.py
cuber12's picture
Publish agent harness research code and paper artifacts
d61821a verified
Raw
History Blame Contribute Delete
9.21 kB
"""Audit the frozen E08 design and generate its balanced reliability cells."""
from __future__ import annotations
from collections import Counter
from hashlib import sha256
import json
from pathlib import Path
from typing import Any
from agent_harness.repository import GitSnapshot
from agent_harness.specs import (
load_agent_systems,
load_embeddings,
load_experiments,
load_harnesses,
load_models,
load_repositories,
load_task_split,
load_tasks,
validate_configuration_tree,
)
from agent_harness.study2_experiment import tokenizer_for
REPEAT_TREATMENTS = ("H000", "H003", "H007", "H011", "A001", "A002")
MODEL_IDS = ("M002", "M003")
REPOSITORY_IDS = ("R001", "R002", "R003")
def validation_record(root: Path, task_id: str) -> dict[str, Any]:
candidates = (
root / "tasks" / "validation" / "study2" / f"{task_id}.json",
root / "tasks" / "validation" / f"{task_id}.json",
)
path = next((item for item in candidates if item.exists()), None)
if path is None:
raise RuntimeError(f"missing hidden-test validation record for {task_id}")
value = json.loads(path.read_text(encoding="utf-8"))
valid = value.get("valid_end_to_end")
if valid is None:
valid = value.get("validation", {}).get("valid")
if valid is not True:
raise RuntimeError(f"task validation is not successful for {task_id}: {path}")
return value
def repository_metrics(
root: Path,
repository_id: str,
models: dict[str, Any],
) -> dict[str, Any]:
repository = load_repositories(root)[repository_id]
snapshot = GitSnapshot(root / repository.local_path)
files = tuple(snapshot.iter_files(repository.pinned_head, repository.source_suffixes))
text = "\n".join(f"FILE: {item.path}\n{item.text}" for item in files)
return {
"repository_id": repository_id,
"name": repository.name,
"language": repository.language,
"pinned_head": repository.pinned_head,
"source_files": len(files),
"source_lines": sum(item.text.count("\n") + 1 for item in files),
"source_bytes": sum(len(item.text.encode("utf-8")) for item in files),
"full_source_tokens": {
model_id: tokenizer_for(models[model_id]).count(text) for model_id in MODEL_IDS
},
}
def balanced_repeat_cells(root: Path, split: tuple[str, ...]) -> list[dict[str, str]]:
tasks = load_tasks(root)
repositories = load_repositories(root)
by_repository = {
repository_id: [
task_id
for task_id in split
if tasks[task_id].repository_url == repositories[repository_id].repository_url
]
for repository_id in REPOSITORY_IDS
}
strata = [
(repository_id, model_id)
for repository_id in REPOSITORY_IDS
for model_id in MODEL_IDS
]
cells: list[dict[str, str]] = []
for treatment_index, treatment_id in enumerate(REPEAT_TREATMENTS):
excluded = {treatment_index % 6, (treatment_index + 1) % 6}
for stratum_index, (repository_id, model_id) in enumerate(strata):
if stratum_index in excluded:
continue
task_id = min(
by_repository[repository_id],
key=lambda value: sha256(
f"E08-repeat|{treatment_id}|{repository_id}|{model_id}|{value}".encode()
).hexdigest(),
)
cells.append(
{
"task_id": task_id,
"treatment_id": treatment_id,
"model_id": model_id,
}
)
if len(cells) != 24 or len({tuple(item.values()) for item in cells}) != 24:
raise RuntimeError("reliability selection did not produce 24 unique cells")
treatment_counts = Counter(item["treatment_id"] for item in cells)
stratum_counts = Counter(
(next(
repository_id
for repository_id in REPOSITORY_IDS
if tasks[item["task_id"]].repository_url
== repositories[repository_id].repository_url
), item["model_id"])
for item in cells
)
if set(treatment_counts.values()) != {4} or set(stratum_counts.values()) != {4}:
raise RuntimeError(
f"reliability balance failed: treatments={treatment_counts}, strata={stratum_counts}"
)
return cells
def audit(root: Path, write: bool) -> dict[str, Any]:
errors, warnings = validate_configuration_tree(root)
if errors or warnings:
raise RuntimeError(f"configuration audit failed: errors={errors}, warnings={warnings}")
experiment = load_experiments(root)["E08"]
models = load_models(root)
tasks = load_tasks(root)
repositories = load_repositories(root)
harnesses = load_harnesses(root)
systems = load_agent_systems(root)
embeddings = load_embeddings(root)
split = load_task_split(root / "tasks" / "splits" / "study2_confirmatory.txt")
if len(split) != 60:
raise RuntimeError(f"Study 2 split has {len(split)} tasks instead of 60")
for task_id in split:
validation_record(root, task_id)
task = tasks[task_id]
for relative in (task.gold_patch, task.test_patch):
if not (root / "tasks" / relative).is_file():
raise RuntimeError(f"missing frozen patch for {task_id}: {relative}")
repository_counts = Counter(
next(
repository_id
for repository_id, repository in repositories.items()
if repository.repository_url == tasks[task_id].repository_url
)
for task_id in split
)
if repository_counts != Counter({"R001": 20, "R002": 20, "R003": 20}):
raise RuntimeError(f"unbalanced repository task counts: {repository_counts}")
metrics = [
repository_metrics(root, repository_id, models)
for repository_id in REPOSITORY_IDS
]
if any(
min(item["full_source_tokens"].values()) <= experiment.context_budgets[0]
for item in metrics
):
raise RuntimeError("at least one repository fits inside the Study 2 context cap")
repeat_cells = balanced_repeat_cells(root, split)
result = {
"schema_version": 1,
"experiment_id": experiment.experiment_id,
"task_count": len(split),
"repository_task_counts": dict(sorted(repository_counts.items())),
"language_task_counts": dict(
sorted(Counter(tasks[item].language for item in split).items())
),
"multi_file_tasks": sum(len(tasks[item].gold_files) > 1 for item in split),
"component_treatments": {
item: harnesses[item].config_hash for item in experiment.harness_ids
},
"controlled_systems": {
item: systems[item].config_hash for item in experiment.agent_system_ids
},
"models": {item: models[item].config_hash for item in experiment.model_ids},
"embedding": {
"embedding_id": experiment.embedding_id,
"config_hash": embeddings[experiment.embedding_id].config_hash,
},
"main_generation_profile": {
"temperature": 0.0,
"top_p": 1.0,
"seeds": list(experiment.seeds),
},
"reliability_generation_profile": {
"temperature": 0.2,
"top_p": 1.0,
"seeds": [0, 1, 2],
},
"repositories": metrics,
"main_cells": experiment.cells_per_task() * len(split),
"repeat_base_cells": len(repeat_cells),
"repeat_additional_cells": len(repeat_cells) * 3,
"planned_live_cells": experiment.cells_per_task() * len(split)
+ len(repeat_cells) * 3,
"repeat_cells": repeat_cells,
}
if write:
reliability_dir = root / "configs" / "reliability"
reliability_dir.mkdir(parents=True, exist_ok=True)
(reliability_dir / "E08_repeat_cells.json").write_text(
json.dumps(
{
"schema_version": 1,
"experiment_id": "E08",
"description": (
"Twenty-four balanced non-oracle E08 cells repeated at "
"temperature 0.2 with seeds 0, 1, and 2."
),
"temperature": 0.2,
"top_p": 1.0,
"seeds": [0, 1, 2],
"cells": repeat_cells,
},
indent=2,
sort_keys=True,
)
+ "\n",
encoding="utf-8",
)
(root / "docs" / "STUDY2_DESIGN_AUDIT.json").write_text(
json.dumps(result, indent=2, sort_keys=True) + "\n", encoding="utf-8"
)
return result
def main() -> None:
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--root", type=Path, default=Path(__file__).resolve().parents[1])
parser.add_argument("--write", action="store_true")
arguments = parser.parse_args()
result = audit(arguments.root.resolve(), arguments.write)
print(json.dumps(result, indent=2, sort_keys=True))
if __name__ == "__main__":
main()