Team Ai
Apppublic

Akjava/open_Deep-Research-DuckDuckGo-Groq

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
2likes
gaia_scorer.py125 linesDownload Raw Back to scripts
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