diff --git a/evaluation/rag_metrics.py b/evaluation/rag_metrics.py index b33a2a6..23b9f33 100644 --- a/evaluation/rag_metrics.py +++ b/evaluation/rag_metrics.py @@ -42,6 +42,13 @@ SEMANTIC_STOPWORDS = { } SEMANTIC_TRIGGER_TOKENS = { + "small", + "slovak", + "question", + "improve", + "context", + "answer", + "accuracy", "annotation", "entity", "evaluation", @@ -84,6 +91,42 @@ def normalize_text( ) +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: @@ -270,6 +313,7 @@ def _semantic_token( if ( token.startswith("pomenov") + or token.startswith("menovan") or token == "named" ): return "named" @@ -308,9 +352,8 @@ def _semantic_token( if ( token.startswith("multiling") - or token.startswith( - "viacjazy" - ) + or token.startswith("viacjazy") + or token.startswith("mnoh") ): return "multilingual" @@ -334,10 +377,13 @@ def _semantic_token( ): return "medical" - if token in { - "data", - "dat", - }: + if ( + token in { + "data", + "dat", + } + or token.startswith("obsah") + ): return "data" if ( @@ -378,15 +424,65 @@ def _semantic_token( ): return "schema" - if token == "hate": + 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" @@ -768,7 +864,14 @@ def evaluate_answer( ) source_required_count = ( - required_source_match_count( + 0 + if not bool( + question.get( + "should_answer", + True, + ) + ) + else required_source_match_count( question, normalized_expected_urls, ) @@ -797,13 +900,24 @@ def evaluate_answer( ) ) - returned_no_answer = ( + 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(