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,169 @@
1
+ import logging
2
+ from typing import Any, Dict, List, Tuple, Optional, Union
3
+
4
+ from pymilvus import (
5
+ DataType,
6
+ FieldSchema,
7
+ CollectionSchema,
8
+ connections,
9
+ Collection,
10
+ MilvusException,
11
+ )
12
+ from pymilvus.orm import utility
13
+
14
+ from autorag.utils.util import apply_recursive
15
+ from autorag.vectordb import BaseVectorStore
16
+
17
+
18
+ logger = logging.getLogger("AutoRAG")
19
+
20
+
21
+ class Milvus(BaseVectorStore):
22
+ def __init__(
23
+ self,
24
+ embedding_model: Union[str, List[dict]],
25
+ collection_name: str,
26
+ embedding_batch: int = 100,
27
+ similarity_metric: str = "cosine",
28
+ index_type: str = "IVF_FLAT",
29
+ uri: str = "http://localhost:19530",
30
+ db_name: str = "",
31
+ token: str = "",
32
+ user: str = "",
33
+ password: str = "",
34
+ timeout: Optional[float] = None,
35
+ params: Dict[str, Any] = {},
36
+ ):
37
+ super().__init__(embedding_model, similarity_metric, embedding_batch)
38
+
39
+ # Connect to Milvus server
40
+ connections.connect(
41
+ "default",
42
+ uri=uri,
43
+ token=token,
44
+ db_name=db_name,
45
+ user=user,
46
+ password=password,
47
+ )
48
+ self.collection_name = collection_name
49
+ self.timeout = timeout
50
+ self.params = params
51
+ self.index_type = index_type
52
+
53
+ # Set Collection
54
+ if not utility.has_collection(collection_name, timeout=timeout):
55
+ # Get the dimension of the embeddings
56
+ test_embedding_result: List[float] = self.embedding.get_query_embedding(
57
+ "test"
58
+ )
59
+ dimension = len(test_embedding_result)
60
+
61
+ pk = FieldSchema(
62
+ name="id",
63
+ dtype=DataType.VARCHAR,
64
+ max_length=128,
65
+ is_primary=True,
66
+ auto_id=False,
67
+ )
68
+ field = FieldSchema(
69
+ name="vector", dtype=DataType.FLOAT_VECTOR, dim=dimension
70
+ )
71
+ schema = CollectionSchema(fields=[pk, field])
72
+
73
+ self.collection = Collection(name=self.collection_name, schema=schema)
74
+ index_params = {
75
+ "metric_type": self.similarity_metric.upper(),
76
+ "index_type": self.index_type.upper(),
77
+ "params": self.params,
78
+ }
79
+ self.collection.create_index(
80
+ field_name="vector", index_params=index_params, timeout=self.timeout
81
+ )
82
+ else:
83
+ self.collection = Collection(name=self.collection_name)
84
+
85
+ async def add(self, ids: List[str], texts: List[str]):
86
+ texts = self.truncated_inputs(texts)
87
+ text_embeddings: List[
88
+ List[float]
89
+ ] = await self.embedding.aget_text_embedding_batch(texts)
90
+ self.add_embedding(ids, text_embeddings)
91
+
92
+ def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
93
+ data = list(
94
+ map(lambda _id, vector: {"id": _id, "vector": vector}, ids, embeddings)
95
+ )
96
+
97
+ # Insert data into the collection
98
+ res = self.collection.insert(data=data, timeout=self.timeout)
99
+ assert res.insert_count == len(ids), (
100
+ f"Insertion failed. Try to insert {len(ids)} but only {res['insert_count']} inserted."
101
+ )
102
+
103
+ self.collection.flush(timeout=self.timeout)
104
+
105
+ async def query(
106
+ self, queries: List[str], top_k: int, **kwargs
107
+ ) -> Tuple[List[List[str]], List[List[float]]]:
108
+ queries = self.truncated_inputs(queries)
109
+ query_embeddings: List[
110
+ List[float]
111
+ ] = await self.embedding.aget_text_embedding_batch(queries)
112
+
113
+ self.collection.load(timeout=self.timeout)
114
+
115
+ # Perform similarity search
116
+ results = self.collection.search(
117
+ data=query_embeddings,
118
+ limit=top_k,
119
+ anns_field="vector",
120
+ param={"metric_type": self.similarity_metric.upper()},
121
+ timeout=self.timeout,
122
+ **kwargs,
123
+ )
124
+
125
+ # Extract IDs and distances
126
+ ids = [[str(hit.id) for hit in result] for result in results]
127
+ distances = [[hit.distance for hit in result] for result in results]
128
+
129
+ if self.similarity_metric in ["l2"]:
130
+ distances = apply_recursive(lambda x: -x, distances)
131
+
132
+ return ids, distances
133
+
134
+ async def fetch(self, ids: List[str]) -> List[List[float]]:
135
+ try:
136
+ self.collection.load(timeout=self.timeout)
137
+ except MilvusException as e:
138
+ logger.warning(f"Failed to load collection: {e}")
139
+ return [[]] * len(ids)
140
+ # Fetch vectors by IDs
141
+ results = self.collection.query(
142
+ expr=f"id in {ids}", output_fields=["id", "vector"], timeout=self.timeout
143
+ )
144
+ id_vector_dict = {str(result["id"]): result["vector"] for result in results}
145
+ result = [id_vector_dict[_id] for _id in ids]
146
+ return result
147
+
148
+ async def is_exist(self, ids: List[str]) -> List[bool]:
149
+ try:
150
+ self.collection.load(timeout=self.timeout)
151
+ except MilvusException:
152
+ return [False] * len(ids)
153
+ # Check the existence of IDs
154
+ results = self.collection.query(
155
+ expr=f"id in {ids}", output_fields=["id"], timeout=self.timeout
156
+ )
157
+ # Determine existence
158
+ existing_ids = {str(result["id"]) for result in results}
159
+ return [str(_id) in existing_ids for _id in ids]
160
+
161
+ async def delete(self, ids: List[str]):
162
+ # Delete entries by IDs
163
+ self.collection.delete(expr=f"id in {ids}", timeout=self.timeout)
164
+
165
+ def delete_collection(self):
166
+ # Delete the collection
167
+ self.collection.release(timeout=self.timeout)
168
+ self.collection.drop_index(timeout=self.timeout)
169
+ self.collection.drop(timeout=self.timeout)
@@ -0,0 +1,121 @@
1
+ import logging
2
+
3
+ from pinecone.grpc import PineconeGRPC as Pinecone_client
4
+ from pinecone import ServerlessSpec
5
+
6
+ from typing import List, Optional, Tuple, Union
7
+
8
+ from autorag.utils.util import make_batch, apply_recursive
9
+ from autorag.vectordb import BaseVectorStore
10
+
11
+ logger = logging.getLogger("AutoRAG")
12
+
13
+
14
+ class Pinecone(BaseVectorStore):
15
+ def __init__(
16
+ self,
17
+ embedding_model: Union[str, List[dict]],
18
+ index_name: str,
19
+ embedding_batch: int = 100,
20
+ dimension: int = 1536,
21
+ similarity_metric: str = "cosine", # "cosine", "dotproduct", "euclidean"
22
+ cloud: Optional[str] = "aws",
23
+ region: Optional[str] = "us-east-1",
24
+ api_key: Optional[str] = None,
25
+ deletion_protection: Optional[str] = "disabled", # "enabled" or "disabled"
26
+ namespace: Optional[str] = "default",
27
+ ingest_batch: int = 200,
28
+ ):
29
+ super().__init__(embedding_model, similarity_metric, embedding_batch)
30
+
31
+ self.index_name = index_name
32
+ self.namespace = namespace
33
+ self.ingest_batch = ingest_batch
34
+
35
+ self.client = Pinecone_client(api_key=api_key)
36
+
37
+ if similarity_metric == "ip":
38
+ similarity_metric = "dotproduct"
39
+ elif similarity_metric == "l2":
40
+ similarity_metric = "euclidean"
41
+
42
+ if not self.client.has_index(index_name):
43
+ self.client.create_index(
44
+ name=index_name,
45
+ dimension=dimension,
46
+ metric=similarity_metric,
47
+ spec=ServerlessSpec(
48
+ cloud=cloud,
49
+ region=region,
50
+ ),
51
+ deletion_protection=deletion_protection,
52
+ )
53
+ self.index = self.client.Index(index_name)
54
+
55
+ async def add(self, ids: List[str], texts: List[str]):
56
+ texts = self.truncated_inputs(texts)
57
+ text_embeddings: List[
58
+ List[float]
59
+ ] = await self.embedding.aget_text_embedding_batch(texts)
60
+ self.add_embedding(ids, text_embeddings)
61
+
62
+ def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
63
+ vector_tuples = list(zip(ids, embeddings))
64
+ batch_vectors = make_batch(vector_tuples, self.ingest_batch)
65
+
66
+ async_res = [
67
+ self.index.upsert(
68
+ vectors=batch_vector_tuples,
69
+ namespace=self.namespace,
70
+ async_req=True,
71
+ )
72
+ for batch_vector_tuples in batch_vectors
73
+ ]
74
+ # Wait for the async requests to finish
75
+ [async_result.result() for async_result in async_res]
76
+
77
+ async def fetch(self, ids: List[str]) -> List[List[float]]:
78
+ results = self.index.fetch(ids=ids, namespace=self.namespace)
79
+ id_vector_dict = {
80
+ str(key): val["values"] for key, val in results["vectors"].items()
81
+ }
82
+ result = [id_vector_dict[_id] for _id in ids]
83
+ return result
84
+
85
+ async def is_exist(self, ids: List[str]) -> List[bool]:
86
+ fetched_result = self.index.fetch(ids=ids, namespace=self.namespace)
87
+ existed_ids = list(map(str, fetched_result.get("vectors", {}).keys()))
88
+ return list(map(lambda x: x in existed_ids, ids))
89
+
90
+ async def query(
91
+ self, queries: List[str], top_k: int, **kwargs
92
+ ) -> Tuple[List[List[str]], List[List[float]]]:
93
+ queries = self.truncated_inputs(queries)
94
+ query_embeddings: List[
95
+ List[float]
96
+ ] = await self.embedding.aget_text_embedding_batch(queries)
97
+
98
+ ids, scores = [], []
99
+ for query_embedding in query_embeddings:
100
+ response = self.index.query(
101
+ vector=query_embedding,
102
+ top_k=top_k,
103
+ include_values=True,
104
+ namespace=self.namespace,
105
+ )
106
+
107
+ ids.append([o.id for o in response.matches])
108
+ scores.append([o.score for o in response.matches])
109
+
110
+ if self.similarity_metric in ["l2"]:
111
+ scores = apply_recursive(lambda x: -x, scores)
112
+
113
+ return ids, scores
114
+
115
+ async def delete(self, ids: List[str]):
116
+ # Delete entries by IDs
117
+ self.index.delete(ids=ids, namespace=self.namespace)
118
+
119
+ def delete_index(self):
120
+ # Delete the index
121
+ self.client.delete_index(self.index_name)
@@ -0,0 +1,155 @@
1
+ import logging
2
+
3
+ from qdrant_client import QdrantClient
4
+ from qdrant_client.models import (
5
+ Distance,
6
+ VectorParams,
7
+ PointStruct,
8
+ PointIdsList,
9
+ HasIdCondition,
10
+ Filter,
11
+ SearchRequest,
12
+ )
13
+
14
+ from typing import List, Tuple, Union
15
+
16
+ from autorag.vectordb import BaseVectorStore
17
+
18
+ logger = logging.getLogger("AutoRAG")
19
+
20
+
21
+ class Qdrant(BaseVectorStore):
22
+ def __init__(
23
+ self,
24
+ embedding_model: Union[str, List[dict]],
25
+ collection_name: str,
26
+ embedding_batch: int = 100,
27
+ similarity_metric: str = "cosine",
28
+ client_type: str = "docker",
29
+ url: str = "http://localhost:6333",
30
+ host: str = "",
31
+ api_key: str = "",
32
+ dimension: int = 1536,
33
+ ingest_batch: int = 64,
34
+ parallel: int = 1,
35
+ max_retries: int = 3,
36
+ ):
37
+ super().__init__(embedding_model, similarity_metric, embedding_batch)
38
+
39
+ self.collection_name = collection_name
40
+ self.ingest_batch = ingest_batch
41
+ self.parallel = parallel
42
+ self.max_retries = max_retries
43
+
44
+ if similarity_metric == "cosine":
45
+ distance = Distance.COSINE
46
+ elif similarity_metric == "ip":
47
+ distance = Distance.DOT
48
+ elif similarity_metric == "l2":
49
+ distance = Distance.EUCLID
50
+ else:
51
+ raise ValueError(
52
+ f"similarity_metric {similarity_metric} is not supported\n"
53
+ "supported similarity metrics are: cosine, ip, l2"
54
+ )
55
+
56
+ if client_type == "docker":
57
+ self.client = QdrantClient(
58
+ url=url,
59
+ )
60
+ elif client_type == "cloud":
61
+ self.client = QdrantClient(
62
+ host=host,
63
+ api_key=api_key,
64
+ )
65
+ else:
66
+ raise ValueError(
67
+ f"client_type {client_type} is not supported\n"
68
+ "supported client types are: docker, cloud"
69
+ )
70
+
71
+ if not self.client.collection_exists(collection_name):
72
+ self.client.create_collection(
73
+ collection_name,
74
+ vectors_config=VectorParams(
75
+ size=dimension,
76
+ distance=distance,
77
+ ),
78
+ )
79
+ self.collection = self.client.get_collection(collection_name)
80
+
81
+ async def add(self, ids: List[str], texts: List[str]):
82
+ texts = self.truncated_inputs(texts)
83
+ text_embeddings = await self.embedding.aget_text_embedding_batch(texts)
84
+ self.add_embedding(ids, text_embeddings)
85
+
86
+ def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
87
+ points = list(
88
+ map(lambda x: PointStruct(id=x[0], vector=x[1]), zip(ids, embeddings))
89
+ )
90
+
91
+ self.client.upload_points(
92
+ collection_name=self.collection_name,
93
+ points=points,
94
+ batch_size=self.ingest_batch,
95
+ parallel=self.parallel,
96
+ max_retries=self.max_retries,
97
+ wait=True,
98
+ )
99
+
100
+ async def fetch(self, ids: List[str]) -> List[List[float]]:
101
+ # Fetch vectors by IDs
102
+ fetched_results = self.client.retrieve(
103
+ collection_name=self.collection_name,
104
+ ids=ids,
105
+ with_vectors=True,
106
+ )
107
+ return list(map(lambda x: x.vector, fetched_results))
108
+
109
+ async def is_exist(self, ids: List[str]) -> List[bool]:
110
+ existed_result = self.client.scroll(
111
+ collection_name=self.collection_name,
112
+ scroll_filter=Filter(
113
+ must=[
114
+ HasIdCondition(has_id=ids),
115
+ ],
116
+ ),
117
+ )
118
+ # existed_result is tuple. So we use existed_result[0] to get list of Record
119
+ existed_ids = list(map(lambda x: x.id, existed_result[0]))
120
+ return list(map(lambda x: x in existed_ids, ids))
121
+
122
+ async def query(
123
+ self, queries: List[str], top_k: int, **kwargs
124
+ ) -> Tuple[List[List[str]], List[List[float]]]:
125
+ queries = self.truncated_inputs(queries)
126
+ query_embeddings: List[
127
+ List[float]
128
+ ] = await self.embedding.aget_text_embedding_batch(queries)
129
+
130
+ search_queries = list(
131
+ map(
132
+ lambda x: SearchRequest(vector=x, limit=top_k, with_vector=True),
133
+ query_embeddings,
134
+ )
135
+ )
136
+
137
+ search_result = self.client.search_batch(
138
+ collection_name=self.collection_name, requests=search_queries
139
+ )
140
+
141
+ # Extract IDs and distances
142
+ ids = [[str(hit.id) for hit in result] for result in search_result]
143
+ scores = [[hit.score for hit in result] for result in search_result]
144
+
145
+ return ids, scores
146
+
147
+ async def delete(self, ids: List[str]):
148
+ self.client.delete(
149
+ collection_name=self.collection_name,
150
+ points_selector=PointIdsList(points=ids),
151
+ )
152
+
153
+ def delete_collection(self):
154
+ # Delete the collection
155
+ self.client.delete_collection(self.collection_name)
@@ -0,0 +1,184 @@
1
+ import logging
2
+
3
+ import weaviate
4
+ from weaviate.classes.init import Auth
5
+ from weaviate.classes.config import Property, DataType
6
+ import weaviate.classes as wvc
7
+ from weaviate.classes.query import MetadataQuery
8
+
9
+ from typing import List, Optional, Tuple, Union
10
+
11
+ from autorag.vectordb import BaseVectorStore
12
+
13
+ logger = logging.getLogger("AutoRAG")
14
+
15
+
16
+ class Weaviate(BaseVectorStore):
17
+ def __init__(
18
+ self,
19
+ embedding_model: Union[str, List[dict]],
20
+ collection_name: str,
21
+ embedding_batch: int = 100,
22
+ similarity_metric: str = "cosine",
23
+ client_type: str = "docker",
24
+ host: str = "localhost",
25
+ port: int = 8080,
26
+ grpc_port: int = 50051,
27
+ url: Optional[str] = None,
28
+ api_key: Optional[str] = None,
29
+ text_key: str = "content",
30
+ ):
31
+ super().__init__(embedding_model, similarity_metric, embedding_batch)
32
+
33
+ self.text_key = text_key
34
+
35
+ if client_type == "docker":
36
+ self.client = weaviate.connect_to_local(
37
+ host=host,
38
+ port=port,
39
+ grpc_port=grpc_port,
40
+ )
41
+ elif client_type == "cloud":
42
+ self.client = weaviate.connect_to_weaviate_cloud(
43
+ cluster_url=url,
44
+ auth_credentials=Auth.api_key(api_key),
45
+ )
46
+ else:
47
+ raise ValueError(
48
+ f"client_type {client_type} is not supported\n"
49
+ "supported client types are: docker, cloud"
50
+ )
51
+ if similarity_metric == "cosine":
52
+ distance_metric = wvc.config.VectorDistances.COSINE
53
+ elif similarity_metric == "ip":
54
+ distance_metric = wvc.config.VectorDistances.DOT
55
+ elif similarity_metric == "l2":
56
+ distance_metric = wvc.config.VectorDistances.L2_SQUARED
57
+ else:
58
+ raise ValueError(
59
+ f"similarity_metric {similarity_metric} is not supported\n"
60
+ "supported similarity metrics are: cosine, ip, l2"
61
+ )
62
+
63
+ if not self.client.collections.exists(collection_name):
64
+ self.client.collections.create(
65
+ collection_name,
66
+ properties=[
67
+ Property(
68
+ name="content", data_type=DataType.TEXT, skip_vectorization=True
69
+ ),
70
+ ],
71
+ vectorizer_config=wvc.config.Configure.Vectorizer.none(),
72
+ vector_index_config=wvc.config.Configure.VectorIndex.hnsw( # hnsw, flat, dynamic,
73
+ distance_metric=distance_metric
74
+ ),
75
+ )
76
+ self.collection = self.client.collections.get(collection_name)
77
+ self.collection_name = collection_name
78
+
79
+ async def add(self, ids: List[str], texts: List[str]):
80
+ texts = self.truncated_inputs(texts)
81
+ text_embeddings = await self.embedding.aget_text_embedding_batch(texts)
82
+
83
+ with self.client.batch.dynamic() as batch:
84
+ for i, text in enumerate(texts):
85
+ data_properties = {self.text_key: text}
86
+
87
+ batch.add_object(
88
+ collection=self.collection_name,
89
+ properties=data_properties,
90
+ uuid=ids[i],
91
+ vector=text_embeddings[i],
92
+ )
93
+
94
+ failed_objs = self.client.batch.failed_objects
95
+ for obj in failed_objs:
96
+ err_message = (
97
+ f"Failed to add object: {obj.original_uuid}\nReason: {obj.message}"
98
+ )
99
+
100
+ logger.error(err_message)
101
+
102
+ def add_embedding(self, ids: List[str], embeddings: List[List[float]]):
103
+ with self.client.batch.dynamic() as batch:
104
+ for i in range(len(ids)):
105
+ batch.add_object(
106
+ collection=self.collection_name,
107
+ uuid=ids[i],
108
+ vector=embeddings[i],
109
+ )
110
+
111
+ failed_objs = self.client.batch.failed_objects
112
+ for obj in failed_objs:
113
+ err_message = (
114
+ f"Failed to add object: {obj.original_uuid}\nReason: {obj.message}"
115
+ )
116
+
117
+ logger.error(err_message)
118
+
119
+ async def fetch(self, ids: List[str]) -> List[List[float]]:
120
+ # Fetch vectors by IDs
121
+ results = self.collection.query.fetch_objects(
122
+ filters=wvc.query.Filter.by_property("_id").contains_any(ids),
123
+ include_vector=True,
124
+ )
125
+ id_vector_dict = {
126
+ str(object.uuid): object.vector["default"] for object in results.objects
127
+ }
128
+ result = [id_vector_dict[_id] for _id in ids]
129
+ return result
130
+
131
+ async def is_exist(self, ids: List[str]) -> List[bool]:
132
+ fetched_result = self.collection.query.fetch_objects(
133
+ filters=wvc.query.Filter.by_property("_id").contains_any(ids),
134
+ )
135
+ existed_ids = [str(result.uuid) for result in fetched_result.objects]
136
+ return list(map(lambda x: x in existed_ids, ids))
137
+
138
+ async def query(
139
+ self, queries: List[str], top_k: int, **kwargs
140
+ ) -> Tuple[List[List[str]], List[List[float]]]:
141
+ queries = self.truncated_inputs(queries)
142
+ query_embeddings: List[
143
+ List[float]
144
+ ] = await self.embedding.aget_text_embedding_batch(queries)
145
+
146
+ ids, scores = [], []
147
+ for query_embedding in query_embeddings:
148
+ response = self.collection.query.near_vector(
149
+ near_vector=query_embedding,
150
+ limit=top_k,
151
+ return_metadata=MetadataQuery(distance=True),
152
+ )
153
+
154
+ ids.append([o.uuid for o in response.objects])
155
+ scores.append(
156
+ [
157
+ distance_to_score(o.metadata.distance, self.similarity_metric)
158
+ for o in response.objects
159
+ ]
160
+ )
161
+
162
+ return ids, scores
163
+
164
+ async def delete(self, ids: List[str]):
165
+ filter = wvc.query.Filter.by_id().contains_any(ids)
166
+ self.collection.data.delete_many(where=filter)
167
+
168
+ def delete_collection(self):
169
+ # Delete the collection
170
+ self.client.collections.delete(self.collection_name)
171
+
172
+
173
+ def distance_to_score(distance: float, similarity_metric) -> float:
174
+ if similarity_metric == "cosine":
175
+ return 1 - distance
176
+ elif similarity_metric == "ip":
177
+ return -distance
178
+ elif similarity_metric == "l2":
179
+ return -distance
180
+ else:
181
+ raise ValueError(
182
+ f"similarity_metric {similarity_metric} is not supported\n"
183
+ "supported similarity metrics are: cosine, ip, l2"
184
+ )
autorag/web.py ADDED
@@ -0,0 +1,81 @@
1
+ from typing import Optional
2
+
3
+ import click
4
+ import streamlit as st
5
+
6
+ from autorag.deploy import Runner
7
+
8
+
9
+ def get_runner(
10
+ yaml_path: Optional[str], project_dir: Optional[str], trial_path: Optional[str]
11
+ ):
12
+ if not yaml_path and not trial_path:
13
+ raise ValueError("yaml_path or trial_path must be given.")
14
+ elif yaml_path and trial_path:
15
+ raise ValueError("yaml_path and trial_path cannot be given at the same time.")
16
+ elif yaml_path:
17
+ return Runner.from_yaml(yaml_path, project_dir=project_dir)
18
+ elif trial_path:
19
+ return Runner.from_trial_folder(trial_path)
20
+
21
+
22
+ def set_initial_state():
23
+ if "messages" not in st.session_state:
24
+ st.session_state["messages"] = [
25
+ {
26
+ "role": "assistant",
27
+ "content": "Welcome !",
28
+ }
29
+ ]
30
+
31
+
32
+ def set_page_config():
33
+ st.set_page_config(
34
+ page_title="AutoRAG",
35
+ page_icon="🤖",
36
+ layout="wide",
37
+ initial_sidebar_state="expanded",
38
+ menu_items={
39
+ "Get help": "https://github.com/Marker-Inc-Korea/AutoRAG/discussions",
40
+ "Report a bug": "https://github.com/Marker-Inc-Korea/AutoRAG/issues",
41
+ },
42
+ )
43
+
44
+
45
+ def set_page_header():
46
+ st.header("📚 AutoRAG", anchor=False)
47
+ st.caption("Input a question and get an answer from the given documents. ")
48
+
49
+
50
+ def chat_box(runner: Runner):
51
+ if query := st.chat_input("How can I help?"):
52
+ # Add the user input to messages state
53
+ st.session_state["messages"].append({"role": "user", "content": query})
54
+ with st.chat_message("user"):
55
+ st.markdown(query)
56
+
57
+ # Generate llama-index stream with user input
58
+ with st.chat_message("assistant"):
59
+ with st.spinner("Processing..."):
60
+ response = st.write(runner.run(query))
61
+
62
+ # Add the final response to messages state
63
+ st.session_state["messages"].append({"role": "assistant", "content": response})
64
+
65
+
66
+ @click.command()
67
+ @click.option("--yaml_path", type=str, help="Path to the YAML file.")
68
+ @click.option("--project_dir", type=str, help="Path to the project directory.")
69
+ @click.option("--trial_path", type=str, help="Path to the trial directory.")
70
+ def run_web_server(
71
+ yaml_path: Optional[str], project_dir: Optional[str], trial_path: Optional[str]
72
+ ):
73
+ runner = get_runner(yaml_path, project_dir, trial_path)
74
+ set_initial_state()
75
+ set_page_config()
76
+ set_page_header()
77
+ chat_box(runner)
78
+
79
+
80
+ if __name__ == "__main__":
81
+ run_web_server()