| 33 | return white_space_fix(remove_articles(remove_punc(lower(s)))) |
| 34 | |
| 35 | class QAExactMatch: |
| 36 | metric_name: str = "qa_exact_match" |
| 37 | |
| 38 | def __init__(self): |
| 39 | self.logger = get_logger(__name__) |
| 40 | |
| 41 | def calculate_metric_scores(self, gold_answers: List[List[str]], predicted_answers: List[str], aggregation_fn: Callable = np.max) -> Tuple[Dict[str, float], List[Dict[str, float]]]: |
| 42 | """ |
| 43 | Calculate exact match (EM) scores |
| 44 | |
| 45 | Args: |
| 46 | gold_answers: List of standard answers, each element is a list of answers |
| 47 | predicted_answers: List of predicted answers |
| 48 | aggregation_fn: Function to aggregate multiple standard answers |
| 49 | |
| 50 | Returns: |
| 51 | Tuple containing: average EM score dictionary, list of EM scores for each sample |
| 52 | """ |
| 53 | assert len(gold_answers) == len(predicted_answers), "Length of gold answers and predicted answers should be the same" |
| 54 | |
| 55 | example_eval_results = [] |
| 56 | total_em = 0 |
| 57 | |
| 58 | for gold_list, predicted in zip(gold_answers, predicted_answers): |
| 59 | em_scores = [1.0 if normalize_answer(gold) == normalize_answer(predicted) else 0.0 for gold in gold_list] |
| 60 | aggregated_em = aggregation_fn(em_scores) |
| 61 | example_eval_results.append({"ExactMatch": aggregated_em}) |
| 62 | total_em += aggregated_em |
| 63 | |
| 64 | avg_em = total_em / len(gold_answers) if gold_answers else 0.0 |
| 65 | pooled_eval_results = {"ExactMatch": avg_em} |
| 66 | |
| 67 | return pooled_eval_results, example_eval_results |
| 68 | |
| 69 | class QAF1Score: |
| 70 | |