zlepsenie semantickeho vyhodnocovania odpovedi

This commit is contained in:
Ján Pták 2026-09-27 14:56:18 +02:00
parent dc592a7b9d
commit 312cb52843

View File

@ -12,7 +12,6 @@ NO_ANSWER_TEXT = (
"informáciu nepodarilo spoľahlivo nájsť." "informáciu nepodarilo spoľahlivo nájsť."
) )
MARKDOWN_URL_RE = re.compile( MARKDOWN_URL_RE = re.compile(
r"\[[^\]]*\]\((https?://[^)]+)\)" r"\[[^\]]*\]\((https?://[^)]+)\)"
) )
@ -22,10 +21,59 @@ WORD_RE = re.compile(
re.UNICODE, 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 = {
"annotation",
"entity",
"evaluation",
"extract",
"hate",
"knowledge",
"llm",
"medical",
"mteb",
"multilingual",
"named",
"ner",
"recognize",
"schema",
"set",
"training",
"triplet",
"unknown",
}
def normalize_text( def normalize_text(
value: str, value: str,
) -> str: ) -> str:
value = UNIT_RE.sub(
lambda match: (
f"{match.group(1)}"
f"{match.group(2).lower()}"
),
value,
)
value = unicodedata.normalize( value = unicodedata.normalize(
"NFKC", "NFKC",
value, value,
@ -69,8 +117,7 @@ def match_tokens(
normalize_match_token( normalize_match_token(
token token
) )
for token for token in WORD_RE.findall(
in WORD_RE.findall(
normalized normalized
) )
] ]
@ -106,16 +153,12 @@ def morphology_token_matches(
) )
return ( return (
expected[ expected[:required_prefix]
:required_prefix == actual[:required_prefix]
]
== actual[
:required_prefix
]
) )
def expected_phrase_matches( def _morphology_phrase_matches(
expected: str, expected: str,
answer: str, answer: str,
) -> bool: ) -> bool:
@ -136,13 +179,13 @@ def expected_phrase_matches(
expected_tokens = ( expected_tokens = (
match_tokens( match_tokens(
expected normalized_expected
) )
) )
answer_tokens = ( answer_tokens = (
match_tokens( match_tokens(
answer normalized_answer
) )
) )
@ -198,6 +241,373 @@ def expected_phrase_matches(
return False return False
def _semantic_token(
token: str,
) -> str:
token = (
normalize_match_token(
token
)
)
if token == "llm":
return "llm"
if (
token.startswith("velk")
or token == "large"
):
return "large"
if (
token.startswith("jazykov")
or token == "language"
):
return "language"
if token.startswith("model"):
return "model"
if (
token.startswith("pomenov")
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("multiling")
or token.startswith(
"viacjazy"
)
):
return "multilingual"
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 in {
"data",
"dat",
}:
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 == "hate":
return "hate"
if (
token.startswith("speech")
or token.startswith("rec")
):
return "speech"
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 _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 "llm" in expected_set:
return (
_contains_llm_concept(
answer_set
)
)
if {
"large",
"language",
"model",
} <= expected_set:
return (
_contains_llm_concept(
answer_set
)
)
# NER je štandardná skratka pre Named Entity Recognition.
# V tomto benchmarku považujeme NER a pomenované/named
# entity za ten istý koncept.
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 (
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( def normalize_url(
value: str, value: str,
) -> str: ) -> str:
@ -210,8 +620,8 @@ def normalize_url(
) )
if match: if match:
value = match.group( value = (
1 match.group(1)
) )
return value.rstrip( return value.rstrip(
@ -219,6 +629,65 @@ def normalize_url(
) )
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( def evaluate_answer(
question: dict[str, Any], question: dict[str, Any],
answer: str, answer: str,
@ -254,17 +723,13 @@ def evaluate_answer(
] ]
answer_contains_score = ( answer_contains_score = (
sum( sum(answer_matches)
answer_matches / len(answer_matches)
)
/ len(
answer_matches
)
if answer_matches if answer_matches
else 1.0 else 1.0
) )
expected_urls = ( raw_expected_urls = (
question.get( question.get(
"expected_source_urls", "expected_source_urls",
[], [],
@ -272,16 +737,23 @@ def evaluate_answer(
) )
if not isinstance( if not isinstance(
expected_urls, raw_expected_urls,
list, list,
): ):
expected_urls = [] raw_expected_urls = []
normalized_expected_urls = [ normalized_expected_urls = [
normalize_url( normalized
str(url) for normalized
in (
valid_expected_url(
url
)
for url
in raw_expected_urls
) )
for url in expected_urls if normalized
is not None
] ]
source_matches = [ source_matches = [
@ -291,17 +763,33 @@ def evaluate_answer(
in normalized_expected_urls in normalized_expected_urls
] ]
source_url_score = ( source_match_count = sum(
sum( source_matches
source_matches
)
/ len(
source_matches
)
if source_matches
else 1.0
) )
source_required_count = (
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( should_answer = bool(
question.get( question.get(
"should_answer", "should_answer",
@ -359,9 +847,7 @@ def evaluate_answer(
and all( and all(
answer_matches answer_matches
) )
and all( and source_ok
source_matches
)
and should_answer_ok and should_answer_ok
and tool_called and tool_called
) )
@ -376,6 +862,12 @@ def evaluate_answer(
"source_matches": ( "source_matches": (
source_matches source_matches
), ),
"source_match_count": (
source_match_count
),
"source_required_count": (
source_required_count
),
"source_url_score": ( "source_url_score": (
source_url_score source_url_score
), ),
@ -420,9 +912,7 @@ def summarize_results(
*, *,
include_groups: bool = True, include_groups: bool = True,
) -> dict[str, Any]: ) -> dict[str, Any]:
total = len( total = len(results)
results
)
errors = [ errors = [
item item
@ -434,9 +924,7 @@ def summarize_results(
completed = ( completed = (
total total
- len( - len(errors)
errors
)
) )
tool_called_values = [ tool_called_values = [
@ -562,9 +1050,7 @@ def summarize_results(
] = { ] = {
"total": total, "total": total,
"completed": completed, "completed": completed,
"errors": len( "errors": len(errors),
errors
),
"tool_call_rate": round( "tool_call_rate": round(
safe_mean( safe_mean(
tool_called_values tool_called_values
@ -653,7 +1139,9 @@ def aggregate_by_field(
]: ]:
grouped: dict[ grouped: dict[
str, str,
list[dict[str, Any]], list[
dict[str, Any]
],
] = defaultdict( ] = defaultdict(
list list
) )