Buckets:
| #!/usr/bin/env python | |
| """CLI to run ReAct, CoT, or Reflexion on HotpotQA or GSM8K.""" | |
| import argparse | |
| import logging | |
| import sys | |
| from pathlib import Path | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.append(str(ROOT)) | |
| from src.config import AgentConfig, EvalConfig, LLMConfig, RunConfig, SearchConfig # noqa: E402 | |
| from src.eval import evaluate_agent # noqa: E402 | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description="Run ReAct, CoT, or Reflexion baselines.") | |
| parser.add_argument("--agent", choices=["react", "cot", "reflexion"], default="react") | |
| parser.add_argument("--backend", choices=["openai", "hf"], default="openai") | |
| parser.add_argument("--model", default="gpt-4o-mini", help="Model name for backend.") | |
| parser.add_argument("--dataset", choices=["hotpot", "gsm8k"], default="hotpot") | |
| parser.add_argument("--split", default="validation[:25]", help="HF dataset split string.") | |
| parser.add_argument("--num-examples", type=int, default=25) | |
| parser.add_argument("--max-steps", type=int, default=6) | |
| parser.add_argument("--temperature", type=float, default=0.7) | |
| parser.add_argument("--max-new-tokens", type=int, default=256) | |
| parser.add_argument("--index-path", default="data/hotpot_index.pkl") | |
| parser.add_argument("--index-split", default="train[:2000]") | |
| parser.add_argument("--log-path", default=None) | |
| return parser.parse_args() | |
| def main(): | |
| args = parse_args() | |
| logging.basicConfig(level=logging.INFO, format="%(levelname)s:%(name)s:%(message)s") | |
| cfg = RunConfig( | |
| llm=LLMConfig( | |
| backend=args.backend, | |
| model=args.model, | |
| temperature=args.temperature, | |
| max_new_tokens=args.max_new_tokens, | |
| ), | |
| agent=AgentConfig(max_steps=args.max_steps, verbose=True), | |
| search=SearchConfig(index_path=args.index_path, dataset_split=args.index_split, k=4), | |
| eval=EvalConfig(dataset=args.dataset, split=args.split, num_examples=args.num_examples, log_path=args.log_path), | |
| ) | |
| result = evaluate_agent(cfg, args.agent) | |
| print( | |
| f"agent={args.agent} dataset={args.dataset} split={args.split} " | |
| f"{result['metric_name']}={result['metric']:.3f} n={result['n']}" | |
| ) | |
| if __name__ == "__main__": | |
| main() | |
| #!/usr/bin/env python | |
| import datetime | |
| import json | |
| import sys | |
| from pathlib import Path | |
| from typing import Optional | |
| import typer | |
| ROOT = Path(__file__).resolve().parents[1] | |
| if str(ROOT) not in sys.path: | |
| sys.path.append(str(ROOT)) | |
| from src.agents import CoTAgent, ReActAgent # noqa: E402 | |
| from src.eval import load_gsm8k, load_hotpot, run_evaluation # noqa: E402 | |
| from src.llm import LLMClient, LLMConfig # noqa: E402 | |
| from src.tools import build_default_tools # noqa: E402 | |
| from src.utils import ensure_dir # noqa: E402 | |
| app = typer.Typer(pretty_exceptions_show_locals=False) | |
| def main( | |
| dataset: str = typer.Option("hotpot", help="hotpot or gsm8k"), | |
| agent: str = typer.Option("react", help="react or cot"), | |
| model: str = typer.Option("gpt-4o-mini", help="OpenAI-compatible model name"), | |
| limit: Optional[int] = typer.Option(50, help="Max examples to evaluate"), | |
| corpus_path: Optional[str] = typer.Option("data/hotpot_corpus.jsonl", help="BM25 corpus path for search"), | |
| max_steps: int = typer.Option(6, help="Max reasoning steps for ReAct"), | |
| temperature: float = typer.Option(0.2, help="Decoding temperature"), | |
| output_dir: str = typer.Option("runs", help="Where to store logs/metrics"), | |
| ) -> None: | |
| """Run evaluation for ReAct or CoT agents.""" | |
| config = LLMConfig(model=model, temperature=temperature) | |
| llm = LLMClient(config) | |
| tools = build_default_tools(corpus_path=corpus_path) | |
| if agent.lower() == "react": | |
| agent_instance = ReActAgent(llm, tools, max_steps=max_steps) | |
| elif agent.lower() == "cot": | |
| agent_instance = CoTAgent(llm) | |
| else: | |
| raise typer.BadParameter("agent must be 'react' or 'cot'") | |
| if dataset.lower() == "hotpot": | |
| data = load_hotpot(limit=limit) | |
| elif dataset.lower() == "gsm8k": | |
| data = load_gsm8k(limit=limit) | |
| else: | |
| raise typer.BadParameter("dataset must be 'hotpot' or 'gsm8k'") | |
| timestamp = datetime.datetime.utcnow().strftime("%Y%m%d-%H%M%S") | |
| run_dir = Path(output_dir) / f"{dataset}-{agent}-{timestamp}" | |
| ensure_dir(str(run_dir)) | |
| metrics = run_evaluation(agent_instance, data, dataset.lower(), str(run_dir)) | |
| metrics_path = run_dir / "metrics.json" | |
| metrics_path.write_text(json.dumps(metrics, indent=2)) | |
| typer.secho(f"Metrics saved to {metrics_path}", fg=typer.colors.CYAN) | |
| if __name__ == "__main__": | |
| app() | |
Xet Storage Details
- Size:
- 4.7 kB
- Xet hash:
- 260213778c6cbbe64b5ec5d6a42cefd79938b560b1a1de5cf03d43969d2dc76d
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.