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.
Files changed (48) hide show
  1. flashrag/__init__.py +0 -0
  2. flashrag/config/__init__.py +2 -0
  3. flashrag/config/config.py +219 -0
  4. flashrag/dataset/__init__.py +2 -0
  5. flashrag/dataset/dataset.py +160 -0
  6. flashrag/dataset/utils.py +60 -0
  7. flashrag/evaluator/__init__.py +2 -0
  8. flashrag/evaluator/_bleu.py +209 -0
  9. flashrag/evaluator/evaluator.py +82 -0
  10. flashrag/evaluator/metrics.py +508 -0
  11. flashrag/evaluator/utils.py +19 -0
  12. flashrag/generator/__init__.py +2 -0
  13. flashrag/generator/fid.py +247 -0
  14. flashrag/generator/generator.py +623 -0
  15. flashrag/generator/openai_generator.py +96 -0
  16. flashrag/generator/stop_word_criteria.py +105 -0
  17. flashrag/judger/__init__.py +1 -0
  18. flashrag/judger/judger.py +184 -0
  19. flashrag/pipeline/__init__.py +3 -0
  20. flashrag/pipeline/active_pipeline.py +980 -0
  21. flashrag/pipeline/branching_pipeline.py +250 -0
  22. flashrag/pipeline/pipeline.py +248 -0
  23. flashrag/pipeline/replug_utils.py +249 -0
  24. flashrag/prompt/__init__.py +1 -0
  25. flashrag/prompt/base_prompt.py +165 -0
  26. flashrag/prompt/selfask_examplars.py +108 -0
  27. flashrag/prompt/trace_examplars.py +4015 -0
  28. flashrag/refiner/__init__.py +2 -0
  29. flashrag/refiner/kg_refiner.py +612 -0
  30. flashrag/refiner/llmlingua_compressor.py +2360 -0
  31. flashrag/refiner/refiner.py +252 -0
  32. flashrag/refiner/selective_context_compressor.py +294 -0
  33. flashrag/retriever/__init__.py +3 -0
  34. flashrag/retriever/__main__.py +4 -0
  35. flashrag/retriever/encoder.py +120 -0
  36. flashrag/retriever/index_builder.py +334 -0
  37. flashrag/retriever/reranker.py +145 -0
  38. flashrag/retriever/retriever.py +346 -0
  39. flashrag/retriever/utils.py +49 -0
  40. flashrag/utils/__init__.py +2 -0
  41. flashrag/utils/constants.py +51 -0
  42. flashrag/utils/pred_parse.py +25 -0
  43. flashrag/utils/utils.py +135 -0
  44. flashrag_dev-0.1.1.dist-info/LICENSE +21 -0
  45. flashrag_dev-0.1.1.dist-info/METADATA +616 -0
  46. flashrag_dev-0.1.1.dist-info/RECORD +48 -0
  47. flashrag_dev-0.1.1.dist-info/WHEEL +5 -0
  48. 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))))
@@ -0,0 +1,2 @@
1
+ from flashrag.generator.generator import *
2
+ from flashrag.generator.openai_generator import *