dp-zp-agent/evaluation/rag_metrics.py

1450 lines
25 KiB
Python

from __future__ import annotations
import re
import statistics
import unicodedata
from collections import defaultdict
from typing import Any
NO_ANSWER_TEXT = (
"V dostupných dokumentoch ZP Wiki sa túto "
"informáciu nepodarilo spoľahlivo nájsť."
)
MARKDOWN_URL_RE = re.compile(
r"\[[^\]]*\]\((https?://[^)]+)\)"
)
WORD_RE = re.compile(
r"[^\W_]+",
re.UNICODE,
)
UNIT_RE = re.compile(
r"(?<!\w)(\d+(?:[.,]\d+)?)\s*(kb|mb|gb|tb)(?!\w)",
re.IGNORECASE,
)
SEMANTIC_STOPWORDS = {
"a",
"aj",
"and",
"from",
"of",
"pomocou",
"pre",
"the",
"v",
"vo",
"z",
"zo",
}
SEMANTIC_TRIGGER_TOKENS = {
"small",
"slovak",
"question",
"improve",
"context",
"answer",
"accuracy",
"annotation",
"cloud",
"entity",
"evaluation",
"extract",
"hate",
"hybrid",
"insert",
"knowledge",
"language",
"llm",
"medical",
"medicine",
"mteb",
"multiple",
"multilingual",
"named",
"ner",
"package",
"recognize",
"schema",
"set",
"training",
"triplet",
"unknown",
}
def normalize_text(
value: str,
) -> str:
value = UNIT_RE.sub(
lambda match: (
f"{match.group(1)}"
f"{match.group(2).lower()}"
),
value,
)
value = unicodedata.normalize(
"NFKC",
value,
).casefold()
return " ".join(
value.split()
)
def looks_like_no_answer(
value: str,
) -> bool:
normalized = normalize_text(
value
)
patterns = (
(
r"\bnepodarilo\b"
r".{0,180}"
r"\b(?:nájsť|zistiť|dohľadať)\b"
),
(
r"\b(?:nie je|nie sú|nebol|nebola|nebolo|neboli)\b"
r".{0,180}"
r"\b(?:uveden|špecifikovan)\w*"
),
(
r"\b(?:neobsahuje|neobsahujú)\b"
r".{0,180}"
r"\b(?:informáci|údaj|špecifikáci)\w*"
),
)
return any(
re.search(
pattern,
normalized,
flags=re.DOTALL,
)
is not None
for pattern in patterns
)
def normalize_match_token(
value: str,
) -> str:
value = unicodedata.normalize(
"NFKD",
value,
)
value = "".join(
character
for character in value
if not unicodedata.combining(
character
)
)
return value.casefold()
def match_tokens(
value: str,
) -> list[str]:
normalized = (
unicodedata.normalize(
"NFKC",
value,
)
)
return [
normalize_match_token(
token
)
for token in WORD_RE.findall(
normalized
)
]
def morphology_token_matches(
expected: str,
actual: str,
) -> bool:
if expected == actual:
return True
if (
expected.isdigit()
or actual.isdigit()
):
return False
shorter = min(
len(expected),
len(actual),
)
if shorter <= 2:
return False
required_prefix = (
3
if shorter <= 4
else 4
if shorter <= 6
else 5
)
return (
expected[:required_prefix]
== actual[:required_prefix]
)
def _morphology_phrase_matches(
expected: str,
answer: str,
) -> bool:
normalized_expected = (
normalize_text(
expected
)
)
normalized_answer = (
normalize_text(
answer
)
)
if not normalized_expected:
return True
expected_tokens = (
match_tokens(
normalized_expected
)
)
answer_tokens = (
match_tokens(
normalized_answer
)
)
has_numeric_token = any(
token.isdigit()
for token in expected_tokens
)
if (
not has_numeric_token
and normalized_expected
in normalized_answer
):
return True
if (
not expected_tokens
or len(answer_tokens)
< len(expected_tokens)
):
return False
window_size = len(
expected_tokens
)
for start in range(
len(answer_tokens)
- window_size
+ 1
):
window = answer_tokens[
start:
start + window_size
]
if all(
morphology_token_matches(
expected_token,
actual_token,
)
for (
expected_token,
actual_token,
)
in zip(
expected_tokens,
window,
)
):
return True
return False
def _semantic_token(
token: str,
) -> str:
token = (
normalize_match_token(
token
)
)
if token == "llm":
return "llm"
number_words = {
"nula": "0",
"jeden": "1",
"jedna": "1",
"jedno": "1",
"dva": "2",
"dve": "2",
"tri": "3",
"styri": "4",
"pat": "5",
"sest": "6",
"sedem": "7",
"osem": "8",
"devat": "9",
"desat": "10",
"zero": "0",
"one": "1",
"two": "2",
"three": "3",
"four": "4",
"five": "5",
"six": "6",
"seven": "7",
"eight": "8",
"nine": "9",
"ten": "10",
}
if token in number_words:
return number_words[token]
if (
token.startswith("velk")
or token == "large"
):
return "large"
if (
token.startswith("jazyk")
or token == "language"
or token == "languages"
):
return "language"
if token.startswith("model"):
return "model"
if (
token.startswith("pomenov")
or token.startswith("menovan")
or token == "named"
):
return "named"
if (
token.startswith("entit")
or token
in {
"entity",
"entities",
}
):
return "entity"
if (
token.startswith("anot")
or token
in {
"annotation",
"annotated",
}
):
return "annotation"
if (
token.startswith("znalost")
or token == "knowledge"
):
return "knowledge"
if (
token.startswith("graf")
or token == "graph"
):
return "graph"
if token.startswith("hybrid"):
return "hybrid"
if (
token.startswith("cloud")
or token.startswith("klaud")
):
return "cloud"
if (
token.startswith("multiling")
or token.startswith("multijaz")
or token.startswith("viacjazy")
or (
token.startswith("viac")
and "jazy" in token
)
or (
token.startswith("mnoho")
and "jazy" in token
)
):
return "multilingual"
if (
token.startswith("viacer")
or token.startswith("mnoh")
or token.startswith("niekolk")
or token == "multiple"
):
return "multiple"
if (
token.startswith("extrak")
or token.startswith("extrah")
or token.startswith("extract")
):
return "extract"
if (
token.startswith("trojic")
or token.startswith("triplet")
):
return "triplet"
if (
token.startswith("medicin")
or token.startswith("lekars")
or token == "medical"
):
return "medical"
if (
token.startswith("liek")
or token
in {
"drug",
"drugs",
"medicine",
"medicines",
"medication",
"medications",
}
):
return "medicine"
if (
token.startswith("packag")
or token.startswith("balick")
or token.startswith("pribal")
):
return "package"
if (
token.startswith("insert")
or token.startswith("letak")
):
return "insert"
if (
token in {
"data",
"dat",
}
or token.startswith("udaj")
or token.startswith("obsah")
):
return "data"
if (
token.startswith("rozpozn")
or token.startswith("rozozn")
or token == "recognize"
):
return "recognize"
if (
token.startswith("neznam")
or token == "unknown"
):
return "unknown"
if (
token.startswith("manual")
or token == "manually"
):
return "manual"
if (
token.startswith("trenovac")
or token == "training"
):
return "training"
if (
token.startswith("mnozin")
or token.startswith("sad")
or token == "set"
):
return "set"
if (
token.startswith("schem")
or token == "schema"
):
return "schema"
if (
token.startswith("sloven")
or token == "slovak"
):
return "slovak"
if (
token.startswith("zleps")
or token == "improve"
):
return "improve"
if (
token.startswith("presn")
or token == "accuracy"
):
return "accuracy"
if (
token.startswith("kratk")
or token.startswith("mal")
or token in {
"short",
"small",
}
):
return "small"
if (
token.startswith("kontext")
or token == "context"
):
return "context"
if (
token == "hate"
or token.startswith("nenavist")
):
return "hate"
if (
token.startswith("speech")
or token.startswith("prejav")
or token.startswith("rec")
):
return "speech"
if (
token.startswith("otaz")
or token == "question"
):
return "question"
if (
token.startswith("odpoved")
or token.startswith("answer")
):
return "answer"
if token == "mteb":
return "mteb"
if (
token.startswith("evalu")
or token.startswith("hodnot")
):
return "evaluation"
if token.startswith("sentence"):
return "sentence"
if token.startswith("transformer"):
return "transformer"
if token == "ner":
return "ner"
if (
token.startswith("korpus")
or token == "corpus"
):
return "set"
return token
def semantic_tokens(
value: str,
) -> list[str]:
tokens: list[str] = []
for token in match_tokens(
value
):
canonical = (
_semantic_token(
token
)
)
if (
canonical
in SEMANTIC_STOPWORDS
):
continue
tokens.append(
canonical
)
return tokens
def _contains_llm_concept(
tokens: set[str],
) -> bool:
return (
"llm" in tokens
or {
"large",
"language",
"model",
}
<= tokens
)
def _contains_multilingual_concept(
tokens: set[str],
) -> bool:
return (
"multilingual" in tokens
or {
"multiple",
"language",
}
<= tokens
)
def _semantic_concept_match(
expected: str,
answer: str,
) -> bool:
expected_tokens = (
semantic_tokens(
expected
)
)
answer_tokens = (
semantic_tokens(
answer
)
)
if (
not expected_tokens
or not answer_tokens
):
return False
expected_set = set(
expected_tokens
)
answer_set = set(
answer_tokens
)
if (
expected_set
and all(
token.isdigit()
for token in expected_set
)
):
return expected_set <= answer_set
if "llm" in expected_set:
return (
_contains_llm_concept(
answer_set
)
)
if {
"large",
"language",
"model",
} <= expected_set:
return (
_contains_llm_concept(
answer_set
)
)
if "ner" in expected_set:
return (
"ner" in answer_set
or {
"named",
"entity",
}
<= answer_set
)
if (
{
"named",
"entity",
}
<= expected_set
and "ner" in answer_set
):
return True
if "multilingual" in expected_set:
required_tokens = (
expected_set
- {
"multilingual",
}
)
if (
_contains_multilingual_concept(
answer_set
)
and required_tokens
<= answer_set
):
return True
if {
"medical",
"package",
"insert",
} <= expected_set:
remaining = (
expected_set
- {
"medical",
"package",
"insert",
}
)
if (
{
"package",
"insert",
}
<= answer_set
and (
{
"medical",
"medicine",
}
& answer_set
)
and remaining
<= answer_set
):
return True
if (
expected_set
& SEMANTIC_TRIGGER_TOKENS
and expected_set
<= answer_set
):
return True
if "mteb" in expected_set:
mteb_supported = (
"mteb"
in answer_set
or {
"sentence",
"transformer",
"evaluation",
}
<= answer_set
)
hate_supported = (
not {
"hate",
"speech",
}
<= expected_set
or {
"hate",
"speech",
}
<= answer_set
)
if (
mteb_supported
and hate_supported
):
return True
return False
def expected_phrase_matches(
expected: str,
answer: str,
) -> bool:
if _morphology_phrase_matches(
expected,
answer,
):
return True
compact_expected = (
normalize_text(
expected
)
)
compact_answer = (
normalize_text(
answer
)
)
if compact_expected:
literal_pattern = re.compile(
rf"(?<!\w)"
rf"{re.escape(compact_expected)}"
rf"(?!\w)"
)
if literal_pattern.search(
compact_answer
):
return True
return _semantic_concept_match(
expected,
answer,
)
def normalize_url(
value: str,
) -> str:
value = value.strip()
match = (
MARKDOWN_URL_RE.search(
value
)
)
if match:
value = (
match.group(1)
)
return value.rstrip(
"/"
)
def valid_expected_url(
value: Any,
) -> str | None:
text = str(
value
or ""
).strip()
if (
not text
or text
in {
"...",
"…",
}
):
return None
normalized = (
normalize_url(
text
)
)
if not normalized.startswith(
(
"http://",
"https://",
)
):
return None
return normalized
def required_source_match_count(
question: dict[str, Any],
expected_urls: list[str],
) -> int:
if not expected_urls:
return 0
if (
str(
question.get(
"category"
)
or ""
)
== "multi_document"
):
return min(
2,
len(expected_urls),
)
return len(expected_urls)
def evaluate_answer(
question: dict[str, Any],
answer: str,
*,
tool_called: bool,
) -> dict[str, Any]:
normalized_answer = (
normalize_text(
answer
)
)
expected_contains = (
question.get(
"expected_answer_contains",
[],
)
)
if not isinstance(
expected_contains,
list,
):
expected_contains = []
answer_matches = [
expected_phrase_matches(
str(expected),
answer,
)
for expected
in expected_contains
]
answer_contains_score = (
sum(answer_matches)
/ len(answer_matches)
if answer_matches
else 1.0
)
raw_expected_urls = (
question.get(
"expected_source_urls",
[],
)
)
if not isinstance(
raw_expected_urls,
list,
):
raw_expected_urls = []
normalized_expected_urls = [
normalized
for normalized
in (
valid_expected_url(
url
)
for url
in raw_expected_urls
)
if normalized
is not None
]
source_matches = [
expected_url
in answer
for expected_url
in normalized_expected_urls
]
source_match_count = sum(
source_matches
)
source_required_count = (
0
if not bool(
question.get(
"should_answer",
True,
)
)
else required_source_match_count(
question,
normalized_expected_urls,
)
)
if source_required_count == 0:
source_url_score = 1.0
source_ok = True
else:
source_url_score = min(
1.0,
source_match_count
/ source_required_count,
)
source_ok = (
source_match_count
>= source_required_count
)
should_answer = bool(
question.get(
"should_answer",
True,
)
)
exact_no_answer = (
normalize_text(
NO_ANSWER_TEXT
)
in normalized_answer
)
returned_no_answer = (
exact_no_answer
if should_answer
else (
exact_no_answer
or looks_like_no_answer(
answer
)
)
)
if should_answer:
should_answer_ok = (
bool(
answer.strip()
)
and not returned_no_answer
)
else:
should_answer_ok = (
returned_no_answer
)
tool_score = (
1.0
if tool_called
else 0.0
)
should_answer_score = (
1.0
if should_answer_ok
else 0.0
)
overall_score = (
statistics.mean(
[
answer_contains_score,
source_url_score,
should_answer_score,
tool_score,
]
)
)
strict_pass = (
bool(
answer.strip()
)
and all(
answer_matches
)
and source_ok
and should_answer_ok
and tool_called
)
return {
"answer_matches": (
answer_matches
),
"answer_contains_score": (
answer_contains_score
),
"source_matches": (
source_matches
),
"source_match_count": (
source_match_count
),
"source_required_count": (
source_required_count
),
"source_url_score": (
source_url_score
),
"should_answer_ok": (
should_answer_ok
),
"should_answer_score": (
should_answer_score
),
"returned_no_answer": (
returned_no_answer
),
"tool_score": (
tool_score
),
"overall_score": (
overall_score
),
"strict_pass": (
strict_pass
),
}
def safe_mean(
values: list[float],
) -> float:
if not values:
return 0.0
return float(
statistics.mean(
values
)
)
def summarize_results(
results: list[
dict[str, Any]
],
*,
include_groups: bool = True,
) -> dict[str, Any]:
total = len(results)
errors = [
item
for item in results
if item.get(
"error"
)
]
completed = (
total
- len(errors)
)
tool_called_values = [
(
1.0
if item.get(
"tool_called"
)
else 0.0
)
for item in results
]
answer_scores = [
float(
item.get(
"answer_contains_score",
0.0,
)
)
for item in results
]
source_scores = [
float(
item.get(
"source_url_score",
0.0,
)
)
for item in results
]
should_answer_scores = [
float(
item.get(
"should_answer_score",
0.0,
)
)
for item in results
]
overall_scores = [
float(
item.get(
"overall_score",
0.0,
)
)
for item in results
]
latencies = [
float(
item.get(
"total_latency_seconds",
0.0,
)
)
for item in results
if not item.get(
"error"
)
]
strict_passes = sum(
1
for item in results
if item.get(
"strict_pass"
)
)
prompt_tokens = sum(
int(
(
item.get(
"usage"
)
or {}
).get(
"prompt_tokens",
0,
)
)
for item in results
)
completion_tokens = sum(
int(
(
item.get(
"usage"
)
or {}
).get(
"completion_tokens",
0,
)
)
for item in results
)
total_tokens = sum(
int(
(
item.get(
"usage"
)
or {}
).get(
"total_tokens",
0,
)
)
for item in results
)
summary: dict[
str,
Any,
] = {
"total": total,
"completed": completed,
"errors": len(errors),
"tool_call_rate": round(
safe_mean(
tool_called_values
),
6,
),
"answer_contains_score": round(
safe_mean(
answer_scores
),
6,
),
"source_url_score": round(
safe_mean(
source_scores
),
6,
),
"should_answer_score": round(
safe_mean(
should_answer_scores
),
6,
),
"overall_score": round(
safe_mean(
overall_scores
),
6,
),
"strict_pass_count": (
strict_passes
),
"strict_pass_rate": round(
(
strict_passes
/ total
if total
else 0.0
),
6,
),
"mean_latency_seconds": round(
safe_mean(
latencies
),
6,
),
"prompt_tokens": (
prompt_tokens
),
"completion_tokens": (
completion_tokens
),
"total_tokens": (
total_tokens
),
}
if include_groups:
summary[
"by_category"
] = aggregate_by_field(
results,
"category",
)
summary[
"by_difficulty"
] = aggregate_by_field(
results,
"difficulty",
)
return summary
def aggregate_by_field(
results: list[
dict[str, Any]
],
field: str,
) -> dict[
str,
dict[str, Any],
]:
grouped: dict[
str,
list[
dict[str, Any]
],
] = defaultdict(
list
)
for result in results:
group_name = str(
result.get(
field,
"unknown",
)
)
grouped[
group_name
].append(
result
)
return {
group_name: summarize_results(
group_rows,
include_groups=False,
)
for (
group_name,
group_rows,
)
in sorted(
grouped.items()
)
}