CoolFace
Modelpublic

admesh/agentic-intent-classifier

sourceHugging Faceotherupdated 14h agoView on Hugging Face
2likes54downloads
regression_suite.py107 linesDownload Raw Back to evaluation
1from __future__ import annotations2 3import json4import sys5from pathlib import Path6 7BASE_DIR = Path(__file__).resolve().parent.parent8if str(BASE_DIR) not in sys.path:9    sys.path.insert(0, str(BASE_DIR))10 11from combined_inference import classify_query12from schemas import validate_classify_response13 14 15def load_cases(path: Path) -> list[dict]:16    return json.loads(path.read_text(encoding="utf-8"))17 18 19def write_json(path: Path, payload: dict | list) -> None:20    path.parent.mkdir(parents=True, exist_ok=True)21    path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n", encoding="utf-8")22 23 24def resolve_path(payload: dict, dotted_path: str):25    value = payload26    for part in dotted_path.split("."):27        if isinstance(value, dict):28            value = value.get(part)29        else:30            return None31    return value32 33 34def evaluate_case_file(cases_path: Path, output_dir: Path, artifact_name: str) -> dict:35    cases = load_cases(cases_path)36    results = []37    counts_by_status: dict[str, dict[str, int]] = {}38 39    for case in cases:40        payload = validate_classify_response(classify_query(case["text"]))41        mismatches = []42        expected = case.get("expected", {})43        actual_snapshot = {}44        for dotted_path, expected_value in expected.items():45            actual_value = resolve_path(payload, dotted_path)46            actual_snapshot[dotted_path] = actual_value47            if actual_value != expected_value:48                mismatches.append(49                    {50                        "path": dotted_path,51                        "expected": expected_value,52                        "actual": actual_value,53                    }54                )55 56        status = case["status"]57        bucket = counts_by_status.setdefault(status, {"total": 0, "passed": 0, "failed": 0})58        bucket["total"] += 159        if mismatches:60            bucket["failed"] += 161        else:62            bucket["passed"] += 163 64        results.append(65            {66                "id": case["id"],67                "status": status,68                "text": case["text"],69                "notes": case.get("notes", ""),70                "pass": not mismatches,71                "mismatches": mismatches,72                "expected": expected,73                "actual": actual_snapshot,74            }75        )76 77    summary = {78        "cases_path": str(cases_path),79        "count": len(results),80        "passed": sum(1 for item in results if item["pass"]),81        "failed": sum(1 for item in results if not item["pass"]),82        "by_status": counts_by_status,83        "results": results,84    }85    write_json(output_dir / artifact_name, summary)86    return summary87 88 89def evaluate_known_failure_cases(cases_path: Path, output_dir: Path) -> dict:90    return evaluate_case_file(cases_path, output_dir, "known_failure_regression.json")91 92 93def evaluate_iab_behavior_lock_cases(cases_path: Path, output_dir: Path) -> dict:94    return evaluate_case_file(cases_path, output_dir, "iab_behavior_lock_regression.json")95 96 97def evaluate_iab_cross_vertical_behavior_lock_cases(cases_path: Path, output_dir: Path) -> dict:98    return evaluate_case_file(cases_path, output_dir, "iab_cross_vertical_behavior_lock_regression.json")99 100 101def evaluate_iab_quality_target_cases(cases_path: Path, output_dir: Path) -> dict:102    return evaluate_case_file(cases_path, output_dir, "iab_quality_target_eval.json")103 104 105def evaluate_iab_cross_vertical_quality_target_cases(cases_path: Path, output_dir: Path) -> dict:106    return evaluate_case_file(cases_path, output_dir, "iab_cross_vertical_quality_target_eval.json")107