AutoRAG 0.0.0__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 (184) hide show
  1. autorag/__init__.py +82 -0
  2. autorag/chunker.py +51 -0
  3. autorag/cli.py +209 -0
  4. autorag/dashboard.py +199 -0
  5. autorag/data/__init__.py +109 -0
  6. autorag/data/chunk/__init__.py +2 -0
  7. autorag/data/chunk/base.py +128 -0
  8. autorag/data/chunk/langchain_chunk.py +76 -0
  9. autorag/data/chunk/llama_index_chunk.py +96 -0
  10. autorag/data/chunk/run.py +38 -0
  11. autorag/data/legacy/__init__.py +0 -0
  12. autorag/data/legacy/corpus/__init__.py +2 -0
  13. autorag/data/legacy/corpus/langchain.py +47 -0
  14. autorag/data/legacy/corpus/llama_index.py +93 -0
  15. autorag/data/legacy/qacreation/__init__.py +6 -0
  16. autorag/data/legacy/qacreation/base.py +239 -0
  17. autorag/data/legacy/qacreation/llama_index.py +253 -0
  18. autorag/data/legacy/qacreation/llama_index_default_prompt.txt +54 -0
  19. autorag/data/legacy/qacreation/ragas.py +75 -0
  20. autorag/data/legacy/qacreation/simple.py +99 -0
  21. autorag/data/parse/__init__.py +1 -0
  22. autorag/data/parse/base.py +79 -0
  23. autorag/data/parse/clova.py +194 -0
  24. autorag/data/parse/langchain_parse.py +87 -0
  25. autorag/data/parse/llamaparse.py +126 -0
  26. autorag/data/parse/run.py +141 -0
  27. autorag/data/parse/table_hybrid_parse.py +134 -0
  28. autorag/data/qa/__init__.py +3 -0
  29. autorag/data/qa/evolve/__init__.py +0 -0
  30. autorag/data/qa/evolve/llama_index_query_evolve.py +64 -0
  31. autorag/data/qa/evolve/openai_query_evolve.py +81 -0
  32. autorag/data/qa/evolve/prompt.py +288 -0
  33. autorag/data/qa/extract_evidence.py +1 -0
  34. autorag/data/qa/filter/__init__.py +0 -0
  35. autorag/data/qa/filter/dontknow.py +117 -0
  36. autorag/data/qa/filter/passage_dependency.py +88 -0
  37. autorag/data/qa/filter/prompt.py +73 -0
  38. autorag/data/qa/generation_gt/__init__.py +0 -0
  39. autorag/data/qa/generation_gt/base.py +16 -0
  40. autorag/data/qa/generation_gt/llama_index_gen_gt.py +41 -0
  41. autorag/data/qa/generation_gt/openai_gen_gt.py +84 -0
  42. autorag/data/qa/generation_gt/prompt.py +27 -0
  43. autorag/data/qa/query/__init__.py +0 -0
  44. autorag/data/qa/query/llama_gen_query.py +82 -0
  45. autorag/data/qa/query/openai_gen_query.py +95 -0
  46. autorag/data/qa/query/prompt.py +201 -0
  47. autorag/data/qa/sample.py +26 -0
  48. autorag/data/qa/schema.py +322 -0
  49. autorag/data/utils/__init__.py +0 -0
  50. autorag/data/utils/util.py +103 -0
  51. autorag/deploy/__init__.py +9 -0
  52. autorag/deploy/api.py +303 -0
  53. autorag/deploy/base.py +235 -0
  54. autorag/deploy/gradio.py +74 -0
  55. autorag/deploy/swagger.yml +202 -0
  56. autorag/embedding/__init__.py +0 -0
  57. autorag/embedding/base.py +144 -0
  58. autorag/embedding/vllm.py +256 -0
  59. autorag/evaluation/__init__.py +3 -0
  60. autorag/evaluation/generation.py +88 -0
  61. autorag/evaluation/metric/__init__.py +22 -0
  62. autorag/evaluation/metric/deepeval_prompt.py +322 -0
  63. autorag/evaluation/metric/g_eval_prompts/coh_detailed.txt +32 -0
  64. autorag/evaluation/metric/g_eval_prompts/con_detailed.txt +33 -0
  65. autorag/evaluation/metric/g_eval_prompts/flu_detailed.txt +26 -0
  66. autorag/evaluation/metric/g_eval_prompts/rel_detailed.txt +33 -0
  67. autorag/evaluation/metric/generation.py +504 -0
  68. autorag/evaluation/metric/retrieval.py +115 -0
  69. autorag/evaluation/metric/retrieval_contents.py +65 -0
  70. autorag/evaluation/metric/util.py +88 -0
  71. autorag/evaluation/retrieval.py +83 -0
  72. autorag/evaluation/retrieval_contents.py +65 -0
  73. autorag/evaluation/util.py +43 -0
  74. autorag/evaluator.py +559 -0
  75. autorag/node_line.py +65 -0
  76. autorag/nodes/__init__.py +0 -0
  77. autorag/nodes/generator/__init__.py +4 -0
  78. autorag/nodes/generator/base.py +103 -0
  79. autorag/nodes/generator/llama_index_llm.py +169 -0
  80. autorag/nodes/generator/openai_llm.py +329 -0
  81. autorag/nodes/generator/run.py +148 -0
  82. autorag/nodes/generator/vllm.py +147 -0
  83. autorag/nodes/generator/vllm_api.py +191 -0
  84. autorag/nodes/hybridretrieval/__init__.py +2 -0
  85. autorag/nodes/hybridretrieval/base.py +58 -0
  86. autorag/nodes/hybridretrieval/hybrid_cc.py +227 -0
  87. autorag/nodes/hybridretrieval/hybrid_rrf.py +149 -0
  88. autorag/nodes/hybridretrieval/run.py +137 -0
  89. autorag/nodes/lexicalretrieval/__init__.py +1 -0
  90. autorag/nodes/lexicalretrieval/bm25.py +381 -0
  91. autorag/nodes/lexicalretrieval/run.py +148 -0
  92. autorag/nodes/passageaugmenter/__init__.py +2 -0
  93. autorag/nodes/passageaugmenter/base.py +76 -0
  94. autorag/nodes/passageaugmenter/pass_passage_augmenter.py +43 -0
  95. autorag/nodes/passageaugmenter/prev_next_augmenter.py +155 -0
  96. autorag/nodes/passageaugmenter/run.py +131 -0
  97. autorag/nodes/passagecompressor/__init__.py +4 -0
  98. autorag/nodes/passagecompressor/base.py +78 -0
  99. autorag/nodes/passagecompressor/longllmlingua.py +115 -0
  100. autorag/nodes/passagecompressor/pass_compressor.py +16 -0
  101. autorag/nodes/passagecompressor/refine.py +54 -0
  102. autorag/nodes/passagecompressor/run.py +186 -0
  103. autorag/nodes/passagecompressor/tree_summarize.py +56 -0
  104. autorag/nodes/passagefilter/__init__.py +6 -0
  105. autorag/nodes/passagefilter/base.py +40 -0
  106. autorag/nodes/passagefilter/pass_passage_filter.py +14 -0
  107. autorag/nodes/passagefilter/percentile_cutoff.py +58 -0
  108. autorag/nodes/passagefilter/recency.py +105 -0
  109. autorag/nodes/passagefilter/run.py +138 -0
  110. autorag/nodes/passagefilter/similarity_percentile_cutoff.py +134 -0
  111. autorag/nodes/passagefilter/similarity_threshold_cutoff.py +112 -0
  112. autorag/nodes/passagefilter/threshold_cutoff.py +78 -0
  113. autorag/nodes/passagereranker/__init__.py +16 -0
  114. autorag/nodes/passagereranker/base.py +44 -0
  115. autorag/nodes/passagereranker/cohere.py +118 -0
  116. autorag/nodes/passagereranker/colbert.py +213 -0
  117. autorag/nodes/passagereranker/flag_embedding.py +112 -0
  118. autorag/nodes/passagereranker/flag_embedding_llm.py +101 -0
  119. autorag/nodes/passagereranker/flashrank.py +245 -0
  120. autorag/nodes/passagereranker/jina.py +115 -0
  121. autorag/nodes/passagereranker/koreranker.py +136 -0
  122. autorag/nodes/passagereranker/mixedbreadai.py +126 -0
  123. autorag/nodes/passagereranker/monot5.py +190 -0
  124. autorag/nodes/passagereranker/openvino.py +191 -0
  125. autorag/nodes/passagereranker/pass_reranker.py +31 -0
  126. autorag/nodes/passagereranker/rankgpt.py +170 -0
  127. autorag/nodes/passagereranker/run.py +145 -0
  128. autorag/nodes/passagereranker/sentence_transformer.py +129 -0
  129. autorag/nodes/passagereranker/tart/__init__.py +1 -0
  130. autorag/nodes/passagereranker/tart/modeling_enc_t5.py +152 -0
  131. autorag/nodes/passagereranker/tart/tart.py +139 -0
  132. autorag/nodes/passagereranker/tart/tokenization_enc_t5.py +112 -0
  133. autorag/nodes/passagereranker/time_reranker.py +72 -0
  134. autorag/nodes/passagereranker/upr.py +160 -0
  135. autorag/nodes/passagereranker/voyageai.py +109 -0
  136. autorag/nodes/promptmaker/__init__.py +12 -0
  137. autorag/nodes/promptmaker/base.py +32 -0
  138. autorag/nodes/promptmaker/chat_fstring.py +73 -0
  139. autorag/nodes/promptmaker/fstring.py +49 -0
  140. autorag/nodes/promptmaker/long_context_reorder.py +83 -0
  141. autorag/nodes/promptmaker/run.py +283 -0
  142. autorag/nodes/promptmaker/window_replacement.py +85 -0
  143. autorag/nodes/queryexpansion/__init__.py +4 -0
  144. autorag/nodes/queryexpansion/base.py +62 -0
  145. autorag/nodes/queryexpansion/hyde.py +43 -0
  146. autorag/nodes/queryexpansion/multi_query_expansion.py +57 -0
  147. autorag/nodes/queryexpansion/pass_query_expansion.py +22 -0
  148. autorag/nodes/queryexpansion/query_decompose.py +111 -0
  149. autorag/nodes/queryexpansion/run.py +308 -0
  150. autorag/nodes/retrieval/__init__.py +0 -0
  151. autorag/nodes/retrieval/base.py +127 -0
  152. autorag/nodes/retrieval/run_util.py +152 -0
  153. autorag/nodes/semanticretrieval/__init__.py +1 -0
  154. autorag/nodes/semanticretrieval/run.py +148 -0
  155. autorag/nodes/semanticretrieval/vectordb.py +339 -0
  156. autorag/nodes/util.py +16 -0
  157. autorag/parser.py +37 -0
  158. autorag/schema/__init__.py +3 -0
  159. autorag/schema/base.py +35 -0
  160. autorag/schema/metricinput.py +99 -0
  161. autorag/schema/module.py +24 -0
  162. autorag/schema/node.py +144 -0
  163. autorag/strategy.py +165 -0
  164. autorag/support.py +235 -0
  165. autorag/utils/__init__.py +8 -0
  166. autorag/utils/cast.py +45 -0
  167. autorag/utils/preprocess.py +149 -0
  168. autorag/utils/util.py +759 -0
  169. autorag/validator.py +98 -0
  170. autorag/vectordb/__init__.py +75 -0
  171. autorag/vectordb/base.py +73 -0
  172. autorag/vectordb/chroma.py +118 -0
  173. autorag/vectordb/couchbase.py +239 -0
  174. autorag/vectordb/milvus.py +169 -0
  175. autorag/vectordb/pinecone.py +121 -0
  176. autorag/vectordb/qdrant.py +155 -0
  177. autorag/vectordb/weaviate.py +184 -0
  178. autorag/web.py +81 -0
  179. autorag-0.0.0.dist-info/METADATA +780 -0
  180. autorag-0.0.0.dist-info/RECORD +184 -0
  181. autorag-0.0.0.dist-info/WHEEL +5 -0
  182. autorag-0.0.0.dist-info/entry_points.txt +2 -0
  183. autorag-0.0.0.dist-info/licenses/LICENSE +201 -0
  184. autorag-0.0.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,127 @@
1
+ import abc
2
+ import logging
3
+ import os
4
+ from typing import List, Union, Tuple
5
+
6
+ import pandas as pd
7
+
8
+ from autorag.schema import BaseModule
9
+ from autorag.support import get_support_modules
10
+ from autorag.utils import fetch_contents, result_to_dataframe, validate_qa_dataset
11
+ from autorag.utils.util import pop_params
12
+
13
+ logger = logging.getLogger("AutoRAG")
14
+
15
+
16
+ class BaseRetrieval(BaseModule, metaclass=abc.ABCMeta):
17
+ def __init__(self, project_dir: str, *args, **kwargs):
18
+ logger.info(f"Initialize retrieval node - {self.__class__.__name__}")
19
+
20
+ self.resources_dir = os.path.join(project_dir, "resources")
21
+ data_dir = os.path.join(project_dir, "data")
22
+ # fetch data from corpus_data
23
+ self.corpus_df = pd.read_parquet(
24
+ os.path.join(data_dir, "corpus.parquet"), engine="pyarrow"
25
+ )
26
+
27
+ def __del__(self):
28
+ logger.info(f"Deleting retrieval node - {self.__class__.__name__} module...")
29
+
30
+ def cast_to_run(self, previous_result: pd.DataFrame, *args, **kwargs):
31
+ logger.info(f"Running retrieval node - {self.__class__.__name__} module...")
32
+ validate_qa_dataset(previous_result)
33
+ # find queries columns & type cast queries
34
+ assert "query" in previous_result.columns, (
35
+ "previous_result must have query column."
36
+ )
37
+ if "queries" not in previous_result.columns:
38
+ previous_result["queries"] = previous_result["query"]
39
+ previous_result.loc[:, "queries"] = previous_result["queries"].apply(
40
+ cast_queries
41
+ )
42
+ queries = previous_result["queries"].tolist()
43
+ return queries
44
+
45
+
46
+ class HybridRetrieval(BaseRetrieval, metaclass=abc.ABCMeta):
47
+ def __init__(
48
+ self, project_dir: str, target_modules, target_module_params, *args, **kwargs
49
+ ):
50
+ super().__init__(project_dir)
51
+ self.target_modules = list(
52
+ map(
53
+ lambda x, y: get_support_modules(x)(
54
+ **y,
55
+ project_dir=project_dir,
56
+ ),
57
+ target_modules,
58
+ target_module_params,
59
+ )
60
+ )
61
+ self.target_module_params = target_module_params
62
+
63
+ @result_to_dataframe(["retrieved_contents", "retrieved_ids", "retrieve_scores"])
64
+ def pure(self, previous_result: pd.DataFrame, *args, **kwargs):
65
+ result_dfs: List[pd.DataFrame] = list(
66
+ map(
67
+ lambda x, y: x.pure(
68
+ **y,
69
+ previous_result=previous_result,
70
+ ),
71
+ self.target_modules,
72
+ self.target_module_params,
73
+ )
74
+ )
75
+ ids = tuple(
76
+ map(lambda df: df["retrieved_ids"].apply(list).tolist(), result_dfs)
77
+ )
78
+ scores = tuple(
79
+ map(
80
+ lambda df: df["retrieve_scores"].apply(list).tolist(),
81
+ result_dfs,
82
+ )
83
+ )
84
+
85
+ _pure_params = pop_params(self._pure, kwargs)
86
+ if "ids" in _pure_params or "scores" in _pure_params:
87
+ raise ValueError(
88
+ "With specifying ids or scores, you must use HybridRRF.run_evaluator instead."
89
+ )
90
+ ids, scores = self._pure(ids=ids, scores=scores, **_pure_params)
91
+ contents = fetch_contents(self.corpus_df, ids)
92
+ return contents, ids, scores
93
+
94
+
95
+ def cast_queries(queries: Union[str, List[str]]) -> List[str]:
96
+ if isinstance(queries, str):
97
+ return [queries]
98
+ elif isinstance(queries, List):
99
+ return queries
100
+ else:
101
+ raise ValueError(f"queries must be str or list, but got {type(queries)}")
102
+
103
+
104
+ def evenly_distribute_passages(
105
+ ids: List[List[str]], scores: List[List[float]], top_k: int
106
+ ) -> Tuple[List[str], List[float]]:
107
+ assert len(ids) == len(scores), "ids and scores must have same length."
108
+ query_cnt = len(ids)
109
+ avg_len = top_k // query_cnt
110
+ remainder = top_k % query_cnt
111
+
112
+ new_ids = []
113
+ new_scores = []
114
+ for i in range(query_cnt):
115
+ if i < remainder:
116
+ new_ids.extend(ids[i][: avg_len + 1])
117
+ new_scores.extend(scores[i][: avg_len + 1])
118
+ else:
119
+ new_ids.extend(ids[i][:avg_len])
120
+ new_scores.extend(scores[i][:avg_len])
121
+
122
+ return new_ids, new_scores
123
+
124
+
125
+ def get_bm25_pkl_name(bm25_tokenizer: str):
126
+ bm25_tokenizer = bm25_tokenizer.replace("/", "")
127
+ return f"bm25_{bm25_tokenizer}.pkl"
@@ -0,0 +1,152 @@
1
+ import os
2
+ import pathlib
3
+ from typing import Tuple, List, Union, Dict
4
+
5
+ import pandas as pd
6
+
7
+ from autorag.evaluation import evaluate_retrieval
8
+ from autorag.schema.metricinput import MetricInput
9
+ from autorag.strategy import measure_speed, filter_by_threshold, select_best
10
+
11
+
12
+ def evaluate_retrieval_node(
13
+ result_df: pd.DataFrame,
14
+ metric_inputs: List[MetricInput],
15
+ metrics: Union[List[str], List[Dict]],
16
+ ) -> pd.DataFrame:
17
+ """
18
+ Evaluate retrieval node from retrieval node result dataframe.
19
+ :param result_df: The result dataframe from a retrieval node.
20
+ :param metric_inputs: List of metric input schema for AutoRAG.
21
+ :param metrics: Metric list from input strategies.
22
+ :return: Return result_df with metrics columns.
23
+ The columns will be 'retrieved_contents', 'retrieved_ids', 'retrieve_scores', and metric names.
24
+ """
25
+
26
+ @evaluate_retrieval(
27
+ metric_inputs=metric_inputs,
28
+ metrics=metrics,
29
+ )
30
+ def evaluate_this_module(df: pd.DataFrame):
31
+ return (
32
+ df["retrieved_contents"].tolist(),
33
+ df["retrieved_ids"].tolist(),
34
+ df["retrieve_scores"].tolist(),
35
+ )
36
+
37
+ return evaluate_this_module(result_df)
38
+
39
+
40
+ def run(
41
+ input_modules,
42
+ input_module_params,
43
+ project_dir: Union[str, pathlib.Path, pathlib.PurePath],
44
+ previous_result: pd.DataFrame,
45
+ strategies,
46
+ metric_inputs: List[MetricInput],
47
+ ) -> Tuple[List[pd.DataFrame], List]:
48
+ """
49
+ Run input modules and parameters.
50
+ :param input_modules: Input modules
51
+ :param input_module_params: Input module parameters
52
+ :param project_dir: Project directory path.
53
+ :param previous_result: Previous result dataframe.
54
+ :param strategies: Strategies for retrieval node.
55
+ :param metric_inputs: List of metric input schema for AutoRAG.
56
+ :return: First, it returns list of result dataframe.
57
+ Second, it returns list of execution times.
58
+ """
59
+ result, execution_times = zip(
60
+ *map(
61
+ lambda task: measure_speed(
62
+ task[0].run_evaluator,
63
+ project_dir=project_dir,
64
+ previous_result=previous_result,
65
+ **task[1],
66
+ ),
67
+ zip(input_modules, input_module_params),
68
+ )
69
+ )
70
+ average_times = list(map(lambda x: x / len(result[0]), execution_times))
71
+
72
+ # run metrics before filtering
73
+ if strategies.get("metrics") is None:
74
+ raise ValueError("You must at least one metrics for retrieval evaluation.")
75
+ result = list(
76
+ map(
77
+ lambda x: evaluate_retrieval_node(
78
+ x,
79
+ metric_inputs,
80
+ strategies.get("metrics"),
81
+ ),
82
+ result,
83
+ )
84
+ )
85
+
86
+ return result, average_times
87
+
88
+
89
+ def save_and_summary(
90
+ input_modules,
91
+ input_module_params,
92
+ result_list,
93
+ execution_time_list,
94
+ filename_start: int,
95
+ save_dir: Union[str, pathlib.Path, pathlib.PurePath],
96
+ strategies,
97
+ ):
98
+ """
99
+ Save the result and make summary file
100
+ :param input_modules: Input modules
101
+ :param input_module_params: Input module parameters
102
+ :param result_list: Result list
103
+ :param execution_time_list: Execution times
104
+ :param filename_start: The first filename to use
105
+ :return: First, it returns list of result dataframe.
106
+ Second, it returns list of execution times.
107
+ """
108
+
109
+ # save results to folder
110
+ filepaths = list(
111
+ map(
112
+ lambda x: os.path.join(save_dir, f"{x}.parquet"),
113
+ range(filename_start, filename_start + len(input_modules)),
114
+ )
115
+ )
116
+ list(
117
+ map(
118
+ lambda x: x[0].to_parquet(x[1], index=False),
119
+ zip(result_list, filepaths),
120
+ )
121
+ ) # execute save to parquet
122
+ filename_list = list(map(lambda x: os.path.basename(x), filepaths))
123
+
124
+ summary_df = pd.DataFrame(
125
+ {
126
+ "filename": filename_list,
127
+ "module_name": list(map(lambda module: module.__name__, input_modules)),
128
+ "module_params": input_module_params,
129
+ "execution_time": execution_time_list,
130
+ **{
131
+ metric: list(map(lambda result: result[metric].mean(), result_list))
132
+ for metric in strategies.get("metrics")
133
+ },
134
+ }
135
+ )
136
+ summary_df.to_csv(os.path.join(save_dir, "summary.csv"), index=False)
137
+ return summary_df
138
+
139
+
140
+ def find_best(results, average_times, filenames, strategies):
141
+ # filter by strategies
142
+ if strategies.get("speed_threshold") is not None:
143
+ results, filenames = filter_by_threshold(
144
+ results, average_times, strategies["speed_threshold"], filenames
145
+ )
146
+ selected_result, selected_filename = select_best(
147
+ results,
148
+ strategies.get("metrics"),
149
+ filenames,
150
+ strategies.get("strategy", "mean"),
151
+ )
152
+ return selected_result, selected_filename
@@ -0,0 +1 @@
1
+ from .vectordb import VectorDB
@@ -0,0 +1,148 @@
1
+ import os
2
+ import pathlib
3
+ from typing import List, Dict, Union
4
+
5
+ import pandas as pd
6
+
7
+ from autorag.evaluation import evaluate_retrieval
8
+ from autorag.evaluation.retrieval import RETRIEVAL_METRIC_FUNC_DICT
9
+ from autorag.nodes.retrieval.run_util import save_and_summary, find_best
10
+ from autorag.schema.metricinput import MetricInput
11
+ from autorag.strategy import measure_speed
12
+ from autorag.utils.util import apply_recursive, to_list
13
+
14
+
15
+ def run_semantic_retrieval_node(
16
+ modules: List,
17
+ module_params: List[Dict],
18
+ previous_result: pd.DataFrame,
19
+ node_line_dir: str,
20
+ strategies: Dict,
21
+ ) -> pd.DataFrame:
22
+ """
23
+ Run the semantic retrieval node.
24
+
25
+ :param modules: Retrieval modules to run.
26
+ :param module_params: Retrieval module parameters.
27
+ :param previous_result: Previous result dataframe.
28
+ Could be query expansion's best result or qa data.
29
+ :param node_line_dir: This node line's directory.
30
+ :param strategies: Strategies for retrieval node.
31
+ :return: The best result dataframe.
32
+ It contains previous result columns and retrieval node's result columns.
33
+ """
34
+ if not os.path.exists(node_line_dir):
35
+ os.makedirs(node_line_dir)
36
+ project_dir = pathlib.PurePath(node_line_dir).parent.parent
37
+ qa_df = pd.read_parquet(
38
+ os.path.join(project_dir, "data", "qa.parquet"), engine="pyarrow"
39
+ )
40
+ retrieval_gt = qa_df["retrieval_gt"].tolist()
41
+ retrieval_gt = apply_recursive(lambda x: str(x), to_list(retrieval_gt))
42
+ # make rows to metric_inputs
43
+ metric_inputs = [
44
+ MetricInput(retrieval_gt=ret_gt, query=query, generation_gt=gen_gt)
45
+ for ret_gt, query, gen_gt in zip(
46
+ retrieval_gt, qa_df["query"].tolist(), qa_df["generation_gt"].tolist()
47
+ )
48
+ ]
49
+
50
+ save_dir = os.path.join(node_line_dir, "semantic_retrieval")
51
+ if not os.path.exists(save_dir):
52
+ os.makedirs(save_dir)
53
+
54
+ # Run the modules
55
+ semantic_results, execution_times = zip(
56
+ *map(
57
+ lambda task: measure_speed(
58
+ task[0].run_evaluator,
59
+ project_dir=project_dir,
60
+ previous_result=previous_result,
61
+ **task[1],
62
+ ),
63
+ zip(modules, module_params),
64
+ )
65
+ )
66
+ semantic_times = list(map(lambda x: x / len(semantic_results[0]), execution_times))
67
+
68
+ # run metrics
69
+ if strategies.get("metrics") is None:
70
+ raise ValueError("You must at least one metrics for retrieval evaluation.")
71
+ semantic_results = list(
72
+ map(
73
+ lambda x: evaluate_semantic_retrieval_node(
74
+ x,
75
+ metric_inputs,
76
+ strategies.get("metrics"),
77
+ ),
78
+ semantic_results,
79
+ )
80
+ )
81
+
82
+ semantic_summary_df = save_and_summary(
83
+ modules,
84
+ module_params,
85
+ semantic_results,
86
+ semantic_times,
87
+ 0,
88
+ save_dir,
89
+ strategies,
90
+ )
91
+ semantic_selected_result, semantic_selected_filename = find_best(
92
+ semantic_results,
93
+ semantic_times,
94
+ semantic_summary_df["filename"].tolist(),
95
+ strategies,
96
+ )
97
+ semantic_summary_df["is_best"] = (
98
+ semantic_summary_df["filename"] == semantic_selected_filename
99
+ )
100
+ previous_result.drop(
101
+ columns=list(RETRIEVAL_METRIC_FUNC_DICT.keys()), inplace=True, errors="ignore"
102
+ )
103
+ semantic_selected_result.rename(
104
+ columns={
105
+ "retrieved_contents": "retrieved_contents_semantic",
106
+ "retrieved_ids": "retrieved_ids_semantic",
107
+ "retrieve_scores": "retrieve_scores_semantic",
108
+ },
109
+ inplace=True,
110
+ )
111
+ best_result = pd.concat([previous_result, semantic_selected_result], axis=1)
112
+ best_result.to_parquet(
113
+ os.path.join(
114
+ save_dir, f"best_{os.path.splitext(semantic_selected_filename)[0]}.parquet"
115
+ ),
116
+ index=False,
117
+ )
118
+ semantic_summary_df.to_csv(os.path.join(save_dir, "summary.csv"), index=False)
119
+ return best_result # The result will be with _semantic suffix
120
+
121
+
122
+ def evaluate_semantic_retrieval_node(
123
+ result_df: pd.DataFrame,
124
+ metric_inputs: List[MetricInput],
125
+ metrics: Union[List[str], List[Dict]],
126
+ ) -> pd.DataFrame:
127
+ """
128
+ Evaluate retrieval node from retrieval node result dataframe.
129
+
130
+ :param result_df: The result dataframe from a retrieval node.
131
+ :param metric_inputs: List of metric input schema for AutoRAG.
132
+ :param metrics: Metric list from input strategies.
133
+ :return: Return result_df with metrics columns.
134
+ The columns will be 'retrieved_contents_semantic', 'retrieved_ids_semantic', 'retrieve_scores_semantic', and metric names.
135
+ """
136
+
137
+ @evaluate_retrieval(
138
+ metric_inputs=metric_inputs,
139
+ metrics=metrics,
140
+ )
141
+ def evaluate_this_module(df: pd.DataFrame):
142
+ return (
143
+ df["retrieved_contents_semantic"].tolist(),
144
+ df["retrieved_ids_semantic"].tolist(),
145
+ df["retrieve_scores_semantic"].tolist(),
146
+ )
147
+
148
+ return evaluate_this_module(result_df)