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
autorag/deploy/api.py ADDED
@@ -0,0 +1,303 @@
1
+ import logging
2
+ import os
3
+ import pathlib
4
+ import uuid
5
+ from typing import Dict, Optional, List, Union, Literal
6
+
7
+ import pandas as pd
8
+ from quart import Quart, request, jsonify
9
+ from quart.helpers import stream_with_context
10
+ from pydantic import BaseModel, ValidationError
11
+
12
+ from autorag.deploy.base import BaseRunner
13
+ from autorag.nodes.generator.base import BaseGenerator
14
+ from autorag.nodes.promptmaker.base import BasePromptMaker
15
+ from autorag.utils.util import fetch_contents, to_list
16
+
17
+ logger = logging.getLogger("AutoRAG")
18
+
19
+ deploy_dir = pathlib.Path(__file__).parent
20
+ root_dir = pathlib.Path(__file__).parent.parent
21
+
22
+ VERSION_PATH = os.path.join(root_dir, "VERSION")
23
+
24
+
25
+ class QueryRequest(BaseModel):
26
+ query: str
27
+ result_column: Optional[str] = "generated_texts"
28
+
29
+
30
+ class RetrievedPassage(BaseModel):
31
+ content: str
32
+ doc_id: str
33
+ score: float
34
+ filepath: Optional[str] = None
35
+ file_page: Optional[int] = None
36
+ start_idx: Optional[int] = None
37
+ end_idx: Optional[int] = None
38
+
39
+
40
+ class RunResponse(BaseModel):
41
+ result: Union[str, List[str]]
42
+ retrieved_passage: List[RetrievedPassage]
43
+
44
+
45
+ class RetrievalResponse(BaseModel):
46
+ passages: List[RetrievedPassage]
47
+
48
+
49
+ class StreamResponse(BaseModel):
50
+ """
51
+ When the type is generated_text, only generated_text is returned. The other fields are None.
52
+ When the type is retrieved_passage, only retrieved_passage and passage_index are returned. The other fields are None.
53
+ """
54
+
55
+ type: Literal["generated_text", "retrieved_passage"]
56
+ generated_text: Optional[str]
57
+ retrieved_passage: Optional[RetrievedPassage]
58
+ passage_index: Optional[int]
59
+
60
+
61
+ class VersionResponse(BaseModel):
62
+ version: str
63
+
64
+
65
+ class ApiRunner(BaseRunner):
66
+ def __init__(self, config: Dict, project_dir: Optional[str] = None):
67
+ super().__init__(config, project_dir)
68
+ self.app = Quart(__name__)
69
+
70
+ data_dir = os.path.join(project_dir, "data")
71
+ self.corpus_df = pd.read_parquet(
72
+ os.path.join(data_dir, "corpus.parquet"), engine="pyarrow"
73
+ )
74
+ self.__add_api_route()
75
+
76
+ def __add_api_route(self):
77
+ @self.app.route("/v1/run", methods=["POST"])
78
+ async def run_query():
79
+ try:
80
+ data = await request.get_json()
81
+ data = QueryRequest(**data)
82
+ except ValidationError as e:
83
+ return jsonify(e.errors()), 400
84
+
85
+ previous_result = pd.DataFrame(
86
+ {
87
+ "qid": str(uuid.uuid4()),
88
+ "query": [data.query],
89
+ "retrieval_gt": [[]],
90
+ "generation_gt": [""],
91
+ }
92
+ ) # pseudo qa data for execution
93
+ for module_instance, module_param in zip(
94
+ self.module_instances, self.module_params
95
+ ):
96
+ new_result = module_instance.pure(
97
+ previous_result=previous_result, **module_param
98
+ )
99
+ duplicated_columns = previous_result.columns.intersection(
100
+ new_result.columns
101
+ )
102
+ drop_previous_result = previous_result.drop(columns=duplicated_columns)
103
+ previous_result = pd.concat([drop_previous_result, new_result], axis=1)
104
+
105
+ # Simulate processing the query
106
+ generated_text = previous_result[data.result_column].tolist()[0]
107
+ retrieved_passage = self.extract_retrieve_passage(previous_result)
108
+
109
+ response = RunResponse(
110
+ result=generated_text, retrieved_passage=retrieved_passage
111
+ )
112
+
113
+ return jsonify(response.model_dump()), 200
114
+
115
+ @self.app.route("/v1/retrieve", methods=["POST"])
116
+ async def run_retrieve_only():
117
+ data = await request.get_json()
118
+ query = data.get("query", None)
119
+ if query is None:
120
+ return jsonify(
121
+ {
122
+ "error": "Invalid request. You need to include 'query' in the request body."
123
+ }
124
+ ), 400
125
+
126
+ previous_result = pd.DataFrame(
127
+ {
128
+ "qid": str(uuid.uuid4()),
129
+ "query": [query],
130
+ "retrieval_gt": [[]],
131
+ "generation_gt": [""],
132
+ }
133
+ ) # pseudo qa data for execution
134
+ for module_instance, module_param in zip(
135
+ self.module_instances, self.module_params
136
+ ):
137
+ if isinstance(module_instance, BasePromptMaker) or isinstance(
138
+ module_instance, BaseGenerator
139
+ ):
140
+ continue
141
+ new_result = module_instance.pure(
142
+ previous_result=previous_result, **module_param
143
+ )
144
+ duplicated_columns = previous_result.columns.intersection(
145
+ new_result.columns
146
+ )
147
+ drop_previous_result = previous_result.drop(columns=duplicated_columns)
148
+ previous_result = pd.concat([drop_previous_result, new_result], axis=1)
149
+
150
+ # Simulate processing the query
151
+ retrieved_passages = self.extract_retrieve_passage(previous_result)
152
+
153
+ retrieval_response = RetrievalResponse(passages=retrieved_passages)
154
+ return jsonify(retrieval_response.model_dump()), 200
155
+
156
+ @self.app.route("/v1/stream", methods=["POST"])
157
+ async def stream_query():
158
+ try:
159
+ data = await request.get_json()
160
+ data = QueryRequest(**data)
161
+ except ValidationError as e:
162
+ return jsonify(e.errors()), 400
163
+
164
+ @stream_with_context
165
+ async def generate():
166
+ previous_result = pd.DataFrame(
167
+ {
168
+ "qid": str(uuid.uuid4()),
169
+ "query": [data.query],
170
+ "retrieval_gt": [[]],
171
+ "generation_gt": [""],
172
+ }
173
+ ) # pseudo qa data for execution
174
+
175
+ for module_instance, module_param in zip(
176
+ self.module_instances, self.module_params
177
+ ):
178
+ if not isinstance(module_instance, BaseGenerator):
179
+ new_result = module_instance.pure(
180
+ previous_result=previous_result, **module_param
181
+ )
182
+ duplicated_columns = previous_result.columns.intersection(
183
+ new_result.columns
184
+ )
185
+ drop_previous_result = previous_result.drop(
186
+ columns=duplicated_columns
187
+ )
188
+ previous_result = pd.concat(
189
+ [drop_previous_result, new_result], axis=1
190
+ )
191
+ else:
192
+ retrieved_passages = self.extract_retrieve_passage(
193
+ previous_result
194
+ )
195
+ for i, retrieved_passage in enumerate(retrieved_passages):
196
+ yield (
197
+ StreamResponse(
198
+ type="retrieved_passage",
199
+ generated_text=None,
200
+ retrieved_passage=retrieved_passage,
201
+ passage_index=i,
202
+ )
203
+ .model_dump_json()
204
+ .encode("utf-8")
205
+ )
206
+ # Start streaming of the result
207
+ assert len(previous_result) == 1
208
+ prompt: str = previous_result["prompts"].tolist()[0]
209
+ async for delta in module_instance.astream(
210
+ prompt=prompt, **module_param
211
+ ):
212
+ response = StreamResponse(
213
+ type="generated_text",
214
+ generated_text=delta,
215
+ retrieved_passage=None,
216
+ passage_index=None,
217
+ )
218
+ yield response.model_dump_json().encode("utf-8")
219
+
220
+ return generate(), 200, {"X-Something": "value"}
221
+
222
+ @self.app.route("/version", methods=["GET"])
223
+ def get_version():
224
+ with open(VERSION_PATH, "r") as f:
225
+ version = f.read().strip()
226
+ response = VersionResponse(version=version)
227
+ return jsonify(response.model_dump()), 200
228
+
229
+ def run_api_server(
230
+ self, host: str = "0.0.0.0", port: int = 8000, remote: bool = True, **kwargs
231
+ ):
232
+ """
233
+ Run the pipeline as an api server.
234
+ Here is api endpoint documentation => https://marker-inc-korea.github.io/AutoRAG/deploy/api_endpoint.html
235
+
236
+ :param host: The host of the api server.
237
+ :param port: The port of the api server.
238
+ :param remote: Whether to expose the api server to the public internet using ngrok.
239
+ :param kwargs: Other arguments for Flask app.run.
240
+ """
241
+ logger.info(f"Run api server at {host}:{port}")
242
+ if remote:
243
+ from pyngrok import ngrok
244
+
245
+ http_tunnel = ngrok.connect(str(port), "http")
246
+ public_url = http_tunnel.public_url
247
+ logger.info(f"Public API URL: {public_url}")
248
+ self.app.run(host=host, port=port, **kwargs)
249
+
250
+ def extract_retrieve_passage(self, df: pd.DataFrame) -> List[RetrievedPassage]:
251
+ if "retrieved_ids" not in df.columns and "retrieved_ids_semantic" in df.columns:
252
+ retrieved_ids: List[str] = df["retrieved_ids_semantic"].tolist()[0]
253
+ scores = df["retrieve_scores_semantic"].tolist()[0]
254
+ elif (
255
+ "retrieved_ids" not in df.columns
256
+ and "retrieved_ids_semantic" not in df.columns
257
+ ):
258
+ retrieved_ids: List[str] = df["retrieved_ids_lexical"].tolist()[0]
259
+ scores = df["retrieve_scores_lexical"].tolist()[0]
260
+ else:
261
+ retrieved_ids: List[str] = df["retrieved_ids"].tolist()[0]
262
+ scores = df["retrieve_scores"].tolist()[0]
263
+ contents = fetch_contents(self.corpus_df, [retrieved_ids])[0]
264
+ if "path" in self.corpus_df.columns:
265
+ paths = fetch_contents(self.corpus_df, [retrieved_ids], column_name="path")[
266
+ 0
267
+ ]
268
+ else:
269
+ paths = [None] * len(retrieved_ids)
270
+ metadatas = fetch_contents(
271
+ self.corpus_df, [retrieved_ids], column_name="metadata"
272
+ )[0]
273
+ if "start_end_idx" in self.corpus_df.columns:
274
+ start_end_indices = fetch_contents(
275
+ self.corpus_df, [retrieved_ids], column_name="start_end_idx"
276
+ )[0]
277
+ else:
278
+ start_end_indices = [None] * len(retrieved_ids)
279
+ start_end_indices = to_list(start_end_indices)
280
+ return list(
281
+ map(
282
+ lambda content,
283
+ doc_id,
284
+ score,
285
+ path,
286
+ metadata,
287
+ start_end_idx: RetrievedPassage(
288
+ content=content,
289
+ doc_id=doc_id,
290
+ score=score,
291
+ filepath=path,
292
+ file_page=metadata.get("page", None),
293
+ start_idx=start_end_idx[0] if start_end_idx else None,
294
+ end_idx=start_end_idx[1] if start_end_idx else None,
295
+ ),
296
+ contents,
297
+ retrieved_ids,
298
+ scores,
299
+ paths,
300
+ metadatas,
301
+ start_end_indices,
302
+ )
303
+ )
autorag/deploy/base.py ADDED
@@ -0,0 +1,235 @@
1
+ import logging
2
+ import os
3
+ import pathlib
4
+ import uuid
5
+ from copy import deepcopy
6
+ from typing import Optional, Dict, List
7
+
8
+ import pandas as pd
9
+ import yaml
10
+
11
+ from autorag.support import get_support_modules
12
+ from autorag.utils.util import load_summary_file, load_yaml_config
13
+
14
+ logger = logging.getLogger("AutoRAG")
15
+
16
+
17
+ def extract_node_line_names(config_dict: Dict) -> List[str]:
18
+ """
19
+ Extract node line names with the given config dictionary order.
20
+
21
+ :param config_dict: The YAML configuration dict for the pipeline.
22
+ You can load this to access trail_folder/config.yaml.
23
+ :return: The list of node line names.
24
+ It is the order of the node line names in the pipeline.
25
+ """
26
+ return [node_line["node_line_name"] for node_line in config_dict["node_lines"]]
27
+
28
+
29
+ def extract_node_strategy(config_dict: Dict) -> Dict:
30
+ """
31
+ Extract node strategies with the given config dictionary.
32
+ The return value is a dictionary of the node type and its strategy.
33
+
34
+ :param config_dict: The YAML configuration dict for the pipeline.
35
+ You can load this to access trail_folder/config.yaml.
36
+ :return: Key is node_type and value is strategy dict.
37
+ """
38
+ return {
39
+ node["node_type"]: node.get("strategy", {})
40
+ for node_line in config_dict["node_lines"]
41
+ for node in node_line["nodes"]
42
+ }
43
+
44
+
45
+ def summary_df_to_yaml(summary_df: pd.DataFrame, config_dict: Dict) -> Dict:
46
+ """
47
+ Convert trial summary dataframe to config yaml file.
48
+
49
+ :param summary_df: The trial summary dataframe of the evaluated trial.
50
+ :param config_dict: The yaml configuration dict for the pipeline.
51
+ You can load this to access trail_folder/config.yaml.
52
+ :return: Dictionary of config yaml file.
53
+ You can save this dictionary to yaml file.
54
+ """
55
+
56
+ # summary_df columns : 'node_line_name', 'node_type', 'best_module_filename',
57
+ # 'best_module_name', 'best_module_params', 'best_execution_time'
58
+ node_line_names = extract_node_line_names(config_dict)
59
+ node_strategies = extract_node_strategy(config_dict)
60
+ strategy_df = pd.DataFrame(
61
+ {
62
+ "node_type": list(node_strategies.keys()),
63
+ "strategy": list(node_strategies.values()),
64
+ }
65
+ )
66
+ summary_df = summary_df.merge(strategy_df, on="node_type", how="left")
67
+ summary_df["categorical_node_line_name"] = pd.Categorical(
68
+ summary_df["node_line_name"], categories=node_line_names, ordered=True
69
+ )
70
+ summary_df = summary_df.sort_values(by="categorical_node_line_name")
71
+ grouped = summary_df.groupby("categorical_node_line_name", observed=False)
72
+
73
+ node_lines = [
74
+ {
75
+ "node_line_name": node_line_name,
76
+ "nodes": [
77
+ {
78
+ "node_type": row["node_type"],
79
+ "strategy": row["strategy"],
80
+ "modules": [
81
+ {
82
+ "module_type": row["best_module_name"],
83
+ **row["best_module_params"],
84
+ }
85
+ ],
86
+ }
87
+ for _, row in node_line.iterrows()
88
+ ],
89
+ }
90
+ for node_line_name, node_line in grouped
91
+ ]
92
+ return {"node_lines": node_lines}
93
+
94
+
95
+ def extract_best_config(trial_path: str, output_path: Optional[str] = None) -> Dict:
96
+ """
97
+ Extract the optimal pipeline from the evaluated trial.
98
+
99
+ :param trial_path: The path to the trial directory that you want to extract the pipeline from.
100
+ Must already be evaluated.
101
+ :param output_path: Output path that pipeline yaml file will be saved.
102
+ Must be .yaml or .yml file.
103
+ If None, it does not save YAML file and just returns dict values.
104
+ Default is None.
105
+ :return: The dictionary of the extracted pipeline.
106
+ """
107
+ summary_path = os.path.join(trial_path, "summary.csv")
108
+ if not os.path.exists(summary_path):
109
+ raise ValueError(f"summary.csv does not exist in {trial_path}.")
110
+ trial_summary_df = load_summary_file(
111
+ summary_path, dict_columns=["best_module_params"]
112
+ )
113
+ config_yaml_path = os.path.join(trial_path, "config.yaml")
114
+ with open(config_yaml_path, "r") as f:
115
+ config_dict = yaml.safe_load(f)
116
+ yaml_dict = summary_df_to_yaml(trial_summary_df, config_dict)
117
+ yaml_dict["vectordb"] = extract_vectordb_config(trial_path)
118
+ if output_path is not None:
119
+ with open(output_path, "w") as f:
120
+ yaml.safe_dump(yaml_dict, f)
121
+ return yaml_dict
122
+
123
+
124
+ def extract_vectordb_config(trial_path: str) -> List[Dict]:
125
+ # get vectordb.yaml file
126
+ project_dir = pathlib.PurePath(os.path.realpath(trial_path)).parent
127
+ vectordb_config_path = os.path.join(project_dir, "resources", "vectordb.yaml")
128
+ if not os.path.exists(vectordb_config_path):
129
+ raise ValueError(f"vectordb.yaml does not exist in {vectordb_config_path}.")
130
+ with open(vectordb_config_path, "r") as f:
131
+ vectordb_dict = yaml.safe_load(f)
132
+ result = vectordb_dict.get("vectordb", [])
133
+ if len(result) != 0:
134
+ return result
135
+ # return default setting of chroma
136
+ return [
137
+ {
138
+ "name": "default",
139
+ "db_type": "chroma",
140
+ "client_type": "persistent",
141
+ "embedding_model": "openai",
142
+ "collection_name": "openai",
143
+ "path": os.path.join(project_dir, "resources", "chroma"),
144
+ }
145
+ ]
146
+
147
+
148
+ class BaseRunner:
149
+ def __init__(self, config: Dict, project_dir: Optional[str] = None):
150
+ self.config = config
151
+ project_dir = os.getcwd() if project_dir is None else project_dir
152
+ os.environ["PROJECT_DIR"] = project_dir
153
+
154
+ # init modules
155
+ node_lines = deepcopy(self.config["node_lines"])
156
+ self.module_instances = []
157
+ self.module_params = []
158
+ for node_line in node_lines:
159
+ for node in node_line["nodes"]:
160
+ if len(node["modules"]) != 1:
161
+ raise ValueError(
162
+ "The number of modules in a node must be 1 for using runner."
163
+ "Please use extract_best_config method for extracting yaml file from evaluated trial."
164
+ )
165
+ module = node["modules"][0]
166
+ module_type = module.pop("module_type")
167
+ module_params = module
168
+ module_instance = get_support_modules(module_type)(
169
+ project_dir=project_dir,
170
+ **module_params,
171
+ )
172
+ self.module_instances.append(module_instance)
173
+ self.module_params.append(module_params)
174
+
175
+ @classmethod
176
+ def from_yaml(cls, yaml_path: str, project_dir: Optional[str] = None):
177
+ """
178
+ Load Runner from the YAML file.
179
+ Must be extracted YAML file from the evaluated trial using the extract_best_config method.
180
+
181
+ :param yaml_path: The path of the YAML file.
182
+ :param project_dir: The path of the project directory.
183
+ Default is the current directory.
184
+ :return: Initialized Runner.
185
+ """
186
+ config = load_yaml_config(yaml_path)
187
+ return cls(config, project_dir=project_dir)
188
+
189
+ @classmethod
190
+ def from_trial_folder(cls, trial_path: str):
191
+ """
192
+ Load Runner from the evaluated trial folder.
193
+ Must already be evaluated using Evaluator class.
194
+ It sets the project_dir as the parent directory of the trial folder.
195
+
196
+ :param trial_path: The path of the trial folder.
197
+ :return: Initialized Runner.
198
+ """
199
+ config = extract_best_config(trial_path)
200
+ return cls(config, project_dir=os.path.dirname(trial_path))
201
+
202
+
203
+ class Runner(BaseRunner):
204
+ def run(self, query: str, result_column: str = "generated_texts"):
205
+ """
206
+ Run the pipeline with query.
207
+ The loaded pipeline must start with a single query,
208
+ so the first module of the pipeline must be `query_expansion` or `retrieval` module.
209
+
210
+ :param query: The query of the user.
211
+ :param result_column: The result column name for the answer.
212
+ Default is `generated_texts`, which is the output of the `generation` module.
213
+ :return: The result of the pipeline.
214
+ """
215
+ previous_result = pd.DataFrame(
216
+ {
217
+ "qid": str(uuid.uuid4()),
218
+ "query": [query],
219
+ "retrieval_gt": [[]],
220
+ "generation_gt": [""],
221
+ }
222
+ ) # pseudo qa data for execution
223
+ for module_instance, module_param in zip(
224
+ self.module_instances, self.module_params
225
+ ):
226
+ new_result = module_instance.pure(
227
+ previous_result=previous_result, **module_param
228
+ )
229
+ duplicated_columns = previous_result.columns.intersection(
230
+ new_result.columns
231
+ )
232
+ drop_previous_result = previous_result.drop(columns=duplicated_columns)
233
+ previous_result = pd.concat([drop_previous_result, new_result], axis=1)
234
+
235
+ return previous_result[result_column].tolist()[0]
@@ -0,0 +1,74 @@
1
+ import logging
2
+ import uuid
3
+
4
+ import pandas as pd
5
+
6
+ from autorag.deploy.base import BaseRunner
7
+
8
+ import gradio as gr
9
+
10
+
11
+ logger = logging.getLogger("AutoRAG")
12
+
13
+
14
+ class GradioRunner(BaseRunner):
15
+ def run_web(
16
+ self,
17
+ server_name: str = "0.0.0.0",
18
+ server_port: int = 7680,
19
+ share: bool = False,
20
+ **kwargs,
21
+ ):
22
+ """
23
+ Run web interface to interact pipeline.
24
+ You can access the web interface at `http://server_name:server_port` in your browser
25
+
26
+ :param server_name: The host of the web. Default is 0.0.0.0.
27
+ :param server_port: The port of the web. Default is 7680.
28
+ :param share: Whether to create a publicly shareable link. Default is False.
29
+ :param kwargs: Other arguments for gr.ChatInterface.launch.
30
+ """
31
+
32
+ logger.info(f"Run web interface at http://{server_name}:{server_port}")
33
+
34
+ def get_response(message, _):
35
+ return self.run(message)
36
+
37
+ gr.ChatInterface(
38
+ get_response, title="📚 AutoRAG", retry_btn=None, undo_btn=None
39
+ ).launch(
40
+ server_name=server_name, server_port=server_port, share=share, **kwargs
41
+ )
42
+
43
+ def run(self, query: str, result_column: str = "generated_texts"):
44
+ """
45
+ Run the pipeline with query.
46
+ The loaded pipeline must start with a single query,
47
+ so the first module of the pipeline must be `query_expansion` or `retrieval` module.
48
+
49
+ :param query: The query of the user.
50
+ :param result_column: The result column name for the answer.
51
+ Default is `generated_texts`, which is the output of the `generation` module.
52
+ :return: The result of the pipeline.
53
+ """
54
+ previous_result = pd.DataFrame(
55
+ {
56
+ "qid": str(uuid.uuid4()),
57
+ "query": [query],
58
+ "retrieval_gt": [[]],
59
+ "generation_gt": [""],
60
+ }
61
+ ) # pseudo qa data for execution
62
+ for module_instance, module_param in zip(
63
+ self.module_instances, self.module_params
64
+ ):
65
+ new_result = module_instance.pure(
66
+ previous_result=previous_result, **module_param
67
+ )
68
+ duplicated_columns = previous_result.columns.intersection(
69
+ new_result.columns
70
+ )
71
+ drop_previous_result = previous_result.drop(columns=duplicated_columns)
72
+ previous_result = pd.concat([drop_previous_result, new_result], axis=1)
73
+
74
+ return previous_result[result_column].tolist()[0]