Cognitive-rag / 03_embedding / run_multihop_agent.py
run_multihop_agent.py
Raw
#!/usr/bin/env python
"""Run the multi-hop agent as a standalone CLI."""

from __future__ import annotations

import argparse
import os
import sys
from pathlib import Path

from dotenv import load_dotenv

PROJECT_ROOT = Path(__file__).resolve().parents[1]
load_dotenv(PROJECT_ROOT / "05_generation" / ".env")

sys.path.insert(0, str(PROJECT_ROOT / "03_embedding"))

from embed_query import load_client
from multihop_agent import MultiHopAgent, MultiHopConfig


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Run CognitiveRAG multi-hop agent")
    parser.add_argument("--query", required=True, help="User query text.")
    parser.add_argument("--strategy", choices=["dense", "sparse", "hybrid"], default="hybrid")
    parser.add_argument("--max-hops", type=int, default=3)
    parser.add_argument("--per-hop-top-k", type=int, default=8)
    parser.add_argument("--total-doc-budget", type=int, default=24)
    parser.add_argument("--sufficiency-threshold", type=float, default=0.75)
    parser.add_argument("--time-limit-sec", type=int, default=60)
    parser.add_argument("--output-dir", default=os.path.join("03_embedding", "output", "multihop_traces"))
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    api_key = os.getenv("NVIDIA_API_KEY")
    if not api_key:
        raise SystemExit("Missing NVIDIA_API_KEY for embeddings.")

    client = load_client("nvidia/llama-3.2-nv-embedqa-1b-v2", api_key, "NONE")
    embed_fn = lambda text: client.embed_query(text)

    config = MultiHopConfig(
        max_hops=args.max_hops,
        per_hop_top_k=args.per_hop_top_k,
        total_doc_budget=args.total_doc_budget,
        sufficiency_threshold=args.sufficiency_threshold,
        time_limit_sec=args.time_limit_sec,
        strategy=args.strategy,
    )

    agent = MultiHopAgent(embed_fn=embed_fn, config=config, output_dir=args.output_dir)
    result = agent.run(args.query)

    print("\nAnswer:")
    print(result.answer)
    print("\nCitations:")
    for cite in result.citations:
        print(f"- {cite}")
    print(f"\nConfidence: {result.confidence:.2f}")
    print(f"Termination: {result.termination_reason}")
    print(f"Trace: {result.trace_path}")


if __name__ == "__main__":
    main()