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