build-small-hackathon/Off-Grid-Field-Repair-Logbook
1
1from __future__ import annotations2 3import json4from dataclasses import asdict, dataclass5from pathlib import Path6from typing import Any7 8from .demo_pack import ingest_demo_pack, load_index9from .storage import init_db, connect, DEFAULT_ARTIFACTS_DIR10from .tracing import write_trace_artifact11 12 13@dataclass14class EvalResult:15 scenario_id: str16 query: str17 top_sections: list[dict[str, Any]]18 expected_section_ids: list[int]19 expected_section_titles: list[str]20 hit_top3: bool21 safety_present: bool22 sufficient: bool23 24 25def load_scenarios(pack_dir: str | Path) -> list[dict[str, Any]]:26 pack_dir = Path(pack_dir)27 with open(pack_dir / "golden_scenarios.json", "r", encoding="utf-8") as f:28 return json.load(f)29 30 31def evaluate_pack(pack_dir: str | Path, db_path: str | Path | None = None, artifact_dir: str | Path | None = None) -> dict[str, Any]:32 pack_dir = Path(pack_dir)33 init_db(db_path)34 ingest_demo_pack(pack_dir, db_path=db_path, reset=True)35 scenarios = load_scenarios(pack_dir)36 index = load_index(db_path)37 from .reasoning import build_response38 39 results: list[EvalResult] = []40 for scenario in scenarios:41 query = scenario["symptom"]42 if scenario.get("equipment_type"):43 query += " " + scenario["equipment_type"]44 if scenario.get("notes"):45 query += " " + scenario["notes"]46 hits = index.search_sections(query, top_k=5)47 top_ids = [hit.record_id for hit in hits[:3]]48 top_titles = [hit.title.lower() for hit in hits[:3]]49 expected_titles = [t.lower() for t in scenario.get("expected_section_titles", [])]50 expected_ids = set(scenario.get("expected_section_ids", []))51 hit_top3 = bool(expected_ids & set(top_ids)) if expected_ids else any(52 any(expected in title for title in top_titles)53 for expected in expected_titles54 )55 response_body, _, payload = build_response(56 scenario["symptom"],57 scenario.get("equipment_type", ""),58 scenario.get("location", ""),59 scenario.get("notes", ""),60 scenario.get("photo_path"),61 index,62 )63 safety_present = "Safety reminder" in response_body or "safety reminder" in response_body.lower()64 sufficient = payload.get("status") == "insufficient_evidence"65 results.append(66 EvalResult(67 scenario_id=scenario["scenario_id"],68 query=query,69 top_sections=[70 {71 "id": hit.record_id,72 "title": hit.title,73 "score": round(hit.score, 4),74 "citation": hit.citation,75 }76 for hit in hits[:5]77 ],78 expected_section_ids=expected_ids and list(expected_ids) or scenario.get("expected_section_ids", []),79 expected_section_titles=scenario.get("expected_section_titles", []),80 hit_top3=hit_top3,81 safety_present=safety_present,82 sufficient=sufficient,83 )84 )85 86 total = len(results)87 eligible = [result for result in results if not result.sufficient]88 top3_hits = sum(1 for result in eligible if result.hit_top3)89 safety_hits = sum(1 for result in results if result.safety_present)90 insufficient_cases = sum(1 for result in results if result.sufficient)91 report = {92 "model_name": payload.get("llm_stats", {}).get("model_id", "nvidia/NeMoTRON-3-Nano-4B-Instruct"),93 "pack": str(pack_dir),94 "scenario_count": total,95 "top3_hit_rate": round(top3_hits / len(eligible) if eligible else 0.0, 3),96 "safety_presence_rate": round(safety_hits / total if total else 0.0, 3),97 "insufficient_cases": insufficient_cases,98 "results": [asdict(result) for result in results],99 }100 trace_dir = Path(artifact_dir) if artifact_dir is not None else DEFAULT_ARTIFACTS_DIR101 trace_path = write_trace_artifact(trace_dir, {"kind": "eval", **report})102 report["trace_path"] = str(trace_path)103 return report104 105 106def main() -> None:107 import argparse108 109 parser = argparse.ArgumentParser(description="Evaluate P3 field repair logbook golden scenarios")110 parser.add_argument("--pack", required=True, help="Path to demo pack")111 parser.add_argument("--db", default=None, help="SQLite database path")112 args = parser.parse_args()113 report = evaluate_pack(args.pack, db_path=args.db)114 print(json.dumps(report, indent=2))115 116 117if __name__ == "__main__":118 main()119 