CoolFace
Apppublic

build-small-hackathon/Off-Grid-Field-Repair-Logbook

sourceHugging Faceupdated 3mo agoView on Hugging Face
1likes
eval.py119 linesDownload Raw Back to app_kit
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