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,127 @@
|
|
|
1
|
+
import abc
|
|
2
|
+
import logging
|
|
3
|
+
import os
|
|
4
|
+
from typing import List, Union, Tuple
|
|
5
|
+
|
|
6
|
+
import pandas as pd
|
|
7
|
+
|
|
8
|
+
from autorag.schema import BaseModule
|
|
9
|
+
from autorag.support import get_support_modules
|
|
10
|
+
from autorag.utils import fetch_contents, result_to_dataframe, validate_qa_dataset
|
|
11
|
+
from autorag.utils.util import pop_params
|
|
12
|
+
|
|
13
|
+
logger = logging.getLogger("AutoRAG")
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class BaseRetrieval(BaseModule, metaclass=abc.ABCMeta):
|
|
17
|
+
def __init__(self, project_dir: str, *args, **kwargs):
|
|
18
|
+
logger.info(f"Initialize retrieval node - {self.__class__.__name__}")
|
|
19
|
+
|
|
20
|
+
self.resources_dir = os.path.join(project_dir, "resources")
|
|
21
|
+
data_dir = os.path.join(project_dir, "data")
|
|
22
|
+
# fetch data from corpus_data
|
|
23
|
+
self.corpus_df = pd.read_parquet(
|
|
24
|
+
os.path.join(data_dir, "corpus.parquet"), engine="pyarrow"
|
|
25
|
+
)
|
|
26
|
+
|
|
27
|
+
def __del__(self):
|
|
28
|
+
logger.info(f"Deleting retrieval node - {self.__class__.__name__} module...")
|
|
29
|
+
|
|
30
|
+
def cast_to_run(self, previous_result: pd.DataFrame, *args, **kwargs):
|
|
31
|
+
logger.info(f"Running retrieval node - {self.__class__.__name__} module...")
|
|
32
|
+
validate_qa_dataset(previous_result)
|
|
33
|
+
# find queries columns & type cast queries
|
|
34
|
+
assert "query" in previous_result.columns, (
|
|
35
|
+
"previous_result must have query column."
|
|
36
|
+
)
|
|
37
|
+
if "queries" not in previous_result.columns:
|
|
38
|
+
previous_result["queries"] = previous_result["query"]
|
|
39
|
+
previous_result.loc[:, "queries"] = previous_result["queries"].apply(
|
|
40
|
+
cast_queries
|
|
41
|
+
)
|
|
42
|
+
queries = previous_result["queries"].tolist()
|
|
43
|
+
return queries
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class HybridRetrieval(BaseRetrieval, metaclass=abc.ABCMeta):
|
|
47
|
+
def __init__(
|
|
48
|
+
self, project_dir: str, target_modules, target_module_params, *args, **kwargs
|
|
49
|
+
):
|
|
50
|
+
super().__init__(project_dir)
|
|
51
|
+
self.target_modules = list(
|
|
52
|
+
map(
|
|
53
|
+
lambda x, y: get_support_modules(x)(
|
|
54
|
+
**y,
|
|
55
|
+
project_dir=project_dir,
|
|
56
|
+
),
|
|
57
|
+
target_modules,
|
|
58
|
+
target_module_params,
|
|
59
|
+
)
|
|
60
|
+
)
|
|
61
|
+
self.target_module_params = target_module_params
|
|
62
|
+
|
|
63
|
+
@result_to_dataframe(["retrieved_contents", "retrieved_ids", "retrieve_scores"])
|
|
64
|
+
def pure(self, previous_result: pd.DataFrame, *args, **kwargs):
|
|
65
|
+
result_dfs: List[pd.DataFrame] = list(
|
|
66
|
+
map(
|
|
67
|
+
lambda x, y: x.pure(
|
|
68
|
+
**y,
|
|
69
|
+
previous_result=previous_result,
|
|
70
|
+
),
|
|
71
|
+
self.target_modules,
|
|
72
|
+
self.target_module_params,
|
|
73
|
+
)
|
|
74
|
+
)
|
|
75
|
+
ids = tuple(
|
|
76
|
+
map(lambda df: df["retrieved_ids"].apply(list).tolist(), result_dfs)
|
|
77
|
+
)
|
|
78
|
+
scores = tuple(
|
|
79
|
+
map(
|
|
80
|
+
lambda df: df["retrieve_scores"].apply(list).tolist(),
|
|
81
|
+
result_dfs,
|
|
82
|
+
)
|
|
83
|
+
)
|
|
84
|
+
|
|
85
|
+
_pure_params = pop_params(self._pure, kwargs)
|
|
86
|
+
if "ids" in _pure_params or "scores" in _pure_params:
|
|
87
|
+
raise ValueError(
|
|
88
|
+
"With specifying ids or scores, you must use HybridRRF.run_evaluator instead."
|
|
89
|
+
)
|
|
90
|
+
ids, scores = self._pure(ids=ids, scores=scores, **_pure_params)
|
|
91
|
+
contents = fetch_contents(self.corpus_df, ids)
|
|
92
|
+
return contents, ids, scores
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def cast_queries(queries: Union[str, List[str]]) -> List[str]:
|
|
96
|
+
if isinstance(queries, str):
|
|
97
|
+
return [queries]
|
|
98
|
+
elif isinstance(queries, List):
|
|
99
|
+
return queries
|
|
100
|
+
else:
|
|
101
|
+
raise ValueError(f"queries must be str or list, but got {type(queries)}")
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def evenly_distribute_passages(
|
|
105
|
+
ids: List[List[str]], scores: List[List[float]], top_k: int
|
|
106
|
+
) -> Tuple[List[str], List[float]]:
|
|
107
|
+
assert len(ids) == len(scores), "ids and scores must have same length."
|
|
108
|
+
query_cnt = len(ids)
|
|
109
|
+
avg_len = top_k // query_cnt
|
|
110
|
+
remainder = top_k % query_cnt
|
|
111
|
+
|
|
112
|
+
new_ids = []
|
|
113
|
+
new_scores = []
|
|
114
|
+
for i in range(query_cnt):
|
|
115
|
+
if i < remainder:
|
|
116
|
+
new_ids.extend(ids[i][: avg_len + 1])
|
|
117
|
+
new_scores.extend(scores[i][: avg_len + 1])
|
|
118
|
+
else:
|
|
119
|
+
new_ids.extend(ids[i][:avg_len])
|
|
120
|
+
new_scores.extend(scores[i][:avg_len])
|
|
121
|
+
|
|
122
|
+
return new_ids, new_scores
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def get_bm25_pkl_name(bm25_tokenizer: str):
|
|
126
|
+
bm25_tokenizer = bm25_tokenizer.replace("/", "")
|
|
127
|
+
return f"bm25_{bm25_tokenizer}.pkl"
|
|
@@ -0,0 +1,152 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import pathlib
|
|
3
|
+
from typing import Tuple, List, Union, Dict
|
|
4
|
+
|
|
5
|
+
import pandas as pd
|
|
6
|
+
|
|
7
|
+
from autorag.evaluation import evaluate_retrieval
|
|
8
|
+
from autorag.schema.metricinput import MetricInput
|
|
9
|
+
from autorag.strategy import measure_speed, filter_by_threshold, select_best
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def evaluate_retrieval_node(
|
|
13
|
+
result_df: pd.DataFrame,
|
|
14
|
+
metric_inputs: List[MetricInput],
|
|
15
|
+
metrics: Union[List[str], List[Dict]],
|
|
16
|
+
) -> pd.DataFrame:
|
|
17
|
+
"""
|
|
18
|
+
Evaluate retrieval node from retrieval node result dataframe.
|
|
19
|
+
:param result_df: The result dataframe from a retrieval node.
|
|
20
|
+
:param metric_inputs: List of metric input schema for AutoRAG.
|
|
21
|
+
:param metrics: Metric list from input strategies.
|
|
22
|
+
:return: Return result_df with metrics columns.
|
|
23
|
+
The columns will be 'retrieved_contents', 'retrieved_ids', 'retrieve_scores', and metric names.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
@evaluate_retrieval(
|
|
27
|
+
metric_inputs=metric_inputs,
|
|
28
|
+
metrics=metrics,
|
|
29
|
+
)
|
|
30
|
+
def evaluate_this_module(df: pd.DataFrame):
|
|
31
|
+
return (
|
|
32
|
+
df["retrieved_contents"].tolist(),
|
|
33
|
+
df["retrieved_ids"].tolist(),
|
|
34
|
+
df["retrieve_scores"].tolist(),
|
|
35
|
+
)
|
|
36
|
+
|
|
37
|
+
return evaluate_this_module(result_df)
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def run(
|
|
41
|
+
input_modules,
|
|
42
|
+
input_module_params,
|
|
43
|
+
project_dir: Union[str, pathlib.Path, pathlib.PurePath],
|
|
44
|
+
previous_result: pd.DataFrame,
|
|
45
|
+
strategies,
|
|
46
|
+
metric_inputs: List[MetricInput],
|
|
47
|
+
) -> Tuple[List[pd.DataFrame], List]:
|
|
48
|
+
"""
|
|
49
|
+
Run input modules and parameters.
|
|
50
|
+
:param input_modules: Input modules
|
|
51
|
+
:param input_module_params: Input module parameters
|
|
52
|
+
:param project_dir: Project directory path.
|
|
53
|
+
:param previous_result: Previous result dataframe.
|
|
54
|
+
:param strategies: Strategies for retrieval node.
|
|
55
|
+
:param metric_inputs: List of metric input schema for AutoRAG.
|
|
56
|
+
:return: First, it returns list of result dataframe.
|
|
57
|
+
Second, it returns list of execution times.
|
|
58
|
+
"""
|
|
59
|
+
result, execution_times = zip(
|
|
60
|
+
*map(
|
|
61
|
+
lambda task: measure_speed(
|
|
62
|
+
task[0].run_evaluator,
|
|
63
|
+
project_dir=project_dir,
|
|
64
|
+
previous_result=previous_result,
|
|
65
|
+
**task[1],
|
|
66
|
+
),
|
|
67
|
+
zip(input_modules, input_module_params),
|
|
68
|
+
)
|
|
69
|
+
)
|
|
70
|
+
average_times = list(map(lambda x: x / len(result[0]), execution_times))
|
|
71
|
+
|
|
72
|
+
# run metrics before filtering
|
|
73
|
+
if strategies.get("metrics") is None:
|
|
74
|
+
raise ValueError("You must at least one metrics for retrieval evaluation.")
|
|
75
|
+
result = list(
|
|
76
|
+
map(
|
|
77
|
+
lambda x: evaluate_retrieval_node(
|
|
78
|
+
x,
|
|
79
|
+
metric_inputs,
|
|
80
|
+
strategies.get("metrics"),
|
|
81
|
+
),
|
|
82
|
+
result,
|
|
83
|
+
)
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
return result, average_times
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def save_and_summary(
|
|
90
|
+
input_modules,
|
|
91
|
+
input_module_params,
|
|
92
|
+
result_list,
|
|
93
|
+
execution_time_list,
|
|
94
|
+
filename_start: int,
|
|
95
|
+
save_dir: Union[str, pathlib.Path, pathlib.PurePath],
|
|
96
|
+
strategies,
|
|
97
|
+
):
|
|
98
|
+
"""
|
|
99
|
+
Save the result and make summary file
|
|
100
|
+
:param input_modules: Input modules
|
|
101
|
+
:param input_module_params: Input module parameters
|
|
102
|
+
:param result_list: Result list
|
|
103
|
+
:param execution_time_list: Execution times
|
|
104
|
+
:param filename_start: The first filename to use
|
|
105
|
+
:return: First, it returns list of result dataframe.
|
|
106
|
+
Second, it returns list of execution times.
|
|
107
|
+
"""
|
|
108
|
+
|
|
109
|
+
# save results to folder
|
|
110
|
+
filepaths = list(
|
|
111
|
+
map(
|
|
112
|
+
lambda x: os.path.join(save_dir, f"{x}.parquet"),
|
|
113
|
+
range(filename_start, filename_start + len(input_modules)),
|
|
114
|
+
)
|
|
115
|
+
)
|
|
116
|
+
list(
|
|
117
|
+
map(
|
|
118
|
+
lambda x: x[0].to_parquet(x[1], index=False),
|
|
119
|
+
zip(result_list, filepaths),
|
|
120
|
+
)
|
|
121
|
+
) # execute save to parquet
|
|
122
|
+
filename_list = list(map(lambda x: os.path.basename(x), filepaths))
|
|
123
|
+
|
|
124
|
+
summary_df = pd.DataFrame(
|
|
125
|
+
{
|
|
126
|
+
"filename": filename_list,
|
|
127
|
+
"module_name": list(map(lambda module: module.__name__, input_modules)),
|
|
128
|
+
"module_params": input_module_params,
|
|
129
|
+
"execution_time": execution_time_list,
|
|
130
|
+
**{
|
|
131
|
+
metric: list(map(lambda result: result[metric].mean(), result_list))
|
|
132
|
+
for metric in strategies.get("metrics")
|
|
133
|
+
},
|
|
134
|
+
}
|
|
135
|
+
)
|
|
136
|
+
summary_df.to_csv(os.path.join(save_dir, "summary.csv"), index=False)
|
|
137
|
+
return summary_df
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def find_best(results, average_times, filenames, strategies):
|
|
141
|
+
# filter by strategies
|
|
142
|
+
if strategies.get("speed_threshold") is not None:
|
|
143
|
+
results, filenames = filter_by_threshold(
|
|
144
|
+
results, average_times, strategies["speed_threshold"], filenames
|
|
145
|
+
)
|
|
146
|
+
selected_result, selected_filename = select_best(
|
|
147
|
+
results,
|
|
148
|
+
strategies.get("metrics"),
|
|
149
|
+
filenames,
|
|
150
|
+
strategies.get("strategy", "mean"),
|
|
151
|
+
)
|
|
152
|
+
return selected_result, selected_filename
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
from .vectordb import VectorDB
|
|
@@ -0,0 +1,148 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import pathlib
|
|
3
|
+
from typing import List, Dict, Union
|
|
4
|
+
|
|
5
|
+
import pandas as pd
|
|
6
|
+
|
|
7
|
+
from autorag.evaluation import evaluate_retrieval
|
|
8
|
+
from autorag.evaluation.retrieval import RETRIEVAL_METRIC_FUNC_DICT
|
|
9
|
+
from autorag.nodes.retrieval.run_util import save_and_summary, find_best
|
|
10
|
+
from autorag.schema.metricinput import MetricInput
|
|
11
|
+
from autorag.strategy import measure_speed
|
|
12
|
+
from autorag.utils.util import apply_recursive, to_list
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def run_semantic_retrieval_node(
|
|
16
|
+
modules: List,
|
|
17
|
+
module_params: List[Dict],
|
|
18
|
+
previous_result: pd.DataFrame,
|
|
19
|
+
node_line_dir: str,
|
|
20
|
+
strategies: Dict,
|
|
21
|
+
) -> pd.DataFrame:
|
|
22
|
+
"""
|
|
23
|
+
Run the semantic retrieval node.
|
|
24
|
+
|
|
25
|
+
:param modules: Retrieval modules to run.
|
|
26
|
+
:param module_params: Retrieval module parameters.
|
|
27
|
+
:param previous_result: Previous result dataframe.
|
|
28
|
+
Could be query expansion's best result or qa data.
|
|
29
|
+
:param node_line_dir: This node line's directory.
|
|
30
|
+
:param strategies: Strategies for retrieval node.
|
|
31
|
+
:return: The best result dataframe.
|
|
32
|
+
It contains previous result columns and retrieval node's result columns.
|
|
33
|
+
"""
|
|
34
|
+
if not os.path.exists(node_line_dir):
|
|
35
|
+
os.makedirs(node_line_dir)
|
|
36
|
+
project_dir = pathlib.PurePath(node_line_dir).parent.parent
|
|
37
|
+
qa_df = pd.read_parquet(
|
|
38
|
+
os.path.join(project_dir, "data", "qa.parquet"), engine="pyarrow"
|
|
39
|
+
)
|
|
40
|
+
retrieval_gt = qa_df["retrieval_gt"].tolist()
|
|
41
|
+
retrieval_gt = apply_recursive(lambda x: str(x), to_list(retrieval_gt))
|
|
42
|
+
# make rows to metric_inputs
|
|
43
|
+
metric_inputs = [
|
|
44
|
+
MetricInput(retrieval_gt=ret_gt, query=query, generation_gt=gen_gt)
|
|
45
|
+
for ret_gt, query, gen_gt in zip(
|
|
46
|
+
retrieval_gt, qa_df["query"].tolist(), qa_df["generation_gt"].tolist()
|
|
47
|
+
)
|
|
48
|
+
]
|
|
49
|
+
|
|
50
|
+
save_dir = os.path.join(node_line_dir, "semantic_retrieval")
|
|
51
|
+
if not os.path.exists(save_dir):
|
|
52
|
+
os.makedirs(save_dir)
|
|
53
|
+
|
|
54
|
+
# Run the modules
|
|
55
|
+
semantic_results, execution_times = zip(
|
|
56
|
+
*map(
|
|
57
|
+
lambda task: measure_speed(
|
|
58
|
+
task[0].run_evaluator,
|
|
59
|
+
project_dir=project_dir,
|
|
60
|
+
previous_result=previous_result,
|
|
61
|
+
**task[1],
|
|
62
|
+
),
|
|
63
|
+
zip(modules, module_params),
|
|
64
|
+
)
|
|
65
|
+
)
|
|
66
|
+
semantic_times = list(map(lambda x: x / len(semantic_results[0]), execution_times))
|
|
67
|
+
|
|
68
|
+
# run metrics
|
|
69
|
+
if strategies.get("metrics") is None:
|
|
70
|
+
raise ValueError("You must at least one metrics for retrieval evaluation.")
|
|
71
|
+
semantic_results = list(
|
|
72
|
+
map(
|
|
73
|
+
lambda x: evaluate_semantic_retrieval_node(
|
|
74
|
+
x,
|
|
75
|
+
metric_inputs,
|
|
76
|
+
strategies.get("metrics"),
|
|
77
|
+
),
|
|
78
|
+
semantic_results,
|
|
79
|
+
)
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
semantic_summary_df = save_and_summary(
|
|
83
|
+
modules,
|
|
84
|
+
module_params,
|
|
85
|
+
semantic_results,
|
|
86
|
+
semantic_times,
|
|
87
|
+
0,
|
|
88
|
+
save_dir,
|
|
89
|
+
strategies,
|
|
90
|
+
)
|
|
91
|
+
semantic_selected_result, semantic_selected_filename = find_best(
|
|
92
|
+
semantic_results,
|
|
93
|
+
semantic_times,
|
|
94
|
+
semantic_summary_df["filename"].tolist(),
|
|
95
|
+
strategies,
|
|
96
|
+
)
|
|
97
|
+
semantic_summary_df["is_best"] = (
|
|
98
|
+
semantic_summary_df["filename"] == semantic_selected_filename
|
|
99
|
+
)
|
|
100
|
+
previous_result.drop(
|
|
101
|
+
columns=list(RETRIEVAL_METRIC_FUNC_DICT.keys()), inplace=True, errors="ignore"
|
|
102
|
+
)
|
|
103
|
+
semantic_selected_result.rename(
|
|
104
|
+
columns={
|
|
105
|
+
"retrieved_contents": "retrieved_contents_semantic",
|
|
106
|
+
"retrieved_ids": "retrieved_ids_semantic",
|
|
107
|
+
"retrieve_scores": "retrieve_scores_semantic",
|
|
108
|
+
},
|
|
109
|
+
inplace=True,
|
|
110
|
+
)
|
|
111
|
+
best_result = pd.concat([previous_result, semantic_selected_result], axis=1)
|
|
112
|
+
best_result.to_parquet(
|
|
113
|
+
os.path.join(
|
|
114
|
+
save_dir, f"best_{os.path.splitext(semantic_selected_filename)[0]}.parquet"
|
|
115
|
+
),
|
|
116
|
+
index=False,
|
|
117
|
+
)
|
|
118
|
+
semantic_summary_df.to_csv(os.path.join(save_dir, "summary.csv"), index=False)
|
|
119
|
+
return best_result # The result will be with _semantic suffix
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
def evaluate_semantic_retrieval_node(
|
|
123
|
+
result_df: pd.DataFrame,
|
|
124
|
+
metric_inputs: List[MetricInput],
|
|
125
|
+
metrics: Union[List[str], List[Dict]],
|
|
126
|
+
) -> pd.DataFrame:
|
|
127
|
+
"""
|
|
128
|
+
Evaluate retrieval node from retrieval node result dataframe.
|
|
129
|
+
|
|
130
|
+
:param result_df: The result dataframe from a retrieval node.
|
|
131
|
+
:param metric_inputs: List of metric input schema for AutoRAG.
|
|
132
|
+
:param metrics: Metric list from input strategies.
|
|
133
|
+
:return: Return result_df with metrics columns.
|
|
134
|
+
The columns will be 'retrieved_contents_semantic', 'retrieved_ids_semantic', 'retrieve_scores_semantic', and metric names.
|
|
135
|
+
"""
|
|
136
|
+
|
|
137
|
+
@evaluate_retrieval(
|
|
138
|
+
metric_inputs=metric_inputs,
|
|
139
|
+
metrics=metrics,
|
|
140
|
+
)
|
|
141
|
+
def evaluate_this_module(df: pd.DataFrame):
|
|
142
|
+
return (
|
|
143
|
+
df["retrieved_contents_semantic"].tolist(),
|
|
144
|
+
df["retrieved_ids_semantic"].tolist(),
|
|
145
|
+
df["retrieve_scores_semantic"].tolist(),
|
|
146
|
+
)
|
|
147
|
+
|
|
148
|
+
return evaluate_this_module(result_df)
|