Wire the guarded conversational RAG answer layer end-to-end
This commit is contained in:
@@ -0,0 +1,93 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
from .models import SubjectScope
|
||||
|
||||
|
||||
class CaseOrigin(StrEnum):
|
||||
EXPERT = "expert"
|
||||
MANUAL_ADVERSARIAL = "manual_adversarial"
|
||||
SOURCE_DERIVED = "source_derived"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class EvaluationCase:
|
||||
case_id: str
|
||||
query: str
|
||||
expected_drug_id: str | None
|
||||
expected_id: str | None
|
||||
origin: CaseOrigin
|
||||
subject_scope: SubjectScope
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class EvaluationOutcome:
|
||||
case: EvaluationCase
|
||||
retrieved_ids: tuple[str, ...]
|
||||
resolved_drug_id: str | None = None
|
||||
drug_resolution_status: str = "not_attempted"
|
||||
|
||||
@property
|
||||
def passed(self) -> bool:
|
||||
if self.case.expected_id is None:
|
||||
return not self.retrieved_ids
|
||||
return self.case.expected_id in self.retrieved_ids
|
||||
|
||||
|
||||
def summarize(outcomes: list[EvaluationOutcome]) -> dict:
|
||||
def metrics(rows: list[EvaluationOutcome]) -> dict:
|
||||
positive = [row for row in rows if row.case.expected_id is not None]
|
||||
negative = [row for row in rows if row.case.expected_id is None]
|
||||
resolution_rows = [
|
||||
row for row in rows
|
||||
if row.case.expected_drug_id is not None
|
||||
and row.case.subject_scope == SubjectScope.HUMAN
|
||||
]
|
||||
return {
|
||||
"cases": len(rows),
|
||||
"positive_cases": len(positive),
|
||||
"negative_cases": len(negative),
|
||||
"recall_at_1": _recall_at(positive, 1),
|
||||
"recall_at_3": _recall_at(positive, 3),
|
||||
"drug_resolution_accuracy": (
|
||||
round(sum(
|
||||
row.resolved_drug_id == row.case.expected_drug_id
|
||||
for row in resolution_rows
|
||||
) / len(resolution_rows), 4)
|
||||
if resolution_rows else None
|
||||
),
|
||||
"drug_resolution_status_counts": {
|
||||
status: sum(
|
||||
row.drug_resolution_status == status for row in resolution_rows
|
||||
)
|
||||
for status in ("resolved", "ambiguous", "not_found", "invalid_state")
|
||||
},
|
||||
"negative_abstain_rate": (
|
||||
round(sum(row.passed for row in negative) / len(negative), 4)
|
||||
if negative else None
|
||||
),
|
||||
}
|
||||
|
||||
return {
|
||||
"expert_release_gate": metrics([
|
||||
row for row in outcomes if row.case.origin == CaseOrigin.EXPERT
|
||||
]),
|
||||
"manual_routing_diagnostic": metrics([
|
||||
row for row in outcomes if row.case.origin == CaseOrigin.MANUAL_ADVERSARIAL
|
||||
]),
|
||||
"source_derived_diagnostic": metrics([
|
||||
row for row in outcomes if row.case.origin == CaseOrigin.SOURCE_DERIVED
|
||||
]),
|
||||
}
|
||||
|
||||
|
||||
def _recall_at(rows: list[EvaluationOutcome], limit: int) -> float | None:
|
||||
if not rows:
|
||||
return None
|
||||
matched = sum(
|
||||
row.case.expected_id in row.retrieved_ids[:limit]
|
||||
for row in rows
|
||||
)
|
||||
return round(matched / len(rows), 4)
|
||||
Reference in New Issue
Block a user