tahamajs's picture
download
raw
4.7 kB
#!/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)
@app.command()
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.