call-qa-processing / ml-services /evaluation /v2 /run_shadow_batch.py
aniketqxp's picture
fix: project evaluator dashboard outcomes
b554e59 verified
Raw
History Blame Contribute Delete
7.29 kB
from __future__ import annotations
import argparse
import json
import time
from pathlib import Path
from .runtime import (
RUNTIME_VERSION,
EvaluationV2Run,
ShadowRunStatus,
safe_run_shadow_evaluation,
)
HERE = Path(__file__).resolve().parent
REPO_ROOT = HERE.parents[2]
TRANSCRIPT_ROOT = REPO_ROOT / "data" / "sentence_segments" / "banking"
LEGACY_ROOT = REPO_ROOT / "ml-services" / "evaluation" / "results"
SENTIMENT_ROOT = (
REPO_ROOT
/ "ml-services"
/ "outputs"
/ "backend"
/ "sentiment_calls_with_features"
/ "banking"
)
DEFAULT_OUTPUT = REPO_ROOT / "frontend" / "public" / "evaluation-v2"
DEFAULT_SUMMARY = HERE / "research" / "shadow_rollout_0_1.json"
def _load(path: Path) -> dict:
return json.loads(path.read_text(encoding="utf-8"))
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT)
parser.add_argument(
"--summary-output",
type=Path,
default=DEFAULT_SUMMARY,
)
parser.add_argument(
"--resume",
action="store_true",
help="Reuse existing succeeded artifacts and retry other calls.",
)
parser.add_argument(
"--rerun-call",
action="append",
default=[],
help="Call ID to rerun even when its existing artifact succeeded.",
)
parser.add_argument(
"--only-call",
action="append",
default=[],
help="Restrict the batch to one or more call IDs.",
)
parser.add_argument(
"--delay-seconds",
type=float,
default=0.0,
help="Pause between provider calls to respect token rate limits.",
)
args = parser.parse_args()
if args.delay_seconds < 0:
parser.error("--delay-seconds cannot be negative")
args.output_dir.mkdir(parents=True, exist_ok=True)
counts: dict[str, int] = {}
decision_statuses: dict[str, int] = {}
comparison_counts = {
"comparable": 0,
"not_comparable": 0,
"attention_agreement": 0,
"attention_disagreement": 0,
}
legacy_attention = 0
v2_attention = 0
rows = []
transcript_paths = sorted(TRANSCRIPT_ROOT.glob("*.json"))
if args.only_call:
selected = set(args.only_call)
transcript_paths = [
path for path in transcript_paths if path.stem in selected
]
missing = sorted(selected - {path.stem for path in transcript_paths})
if missing:
parser.error(f"unknown --only-call values: {missing}")
for transcript_index, transcript_path in enumerate(transcript_paths):
call_id = transcript_path.stem
legacy_path = LEGACY_ROOT / f"{call_id}_graph.json"
sentiment_path = (
SENTIMENT_ROOT
/ f"{call_id}_backend_sentiment_with_features.json"
)
output = args.output_dir / f"{call_id}.json"
run = None
if (
args.resume
and call_id not in args.rerun_call
and output.exists()
):
existing = EvaluationV2Run.model_validate(_load(output))
if existing.status == ShadowRunStatus.SUCCEEDED:
run = existing
if run is None:
run = safe_run_shadow_evaluation(
transcript=_load(transcript_path),
sentiment=(
_load(sentiment_path)
if sentiment_path.exists()
else None
),
legacy_evaluation=(
_load(legacy_path) if legacy_path.exists() else None
),
transcript_source=str(
transcript_path.relative_to(REPO_ROOT)
),
sentiment_source=(
str(sentiment_path.relative_to(REPO_ROOT))
if sentiment_path.exists()
else None
),
)
output.write_text(
json.dumps(run.model_dump(mode="json"), indent=2) + "\n",
encoding="utf-8",
)
if (
args.delay_seconds
and transcript_index < len(transcript_paths) - 1
):
time.sleep(args.delay_seconds)
counts[run.status.value] = counts.get(run.status.value, 0) + 1
if run.legacy_proxy and run.legacy_proxy.attention_required:
legacy_attention += 1
if run.decision:
status = run.decision.decision_status.value
decision_statuses[status] = decision_statuses.get(status, 0) + 1
v2_attention += int(run.decision.attention_required)
if run.comparison:
if run.comparison.comparable:
comparison_counts["comparable"] += 1
key = (
"attention_agreement"
if run.comparison.attention_agreement
else "attention_disagreement"
)
comparison_counts[key] += 1
else:
comparison_counts["not_comparable"] += 1
rows.append({
"call_id": call_id,
"status": run.status.value,
"decision_status": (
run.decision.decision_status.value
if run.decision
else None
),
"legacy_attention_proxy": (
run.legacy_proxy.attention_required
if run.legacy_proxy
else None
),
"v2_attention": (
run.decision.attention_required
if run.decision
else None
),
"comparable": (
run.comparison.comparable
if run.comparison
else False
),
"limitations": run.limitations,
})
summary = {
"schema_version": "1.0",
"runtime_version": RUNTIME_VERSION,
"population": "ten_banking_calls",
"call_count": len(rows),
"run_statuses": dict(sorted(counts.items())),
"decision_statuses": dict(sorted(decision_statuses.items())),
"legacy_attention_proxy_count": legacy_attention,
"v2_attention_count": v2_attention,
"comparison": comparison_counts,
"conclusion": (
f"{comparison_counts['comparable']} calls produced comparable "
"attention decisions: "
f"{comparison_counts['attention_agreement']} agreements and "
f"{comparison_counts['attention_disagreement']} disagreements. "
f"{comparison_counts['not_comparable']} calls remain "
"non-comparable because requirement coverage is partial."
),
"calls": rows,
}
args.summary_output.parent.mkdir(parents=True, exist_ok=True)
args.summary_output.write_text(
json.dumps(summary, indent=2) + "\n",
encoding="utf-8",
)
print(
f"Wrote {sum(counts.values())} shadow runs: "
+ ", ".join(
f"{status}={count}"
for status, count in sorted(counts.items())
)
)
print(f"Summary -> {args.summary_output}")
return 0
if __name__ == "__main__":
raise SystemExit(main())