90shikhar08's picture
feat(eval): async harness with Groq/Gemini/HF adapters and 5-bias matrix
a8f17de
Raw
History Blame Contribute Delete
5.91 kB
#!/usr/bin/env python3
"""Run procedural bias eval across splits, models, and biases; write JSON results."""
from __future__ import annotations
import argparse
import asyncio
import sys
from collections import defaultdict
from itertools import islice
from pathlib import Path
from typing import Any
ROOT = Path(__file__).resolve().parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
from bias_types import BIAS_TYPES, validate_bias # noqa: E402
from eval.adapters import GeminiAdapter, GroqAdapter, HFAdapter # noqa: E402
from eval.harness import evaluate_batch # noqa: E402
from eval.results_writer import write_result # noqa: E402
from splits import iter_split # noqa: E402
_ALLOWED_MODELS = ("gemini", "groq", "hf")
_ALLOWED_ARCHETYPES = ("easy", "medium", "hard")
_ALLOWED_SPLITS = ("train", "val", "test")
def _parse_csv(raw: str) -> list[str]:
return [x.strip() for x in raw.split(",") if x.strip()]
def _build_tasks(
split: str,
archetypes: list[str],
bias_types: list[str],
max_seeds: int,
) -> list[dict[str, Any]]:
tasks: list[dict[str, Any]] = []
for arch in archetypes:
for bias in bias_types:
gen = iter_split(split, arch, bias_type=bias)
for task in islice(gen, max_seeds):
meta = dict(task.get("_meta") or {})
meta["bias_type"] = bias
task["_meta"] = meta
tasks.append(task)
return tasks
def _make_adapter(kind: str, *, groq_model: str, hf_model: str) -> Any:
k = kind.lower()
if k == "gemini":
return GeminiAdapter()
if k == "groq":
return GroqAdapter(groq_model)
if k == "hf":
return HFAdapter(hf_model)
raise ValueError(f"unknown model kind {kind!r}")
def _print_summary(results: list[dict[str, Any]]) -> None:
cell: dict[tuple[str, str, str], list[int]] = defaultdict(lambda: [0, 0])
for r in results:
key = (str(r["model"]), str(r["bias_type"]), str(r["archetype"]))
cell[key][1] += 1
cell[key][0] += int(bool(r.get("is_optimal")))
keys = sorted(cell.keys())
w_m = max(len("model"), max((len(k[0]) for k in keys), default=6))
w_b = max(len("bias_type"), max((len(k[1]) for k in keys), default=10))
w_a = max(len("archetype"), max((len(k[2]) for k in keys), default=8))
header = (
f"{'model':<{w_m}} {'bias_type':<{w_b}} {'archetype':<{w_a}} "
"optimal_rate (opt/n)"
)
print("\n=== Summary (optimal rate) ===")
print(header)
print("-" * len(header))
for model, bias, arch in keys:
opt, n = cell[(model, bias, arch)]
rate = opt / n if n else 0.0
print(
f"{model:<{w_m}} {bias:<{w_b}} {arch:<{w_a}} "
f"{rate:>12.4f} ({opt}/{n})"
)
async def _async_main(args: argparse.Namespace) -> None:
archetypes = _parse_csv(args.archetypes)
for a in archetypes:
if a not in _ALLOWED_ARCHETYPES:
raise SystemExit(f"unknown archetype {a!r}; allowed {_ALLOWED_ARCHETYPES}")
bias_types = _parse_csv(args.bias_types)
for b in bias_types:
validate_bias(b)
models = _parse_csv(args.models)
for m in models:
if m.lower() not in _ALLOWED_MODELS:
raise SystemExit(f"unknown model {m!r}; allowed {_ALLOWED_MODELS}")
tasks = _build_tasks(args.split, archetypes, bias_types, args.max_seeds)
if not tasks:
print("No tasks built (check split files and max-seeds).", file=sys.stderr)
return
results_dir = str(ROOT / "results")
all_results: list[dict[str, Any]] = []
for mk in models:
adapter = _make_adapter(
mk.lower(),
groq_model=args.groq_model,
hf_model=args.hf_model,
)
batch = await evaluate_batch(
adapter,
tasks,
max_concurrent=args.max_concurrent,
)
for r in batch:
write_result(r, results_dir=results_dir)
all_results.extend(batch)
_print_summary(all_results)
def main() -> None:
parser = argparse.ArgumentParser(description="Bias eval harness over data splits.")
parser.add_argument(
"--models",
default="gemini",
help=f"Comma-separated: {', '.join(_ALLOWED_MODELS)} (default: gemini)",
)
parser.add_argument(
"--bias-types",
default=",".join(BIAS_TYPES),
help="Comma-separated bias names (default: all BIAS_TYPES)",
)
parser.add_argument(
"--archetypes",
default="easy,medium,hard",
help="Comma-separated archetypes (default: easy,medium,hard)",
)
parser.add_argument(
"--split",
default="val",
choices=list(_ALLOWED_SPLITS),
help="Seed split (default: val)",
)
parser.add_argument(
"--max-seeds",
type=int,
default=10,
metavar="N",
help="Cap seeds per (archetype × bias_type) combo (default: 10)",
)
parser.add_argument(
"--max-concurrent",
type=int,
default=5,
metavar="N",
help="Concurrent API calls per model batch (default: 5)",
)
parser.add_argument(
"--groq-model",
default="llama-3.3-70b-versatile",
help="Groq chat model id when --models includes groq",
)
parser.add_argument(
"--hf-model",
default="Qwen/Qwen2.5-0.5B-Instruct",
help="Hugging Face model id when --models includes hf",
)
parser.add_argument(
"--confirm-test",
action="store_true",
help="Required when --split test (intentional held-out evaluation)",
)
args = parser.parse_args()
if args.split == "test" and not args.confirm_test:
parser.error("--split test requires --confirm-test")
asyncio.run(_async_main(args))
if __name__ == "__main__":
main()