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,141 @@
1
+ import os
2
+ from typing import List, Callable, Dict
3
+ import pandas as pd
4
+ from glob import glob
5
+
6
+ from autorag.strategy import measure_speed
7
+ from autorag.data.utils.util import get_param_combinations
8
+
9
+ default_map = {
10
+ "pdf": {
11
+ "file_type": "pdf",
12
+ "module_type": "langchain_parse",
13
+ "parse_method": "pdfminer",
14
+ },
15
+ "csv": {
16
+ "file_type": "csv",
17
+ "module_type": "langchain_parse",
18
+ "parse_method": "csv",
19
+ },
20
+ "md": {
21
+ "file_type": "md",
22
+ "module_type": "langchain_parse",
23
+ "parse_method": "unstructuredmarkdown",
24
+ },
25
+ "html": {
26
+ "file_type": "html",
27
+ "module_type": "langchain_parse",
28
+ "parse_method": "bshtml",
29
+ },
30
+ "xml": {
31
+ "file_type": "xml",
32
+ "module_type": "langchain_parse",
33
+ "parse_method": "unstructuredxml",
34
+ },
35
+ }
36
+
37
+
38
+ def run_parser(
39
+ modules: List[Callable],
40
+ module_params: List[Dict],
41
+ data_path_glob: str,
42
+ project_dir: str,
43
+ all_files: bool,
44
+ ):
45
+ if not all_files:
46
+ # Set the parsing module to default if it is a file type in paths but not set in YAML.
47
+ data_path_list = glob(data_path_glob)
48
+ if not data_path_list:
49
+ raise FileNotFoundError(f"data does not exits in {data_path_glob}")
50
+
51
+ file_types = set(
52
+ [os.path.basename(data_path).split(".")[-1] for data_path in data_path_list]
53
+ )
54
+ set_file_types = set([module["file_type"] for module in module_params])
55
+
56
+ # Calculate the set difference once
57
+ file_types_to_remove = set_file_types - file_types
58
+
59
+ # Use list comprehension to filter out unwanted elements
60
+ module_params = [
61
+ param
62
+ for param in module_params
63
+ if param["file_type"] not in file_types_to_remove
64
+ ]
65
+ modules = [
66
+ module
67
+ for module, param in zip(modules, module_params)
68
+ if param["file_type"] not in file_types_to_remove
69
+ ]
70
+
71
+ # create a list of only those file_types that are in file_types but not in set_file_types
72
+ missing_file_types = list(file_types - set_file_types)
73
+
74
+ if missing_file_types:
75
+ add_modules_list = []
76
+ for missing_file_type in missing_file_types:
77
+ if missing_file_type == "json":
78
+ raise ValueError(
79
+ "JSON file type must have a jq_schema so you must set it in the YAML file."
80
+ )
81
+
82
+ add_modules_list.append(default_map[missing_file_type])
83
+
84
+ add_modules, add_params = get_param_combinations(add_modules_list)
85
+ modules.extend(add_modules)
86
+ module_params.extend(add_params)
87
+
88
+ results, execution_times = zip(
89
+ *map(
90
+ lambda x: measure_speed(x[0], data_path_glob=data_path_glob, **x[1]),
91
+ zip(modules, module_params),
92
+ )
93
+ )
94
+ average_times = list(map(lambda x: x / len(results[0]), execution_times))
95
+
96
+ # save results to parquet files
97
+ if all_files:
98
+ if len(module_params) > 1:
99
+ raise ValueError(
100
+ "All files is set to True, You can only use one parsing module."
101
+ )
102
+ filepaths = [os.path.join(project_dir, "parsed_result.parquet")]
103
+ else:
104
+ filepaths = list(
105
+ map(
106
+ lambda x: os.path.join(project_dir, f"{x['file_type']}.parquet"),
107
+ module_params,
108
+ )
109
+ )
110
+
111
+ _files = {}
112
+ for result, filepath in zip(results, filepaths):
113
+ _files[filepath].append(result) if filepath in _files.keys() else _files.update(
114
+ {filepath: [result]}
115
+ )
116
+ # Save files with a specific file type as Parquet files.
117
+ for filepath, value in _files.items():
118
+ pd.concat(value).to_parquet(filepath, index=False)
119
+
120
+ filenames = list(map(lambda x: os.path.basename(x), filepaths))
121
+
122
+ summary_df = pd.DataFrame(
123
+ {
124
+ "filename": filenames,
125
+ "module_name": list(map(lambda module: module.__name__, modules)),
126
+ "module_params": module_params,
127
+ "execution_time": average_times,
128
+ }
129
+ )
130
+ summary_df.to_csv(os.path.join(project_dir, "summary.csv"), index=False)
131
+
132
+ # concat all parquet files here if not all_files.
133
+ _filepaths = list(_files.keys())
134
+ if not all_files:
135
+ dataframes = [pd.read_parquet(file) for file in _filepaths]
136
+ combined_df = pd.concat(dataframes, ignore_index=True)
137
+ combined_df.to_parquet(
138
+ os.path.join(project_dir, "parsed_result.parquet"), index=False
139
+ )
140
+
141
+ return summary_df
@@ -0,0 +1,134 @@
1
+ import os
2
+ import tempfile
3
+ from glob import glob
4
+ from typing import List, Tuple, Dict
5
+
6
+ from PyPDF2 import PdfFileReader, PdfFileWriter
7
+ import pdfplumber
8
+
9
+ from autorag.support import get_support_modules
10
+ from autorag.data.parse.base import parser_node
11
+
12
+
13
+ @parser_node
14
+ def table_hybrid_parse(
15
+ data_path_list: List[str],
16
+ text_parse_module: str,
17
+ text_params: Dict,
18
+ table_parse_module: str,
19
+ table_params: Dict,
20
+ ) -> Tuple[List[str], List[str], List[int]]:
21
+ """
22
+ Parse documents to use table_hybrid_parse method.
23
+ The table_hybrid_parse method is a hybrid method that combines the parsing results of PDFs with and without tables.
24
+ It splits the PDF file into pages, separates pages with and without tables, and then parses and merges the results.
25
+
26
+ :param data_path_list: The list of data paths to parse.
27
+ :param text_parse_module: The text parsing module to use. The type should be a string.
28
+ :param text_params: The extra parameters for the text parsing module. The type should be a dictionary.
29
+ :param table_parse_module: The table parsing module to use. The type should be a string.
30
+ :param table_params: The extra parameters for the table parsing module. The type should be a dictionary.
31
+ :return: tuple of lists containing the parsed texts, path and pages.
32
+ """
33
+ # make save folder directory
34
+ with tempfile.TemporaryDirectory(ignore_cleanup_errors=True) as save_dir:
35
+ text_dir = os.path.join(save_dir, "text")
36
+ table_dir = os.path.join(save_dir, "table")
37
+
38
+ os.makedirs(text_dir, exist_ok=True)
39
+ os.makedirs(table_dir, exist_ok=True)
40
+
41
+ # Split PDF file into pages and Save PDFs with and without tables
42
+ path_map_dict_lst = [
43
+ save_page_by_table(data_path, text_dir, table_dir)
44
+ for data_path in data_path_list
45
+ ]
46
+ path_map_dict = {k: v for d in path_map_dict_lst for k, v in d.items()}
47
+
48
+ # Extract text pages
49
+ table_results, table_file_path = get_each_module_result(
50
+ table_parse_module, table_params, os.path.join(table_dir, "*")
51
+ )
52
+
53
+ # Extract table pages
54
+ text_results, text_file_path = get_each_module_result(
55
+ text_parse_module, text_params, os.path.join(text_dir, "*")
56
+ )
57
+
58
+ # Merge parsing results of PDFs with and without tables
59
+ texts = table_results + text_results
60
+ temp_path_lst = table_file_path + text_file_path
61
+
62
+ # Sort by file names
63
+ temp_path_lst, texts = zip(*sorted(zip(temp_path_lst, texts)))
64
+
65
+ # get original file path
66
+ path = list(map(lambda temp_path: path_map_dict[temp_path], temp_path_lst))
67
+
68
+ # get pages
69
+ pages = list(map(lambda x: get_page_from_path(x), temp_path_lst))
70
+
71
+ return list(texts), path, pages
72
+
73
+
74
+ # Save PDFs with and without tables
75
+ def save_page_by_table(data_path: str, text_dir: str, table_dir: str) -> Dict[str, str]:
76
+ file_name = os.path.basename(data_path).split(".pdf")[0]
77
+
78
+ with open(data_path, "rb") as input_data:
79
+ pdf_reader = PdfFileReader(input_data)
80
+ num_pages = pdf_reader.getNumPages()
81
+
82
+ path_map_dict = {}
83
+ for page_num in range(num_pages):
84
+ output_pdf_path = _get_output_path(
85
+ data_path, page_num, file_name, text_dir, table_dir
86
+ )
87
+ _save_single_page(pdf_reader, page_num, output_pdf_path)
88
+ path_map_dict.update({output_pdf_path: data_path})
89
+
90
+ return path_map_dict
91
+
92
+
93
+ def _get_output_path(
94
+ data_path: str, page_num: int, file_name: str, text_dir: str, table_dir: str
95
+ ) -> str:
96
+ with pdfplumber.open(data_path) as pdf:
97
+ page = pdf.pages[page_num]
98
+ tables = page.extract_tables()
99
+ directory = table_dir if tables else text_dir
100
+ return os.path.join(directory, f"{file_name}_page_{page_num + 1}.pdf")
101
+
102
+
103
+ def _save_single_page(pdf_reader: PdfFileReader, page_num: int, output_pdf_path: str):
104
+ pdf_writer = PdfFileWriter()
105
+ pdf_writer.addPage(pdf_reader.getPage(page_num))
106
+
107
+ with open(output_pdf_path, "wb") as output_file:
108
+ pdf_writer.write(output_file)
109
+
110
+
111
+ def get_each_module_result(
112
+ module: str, module_params: Dict, data_path_glob: str
113
+ ) -> Tuple[List[str], List[str]]:
114
+ module_params["module_type"] = module
115
+
116
+ data_path_list = glob(data_path_glob)
117
+ if not data_path_list:
118
+ return [], []
119
+
120
+ module_name = module_params.pop("module_type")
121
+ module_callable = get_support_modules(module_name)
122
+ module_original = module_callable.__wrapped__
123
+ texts, path, _ = module_original(data_path_list, **module_params)
124
+
125
+ return texts, path
126
+
127
+
128
+ def get_page_from_path(file_path: str) -> int:
129
+ file_name = os.path.basename(file_path)
130
+ split_result = file_name.rsplit("_page_", -1)
131
+ page_number_with_extension = split_result[1]
132
+ page_number, _ = page_number_with_extension.split(".")
133
+
134
+ return int(page_number)
@@ -0,0 +1,3 @@
1
+ # This is v2 version, the next version of data creation
2
+ # The legacy (v1) version will be deprecated on AutoRAG version 0.3
3
+ # The legacy (v1) version and new v2 data creation is not compatible with each other
File without changes
@@ -0,0 +1,64 @@
1
+ import itertools
2
+ from typing import Dict, List
3
+
4
+ from llama_index.core.base.llms.base import BaseLLM
5
+ from llama_index.core.base.llms.types import ChatResponse, ChatMessage, MessageRole
6
+
7
+ from autorag.data.qa.evolve.prompt import QUERY_EVOLVE_PROMPT
8
+
9
+
10
+ async def llama_index_generate_base(
11
+ row: Dict,
12
+ llm: BaseLLM,
13
+ messages: List[ChatMessage],
14
+ ) -> Dict:
15
+ original_query = row["query"]
16
+ context = list(itertools.chain.from_iterable(row["retrieval_gt_contents"]))
17
+ context_str = "Text:\n" + "\n".join(
18
+ [f"{i + 1}. {c}" for i, c in enumerate(context)]
19
+ )
20
+ user_prompt = f"Question: {original_query}\nContext: {context_str}\nOutput: "
21
+ messages.append(ChatMessage(role=MessageRole.USER, content=user_prompt))
22
+
23
+ chat_response: ChatResponse = await llm.achat(messages=messages)
24
+ row["query"] = chat_response.message.content
25
+ return row
26
+
27
+
28
+ async def conditional_evolve_ragas(
29
+ row: Dict,
30
+ llm: BaseLLM,
31
+ lang: str = "en",
32
+ ) -> Dict:
33
+ return await llama_index_generate_base(
34
+ row,
35
+ llm,
36
+ QUERY_EVOLVE_PROMPT["conditional_evolve_ragas"][lang],
37
+ )
38
+
39
+
40
+ async def reasoning_evolve_ragas(
41
+ row: Dict,
42
+ llm: BaseLLM,
43
+ lang: str = "en",
44
+ ) -> Dict:
45
+ return await llama_index_generate_base(
46
+ row,
47
+ llm,
48
+ QUERY_EVOLVE_PROMPT["reasoning_evolve_ragas"][lang],
49
+ )
50
+
51
+
52
+ async def compress_ragas(
53
+ row: Dict,
54
+ llm: BaseLLM,
55
+ lang: str = "en",
56
+ ) -> Dict:
57
+ original_query = row["query"]
58
+ user_prompt = f"Question: {original_query}\nOutput: "
59
+ messages = QUERY_EVOLVE_PROMPT["compress_ragas"][lang]
60
+ messages.append(ChatMessage(role=MessageRole.USER, content=user_prompt))
61
+
62
+ chat_response: ChatResponse = await llm.achat(messages=messages)
63
+ row["query"] = chat_response.message.content
64
+ return row
@@ -0,0 +1,81 @@
1
+ import itertools
2
+ from typing import Dict, List
3
+
4
+ from llama_index.core.base.llms.types import ChatMessage, MessageRole
5
+ from llama_index.llms.openai.utils import to_openai_message_dicts
6
+ from openai import AsyncClient
7
+ from pydantic import BaseModel
8
+
9
+ from autorag.data.qa.evolve.prompt import QUERY_EVOLVE_PROMPT
10
+
11
+
12
+ class Response(BaseModel):
13
+ evolved_query: str
14
+
15
+
16
+ async def query_evolve_openai_base(
17
+ row: Dict,
18
+ client: AsyncClient,
19
+ messages: List[ChatMessage],
20
+ model_name: str = "gpt-4o-2024-08-06",
21
+ ):
22
+ """
23
+ Evolve the original query to a new evolved query using OpenAI structured outputs.
24
+ """
25
+ original_query = row["query"]
26
+ context = list(itertools.chain.from_iterable(row["retrieval_gt_contents"]))
27
+ context_str = "Text:\n" + "\n".join(
28
+ [f"{i + 1}. {c}" for i, c in enumerate(context)]
29
+ )
30
+ user_prompt = f"Question: {original_query}\nContext: {context_str}\nOutput: "
31
+ messages.append(ChatMessage(role=MessageRole.USER, content=user_prompt))
32
+
33
+ completion = await client.beta.chat.completions.parse(
34
+ model=model_name,
35
+ messages=to_openai_message_dicts(messages),
36
+ response_format=Response,
37
+ )
38
+ row["query"] = completion.choices[0].message.parsed.evolved_query
39
+ return row
40
+
41
+
42
+ async def conditional_evolve_ragas(
43
+ row: Dict,
44
+ client: AsyncClient,
45
+ model_name: str = "gpt-4o-2024-08-06",
46
+ lang: str = "en",
47
+ ) -> Dict:
48
+ return await query_evolve_openai_base(
49
+ row, client, QUERY_EVOLVE_PROMPT["conditional_evolve_ragas"][lang], model_name
50
+ )
51
+
52
+
53
+ async def reasoning_evolve_ragas(
54
+ row: Dict,
55
+ client: AsyncClient,
56
+ model_name: str = "gpt-4o-2024-08-06",
57
+ lang: str = "en",
58
+ ) -> Dict:
59
+ return await query_evolve_openai_base(
60
+ row, client, QUERY_EVOLVE_PROMPT["reasoning_evolve_ragas"][lang], model_name
61
+ )
62
+
63
+
64
+ async def compress_ragas(
65
+ row: Dict,
66
+ client: AsyncClient,
67
+ model_name: str = "gpt-4o-2024-08-06",
68
+ lang: str = "en",
69
+ ) -> Dict:
70
+ original_query = row["query"]
71
+ messages = QUERY_EVOLVE_PROMPT["compress_ragas"][lang]
72
+ user_prompt = f"Question: {original_query}\nOutput: "
73
+ messages.append(ChatMessage(role=MessageRole.USER, content=user_prompt))
74
+
75
+ completion = await client.beta.chat.completions.parse(
76
+ model=model_name,
77
+ messages=to_openai_message_dicts(messages),
78
+ response_format=Response,
79
+ )
80
+ row["query"] = completion.choices[0].message.parsed.evolved_query
81
+ return row