| 77 | |
| 78 | @TEXT_POSTPROCESSORS.register_module() |
| 79 | def answer_cleansing( |
| 80 | method: str, |
| 81 | prediction: str, |
| 82 | options: list, |
| 83 | label: str, |
| 84 | ) -> str: |
| 85 | |
| 86 | # Clean up unwanted phrases in the prediction |
| 87 | for unwanted_phrase in [ |
| 88 | 'I understand', |
| 89 | 'A through J', |
| 90 | 'A through E', |
| 91 | 'A through D', |
| 92 | ]: |
| 93 | prediction = prediction.replace(unwanted_phrase, '') |
| 94 | |
| 95 | options_num = len(options) |
| 96 | options = [chr(65 + i) for i in range(options_num)] |
| 97 | options_str = r'\b(' + '|'.join(options) + r')\b' |
| 98 | prediction = re.findall(options_str, prediction) |
| 99 | |
| 100 | if len(prediction) == 0: |
| 101 | prediction = [] |
| 102 | return prediction |
| 103 | else: |
| 104 | # If there is a "label" and its length is 1, |
| 105 | # process prediction accordingly |
| 106 | if len(label) == 1: |
| 107 | if method == 'few-shot': |
| 108 | answer_flag = True if len(prediction) > 1 else False |
| 109 | # choose the first or last element based on the answer_flag |
| 110 | if answer_flag: |
| 111 | prediction = [prediction[0]] |
| 112 | else: |
| 113 | prediction = [prediction[-1]] |
| 114 | elif method == 'zero-shot': |
| 115 | # choose the first element in list |
| 116 | prediction = [prediction[0]] |
| 117 | else: |
| 118 | raise ValueError('Method is not properly defined ...') |
| 119 | |
| 120 | # Remove trailing period if it exists |
| 121 | if prediction[0] and prediction[0].endswith('.'): |
| 122 | prediction[0] = prediction[0][:-1] |
| 123 | |
| 124 | return prediction[0] |
| 125 | |
| 126 | |
| 127 | def _generic_llmjudge_postprocess(judgement: str): |