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.
- autorag/__init__.py +82 -0
- autorag/chunker.py +51 -0
- autorag/cli.py +209 -0
- autorag/dashboard.py +199 -0
- autorag/data/__init__.py +109 -0
- autorag/data/chunk/__init__.py +2 -0
- autorag/data/chunk/base.py +128 -0
- autorag/data/chunk/langchain_chunk.py +76 -0
- autorag/data/chunk/llama_index_chunk.py +96 -0
- autorag/data/chunk/run.py +38 -0
- autorag/data/legacy/__init__.py +0 -0
- autorag/data/legacy/corpus/__init__.py +2 -0
- autorag/data/legacy/corpus/langchain.py +47 -0
- autorag/data/legacy/corpus/llama_index.py +93 -0
- autorag/data/legacy/qacreation/__init__.py +6 -0
- autorag/data/legacy/qacreation/base.py +239 -0
- autorag/data/legacy/qacreation/llama_index.py +253 -0
- autorag/data/legacy/qacreation/llama_index_default_prompt.txt +54 -0
- autorag/data/legacy/qacreation/ragas.py +75 -0
- autorag/data/legacy/qacreation/simple.py +99 -0
- autorag/data/parse/__init__.py +1 -0
- autorag/data/parse/base.py +79 -0
- autorag/data/parse/clova.py +194 -0
- autorag/data/parse/langchain_parse.py +87 -0
- autorag/data/parse/llamaparse.py +126 -0
- autorag/data/parse/run.py +141 -0
- autorag/data/parse/table_hybrid_parse.py +134 -0
- autorag/data/qa/__init__.py +3 -0
- autorag/data/qa/evolve/__init__.py +0 -0
- autorag/data/qa/evolve/llama_index_query_evolve.py +64 -0
- autorag/data/qa/evolve/openai_query_evolve.py +81 -0
- autorag/data/qa/evolve/prompt.py +288 -0
- autorag/data/qa/extract_evidence.py +1 -0
- autorag/data/qa/filter/__init__.py +0 -0
- autorag/data/qa/filter/dontknow.py +117 -0
- autorag/data/qa/filter/passage_dependency.py +88 -0
- autorag/data/qa/filter/prompt.py +73 -0
- autorag/data/qa/generation_gt/__init__.py +0 -0
- autorag/data/qa/generation_gt/base.py +16 -0
- autorag/data/qa/generation_gt/llama_index_gen_gt.py +41 -0
- autorag/data/qa/generation_gt/openai_gen_gt.py +84 -0
- autorag/data/qa/generation_gt/prompt.py +27 -0
- autorag/data/qa/query/__init__.py +0 -0
- autorag/data/qa/query/llama_gen_query.py +82 -0
- autorag/data/qa/query/openai_gen_query.py +95 -0
- autorag/data/qa/query/prompt.py +201 -0
- autorag/data/qa/sample.py +26 -0
- autorag/data/qa/schema.py +322 -0
- autorag/data/utils/__init__.py +0 -0
- autorag/data/utils/util.py +103 -0
- autorag/deploy/__init__.py +9 -0
- autorag/deploy/api.py +303 -0
- autorag/deploy/base.py +235 -0
- autorag/deploy/gradio.py +74 -0
- autorag/deploy/swagger.yml +202 -0
- autorag/embedding/__init__.py +0 -0
- autorag/embedding/base.py +144 -0
- autorag/embedding/vllm.py +256 -0
- autorag/evaluation/__init__.py +3 -0
- autorag/evaluation/generation.py +88 -0
- autorag/evaluation/metric/__init__.py +22 -0
- autorag/evaluation/metric/deepeval_prompt.py +322 -0
- autorag/evaluation/metric/g_eval_prompts/coh_detailed.txt +32 -0
- autorag/evaluation/metric/g_eval_prompts/con_detailed.txt +33 -0
- autorag/evaluation/metric/g_eval_prompts/flu_detailed.txt +26 -0
- autorag/evaluation/metric/g_eval_prompts/rel_detailed.txt +33 -0
- autorag/evaluation/metric/generation.py +504 -0
- autorag/evaluation/metric/retrieval.py +115 -0
- autorag/evaluation/metric/retrieval_contents.py +65 -0
- autorag/evaluation/metric/util.py +88 -0
- autorag/evaluation/retrieval.py +83 -0
- autorag/evaluation/retrieval_contents.py +65 -0
- autorag/evaluation/util.py +43 -0
- autorag/evaluator.py +559 -0
- autorag/node_line.py +65 -0
- autorag/nodes/__init__.py +0 -0
- autorag/nodes/generator/__init__.py +4 -0
- autorag/nodes/generator/base.py +103 -0
- autorag/nodes/generator/llama_index_llm.py +169 -0
- autorag/nodes/generator/openai_llm.py +329 -0
- autorag/nodes/generator/run.py +148 -0
- autorag/nodes/generator/vllm.py +147 -0
- autorag/nodes/generator/vllm_api.py +191 -0
- autorag/nodes/hybridretrieval/__init__.py +2 -0
- autorag/nodes/hybridretrieval/base.py +58 -0
- autorag/nodes/hybridretrieval/hybrid_cc.py +227 -0
- autorag/nodes/hybridretrieval/hybrid_rrf.py +149 -0
- autorag/nodes/hybridretrieval/run.py +137 -0
- autorag/nodes/lexicalretrieval/__init__.py +1 -0
- autorag/nodes/lexicalretrieval/bm25.py +381 -0
- autorag/nodes/lexicalretrieval/run.py +148 -0
- autorag/nodes/passageaugmenter/__init__.py +2 -0
- autorag/nodes/passageaugmenter/base.py +76 -0
- autorag/nodes/passageaugmenter/pass_passage_augmenter.py +43 -0
- autorag/nodes/passageaugmenter/prev_next_augmenter.py +155 -0
- autorag/nodes/passageaugmenter/run.py +131 -0
- autorag/nodes/passagecompressor/__init__.py +4 -0
- autorag/nodes/passagecompressor/base.py +78 -0
- autorag/nodes/passagecompressor/longllmlingua.py +115 -0
- autorag/nodes/passagecompressor/pass_compressor.py +16 -0
- autorag/nodes/passagecompressor/refine.py +54 -0
- autorag/nodes/passagecompressor/run.py +186 -0
- autorag/nodes/passagecompressor/tree_summarize.py +56 -0
- autorag/nodes/passagefilter/__init__.py +6 -0
- autorag/nodes/passagefilter/base.py +40 -0
- autorag/nodes/passagefilter/pass_passage_filter.py +14 -0
- autorag/nodes/passagefilter/percentile_cutoff.py +58 -0
- autorag/nodes/passagefilter/recency.py +105 -0
- autorag/nodes/passagefilter/run.py +138 -0
- autorag/nodes/passagefilter/similarity_percentile_cutoff.py +134 -0
- autorag/nodes/passagefilter/similarity_threshold_cutoff.py +112 -0
- autorag/nodes/passagefilter/threshold_cutoff.py +78 -0
- autorag/nodes/passagereranker/__init__.py +16 -0
- autorag/nodes/passagereranker/base.py +44 -0
- autorag/nodes/passagereranker/cohere.py +118 -0
- autorag/nodes/passagereranker/colbert.py +213 -0
- autorag/nodes/passagereranker/flag_embedding.py +112 -0
- autorag/nodes/passagereranker/flag_embedding_llm.py +101 -0
- autorag/nodes/passagereranker/flashrank.py +245 -0
- autorag/nodes/passagereranker/jina.py +115 -0
- autorag/nodes/passagereranker/koreranker.py +136 -0
- autorag/nodes/passagereranker/mixedbreadai.py +126 -0
- autorag/nodes/passagereranker/monot5.py +190 -0
- autorag/nodes/passagereranker/openvino.py +191 -0
- autorag/nodes/passagereranker/pass_reranker.py +31 -0
- autorag/nodes/passagereranker/rankgpt.py +170 -0
- autorag/nodes/passagereranker/run.py +145 -0
- autorag/nodes/passagereranker/sentence_transformer.py +129 -0
- autorag/nodes/passagereranker/tart/__init__.py +1 -0
- autorag/nodes/passagereranker/tart/modeling_enc_t5.py +152 -0
- autorag/nodes/passagereranker/tart/tart.py +139 -0
- autorag/nodes/passagereranker/tart/tokenization_enc_t5.py +112 -0
- autorag/nodes/passagereranker/time_reranker.py +72 -0
- autorag/nodes/passagereranker/upr.py +160 -0
- autorag/nodes/passagereranker/voyageai.py +109 -0
- autorag/nodes/promptmaker/__init__.py +12 -0
- autorag/nodes/promptmaker/base.py +32 -0
- autorag/nodes/promptmaker/chat_fstring.py +73 -0
- autorag/nodes/promptmaker/fstring.py +49 -0
- autorag/nodes/promptmaker/long_context_reorder.py +83 -0
- autorag/nodes/promptmaker/run.py +283 -0
- autorag/nodes/promptmaker/window_replacement.py +85 -0
- autorag/nodes/queryexpansion/__init__.py +4 -0
- autorag/nodes/queryexpansion/base.py +62 -0
- autorag/nodes/queryexpansion/hyde.py +43 -0
- autorag/nodes/queryexpansion/multi_query_expansion.py +57 -0
- autorag/nodes/queryexpansion/pass_query_expansion.py +22 -0
- autorag/nodes/queryexpansion/query_decompose.py +111 -0
- autorag/nodes/queryexpansion/run.py +308 -0
- autorag/nodes/retrieval/__init__.py +0 -0
- autorag/nodes/retrieval/base.py +127 -0
- autorag/nodes/retrieval/run_util.py +152 -0
- autorag/nodes/semanticretrieval/__init__.py +1 -0
- autorag/nodes/semanticretrieval/run.py +148 -0
- autorag/nodes/semanticretrieval/vectordb.py +339 -0
- autorag/nodes/util.py +16 -0
- autorag/parser.py +37 -0
- autorag/schema/__init__.py +3 -0
- autorag/schema/base.py +35 -0
- autorag/schema/metricinput.py +99 -0
- autorag/schema/module.py +24 -0
- autorag/schema/node.py +144 -0
- autorag/strategy.py +165 -0
- autorag/support.py +235 -0
- autorag/utils/__init__.py +8 -0
- autorag/utils/cast.py +45 -0
- autorag/utils/preprocess.py +149 -0
- autorag/utils/util.py +759 -0
- autorag/validator.py +98 -0
- autorag/vectordb/__init__.py +75 -0
- autorag/vectordb/base.py +73 -0
- autorag/vectordb/chroma.py +118 -0
- autorag/vectordb/couchbase.py +239 -0
- autorag/vectordb/milvus.py +169 -0
- autorag/vectordb/pinecone.py +121 -0
- autorag/vectordb/qdrant.py +155 -0
- autorag/vectordb/weaviate.py +184 -0
- autorag/web.py +81 -0
- autorag-0.0.0.dist-info/METADATA +780 -0
- autorag-0.0.0.dist-info/RECORD +184 -0
- autorag-0.0.0.dist-info/WHEEL +5 -0
- autorag-0.0.0.dist-info/entry_points.txt +2 -0
- autorag-0.0.0.dist-info/licenses/LICENSE +201 -0
- 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)
|
|
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
|