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,99 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import pathlib
|
|
3
|
+
import uuid
|
|
4
|
+
from typing import Callable
|
|
5
|
+
|
|
6
|
+
import pandas as pd
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def generate_qa_row(llm, corpus_data_row):
|
|
10
|
+
"""
|
|
11
|
+
this sample code to generate rag dataset using OpenAI chat model
|
|
12
|
+
|
|
13
|
+
:param llm: guidance model
|
|
14
|
+
:param corpus_data_row: need "contents" column
|
|
15
|
+
:return: should to be dict which has "query", "generation_gt" columns at least.
|
|
16
|
+
"""
|
|
17
|
+
from guidance import gen
|
|
18
|
+
import guidance
|
|
19
|
+
|
|
20
|
+
temp_llm = llm
|
|
21
|
+
with guidance.user():
|
|
22
|
+
temp_llm += f"""
|
|
23
|
+
You have to found a passge to solve "the problem".
|
|
24
|
+
You need to build a clean and clear set of (problem, passage, answer) in json format
|
|
25
|
+
so that you don't have to ask about "the problem" again.
|
|
26
|
+
problem need to end with question mark("?").
|
|
27
|
+
The process of approaching the answer based on the information of the given passage
|
|
28
|
+
must be clearly and neatly displayed in the answer.\n
|
|
29
|
+
\n
|
|
30
|
+
Here is set of (problem, passage, answer) in JSON format:\n
|
|
31
|
+
{{\n
|
|
32
|
+
"passage": {corpus_data_row["contents"]}\n
|
|
33
|
+
"problem":
|
|
34
|
+
"""
|
|
35
|
+
|
|
36
|
+
with guidance.assistant():
|
|
37
|
+
temp_llm += gen("query", stop="?")
|
|
38
|
+
with guidance.user():
|
|
39
|
+
temp_llm += """
|
|
40
|
+
"answer":
|
|
41
|
+
"""
|
|
42
|
+
with guidance.assistant():
|
|
43
|
+
temp_llm += gen("generation_gt")
|
|
44
|
+
|
|
45
|
+
corpus_data_row["metadata"]["qa_generation"] = "simple"
|
|
46
|
+
|
|
47
|
+
response = {"query": temp_llm["query"], "generation_gt": temp_llm["generation_gt"]}
|
|
48
|
+
return response
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def generate_simple_qa_dataset(
|
|
52
|
+
llm,
|
|
53
|
+
corpus_data: pd.DataFrame,
|
|
54
|
+
output_filepath: str,
|
|
55
|
+
generate_row_function: Callable,
|
|
56
|
+
**kwargs,
|
|
57
|
+
):
|
|
58
|
+
"""
|
|
59
|
+
corpus_data to qa_dataset
|
|
60
|
+
qa_dataset will be saved to filepath(file_dir/filename)
|
|
61
|
+
|
|
62
|
+
:param llm: guidance.models.Model
|
|
63
|
+
:param corpus_data: pd.DataFrame. refer to the basic structure
|
|
64
|
+
:param output_filepath: file_dir must exist, filepath must not exist. file extension must be .parquet
|
|
65
|
+
:param generate_row_function: input(llm, corpus_data_row, kwargs) output(dict[columns contain "query" and "generation_gt"])
|
|
66
|
+
:param kwargs: if generate_row_function requires more args, use kwargs
|
|
67
|
+
:return: qa_dataset as pd.DataFrame
|
|
68
|
+
"""
|
|
69
|
+
output_file_dir = pathlib.PurePath(output_filepath).parent
|
|
70
|
+
if not os.path.isdir(output_file_dir):
|
|
71
|
+
raise NotADirectoryError(f"directory {output_file_dir} not found.")
|
|
72
|
+
if not output_filepath.endswith("parquet"):
|
|
73
|
+
raise NameError(
|
|
74
|
+
f'file path: {output_filepath} filename extension need to be ".parquet"'
|
|
75
|
+
)
|
|
76
|
+
if os.path.exists(output_filepath):
|
|
77
|
+
raise FileExistsError(
|
|
78
|
+
f"{output_filepath.split('/')[-1]} already exists in {output_file_dir}."
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
qa_data_lst = []
|
|
82
|
+
for _, corpus_data_row in corpus_data.iterrows():
|
|
83
|
+
response = generate_row_function(
|
|
84
|
+
llm=llm, corpus_data_row=corpus_data_row, **kwargs
|
|
85
|
+
)
|
|
86
|
+
qa_data_lst.append(
|
|
87
|
+
{
|
|
88
|
+
"qid": str(uuid.uuid4()),
|
|
89
|
+
"query": response["query"],
|
|
90
|
+
"retrieval_gt": [[corpus_data_row["doc_id"]]],
|
|
91
|
+
"generation_gt": [response["generation_gt"]],
|
|
92
|
+
"metadata": corpus_data_row["metadata"],
|
|
93
|
+
}
|
|
94
|
+
)
|
|
95
|
+
|
|
96
|
+
qa_dataset = pd.DataFrame(qa_data_lst)
|
|
97
|
+
qa_dataset.to_parquet(output_filepath, index=False)
|
|
98
|
+
|
|
99
|
+
return qa_dataset
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
from .langchain_parse import langchain_parse
|
|
@@ -0,0 +1,79 @@
|
|
|
1
|
+
import functools
|
|
2
|
+
import logging
|
|
3
|
+
from datetime import datetime
|
|
4
|
+
from glob import glob
|
|
5
|
+
from typing import Tuple, List, Optional
|
|
6
|
+
import os
|
|
7
|
+
|
|
8
|
+
from autorag.utils import result_to_dataframe
|
|
9
|
+
from autorag.data.utils.util import get_file_metadata
|
|
10
|
+
|
|
11
|
+
logger = logging.getLogger("AutoRAG")
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def parser_node(func):
|
|
15
|
+
@functools.wraps(func)
|
|
16
|
+
@result_to_dataframe(["texts", "path", "page", "last_modified_datetime"])
|
|
17
|
+
def wrapper(
|
|
18
|
+
data_path_glob: str,
|
|
19
|
+
file_type: str,
|
|
20
|
+
parse_method: Optional[str] = None,
|
|
21
|
+
**kwargs,
|
|
22
|
+
) -> Tuple[List[str], List[str], List[int], List[datetime]]:
|
|
23
|
+
logger.info(f"Running parser - {func.__name__} module...")
|
|
24
|
+
|
|
25
|
+
data_path_list = glob(data_path_glob)
|
|
26
|
+
if not data_path_list:
|
|
27
|
+
raise FileNotFoundError(f"data does not exits in {data_path_glob}")
|
|
28
|
+
|
|
29
|
+
assert file_type in [
|
|
30
|
+
"pdf",
|
|
31
|
+
"csv",
|
|
32
|
+
"json",
|
|
33
|
+
"md",
|
|
34
|
+
"html",
|
|
35
|
+
"xml",
|
|
36
|
+
"all_files",
|
|
37
|
+
], f"search type {file_type} is not supported"
|
|
38
|
+
|
|
39
|
+
# extract only files from data_path_list based on the file_type set in the YAML file
|
|
40
|
+
data_paths = (
|
|
41
|
+
[
|
|
42
|
+
data_path
|
|
43
|
+
for data_path in data_path_list
|
|
44
|
+
if os.path.basename(data_path).split(".")[-1] == file_type
|
|
45
|
+
]
|
|
46
|
+
if file_type != "all_files"
|
|
47
|
+
else data_path_list
|
|
48
|
+
)
|
|
49
|
+
|
|
50
|
+
if func.__name__ == "langchain_parse":
|
|
51
|
+
parse_method = parse_method.lower()
|
|
52
|
+
if parse_method == "directory":
|
|
53
|
+
path_split_list = data_path_glob.split("/")
|
|
54
|
+
glob_path = path_split_list.pop()
|
|
55
|
+
folder_path = "/".join(path_split_list)
|
|
56
|
+
kwargs.update({"glob": glob_path, "path": folder_path})
|
|
57
|
+
result = func(
|
|
58
|
+
data_path_list=data_paths, parse_method=parse_method, **kwargs
|
|
59
|
+
)
|
|
60
|
+
else:
|
|
61
|
+
result = func(
|
|
62
|
+
data_path_list=data_paths, parse_method=parse_method, **kwargs
|
|
63
|
+
)
|
|
64
|
+
elif func.__name__ in ["clova_ocr", "llama_parse", "table_hybrid_parse"]:
|
|
65
|
+
result = func(data_path_list=data_paths, **kwargs)
|
|
66
|
+
else:
|
|
67
|
+
raise ValueError(f"Unsupported module_type: {func.__name__}")
|
|
68
|
+
result = _add_last_modified_datetime(result)
|
|
69
|
+
return result
|
|
70
|
+
|
|
71
|
+
return wrapper
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def _add_last_modified_datetime(result):
|
|
75
|
+
last_modified_datetime_lst = list(
|
|
76
|
+
map(lambda x: get_file_metadata(x)["last_modified_datetime"], result[1])
|
|
77
|
+
)
|
|
78
|
+
result_with_dates = result + (last_modified_datetime_lst,)
|
|
79
|
+
return result_with_dates
|
|
@@ -0,0 +1,194 @@
|
|
|
1
|
+
import base64
|
|
2
|
+
import itertools
|
|
3
|
+
import json
|
|
4
|
+
import os
|
|
5
|
+
from typing import List, Optional, Tuple
|
|
6
|
+
|
|
7
|
+
import aiohttp
|
|
8
|
+
import fitz # PyMuPDF
|
|
9
|
+
|
|
10
|
+
from autorag.data.parse.base import parser_node
|
|
11
|
+
from autorag.utils.util import process_batch, get_event_loop
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@parser_node
|
|
15
|
+
def clova_ocr(
|
|
16
|
+
data_path_list: List[str],
|
|
17
|
+
url: Optional[str] = None,
|
|
18
|
+
api_key: Optional[str] = None,
|
|
19
|
+
batch: int = 5,
|
|
20
|
+
table_detection: bool = False,
|
|
21
|
+
) -> Tuple[List[str], List[str], List[int]]:
|
|
22
|
+
"""
|
|
23
|
+
Parse documents to use Naver Clova OCR.
|
|
24
|
+
|
|
25
|
+
:param data_path_list: The list of data paths to parse.
|
|
26
|
+
:param url: The URL for Clova OCR.
|
|
27
|
+
You can get the URL with the guide at https://guide.ncloud-docs.com/docs/clovaocr-example01
|
|
28
|
+
You can set the environment variable CLOVA_URL, or you can set it directly as a parameter.
|
|
29
|
+
:param api_key: The API key for Clova OCR.
|
|
30
|
+
You can get the API key with the guide at https://guide.ncloud-docs.com/docs/clovaocr-example01
|
|
31
|
+
You can set the environment variable CLOVA_API_KEY, or you can set it directly as a parameter.
|
|
32
|
+
:param batch: The batch size for parse documents. Default is 8.
|
|
33
|
+
:param table_detection: Whether to enable table detection. Default is False.
|
|
34
|
+
:return: tuple of lists containing the parsed texts, path and pages.
|
|
35
|
+
"""
|
|
36
|
+
url = os.getenv("CLOVA_URL", None) if url is None else url
|
|
37
|
+
if url is None:
|
|
38
|
+
raise KeyError(
|
|
39
|
+
"Please set the URL for Clova OCR in the environment variable CLOVA_URL "
|
|
40
|
+
"or directly set it on the config YAML file."
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
api_key = os.getenv("CLOVA_API_KEY", None) if api_key is None else api_key
|
|
44
|
+
if api_key is None:
|
|
45
|
+
raise KeyError(
|
|
46
|
+
"Please set the API key for Clova OCR in the environment variable CLOVA_API_KEY "
|
|
47
|
+
"or directly set it on the config YAML file."
|
|
48
|
+
)
|
|
49
|
+
if batch > 5:
|
|
50
|
+
raise ValueError("The batch size should be less than or equal to 5.")
|
|
51
|
+
|
|
52
|
+
image_data_lst = list(
|
|
53
|
+
map(lambda data_path: pdf_to_images(data_path), data_path_list)
|
|
54
|
+
)
|
|
55
|
+
image_info_lst = [
|
|
56
|
+
generate_image_info(pdf_path, len(image_data))
|
|
57
|
+
for pdf_path, image_data in zip(data_path_list, image_data_lst)
|
|
58
|
+
]
|
|
59
|
+
|
|
60
|
+
image_data_list = list(itertools.chain(*image_data_lst))
|
|
61
|
+
image_info_list = list(itertools.chain(*image_info_lst))
|
|
62
|
+
|
|
63
|
+
tasks = [
|
|
64
|
+
clova_ocr_pure(image_data, image_info, url, api_key, table_detection)
|
|
65
|
+
for image_data, image_info in zip(image_data_list, image_info_list)
|
|
66
|
+
]
|
|
67
|
+
loop = get_event_loop()
|
|
68
|
+
results = loop.run_until_complete(process_batch(tasks, batch))
|
|
69
|
+
|
|
70
|
+
texts, path, pages = zip(*results)
|
|
71
|
+
return list(texts), list(path), list(pages)
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
async def clova_ocr_pure(
|
|
75
|
+
image_data: bytes,
|
|
76
|
+
image_info: dict,
|
|
77
|
+
url: str,
|
|
78
|
+
api_key: str,
|
|
79
|
+
table_detection: bool = False,
|
|
80
|
+
) -> Tuple[str, str, int]:
|
|
81
|
+
session = aiohttp.ClientSession()
|
|
82
|
+
table_html = ""
|
|
83
|
+
headers = {"X-OCR-SECRET": api_key, "Content-Type": "application/json"}
|
|
84
|
+
|
|
85
|
+
# Convert image data to base64
|
|
86
|
+
image_base64 = base64.b64encode(image_data).decode("utf-8")
|
|
87
|
+
|
|
88
|
+
# Set data
|
|
89
|
+
data = {
|
|
90
|
+
"version": "V2",
|
|
91
|
+
"requestId": "sample_id",
|
|
92
|
+
"timestamp": 0,
|
|
93
|
+
"images": [{"format": "png", "name": "sample_image", "data": image_base64}],
|
|
94
|
+
"enableTableDetection": table_detection,
|
|
95
|
+
}
|
|
96
|
+
|
|
97
|
+
async with session.post(url, headers=headers, data=json.dumps(data)) as response:
|
|
98
|
+
resp_json = await response.json()
|
|
99
|
+
if "images" not in resp_json:
|
|
100
|
+
raise RuntimeError(
|
|
101
|
+
f"Invalid response from Clova API: {resp_json['detail']}"
|
|
102
|
+
)
|
|
103
|
+
if "tables" in resp_json["images"][0].keys():
|
|
104
|
+
table_html = json_to_html_table(
|
|
105
|
+
resp_json["images"][0]["tables"][0]["cells"]
|
|
106
|
+
)
|
|
107
|
+
page_text = extract_text_from_fields(resp_json["images"][0]["fields"])
|
|
108
|
+
|
|
109
|
+
if table_html:
|
|
110
|
+
page_text += f"\n\ntable html:\n{table_html}"
|
|
111
|
+
|
|
112
|
+
await session.close()
|
|
113
|
+
return page_text, image_info["pdf_path"], image_info["pdf_page"]
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
def pdf_to_images(pdf_path: str) -> List[bytes]:
|
|
117
|
+
"""Convert each page of the PDF to an image and return the image data."""
|
|
118
|
+
pdf_document = fitz.open(pdf_path)
|
|
119
|
+
image_data_lst = []
|
|
120
|
+
for page_num in range(len(pdf_document)):
|
|
121
|
+
page = pdf_document.load_page(page_num)
|
|
122
|
+
pix = page.get_pixmap()
|
|
123
|
+
img_data = pix.tobytes("png")
|
|
124
|
+
image_data_lst.append(img_data)
|
|
125
|
+
return image_data_lst
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
def generate_image_info(pdf_path: str, num_pages: int) -> List[dict]:
|
|
129
|
+
"""Generate image names based on the PDF file name and the number of pages."""
|
|
130
|
+
image_info_lst = [
|
|
131
|
+
{"pdf_path": pdf_path, "pdf_page": page_num + 1}
|
|
132
|
+
for page_num in range(num_pages)
|
|
133
|
+
]
|
|
134
|
+
return image_info_lst
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
def extract_text_from_fields(fields):
|
|
138
|
+
text = ""
|
|
139
|
+
for field in fields:
|
|
140
|
+
text += field["inferText"]
|
|
141
|
+
if field["lineBreak"]:
|
|
142
|
+
text += "\n"
|
|
143
|
+
else:
|
|
144
|
+
text += " "
|
|
145
|
+
return text.strip()
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def json_to_html_table(json_data):
|
|
149
|
+
# Initialize the HTML table
|
|
150
|
+
html = '<table border="1">\n'
|
|
151
|
+
# Determine the number of rows and columns
|
|
152
|
+
max_row = max(cell["rowIndex"] + cell["rowSpan"] for cell in json_data)
|
|
153
|
+
max_col = max(cell["columnIndex"] + cell["columnSpan"] for cell in json_data)
|
|
154
|
+
# Create a 2D array to keep track of merged cells
|
|
155
|
+
table = [["" for _ in range(max_col)] for _ in range(max_row)]
|
|
156
|
+
# Fill the table with cell data
|
|
157
|
+
for cell in json_data:
|
|
158
|
+
row = cell["rowIndex"]
|
|
159
|
+
col = cell["columnIndex"]
|
|
160
|
+
row_span = cell["rowSpan"]
|
|
161
|
+
col_span = cell["columnSpan"]
|
|
162
|
+
cell_text = (
|
|
163
|
+
" ".join(
|
|
164
|
+
line["inferText"] for line in cell["cellTextLines"][0]["cellWords"]
|
|
165
|
+
)
|
|
166
|
+
if cell["cellTextLines"]
|
|
167
|
+
else ""
|
|
168
|
+
)
|
|
169
|
+
# Place the cell in the table
|
|
170
|
+
table[row][col] = {"text": cell_text, "rowSpan": row_span, "colSpan": col_span}
|
|
171
|
+
# Mark merged cells as occupied
|
|
172
|
+
for r in range(row, row + row_span):
|
|
173
|
+
for c in range(col, col + col_span):
|
|
174
|
+
if r != row or c != col:
|
|
175
|
+
table[r][c] = None
|
|
176
|
+
# Generate HTML from the table array
|
|
177
|
+
for row in table:
|
|
178
|
+
html += " <tr>\n"
|
|
179
|
+
for cell in row:
|
|
180
|
+
if cell is None:
|
|
181
|
+
continue
|
|
182
|
+
if cell == "":
|
|
183
|
+
html += " <td></td>\n"
|
|
184
|
+
else:
|
|
185
|
+
row_span_attr = (
|
|
186
|
+
f' rowspan="{cell["rowSpan"]}"' if cell["rowSpan"] > 1 else ""
|
|
187
|
+
)
|
|
188
|
+
col_span_attr = (
|
|
189
|
+
f' colspan="{cell["colSpan"]}"' if cell["colSpan"] > 1 else ""
|
|
190
|
+
)
|
|
191
|
+
html += f" <td{row_span_attr}{col_span_attr}>{cell['text']}</td>\n"
|
|
192
|
+
html += " </tr>\n"
|
|
193
|
+
html += "</table>"
|
|
194
|
+
return html
|
|
@@ -0,0 +1,87 @@
|
|
|
1
|
+
import multiprocessing as mp
|
|
2
|
+
from itertools import chain
|
|
3
|
+
from typing import List, Tuple
|
|
4
|
+
|
|
5
|
+
from autorag.data import parse_modules
|
|
6
|
+
from autorag.data.parse.base import parser_node
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@parser_node
|
|
10
|
+
def langchain_parse(
|
|
11
|
+
data_path_list: List[str], parse_method: str, **kwargs
|
|
12
|
+
) -> Tuple[List[str], List[str], List[int]]:
|
|
13
|
+
"""
|
|
14
|
+
Parse documents to use langchain document_loaders(parse) method
|
|
15
|
+
|
|
16
|
+
:param data_path_list: The list of data paths to parse.
|
|
17
|
+
:param parse_method: A langchain document_loaders(parse) method to use.
|
|
18
|
+
:param kwargs: The extra parameters for creating the langchain document_loaders(parse) instance.
|
|
19
|
+
:return: tuple of lists containing the parsed texts, path and pages.
|
|
20
|
+
"""
|
|
21
|
+
if parse_method in ["directory", "unstructured"]:
|
|
22
|
+
results = parse_all_files(data_path_list, parse_method, **kwargs)
|
|
23
|
+
texts, path = results[0], results[1]
|
|
24
|
+
pages = [-1] * len(texts)
|
|
25
|
+
|
|
26
|
+
else:
|
|
27
|
+
num_workers = mp.cpu_count()
|
|
28
|
+
# Execute parallel processing
|
|
29
|
+
with mp.Pool(num_workers) as pool:
|
|
30
|
+
results = pool.starmap(
|
|
31
|
+
langchain_parse_pure,
|
|
32
|
+
[(data_path, parse_method, kwargs) for data_path in data_path_list],
|
|
33
|
+
)
|
|
34
|
+
|
|
35
|
+
texts, path, pages = (list(chain.from_iterable(item)) for item in zip(*results))
|
|
36
|
+
|
|
37
|
+
return texts, path, pages
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def langchain_parse_pure(
|
|
41
|
+
data_path: str, parse_method: str, kwargs
|
|
42
|
+
) -> Tuple[List[str], List[str], List[int]]:
|
|
43
|
+
"""
|
|
44
|
+
Parses a single file using the specified parse method.
|
|
45
|
+
|
|
46
|
+
Args:
|
|
47
|
+
data_path (str): The file path to parse.
|
|
48
|
+
parse_method (str): The parsing method to use.
|
|
49
|
+
kwargs (Dict): Additional keyword arguments for the parsing method.
|
|
50
|
+
|
|
51
|
+
Returns:
|
|
52
|
+
Tuple[str, str]: A tuple containing the parsed text and the file path.
|
|
53
|
+
"""
|
|
54
|
+
|
|
55
|
+
parse_instance = parse_modules[parse_method](data_path, **kwargs)
|
|
56
|
+
|
|
57
|
+
# Load the text from the file
|
|
58
|
+
documents = parse_instance.load()
|
|
59
|
+
|
|
60
|
+
texts = list(map(lambda x: x.page_content, documents))
|
|
61
|
+
path = [data_path] * len(texts)
|
|
62
|
+
if parse_method in ["pymupdf", "pdfplumber", "pypdf", "pypdfium2"]:
|
|
63
|
+
pages = list(range(1, len(documents) + 1))
|
|
64
|
+
else:
|
|
65
|
+
pages = [-1] * len(texts)
|
|
66
|
+
|
|
67
|
+
# Clean up the parse instance
|
|
68
|
+
del parse_instance
|
|
69
|
+
|
|
70
|
+
return texts, path, pages
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def parse_all_files(
|
|
74
|
+
data_path_list: List[str], parse_method: str, **kwargs
|
|
75
|
+
) -> Tuple[List[str], List[str]]:
|
|
76
|
+
if parse_method == "unstructured":
|
|
77
|
+
parse_instance = parse_modules[parse_method](data_path_list, **kwargs)
|
|
78
|
+
elif parse_method == "directory":
|
|
79
|
+
parse_instance = parse_modules[parse_method](**kwargs)
|
|
80
|
+
else:
|
|
81
|
+
raise ValueError(f"Unsupported parse method: {parse_method}")
|
|
82
|
+
docs = parse_instance.load()
|
|
83
|
+
texts = [doc.page_content for doc in docs]
|
|
84
|
+
file_names = [doc.metadata["source"] for doc in docs]
|
|
85
|
+
|
|
86
|
+
del parse_instance
|
|
87
|
+
return texts, file_names
|
|
@@ -0,0 +1,126 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from typing import List, Tuple
|
|
3
|
+
from itertools import chain
|
|
4
|
+
|
|
5
|
+
from llama_parse import LlamaParse
|
|
6
|
+
|
|
7
|
+
from autorag.data.parse.base import parser_node
|
|
8
|
+
from autorag.utils.util import process_batch, get_event_loop
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@parser_node
|
|
12
|
+
def llama_parse(
|
|
13
|
+
data_path_list: List[str],
|
|
14
|
+
batch: int = 8,
|
|
15
|
+
use_vendor_multimodal_model: bool = False,
|
|
16
|
+
vendor_multimodal_model_name: str = "openai-gpt4o",
|
|
17
|
+
use_own_key: bool = False,
|
|
18
|
+
vendor_multimodal_api_key: str = None,
|
|
19
|
+
**kwargs,
|
|
20
|
+
) -> Tuple[List[str], List[str], List[int]]:
|
|
21
|
+
"""
|
|
22
|
+
Parse documents to use llama_parse.
|
|
23
|
+
LLAMA_CLOUD_API_KEY environment variable should be set.
|
|
24
|
+
You can get the key from https://cloud.llamaindex.ai/api-key
|
|
25
|
+
|
|
26
|
+
:param data_path_list: The list of data paths to parse.
|
|
27
|
+
:param batch: The batch size for parse documents. Default is 8.
|
|
28
|
+
:param use_vendor_multimodal_model: Whether to use the vendor multimodal model. Default is False.
|
|
29
|
+
:param vendor_multimodal_model_name: The name of the vendor multimodal model. Default is "openai-gpt4o".
|
|
30
|
+
:param use_own_key: Whether to use the own API key. Default is False.
|
|
31
|
+
:param vendor_multimodal_api_key: The API key for the vendor multimodal model.
|
|
32
|
+
:param kwargs: The extra parameters for creating the llama_parse instance.
|
|
33
|
+
:return: tuple of lists containing the parsed texts, path and pages.
|
|
34
|
+
"""
|
|
35
|
+
if use_vendor_multimodal_model:
|
|
36
|
+
kwargs = _add_multimodal_params(
|
|
37
|
+
kwargs,
|
|
38
|
+
use_vendor_multimodal_model,
|
|
39
|
+
vendor_multimodal_model_name,
|
|
40
|
+
use_own_key,
|
|
41
|
+
vendor_multimodal_api_key,
|
|
42
|
+
)
|
|
43
|
+
|
|
44
|
+
parse_instance = LlamaParse(**kwargs)
|
|
45
|
+
|
|
46
|
+
tasks = [
|
|
47
|
+
llama_parse_pure(data_path, parse_instance) for data_path in data_path_list
|
|
48
|
+
]
|
|
49
|
+
loop = get_event_loop()
|
|
50
|
+
results = loop.run_until_complete(process_batch(tasks, batch))
|
|
51
|
+
|
|
52
|
+
del parse_instance
|
|
53
|
+
|
|
54
|
+
texts, path, pages = (list(chain.from_iterable(item)) for item in zip(*results))
|
|
55
|
+
|
|
56
|
+
return texts, path, pages
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
async def llama_parse_pure(
|
|
60
|
+
data_path: str, parse_instance
|
|
61
|
+
) -> Tuple[List[str], List[str], List[int]]:
|
|
62
|
+
documents = await parse_instance.aload_data(data_path)
|
|
63
|
+
|
|
64
|
+
texts = list(map(lambda x: x.text, documents))
|
|
65
|
+
path = [data_path] * len(texts)
|
|
66
|
+
pages = list(range(1, len(documents) + 1))
|
|
67
|
+
|
|
68
|
+
return texts, path, pages
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def _add_multimodal_params(
|
|
72
|
+
kwargs,
|
|
73
|
+
use_vendor_multimodal_model,
|
|
74
|
+
vendor_multimodal_model_name,
|
|
75
|
+
use_own_key,
|
|
76
|
+
vendor_multimodal_api_key,
|
|
77
|
+
) -> dict:
|
|
78
|
+
kwargs["use_vendor_multimodal_model"] = use_vendor_multimodal_model
|
|
79
|
+
kwargs["vendor_multimodal_model_name"] = vendor_multimodal_model_name
|
|
80
|
+
|
|
81
|
+
def set_multimodal_api_key(
|
|
82
|
+
multimodal_model_name: str = "openai-gpt4o", _api_key: str = None
|
|
83
|
+
) -> str:
|
|
84
|
+
if multimodal_model_name in ["openai-gpt4o", "openai-gpt-4o-mini"]:
|
|
85
|
+
_api_key = (
|
|
86
|
+
os.getenv("OPENAI_API_KEY", None) if _api_key is None else _api_key
|
|
87
|
+
)
|
|
88
|
+
if _api_key is None:
|
|
89
|
+
raise KeyError(
|
|
90
|
+
"Please set the OPENAI_API_KEY in the environment variable OPENAI_API_KEY "
|
|
91
|
+
"or directly set it on the config YAML file."
|
|
92
|
+
)
|
|
93
|
+
elif multimodal_model_name in ["anthropic-sonnet-3.5"]:
|
|
94
|
+
_api_key = (
|
|
95
|
+
os.getenv("ANTHROPIC_API_KEY", None) if _api_key is None else _api_key
|
|
96
|
+
)
|
|
97
|
+
if _api_key is None:
|
|
98
|
+
raise KeyError(
|
|
99
|
+
"Please set the ANTHROPIC_API_KEY in the environment variable ANTHROPIC_API_KEY "
|
|
100
|
+
"or directly set it on the config YAML file."
|
|
101
|
+
)
|
|
102
|
+
elif multimodal_model_name in ["gemini-1.5-flash", "gemini-1.5-pro"]:
|
|
103
|
+
_api_key = (
|
|
104
|
+
os.getenv("GEMINI_API_KEY", None) if _api_key is None else _api_key
|
|
105
|
+
)
|
|
106
|
+
if _api_key is None:
|
|
107
|
+
raise KeyError(
|
|
108
|
+
"Please set the GEMINI_API_KEY in the environment variable GEMINI_API_KEY "
|
|
109
|
+
"or directly set it on the config YAML file."
|
|
110
|
+
)
|
|
111
|
+
elif multimodal_model_name in ["custom-azure-model"]:
|
|
112
|
+
raise NotImplementedError(
|
|
113
|
+
"Custom Azure multimodal model is not supported yet."
|
|
114
|
+
)
|
|
115
|
+
else:
|
|
116
|
+
raise ValueError("Invalid multimodal model name.")
|
|
117
|
+
|
|
118
|
+
return _api_key
|
|
119
|
+
|
|
120
|
+
if use_own_key:
|
|
121
|
+
api_key = set_multimodal_api_key(
|
|
122
|
+
vendor_multimodal_model_name, vendor_multimodal_api_key
|
|
123
|
+
)
|
|
124
|
+
kwargs["vendor_multimodal_api_key"] = api_key
|
|
125
|
+
|
|
126
|
+
return kwargs
|