Akjava/open_Deep-Research-DuckDuckGo-Groq
2
1import re2import string3import warnings4 5 6def normalize_number_str(number_str: str) -> float:7 # we replace these common units and commas to allow8 # conversion to float9 for char in ["$", "%", ","]:10 number_str = number_str.replace(char, "")11 try:12 return float(number_str)13 except ValueError:14 print(f"String {number_str} cannot be normalized to number str.")15 return float("inf")16 17 18def split_string(19 s: str,20 char_list: list[str] = [",", ";"],21) -> list[str]:22 pattern = f"[{''.join(char_list)}]"23 return re.split(pattern, s)24 25 26def is_float(element: any) -> bool:27 try:28 float(element)29 return True30 except ValueError:31 return False32 33 34def question_scorer(35 model_answer: str,36 ground_truth: str,37) -> bool:38 # if gt is a number39 if is_float(ground_truth):40 normalized_answer = normalize_number_str(str(model_answer))41 return normalized_answer == float(ground_truth)42 43 # if gt is a list44 elif any(char in ground_truth for char in [",", ";"]):45 # question with the fish: normalization removes punct46 47 gt_elems = split_string(ground_truth)48 ma_elems = split_string(model_answer)49 50 # check length is the same51 if len(gt_elems) != len(ma_elems):52 warnings.warn("Answer lists have different lengths, returning False.", UserWarning)53 return False54 55 # compare each element as float or str56 comparisons = []57 for ma_elem, gt_elem in zip(ma_elems, gt_elems):58 if is_float(gt_elem):59 normalized_ma_elem = normalize_number_str(ma_elem)60 comparisons.append(normalized_ma_elem == float(gt_elem))61 else:62 # we do not remove punct since comparisons can include punct63 comparisons.append(64 normalize_str(ma_elem, remove_punct=False) == normalize_str(gt_elem, remove_punct=False)65 )66 return all(comparisons)67 68 # if gt is a str69 else:70 return normalize_str(model_answer) == normalize_str(ground_truth)71 72 73def check_prediction_contains_answer_letters_in_order(prediction, true_answer):74 prediction = prediction.lower()75 true_answer = true_answer.lower()76 if len(prediction) > len(true_answer) * 3:77 return False78 i = 079 for letter in true_answer:80 if letter in prediction[i:]:81 i += prediction[i:].index(letter)82 else:83 return False84 return True85 86 87def check_close_call(prediction, true_answer, is_correct):88 if is_correct:89 return True90 else:91 if is_float(true_answer):92 return is_correct93 else:94 if (95 check_prediction_contains_answer_letters_in_order(str(prediction), str(true_answer))96 and len(str(true_answer)) * 0.5 <= len(str(prediction)) <= len(str(true_answer)) * 297 ):98 print(f"Close call: {prediction} vs {true_answer}")99 return True100 else:101 return False102 103 104def normalize_str(input_str, remove_punct=True) -> str:105 """106 Normalize a string by:107 - Removing all white spaces108 - Optionally removing punctuation (if remove_punct is True)109 - Converting to lowercase110 Parameters:111 - input_str: str, the string to normalize112 - remove_punct: bool, whether to remove punctuation (default: True)113 Returns:114 - str, the normalized string115 """116 # Remove all white spaces. Required e.g for seagull vs. sea gull117 no_spaces = re.sub(r"\s", "", input_str)118 119 # Remove punctuation, if specified.120 if remove_punct:121 translator = str.maketrans("", "", string.punctuation)122 return no_spaces.lower().translate(translator)123 else:124 return no_spaces.lower()125 