flashrag-dev 0.1.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- flashrag/__init__.py +0 -0
- flashrag/config/__init__.py +2 -0
- flashrag/config/config.py +219 -0
- flashrag/dataset/__init__.py +2 -0
- flashrag/dataset/dataset.py +160 -0
- flashrag/dataset/utils.py +60 -0
- flashrag/evaluator/__init__.py +2 -0
- flashrag/evaluator/_bleu.py +209 -0
- flashrag/evaluator/evaluator.py +82 -0
- flashrag/evaluator/metrics.py +508 -0
- flashrag/evaluator/utils.py +19 -0
- flashrag/generator/__init__.py +2 -0
- flashrag/generator/fid.py +247 -0
- flashrag/generator/generator.py +623 -0
- flashrag/generator/openai_generator.py +96 -0
- flashrag/generator/stop_word_criteria.py +105 -0
- flashrag/judger/__init__.py +1 -0
- flashrag/judger/judger.py +184 -0
- flashrag/pipeline/__init__.py +3 -0
- flashrag/pipeline/active_pipeline.py +980 -0
- flashrag/pipeline/branching_pipeline.py +250 -0
- flashrag/pipeline/pipeline.py +248 -0
- flashrag/pipeline/replug_utils.py +249 -0
- flashrag/prompt/__init__.py +1 -0
- flashrag/prompt/base_prompt.py +165 -0
- flashrag/prompt/selfask_examplars.py +108 -0
- flashrag/prompt/trace_examplars.py +4015 -0
- flashrag/refiner/__init__.py +2 -0
- flashrag/refiner/kg_refiner.py +612 -0
- flashrag/refiner/llmlingua_compressor.py +2360 -0
- flashrag/refiner/refiner.py +252 -0
- flashrag/refiner/selective_context_compressor.py +294 -0
- flashrag/retriever/__init__.py +3 -0
- flashrag/retriever/__main__.py +4 -0
- flashrag/retriever/encoder.py +120 -0
- flashrag/retriever/index_builder.py +334 -0
- flashrag/retriever/reranker.py +145 -0
- flashrag/retriever/retriever.py +346 -0
- flashrag/retriever/utils.py +49 -0
- flashrag/utils/__init__.py +2 -0
- flashrag/utils/constants.py +51 -0
- flashrag/utils/pred_parse.py +25 -0
- flashrag/utils/utils.py +135 -0
- flashrag_dev-0.1.1.dist-info/LICENSE +21 -0
- flashrag_dev-0.1.1.dist-info/METADATA +616 -0
- flashrag_dev-0.1.1.dist-info/RECORD +48 -0
- flashrag_dev-0.1.1.dist-info/WHEEL +5 -0
- flashrag_dev-0.1.1.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from flashrag.evaluator.metrics import BaseMetric
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class Evaluator:
|
|
6
|
+
"""Evaluator is used to summarize the results of all metrics."""
|
|
7
|
+
|
|
8
|
+
def __init__(self, config):
|
|
9
|
+
self.config = config
|
|
10
|
+
self.save_dir = config["save_dir"]
|
|
11
|
+
|
|
12
|
+
self.save_metric_flag = config["save_metric_score"]
|
|
13
|
+
self.save_data_flag = config["save_intermediate_data"]
|
|
14
|
+
self.metrics = [metric.lower() for metric in self.config["metrics"]]
|
|
15
|
+
|
|
16
|
+
self.avaliable_metrics = self._collect_metrics()
|
|
17
|
+
|
|
18
|
+
self.metric_class = {}
|
|
19
|
+
for metric in self.metrics:
|
|
20
|
+
if metric in self.avaliable_metrics:
|
|
21
|
+
self.metric_class[metric] = self.avaliable_metrics[metric](self.config)
|
|
22
|
+
else:
|
|
23
|
+
print(f"{metric} has not been implemented!")
|
|
24
|
+
raise NotImplementedError
|
|
25
|
+
|
|
26
|
+
def _collect_metrics(self):
|
|
27
|
+
"""Collect all classes based on ```BaseMetric```."""
|
|
28
|
+
|
|
29
|
+
def find_descendants(base_class, subclasses=None):
|
|
30
|
+
if subclasses is None:
|
|
31
|
+
subclasses = set()
|
|
32
|
+
|
|
33
|
+
direct_subclasses = base_class.__subclasses__()
|
|
34
|
+
for subclass in direct_subclasses:
|
|
35
|
+
if subclass not in subclasses:
|
|
36
|
+
subclasses.add(subclass)
|
|
37
|
+
find_descendants(subclass, subclasses)
|
|
38
|
+
return subclasses
|
|
39
|
+
|
|
40
|
+
avaliable_metrics = {}
|
|
41
|
+
for cls in find_descendants(BaseMetric):
|
|
42
|
+
metric_name = cls.metric_name
|
|
43
|
+
avaliable_metrics[metric_name] = cls
|
|
44
|
+
return avaliable_metrics
|
|
45
|
+
|
|
46
|
+
def evaluate(self, data):
|
|
47
|
+
"""Calculate all metric indicators and summarize them."""
|
|
48
|
+
|
|
49
|
+
result_dict = {}
|
|
50
|
+
for metric in self.metrics:
|
|
51
|
+
try:
|
|
52
|
+
metric_result, metric_scores = self.metric_class[metric].calculate_metric(data)
|
|
53
|
+
result_dict.update(metric_result)
|
|
54
|
+
|
|
55
|
+
for metric_score, item in zip(metric_scores, data):
|
|
56
|
+
item.update_evaluation_score(metric, metric_score)
|
|
57
|
+
except Exception as e:
|
|
58
|
+
print(f"Error in {metric}!")
|
|
59
|
+
print(e)
|
|
60
|
+
continue
|
|
61
|
+
|
|
62
|
+
if self.save_metric_flag:
|
|
63
|
+
self.save_metric_score(result_dict)
|
|
64
|
+
|
|
65
|
+
if self.save_data_flag:
|
|
66
|
+
self.save_data(data)
|
|
67
|
+
|
|
68
|
+
return result_dict
|
|
69
|
+
|
|
70
|
+
def save_metric_score(self, result_dict, file_name="metric_score.txt"):
|
|
71
|
+
save_path = os.path.join(self.save_dir, file_name)
|
|
72
|
+
with open(save_path, "w", encoding="utf-8") as f:
|
|
73
|
+
for k, v in result_dict.items():
|
|
74
|
+
f.write(f"{k}: {v}\n")
|
|
75
|
+
|
|
76
|
+
def save_data(self, data, file_name="intermediate_data.json"):
|
|
77
|
+
"""Save the evaluated data, including the raw data and the score of each data
|
|
78
|
+
sample on each metric."""
|
|
79
|
+
|
|
80
|
+
save_path = os.path.join(self.save_dir, file_name)
|
|
81
|
+
|
|
82
|
+
data.save(save_path)
|
|
@@ -0,0 +1,508 @@
|
|
|
1
|
+
import re
|
|
2
|
+
import warnings
|
|
3
|
+
from collections import Counter
|
|
4
|
+
from flashrag.evaluator.utils import normalize_answer
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class BaseMetric:
|
|
8
|
+
"""`BaseMetric` serves as the base object of all metrics. Implemented metric should
|
|
9
|
+
inherit this class.
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
metric_name = "base"
|
|
13
|
+
|
|
14
|
+
def __init__(self, config):
|
|
15
|
+
self.config = config
|
|
16
|
+
self.dataset_name = config["dataset_name"]
|
|
17
|
+
|
|
18
|
+
def calculate_metric(self, data):
|
|
19
|
+
"""Get the total score of this metric and score for each sample.
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
data object: it contains basic information and generated information.
|
|
23
|
+
|
|
24
|
+
Returns:
|
|
25
|
+
(metric_score: dict, metric_score_list: list)
|
|
26
|
+
metric_score: such as ``{'em': 0.53}``.
|
|
27
|
+
metric_score_list: score for each sample.
|
|
28
|
+
|
|
29
|
+
"""
|
|
30
|
+
return {}, []
|
|
31
|
+
|
|
32
|
+
def get_dataset_answer(self, data):
|
|
33
|
+
if any(choice == [] for choice in data.choices):
|
|
34
|
+
golden_answers_list = data.golden_answers
|
|
35
|
+
else:
|
|
36
|
+
# multi-choice dataset
|
|
37
|
+
all_choices_list = data.choices
|
|
38
|
+
golden_choice_idx_list = data.golden_answers
|
|
39
|
+
golden_answers_list = [
|
|
40
|
+
[choices[idx] for idx in idx_list]
|
|
41
|
+
for choices, idx_list in zip(all_choices_list, golden_choice_idx_list)
|
|
42
|
+
]
|
|
43
|
+
|
|
44
|
+
return golden_answers_list
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class F1_Score(BaseMetric):
|
|
48
|
+
"""Token-level F1 score"""
|
|
49
|
+
|
|
50
|
+
metric_name = "f1"
|
|
51
|
+
|
|
52
|
+
def __init__(self, config):
|
|
53
|
+
super().__init__(config)
|
|
54
|
+
|
|
55
|
+
def token_level_scores(self, prediction: str, ground_truths: str):
|
|
56
|
+
final_metric = {"f1": 0, "precision": 0, "recall": 0}
|
|
57
|
+
if isinstance(ground_truths, str):
|
|
58
|
+
ground_truths = [ground_truths]
|
|
59
|
+
for ground_truth in ground_truths:
|
|
60
|
+
normalized_prediction = normalize_answer(prediction)
|
|
61
|
+
normalized_ground_truth = normalize_answer(ground_truth)
|
|
62
|
+
|
|
63
|
+
if normalized_prediction in ["yes", "no", "noanswer"] and normalized_prediction != normalized_ground_truth:
|
|
64
|
+
continue
|
|
65
|
+
if (
|
|
66
|
+
normalized_ground_truth in ["yes", "no", "noanswer"]
|
|
67
|
+
and normalized_prediction != normalized_ground_truth
|
|
68
|
+
):
|
|
69
|
+
continue
|
|
70
|
+
prediction_tokens = normalized_prediction.split()
|
|
71
|
+
ground_truth_tokens = normalized_ground_truth.split()
|
|
72
|
+
common = Counter(prediction_tokens) & Counter(ground_truth_tokens)
|
|
73
|
+
num_same = sum(common.values())
|
|
74
|
+
if num_same == 0:
|
|
75
|
+
continue
|
|
76
|
+
precision = 1.0 * num_same / len(prediction_tokens)
|
|
77
|
+
recall = 1.0 * num_same / len(ground_truth_tokens)
|
|
78
|
+
f1 = (2 * precision * recall) / (precision + recall)
|
|
79
|
+
for k in ["f1", "precision", "recall"]:
|
|
80
|
+
final_metric[k] = max(eval(k), final_metric[k])
|
|
81
|
+
return final_metric
|
|
82
|
+
|
|
83
|
+
def calculate_metric(self, data):
|
|
84
|
+
pred_list = data.pred
|
|
85
|
+
golden_answers_list = self.get_dataset_answer(data)
|
|
86
|
+
|
|
87
|
+
metric_score_list = [
|
|
88
|
+
self.token_level_scores(pred, golden_answers)["f1"]
|
|
89
|
+
for pred, golden_answers in zip(pred_list, golden_answers_list)
|
|
90
|
+
]
|
|
91
|
+
f1 = sum(metric_score_list) / len(metric_score_list)
|
|
92
|
+
return {"f1": f1}, metric_score_list
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
class Recall_Score(F1_Score):
|
|
96
|
+
"""Token-level Recall score"""
|
|
97
|
+
|
|
98
|
+
metric_name = "recall"
|
|
99
|
+
|
|
100
|
+
def __init__(self, config):
|
|
101
|
+
super().__init__(config)
|
|
102
|
+
|
|
103
|
+
def calculate_metric(self, data):
|
|
104
|
+
pred_list = data.pred
|
|
105
|
+
golden_answers_list = self.get_dataset_answer(data)
|
|
106
|
+
metric_score_list = [
|
|
107
|
+
self.token_level_scores(pred, golden_answers)["recall"]
|
|
108
|
+
for pred, golden_answers in zip(pred_list, golden_answers_list)
|
|
109
|
+
]
|
|
110
|
+
precision = sum(metric_score_list) / len(metric_score_list)
|
|
111
|
+
return {"recall": precision}, metric_score_list
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
class Precision_Score(F1_Score):
|
|
115
|
+
"""Token-level Precision score"""
|
|
116
|
+
|
|
117
|
+
metric_name = "precision"
|
|
118
|
+
|
|
119
|
+
def __init__(self, config):
|
|
120
|
+
super().__init__(config)
|
|
121
|
+
|
|
122
|
+
def calculate_metric(self, data):
|
|
123
|
+
pred_list = data.pred
|
|
124
|
+
golden_answers_list = self.get_dataset_answer(data)
|
|
125
|
+
metric_score_list = [
|
|
126
|
+
self.token_level_scores(pred, golden_answers)["precision"]
|
|
127
|
+
for pred, golden_answers in zip(pred_list, golden_answers_list)
|
|
128
|
+
]
|
|
129
|
+
precision = sum(metric_score_list) / len(metric_score_list)
|
|
130
|
+
return {"precision": precision}, metric_score_list
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
class ExactMatch(BaseMetric):
|
|
134
|
+
r"""Exact match measure whether the predicted answer is completely consistent
|
|
135
|
+
with the standard answer.
|
|
136
|
+
|
|
137
|
+
"""
|
|
138
|
+
|
|
139
|
+
metric_name = "em"
|
|
140
|
+
|
|
141
|
+
def __init__(self, config):
|
|
142
|
+
super().__init__(config)
|
|
143
|
+
self.is_regex = self.dataset_name == "curatedtrec"
|
|
144
|
+
|
|
145
|
+
def calculate_em(self, prediction: str, golden_answers: list) -> float:
|
|
146
|
+
if isinstance(golden_answers, str):
|
|
147
|
+
golden_answers = [golden_answers]
|
|
148
|
+
normalized_prediction = normalize_answer(prediction)
|
|
149
|
+
score = 0.0
|
|
150
|
+
for golden_answer in golden_answers:
|
|
151
|
+
if self.is_regex:
|
|
152
|
+
print("Consider answer as regex!")
|
|
153
|
+
golden_answer = re.compile(golden_answer, re.IGNORECASE)
|
|
154
|
+
match = re.fullmatch(golden_answer, normalized_prediction)
|
|
155
|
+
if match is not None:
|
|
156
|
+
score = 1.0
|
|
157
|
+
break
|
|
158
|
+
else:
|
|
159
|
+
golden_answer = normalize_answer(golden_answer)
|
|
160
|
+
if golden_answer == normalized_prediction:
|
|
161
|
+
score = 1.0
|
|
162
|
+
break
|
|
163
|
+
return score
|
|
164
|
+
|
|
165
|
+
def calculate_metric(self, data):
|
|
166
|
+
pred_list = data.pred
|
|
167
|
+
golden_answers_list = self.get_dataset_answer(data)
|
|
168
|
+
|
|
169
|
+
metric_score_list = [
|
|
170
|
+
self.calculate_em(pred, golden_answers) for pred, golden_answers in zip(pred_list, golden_answers_list)
|
|
171
|
+
]
|
|
172
|
+
em_score = sum(metric_score_list) / len(metric_score_list)
|
|
173
|
+
|
|
174
|
+
return {"em": em_score}, metric_score_list
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
class Sub_ExactMatch(BaseMetric):
|
|
178
|
+
r"""Sub-Exact match measure whether the predicted answer contains the standard answer."""
|
|
179
|
+
|
|
180
|
+
metric_name = "acc"
|
|
181
|
+
|
|
182
|
+
def __init__(self, config):
|
|
183
|
+
super().__init__(config)
|
|
184
|
+
self.is_regex = self.dataset_name == "curatedtrec"
|
|
185
|
+
|
|
186
|
+
def calculate_sub_em(self, prediction: str, golden_answers: list) -> float:
|
|
187
|
+
if isinstance(golden_answers, str):
|
|
188
|
+
golden_answers = [golden_answers]
|
|
189
|
+
normalized_prediction = normalize_answer(prediction)
|
|
190
|
+
score = 0.0
|
|
191
|
+
for golden_answer in golden_answers:
|
|
192
|
+
if self.is_regex:
|
|
193
|
+
print("Consider answer as regex!")
|
|
194
|
+
golden_answer = re.compile(golden_answer, re.IGNORECASE)
|
|
195
|
+
match = re.search(golden_answer, normalized_prediction)
|
|
196
|
+
if match is not None:
|
|
197
|
+
score = 1.0
|
|
198
|
+
break
|
|
199
|
+
else:
|
|
200
|
+
golden_answer = normalize_answer(golden_answer)
|
|
201
|
+
if golden_answer in normalized_prediction:
|
|
202
|
+
score = 1.0
|
|
203
|
+
break
|
|
204
|
+
return score
|
|
205
|
+
|
|
206
|
+
def calculate_metric(self, data):
|
|
207
|
+
golden_answers_list = self.get_dataset_answer(data)
|
|
208
|
+
pred_list = data.pred
|
|
209
|
+
|
|
210
|
+
metric_score_list = [
|
|
211
|
+
self.calculate_sub_em(pred, golden_answers) for pred, golden_answers in zip(pred_list, golden_answers_list)
|
|
212
|
+
]
|
|
213
|
+
sub_em_score = sum(metric_score_list) / len(metric_score_list)
|
|
214
|
+
|
|
215
|
+
return {"acc": sub_em_score}, metric_score_list
|
|
216
|
+
|
|
217
|
+
|
|
218
|
+
class Retrieval_Recall(BaseMetric):
|
|
219
|
+
r"""The recall of the top-k retreived passages, we measure if any of the passage contain the answer string."""
|
|
220
|
+
|
|
221
|
+
metric_name = "retrieval_recall"
|
|
222
|
+
|
|
223
|
+
def __init__(self, config):
|
|
224
|
+
super().__init__(config)
|
|
225
|
+
self.topk = config["metric_setting"]["retrieval_recall_topk"]
|
|
226
|
+
|
|
227
|
+
def calculate_metric(self, data):
|
|
228
|
+
golden_answers_list = self.get_dataset_answer(data)
|
|
229
|
+
retrieve_docs = data.retrieval_result
|
|
230
|
+
recall_score_list = []
|
|
231
|
+
for doc_list, golden_answers in zip(retrieve_docs, golden_answers_list):
|
|
232
|
+
if len(doc_list) < self.topk:
|
|
233
|
+
warnings.warn(f"Length of retrieved docs is smaller than topk ({self.topk})")
|
|
234
|
+
doc_list = [doc["contents"] for doc in doc_list[: self.topk]]
|
|
235
|
+
hit_list = []
|
|
236
|
+
for doc in doc_list:
|
|
237
|
+
for golden_answer in golden_answers:
|
|
238
|
+
if normalize_answer(golden_answer) in normalize_answer(doc):
|
|
239
|
+
hit_list.append(True)
|
|
240
|
+
break
|
|
241
|
+
else:
|
|
242
|
+
hit_list.append(False)
|
|
243
|
+
score = 1 if any(hit_list) else 0
|
|
244
|
+
recall_score_list.append(score)
|
|
245
|
+
recall_score = sum(recall_score_list) / len(recall_score_list)
|
|
246
|
+
|
|
247
|
+
return {f"retrieval_recall_top{self.topk}": recall_score}, recall_score_list
|
|
248
|
+
|
|
249
|
+
|
|
250
|
+
class Retrieval_Precision(BaseMetric):
|
|
251
|
+
r"""The precision of the top-k retreived passages, we measure if any of the passage contain the answer string."""
|
|
252
|
+
|
|
253
|
+
metric_name = "retrieval_precision"
|
|
254
|
+
|
|
255
|
+
def __init__(self, config):
|
|
256
|
+
super().__init__(config)
|
|
257
|
+
self.topk = config["metric_setting"]["retrieval_recall_topk"]
|
|
258
|
+
|
|
259
|
+
def calculate_metric(self, data):
|
|
260
|
+
golden_answers_list = self.get_dataset_answer(data)
|
|
261
|
+
retrieve_docs = data.retrieval_result
|
|
262
|
+
precision_score_list = []
|
|
263
|
+
for doc_list, golden_answers in zip(retrieve_docs, golden_answers_list):
|
|
264
|
+
if len(doc_list) < self.topk:
|
|
265
|
+
warnings.warn(f"Length of retrieved docs is smaller than topk ({self.topk})")
|
|
266
|
+
doc_list = [doc["contents"] for doc in doc_list[: self.topk]]
|
|
267
|
+
hit_list = []
|
|
268
|
+
for doc in doc_list:
|
|
269
|
+
for golden_answer in golden_answers:
|
|
270
|
+
if normalize_answer(golden_answer) in normalize_answer(doc):
|
|
271
|
+
hit_list.append(True)
|
|
272
|
+
break
|
|
273
|
+
else:
|
|
274
|
+
hit_list.append(False)
|
|
275
|
+
score = sum(hit_list) / len(hit_list)
|
|
276
|
+
precision_score_list.append(score)
|
|
277
|
+
precision_score = sum(precision_score_list) / len(precision_score_list)
|
|
278
|
+
|
|
279
|
+
return {f"retrieval_precision_top{self.topk}": precision_score}, precision_score_list
|
|
280
|
+
|
|
281
|
+
|
|
282
|
+
class Rouge_Score(BaseMetric):
|
|
283
|
+
metric_name = "rouge_score"
|
|
284
|
+
|
|
285
|
+
def __init__(self, config):
|
|
286
|
+
super().__init__(config)
|
|
287
|
+
from rouge import Rouge
|
|
288
|
+
|
|
289
|
+
self.scorer = Rouge()
|
|
290
|
+
|
|
291
|
+
def calculate_rouge(self, pred, golden_answers):
|
|
292
|
+
output = {}
|
|
293
|
+
for answer in golden_answers:
|
|
294
|
+
scores = self.scorer.get_scores(pred, answer)
|
|
295
|
+
for key in ["rouge-1", "rouge-2", "rouge-l"]:
|
|
296
|
+
if key not in output:
|
|
297
|
+
output[key] = []
|
|
298
|
+
output[key].append(scores[0][key]["f"])
|
|
299
|
+
for k, v in output.items():
|
|
300
|
+
output[k] = max(v)
|
|
301
|
+
|
|
302
|
+
return output
|
|
303
|
+
|
|
304
|
+
|
|
305
|
+
class Rouge_1(Rouge_Score):
|
|
306
|
+
metric_name = "rouge-1"
|
|
307
|
+
|
|
308
|
+
def __init__(self, config):
|
|
309
|
+
super().__init__(config)
|
|
310
|
+
|
|
311
|
+
def calculate_metric(self, data):
|
|
312
|
+
golden_answers_list = self.get_dataset_answer(data)
|
|
313
|
+
pred_list = data.pred
|
|
314
|
+
|
|
315
|
+
metric_score_list = [
|
|
316
|
+
self.calculate_rouge(pred, golden_answers)["rouge-1"]
|
|
317
|
+
for pred, golden_answers in zip(pred_list, golden_answers_list)
|
|
318
|
+
]
|
|
319
|
+
score = sum(metric_score_list) / len(metric_score_list)
|
|
320
|
+
|
|
321
|
+
return {"rouge-1": score}, metric_score_list
|
|
322
|
+
|
|
323
|
+
|
|
324
|
+
class Rouge_2(Rouge_Score):
|
|
325
|
+
metric_name = "rouge-2"
|
|
326
|
+
|
|
327
|
+
def __init__(self, config):
|
|
328
|
+
super().__init__(config)
|
|
329
|
+
|
|
330
|
+
def calculate_metric(self, data):
|
|
331
|
+
golden_answers_list = self.get_dataset_answer(data)
|
|
332
|
+
pred_list = data.pred
|
|
333
|
+
|
|
334
|
+
metric_score_list = [
|
|
335
|
+
self.calculate_rouge(pred, golden_answers)["rouge-2"]
|
|
336
|
+
for pred, golden_answers in zip(pred_list, golden_answers_list)
|
|
337
|
+
]
|
|
338
|
+
score = sum(metric_score_list) / len(metric_score_list)
|
|
339
|
+
|
|
340
|
+
return {"rouge-2": score}, metric_score_list
|
|
341
|
+
|
|
342
|
+
|
|
343
|
+
class Rouge_L(Rouge_Score):
|
|
344
|
+
metric_name = "rouge-l"
|
|
345
|
+
|
|
346
|
+
def __init__(self, config):
|
|
347
|
+
super().__init__(config)
|
|
348
|
+
|
|
349
|
+
def calculate_metric(self, data):
|
|
350
|
+
golden_answers_list = self.get_dataset_answer(data)
|
|
351
|
+
pred_list = data.pred
|
|
352
|
+
|
|
353
|
+
metric_score_list = [
|
|
354
|
+
self.calculate_rouge(pred, golden_answers)["rouge-l"]
|
|
355
|
+
for pred, golden_answers in zip(pred_list, golden_answers_list)
|
|
356
|
+
]
|
|
357
|
+
score = sum(metric_score_list) / len(metric_score_list)
|
|
358
|
+
|
|
359
|
+
return {"rouge-l": score}, metric_score_list
|
|
360
|
+
|
|
361
|
+
|
|
362
|
+
class BLEU(BaseMetric):
|
|
363
|
+
metric_name = "bleu"
|
|
364
|
+
|
|
365
|
+
def __init__(self, config):
|
|
366
|
+
super().__init__(config)
|
|
367
|
+
from ._bleu import Tokenizer13a
|
|
368
|
+
|
|
369
|
+
self.tokenizer = Tokenizer13a()
|
|
370
|
+
self.max_order = config["metric_setting"].get("bleu_max_order", 4)
|
|
371
|
+
self.smooth = config["metric_setting"].get("bleu_smooth", False)
|
|
372
|
+
|
|
373
|
+
def calculate_metric(self, data):
|
|
374
|
+
from ._bleu import compute_bleu
|
|
375
|
+
|
|
376
|
+
golden_answers_list = self.get_dataset_answer(data)
|
|
377
|
+
pred_list = data.pred
|
|
378
|
+
|
|
379
|
+
pred_list = [self.tokenizer(pred) for pred in pred_list]
|
|
380
|
+
golden_answers_list = [
|
|
381
|
+
[self.tokenizer(ans) for ans in golden_answers] for golden_answers in golden_answers_list
|
|
382
|
+
]
|
|
383
|
+
score = compute_bleu(
|
|
384
|
+
reference_corpus=golden_answers_list,
|
|
385
|
+
translation_corpus=pred_list,
|
|
386
|
+
max_order=self.max_order,
|
|
387
|
+
smooth=self.smooth,
|
|
388
|
+
)
|
|
389
|
+
(total_bleu, precisions, bp, ratio, translation_length, reference_length) = score
|
|
390
|
+
|
|
391
|
+
score_list = []
|
|
392
|
+
for pred, golden_answers in zip(pred_list, golden_answers_list):
|
|
393
|
+
pred = [pred]
|
|
394
|
+
golden_answers = [golden_answers]
|
|
395
|
+
score = compute_bleu(
|
|
396
|
+
reference_corpus=golden_answers_list,
|
|
397
|
+
translation_corpus=pred_list,
|
|
398
|
+
max_order=self.max_order,
|
|
399
|
+
smooth=self.smooth,
|
|
400
|
+
)
|
|
401
|
+
(bleu, precisions, bp, ratio, translation_length, reference_length) = score
|
|
402
|
+
score_list.append(bleu)
|
|
403
|
+
|
|
404
|
+
return {"bleu": total_bleu}, score_list
|
|
405
|
+
|
|
406
|
+
|
|
407
|
+
class LLMJudge(BaseMetric):
|
|
408
|
+
metric_name = "llm_judge"
|
|
409
|
+
JUDGE_PROMPT = """
|
|
410
|
+
You will be given a user_question and system_answer couple.
|
|
411
|
+
Your task is to provide a 'total rating' scoring how well the system_answer answers the user concerns expressed in the user_question.
|
|
412
|
+
Give your answer as a float on a scale of 0 to 10, where 0 means that the system_answer is not helpful at all, and 10 means that the answer completely and helpfully addresses the question.
|
|
413
|
+
|
|
414
|
+
Provide your feedback as follows:
|
|
415
|
+
|
|
416
|
+
Feedback:::
|
|
417
|
+
Total rating: (your rating, as a float between 0 and 10)
|
|
418
|
+
|
|
419
|
+
Now here are the question and answer.
|
|
420
|
+
|
|
421
|
+
Question: {question}
|
|
422
|
+
Answer: {answer}
|
|
423
|
+
|
|
424
|
+
Feedback:::
|
|
425
|
+
Total rating: """
|
|
426
|
+
|
|
427
|
+
def __init__(self, config):
|
|
428
|
+
super().__init__(config)
|
|
429
|
+
if "llm_judge_setting" in config["metric_setting"]:
|
|
430
|
+
llm_setting = config["metric_setting"]["llm_judge_setting"]
|
|
431
|
+
else:
|
|
432
|
+
assert False, "No available LLM settings!"
|
|
433
|
+
# TODO: integrate generator class
|
|
434
|
+
llm_name = llm_setting["model_name"]
|
|
435
|
+
if "model_path" not in llm_setting:
|
|
436
|
+
model_path = config["model2path"].get(llm_name, None)
|
|
437
|
+
else:
|
|
438
|
+
model_path = llm_setting["model_path"]
|
|
439
|
+
if model_path is None:
|
|
440
|
+
assert False, "None model path "
|
|
441
|
+
|
|
442
|
+
from transformers import pipeline
|
|
443
|
+
|
|
444
|
+
self.llm_pipeline = pipeline("text2text-generation", model=model_path, device=0)
|
|
445
|
+
|
|
446
|
+
def extract_judge_score(answer: str, split_str: str = "Total rating:") -> int:
|
|
447
|
+
try:
|
|
448
|
+
if split_str in answer:
|
|
449
|
+
rating = answer.split(split_str)[1]
|
|
450
|
+
else:
|
|
451
|
+
rating = answer
|
|
452
|
+
digit_groups = [el.strip() for el in re.findall(r"\d+(?:\.\d+)?", rating)]
|
|
453
|
+
return float(digit_groups[0])
|
|
454
|
+
except Exception as e:
|
|
455
|
+
print(e)
|
|
456
|
+
return 0
|
|
457
|
+
|
|
458
|
+
def calculate_metric(self, data):
|
|
459
|
+
question_list = data.question
|
|
460
|
+
pred_list = data.pred
|
|
461
|
+
|
|
462
|
+
judge_input_prompt = [self.JUDGE_PROMPT.format(question=q, answer=a) for q, a in zip(question_list, pred_list)]
|
|
463
|
+
judge_output = self.llm_pipeline(judge_input_prompt, max_new_tokens=100, batch_size=8)
|
|
464
|
+
judge_output = [item["generated_text"] for item in judge_output]
|
|
465
|
+
|
|
466
|
+
metric_score_list = [self.extract_judge_score(o) for o in judge_output]
|
|
467
|
+
# rescale score
|
|
468
|
+
metric_score_list = [score / 10 + 1 for score in metric_score_list]
|
|
469
|
+
|
|
470
|
+
score = sum(metric_score_list) / len(metric_score_list)
|
|
471
|
+
|
|
472
|
+
return {"llm_judge_score": score}, metric_score_list
|
|
473
|
+
|
|
474
|
+
|
|
475
|
+
class CountToken(BaseMetric):
|
|
476
|
+
metric_name = "input_tokens"
|
|
477
|
+
|
|
478
|
+
def __init__(self, config):
|
|
479
|
+
super().__init__(config)
|
|
480
|
+
tokenizer_name = config["metric_setting"].get("tokenizer_name", None)
|
|
481
|
+
is_hf_tokenizer = True
|
|
482
|
+
from flashrag.utils.constants import OPENAI_MODEL_DICT
|
|
483
|
+
|
|
484
|
+
if tokenizer_name is None or tokenizer_name in OPENAI_MODEL_DICT:
|
|
485
|
+
# use gpt4 tokenizer
|
|
486
|
+
import tiktoken
|
|
487
|
+
|
|
488
|
+
if tokenizer_name is None:
|
|
489
|
+
tokenizer_name = "gpt-4"
|
|
490
|
+
tokenizer = tiktoken.encoding_for_model(tokenizer_name)
|
|
491
|
+
is_hf_tokenizer = False
|
|
492
|
+
else:
|
|
493
|
+
from transformers import AutoTokenizer
|
|
494
|
+
|
|
495
|
+
tokenizer = AutoTokenizer.from_pretrained(tokenizer_name)
|
|
496
|
+
|
|
497
|
+
self.tokenizer = tokenizer
|
|
498
|
+
self.is_hf_tokenizer = is_hf_tokenizer
|
|
499
|
+
|
|
500
|
+
def calculate_metric(self, data):
|
|
501
|
+
input_prompts = data.prompt
|
|
502
|
+
if self.is_hf_tokenizer:
|
|
503
|
+
token_counts = [len(self.tokenizer.tokenize(text)) for text in input_prompts]
|
|
504
|
+
else:
|
|
505
|
+
token_counts = [len(self.tokenizer.encode(text)) for text in input_prompts]
|
|
506
|
+
avg_tokens = sum(token_counts) / len(token_counts)
|
|
507
|
+
|
|
508
|
+
return {"avg_input_tokens": avg_tokens}, token_counts
|
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
import re
|
|
2
|
+
import string
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
def normalize_answer(s):
|
|
6
|
+
def remove_articles(text):
|
|
7
|
+
return re.sub(r"\b(a|an|the)\b", " ", text)
|
|
8
|
+
|
|
9
|
+
def white_space_fix(text):
|
|
10
|
+
return " ".join(text.split())
|
|
11
|
+
|
|
12
|
+
def remove_punc(text):
|
|
13
|
+
exclude = set(string.punctuation)
|
|
14
|
+
return "".join(ch for ch in text if ch not in exclude)
|
|
15
|
+
|
|
16
|
+
def lower(text):
|
|
17
|
+
return text.lower()
|
|
18
|
+
|
|
19
|
+
return white_space_fix(remove_articles(remove_punc(lower(s))))
|