This commit is contained in:
Ján Pták 2026-08-16 22:24:29 +02:00
parent 2af4d6aec3
commit c13f5a39ee

View File

@ -17,6 +17,11 @@ MARKDOWN_URL_RE = re.compile(
r"\[[^\]]*\]\((https?://[^)]+)\)" r"\[[^\]]*\]\((https?://[^)]+)\)"
) )
WORD_RE = re.compile(
r"[^\W_]+",
re.UNICODE,
)
def normalize_text( def normalize_text(
value: str, value: str,
@ -24,15 +29,175 @@ def normalize_text(
value = unicodedata.normalize( value = unicodedata.normalize(
"NFKC", "NFKC",
value, value,
) ).casefold()
value = value.casefold()
return " ".join( return " ".join(
value.split() value.split()
) )
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 expected_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(
expected
)
)
answer_tokens = (
match_tokens(
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 normalize_url( def normalize_url(
value: str, value: str,
) -> str: ) -> str:
@ -79,25 +244,15 @@ def evaluate_answer(
): ):
expected_contains = [] expected_contains = []
answer_matches: list[ answer_matches = [
bool expected_phrase_matches(
] = [] str(expected),
answer,
)
for expected
in expected_contains
]
for expected in expected_contains:
expected_text = (
normalize_text(
str(
expected
)
)
)
answer_matches.append(
expected_text
in normalized_answer
)
if answer_matches:
answer_contains_score = ( answer_contains_score = (
sum( sum(
answer_matches answer_matches
@ -105,11 +260,10 @@ def evaluate_answer(
/ len( / len(
answer_matches answer_matches
) )
if answer_matches
else 1.0
) )
else:
answer_contains_score = 1.0
expected_urls = ( expected_urls = (
question.get( question.get(
"expected_source_urls", "expected_source_urls",
@ -125,26 +279,18 @@ def evaluate_answer(
normalized_expected_urls = [ normalized_expected_urls = [
normalize_url( normalize_url(
str( str(url)
url
)
) )
for url in expected_urls for url in expected_urls
] ]
source_matches: list[ source_matches = [
bool
] = []
for expected_url in (
normalized_expected_urls
):
source_matches.append(
expected_url expected_url
in answer in answer
) for expected_url
in normalized_expected_urls
]
if source_matches:
source_url_score = ( source_url_score = (
sum( sum(
source_matches source_matches
@ -152,11 +298,10 @@ def evaluate_answer(
/ len( / len(
source_matches source_matches
) )
if source_matches
else 1.0
) )
else:
source_url_score = 1.0
should_answer = bool( should_answer = bool(
question.get( question.get(
"should_answer", "should_answer",
@ -164,14 +309,10 @@ def evaluate_answer(
) )
) )
normalized_no_answer = ( returned_no_answer = (
normalize_text( normalize_text(
NO_ANSWER_TEXT NO_ANSWER_TEXT
) )
)
returned_no_answer = (
normalized_no_answer
in normalized_answer in normalized_answer
) )
@ -539,7 +680,8 @@ def aggregate_by_field(
for ( for (
group_name, group_name,
group_rows, group_rows,
) in sorted( )
in sorted(
grouped.items() grouped.items()
) )
} }