diff --git a/evaluation/rag_metrics.py b/evaluation/rag_metrics.py index 4db6c7d..1d4ffd8 100644 --- a/evaluation/rag_metrics.py +++ b/evaluation/rag_metrics.py @@ -17,6 +17,11 @@ MARKDOWN_URL_RE = re.compile( r"\[[^\]]*\]\((https?://[^)]+)\)" ) +WORD_RE = re.compile( + r"[^\W_]+", + re.UNICODE, +) + def normalize_text( value: str, @@ -24,15 +29,175 @@ def normalize_text( value = unicodedata.normalize( "NFKC", value, - ) - - value = value.casefold() + ).casefold() return " ".join( 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( value: str, ) -> str: @@ -79,36 +244,25 @@ def evaluate_answer( ): expected_contains = [] - answer_matches: list[ - bool - ] = [] - - for expected in expected_contains: - expected_text = ( - normalize_text( - str( - expected - ) - ) + answer_matches = [ + expected_phrase_matches( + str(expected), + answer, ) + for expected + in expected_contains + ] - answer_matches.append( - expected_text - in normalized_answer + answer_contains_score = ( + sum( + answer_matches ) - - if answer_matches: - answer_contains_score = ( - sum( - answer_matches - ) - / len( - answer_matches - ) + / len( + answer_matches ) - - else: - answer_contains_score = 1.0 + if answer_matches + else 1.0 + ) expected_urls = ( question.get( @@ -125,37 +279,28 @@ def evaluate_answer( normalized_expected_urls = [ normalize_url( - str( - url - ) + str(url) ) for url in expected_urls ] - source_matches: list[ - bool - ] = [] + source_matches = [ + expected_url + in answer + for expected_url + in normalized_expected_urls + ] - for expected_url in ( - normalized_expected_urls - ): - source_matches.append( - expected_url - in answer + source_url_score = ( + sum( + source_matches ) - - if source_matches: - source_url_score = ( - sum( - source_matches - ) - / len( - source_matches - ) + / len( + source_matches ) - - else: - source_url_score = 1.0 + if source_matches + else 1.0 + ) should_answer = bool( question.get( @@ -164,14 +309,10 @@ def evaluate_answer( ) ) - normalized_no_answer = ( + returned_no_answer = ( normalize_text( NO_ANSWER_TEXT ) - ) - - returned_no_answer = ( - normalized_no_answer in normalized_answer ) @@ -539,7 +680,8 @@ def aggregate_by_field( for ( group_name, group_rows, - ) in sorted( + ) + in sorted( grouped.items() ) }