diff --git a/evaluation/rag_metrics.py b/evaluation/rag_metrics.py index 1d4ffd8..b33a2a6 100644 --- a/evaluation/rag_metrics.py +++ b/evaluation/rag_metrics.py @@ -12,7 +12,6 @@ NO_ANSWER_TEXT = ( "informáciu nepodarilo spoľahlivo nájsť." ) - MARKDOWN_URL_RE = re.compile( r"\[[^\]]*\]\((https?://[^)]+)\)" ) @@ -22,10 +21,59 @@ WORD_RE = re.compile( re.UNICODE, ) +UNIT_RE = re.compile( + r"(? str: + value = UNIT_RE.sub( + lambda match: ( + f"{match.group(1)}" + f"{match.group(2).lower()}" + ), + value, + ) + value = unicodedata.normalize( "NFKC", value, @@ -69,8 +117,7 @@ def match_tokens( normalize_match_token( token ) - for token - in WORD_RE.findall( + for token in WORD_RE.findall( normalized ) ] @@ -106,16 +153,12 @@ def morphology_token_matches( ) return ( - expected[ - :required_prefix - ] - == actual[ - :required_prefix - ] + expected[:required_prefix] + == actual[:required_prefix] ) -def expected_phrase_matches( +def _morphology_phrase_matches( expected: str, answer: str, ) -> bool: @@ -136,13 +179,13 @@ def expected_phrase_matches( expected_tokens = ( match_tokens( - expected + normalized_expected ) ) answer_tokens = ( match_tokens( - answer + normalized_answer ) ) @@ -198,6 +241,373 @@ def expected_phrase_matches( 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"(? str: @@ -210,8 +620,8 @@ def normalize_url( ) if match: - value = match.group( - 1 + value = ( + match.group(1) ) 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( question: dict[str, Any], answer: str, @@ -254,17 +723,13 @@ def evaluate_answer( ] answer_contains_score = ( - sum( - answer_matches - ) - / len( - answer_matches - ) + sum(answer_matches) + / len(answer_matches) if answer_matches else 1.0 ) - expected_urls = ( + raw_expected_urls = ( question.get( "expected_source_urls", [], @@ -272,16 +737,23 @@ def evaluate_answer( ) if not isinstance( - expected_urls, + raw_expected_urls, list, ): - expected_urls = [] + raw_expected_urls = [] normalized_expected_urls = [ - normalize_url( - str(url) + normalized + 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 = [ @@ -291,17 +763,33 @@ def evaluate_answer( in normalized_expected_urls ] - source_url_score = ( - sum( - source_matches - ) - / len( - source_matches - ) - if source_matches - else 1.0 + source_match_count = sum( + source_matches ) + 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( question.get( "should_answer", @@ -359,9 +847,7 @@ def evaluate_answer( and all( answer_matches ) - and all( - source_matches - ) + and source_ok and should_answer_ok and tool_called ) @@ -376,6 +862,12 @@ def evaluate_answer( "source_matches": ( source_matches ), + "source_match_count": ( + source_match_count + ), + "source_required_count": ( + source_required_count + ), "source_url_score": ( source_url_score ), @@ -420,9 +912,7 @@ def summarize_results( *, include_groups: bool = True, ) -> dict[str, Any]: - total = len( - results - ) + total = len(results) errors = [ item @@ -434,9 +924,7 @@ def summarize_results( completed = ( total - - len( - errors - ) + - len(errors) ) tool_called_values = [ @@ -562,9 +1050,7 @@ def summarize_results( ] = { "total": total, "completed": completed, - "errors": len( - errors - ), + "errors": len(errors), "tool_call_rate": round( safe_mean( tool_called_values @@ -653,7 +1139,9 @@ def aggregate_by_field( ]: grouped: dict[ str, - list[dict[str, Any]], + list[ + dict[str, Any] + ], ] = defaultdict( list )