Compare commits

...

49 Commits

Author SHA1 Message Date
Deshraj Yadav 2f285ea00a [Bug fix] Fix history sequence in prompt (#1254) 2024-02-11 16:07:36 -08:00
Dhravya Shah d38120c839 [Docs] Added documentation to deploy to Railway.app (#1250) 2024-02-11 15:58:42 -08:00
Michael d94aee812b [Improvements] Fixes to null data results and OpenAI embedding limits (#1238) 2024-02-11 15:45:02 -08:00
Rishiraj2594 68d650ec40 [Docs] Typo fixed youtube-video.mdx (#1253) 2024-02-09 16:34:08 -08:00
Rishiraj2594 769d926f5a [Docs] Typo fixed in youtube-channel.mdx (#1252) 2024-02-09 16:33:48 -08:00
Oskar 9478bab04e Fix links to the Discourse docs in the Discourse Loader (#1251) 2024-02-09 08:21:52 -08:00
Deshraj Yadav 7ad4af250f [Feature] Add support for optionally fetch all chat history for app (#1249) 2024-02-07 14:52:39 -08:00
Deshraj Yadav 9fa368b114 [Refactor] Remove usage of 'Pipeline' in favor of 'App' (#1246) 2024-02-06 19:00:33 -08:00
Deshraj Yadav 4afef04f26 [Feature] Add support for metadata filtering on search API (#1245) 2024-02-06 15:42:51 -08:00
Thomas T 8fe2c3effc [Bug Fix] Add support for AWS_REGION override (#1237) 2024-02-06 11:25:58 -08:00
Deshraj Yadav fa78c972be [Bug Fix] Fix issue related to using embedding model from huggingface (#1242) 2024-02-06 10:54:58 -08:00
Deshraj Yadav 0e66261644 Update docs (#1240) 2024-02-05 18:56:05 -08:00
Juanan Pereira 819650a254 Update URL Validation Regex to Support IP Addresses and Port Numbers (#1233) 2024-02-02 09:06:56 +05:30
Taranjeet Singh 34c41c87dc Docs: Update full stack docs (#1230) 2024-01-30 09:51:32 +05:30
Deshraj Yadav 2985b667b0 [Bug fix] Fix issue with gmail loader (#1228) 2024-01-29 18:36:02 +05:30
Taranjeet Singh 31bb0e7f0f Bump version to 0.1.71 (#1223) 2024-01-27 13:34:34 +05:30
Taranjeet Singh 8f28264aec feat: add UA header for pdf and sitemap (#1222) 2024-01-27 13:29:09 +05:30
Taranjeet Singh ec4fb11aa5 bump version to 0.1.70 (#1221) 2024-01-27 09:31:32 +05:30
Deven Patel b210723de1 [Improvement] add default user-agent header in webpage loader (#1219) 2024-01-26 11:04:25 +05:30
Deven Patel 433f99dd78 [Bugfix] fix typo in opensearch db (#1218) 2024-01-26 10:08:47 +05:30
Deven Patel e75c05112e [Improvement] update pinecone client v3 (#1200)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-26 09:08:37 +05:30
Taranjeet Singh d2a5b50ff8 add support for openai embedding models - text-em-3 (#1216) 2024-01-26 00:46:17 +05:30
Deven Patel 120690afd4 [Docs] Update mistral model in quickstart example (#1215) 2024-01-25 22:29:24 +05:30
Deven Patel 344dbeee42 [Bugfix] openai assistant (#1213)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-25 15:22:18 +05:30
Taranjeet Singh 3fe3b0320a Bump version to 0.1.69 (#1212) 2024-01-25 13:42:12 +05:30
Deven Patel 75896b647f [Docs] add docs for getting the list of added data sources (#1209)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-25 13:33:09 +05:30
Peter Jausovec 446d0975aa enable using custom Pinecone index name (#1172)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-25 13:30:10 +05:30
Deven Patel b7d365119c [Feature] add app.delete() method (#1187)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-23 14:27:30 +05:30
Deven Patel 2d9fbd4e49 [Bugfix] fix qdrant and weaviate db integration (#1181)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-23 14:24:29 +05:30
Deven Patel 22e14b5e65 [Bugfix] update zilliz db (#1186)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-23 14:23:57 +05:30
Deven Patel 1a654beea4 [Bugfix] fix pinecone db (#1185)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-23 14:23:30 +05:30
Deven Patel f50f8a444a [Bugfix] fix opensearch db (#1184)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-23 14:22:58 +05:30
Deven Patel 3cc3a0058d [bugfix] fix elasticsearch db (#1183)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-23 14:22:22 +05:30
Taranjeet Singh ae473b5e3c Bump version to 0.1.68 (#1206) 2024-01-23 14:19:11 +05:30
Deven Patel efb7e31565 [Docs] fix slack join link (#1205) 2024-01-22 20:54:56 -08:00
Deven Patel 069d265338 [Feature] Add support for AWS Bedrock LLM (#1189)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-21 14:09:08 +05:30
Taranjeet Singh 751a3a4bd1 Bump version to 0.1.67 (#1198) 2024-01-20 12:40:43 +05:30
Deven Patel cb0499407e [Feature] Add support for Mistral API (#1194)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-20 12:31:50 +05:30
Deven Patel 9afc6878c8 [Update] add test for passing vector dimension in embedder config (#1196)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-19 23:38:09 +05:30
Deven Patel 0b5b12575a [Bugfix] fix google ai embedding function (#1195)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-19 21:35:44 +05:30
aryankhanna475 d79d30bf0c Update Askabraham showcase (#1190) 2024-01-19 13:24:06 +05:30
Deven Patel 59600e2a5b [Improvement] add vector_dimension configuration in embedder config (#1192)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-19 10:31:41 +05:30
Deven Patel e572b5a3dc [Bugfix] fix import youtube allowed netlocks by defining them locally (#1191)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-19 09:35:46 +05:30
Juanan Pereira 5b46daaee4 Fix #1176 (a bug in the chromadb provider definition example) (#1177) 2024-01-18 02:28:55 +05:30
Deven Patel 2784bae772 [Tests] add tests for evaluation metrics (#1174)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-15 16:05:58 +05:30
Deshraj Yadav 325e11f0de Update docs (#1170) 2024-01-14 12:09:40 +05:30
Deven Patel 7444f59e3c [Bugfix] fix ec dev command for hf spaces (#1168)
Co-authored-by: Deven Patel <deven298@yahoo.com>
2024-01-14 08:36:23 +05:30
Deshraj Yadav affe319460 [Refactor] Change evaluation script path (#1165) 2024-01-12 21:29:59 +05:30
Deshraj Yadav 862ff6cca6 [Bug fix] Fix embedding issue for opensearch and some other vector databases (#1163) 2024-01-12 14:15:39 +05:30
104 changed files with 2293 additions and 727 deletions
+1 -4
View File
@@ -32,9 +32,6 @@
<hr />
> ### Checkout our latest [Sadhguru AI app](https://sadhguru-ai.streamlit.app/) built using Embedchain.
## What is Embedchain?
Embedchain is an Open Source RAG Framework that makes it easy to create and deploy AI apps. At its core, Embedchain follows the design principle of being *"Conventional but Configurable"* to serve both software engineers and machine learning engineers.
@@ -64,7 +61,7 @@ For example, you can create an Elon Musk bot using the following code:
```python
import os
from embedchain import Pipeline as App
from embedchain import App
# Create a bot instance
os.environ["OPENAI_API_KEY"] = "YOUR API KEY"
+1 -1
View File
@@ -2,7 +2,7 @@
<Card title="Talk to founders" icon="calendar" href="https://cal.com/taranjeetio/ec">
Schedule a call
</Card>
<Card title="Slack" icon="slack" href="https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw" color="#4A154B">
<Card title="Slack" icon="slack" href="https://embedchain.ai/slack" color="#4A154B">
Join our slack community
</Card>
<Card title="Discord" icon="discord" href="https://discord.gg/6PzXDgEjG5" color="#7289DA">
+1 -1
View File
@@ -4,7 +4,7 @@
<Card title="Google Form" icon="file" href="https://forms.gle/NDRCKsRpUHsz2Wcm8" color="#7387d0">
Fill out this form
</Card>
<Card title="Slack" icon="slack" href="https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw" color="#4A154B">
<Card title="Slack" icon="slack" href="https://embedchain.ai/slack" color="#4A154B">
Let us know on our slack community
</Card>
<Card title="Discord" icon="discord" href="https://discord.gg/6PzXDgEjG5" color="#7289DA">
+1 -1
View File
@@ -1,7 +1,7 @@
<p>If you can't find the specific LLM you need, no need to fret. We're continuously expanding our support for additional LLMs, and you can help us prioritize by opening an issue on our GitHub or simply reaching out to us on our Slack or Discord community.</p>
<CardGroup cols={2}>
<Card title="Slack" icon="slack" href="https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw" color="#4A154B">
<Card title="Slack" icon="slack" href="https://embedchain.ai/slack" color="#4A154B">
Let us know on our slack community
</Card>
<Card title="Discord" icon="discord" href="https://discord.gg/6PzXDgEjG5" color="#7289DA">
+1 -1
View File
@@ -3,7 +3,7 @@
<p>If you can't find the specific vector database, please feel free to request through one of the following channels and help us prioritize.</p>
<CardGroup cols={2}>
<Card title="Slack" icon="slack" href="https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw" color="#4A154B">
<Card title="Slack" icon="slack" href="https://embedchain.ai/slack" color="#4A154B">
Let us know on our slack community
</Card>
<Card title="Discord" icon="discord" href="https://discord.gg/6PzXDgEjG5" color="#7289DA">
@@ -8,7 +8,7 @@ You can configure different components of your app (`llm`, `embedding model`, or
<Tip>
Embedchain applications are configurable using YAML file, JSON file or by directly passing the config dictionary. Checkout the [docs here](/api-reference/pipeline/overview#usage) on how to use other formats.
Embedchain applications are configurable using YAML file, JSON file or by directly passing the config dictionary. Checkout the [docs here](/api-reference/app/overview#usage) on how to use other formats.
</Tip>
<CodeGroup>
@@ -200,9 +200,10 @@ Alright, let's dive into what each key means in the yaml config above:
- `stream` (Boolean): Controls if the response is streamed back to the user (set to false).
- `prompt` (String): A prompt for the model to follow when generating responses, requires `$context` and `$query` variables.
- `system_prompt` (String): A system prompt for the model to follow when generating responses, in this case, it's set to the style of William Shakespeare.
- `stream` (Boolean): Controls if the response is streamed back to the user (set to false).
- `stream` (Boolean): Controls if the response is streamed back to the user (set to false).
- `number_documents` (Integer): Number of documents to pull from the vectordb as context, defaults to 1
- `api_key` (String): The API key for the language model.
- `model_kwargs` (Dict): Keyword arguments to pass to the language model. Used for `aws_bedrock` provider, since it requires different arguments for each model.
3. `vectordb` Section:
- `provider` (String): The provider for the vector database, set to 'chroma'. You can find the full list of vector database providers in [our docs](/components/vector-databases).
- `config`:
@@ -214,7 +215,11 @@ Alright, let's dive into what each key means in the yaml config above:
- `provider` (String): The provider for the embedder, set to 'openai'. You can find the full list of embedding model providers in [our docs](/components/embedding-models).
- `config`:
- `model` (String): The specific model used for text embedding, 'text-embedding-ada-002'.
- `vector_dimension` (Integer): The vector dimension of the embedding model. [Defaults](https://github.com/embedchain/embedchain/blob/e572b5a3dc1b66f1e9b3357d11a88c63b5ce06e3/embedchain/models/vector_dimensions.py)
- `api_key` (String): The API key for the embedding model.
- `deployment_name` (String): The deployment name for the embedding model.
- `title` (String): The title for the embedding model for Google Embedder.
- `task_type` (String): The task type for the embedding model for Google Embedder.
5. `chunker` Section:
- `chunk_size` (Integer): The size of each chunk of text that is sent to the language model.
- `chunk_overlap` (Integer): The amount of overlap between each chunk of text.
@@ -129,3 +129,18 @@ app.chat("What is the net worth of Bill Gates?", session_id="user2")
app.chat("What was my last question", session_id="user1")
# 'Your last question was "What is the net worth of Elon Musk?"'
```
### With custom context window
If you want to customize the context window that you want to use during chat (default context window is 3 document chunks), you can do using the following code snippet:
```python with custom chunks size
from embedchain import App
from embedchain.config import BaseLlmConfig
app = App()
app.add("https://www.forbes.com/profile/elon-musk")
query_config = BaseLlmConfig(number_documents=5)
app.chat("What is the net worth of Elon Musk?", config=query_config)
```
+48
View File
@@ -0,0 +1,48 @@
---
title: 🗑 delete
---
## Delete Document
`delete()` method allows you to delete a document previously added to the app.
### Usage
```python
from embedchain import App
app = App()
forbes_doc_id = app.add("https://www.forbes.com/profile/elon-musk")
wiki_doc_id = app.add("https://en.wikipedia.org/wiki/Elon_Musk")
app.delete(forbes_doc_id) # deletes the forbes document
```
<Note>
If you do not have the document id, you can use `app.db.get()` method to get the document and extract the `hash` key from `metadatas` dictionary object, which serves as the document id.
</Note>
## Delete Chat Session History
`delete_session_chat_history()` method allows you to delete all previous messages in a chat history.
### Usage
```python
from embedchain import App
app = App()
app.add("https://www.forbes.com/profile/elon-musk")
app.chat("What is the net worth of Elon Musk?")
app.delete_session_chat_history()
```
<Note>
`delete_session_chat_history(session_id="session_1")` method also accepts `session_id` optional param for deleting chat history of a specific session.
It assumes the default session if no `session_id` is provided.
</Note>
+5
View File
@@ -0,0 +1,5 @@
---
title: 🚀 deploy
---
The `deploy()` method is currently available on an invitation-only basis. To request access, please submit your information via the provided [Google Form](https://forms.gle/vigN11h7b4Ywat668). We will review your request and respond promptly.
+33
View File
@@ -0,0 +1,33 @@
---
title: 📄 get
---
## Get data sources
`get_data_sources()` returns a list of all the data sources added in the app.
### Usage
```python
from embedchain import App
app = App()
app.add("https://www.forbes.com/profile/elon-musk")
app.add("https://en.wikipedia.org/wiki/Elon_Musk")
data_sources = app.get_data_sources()
# [
# {
# 'data_type': 'web_page',
# 'data_value': 'https://en.wikipedia.org/wiki/Elon_Musk',
# 'metadata': 'null'
# },
# {
# 'data_type': 'web_page',
# 'data_value': 'https://www.forbes.com/profile/elon-musk',
# 'metadata': 'null'
# }
# ]
```
@@ -1,34 +1,34 @@
---
title: "Pipeline"
title: "App"
---
Create a RAG pipeline object on Embedchain. This is the main entrypoint for a developer to interact with Embedchain APIs. A pipeline configures the llm, vector database, embedding model, and retrieval strategy of your choice.
Create a RAG app object on Embedchain. This is the main entrypoint for a developer to interact with Embedchain APIs. An app configures the llm, vector database, embedding model, and retrieval strategy of your choice.
### Attributes
<ParamField path="local_id" type="str">
Pipeline ID
App ID
</ParamField>
<ParamField path="name" type="str" optional>
Name of the pipeline
Name of the app
</ParamField>
<ParamField path="config" type="BaseConfig">
Configuration of the pipeline
Configuration of the app
</ParamField>
<ParamField path="llm" type="BaseLlm">
Configured LLM for the RAG pipeline
Configured LLM for the RAG app
</ParamField>
<ParamField path="db" type="BaseVectorDB">
Configured vector database for the RAG pipeline
Configured vector database for the RAG app
</ParamField>
<ParamField path="embedding_model" type="BaseEmbedder">
Configured embedding model for the RAG pipeline
Configured embedding model for the RAG app
</ParamField>
<ParamField path="chunker" type="ChunkerConfig">
Chunker configuration
</ParamField>
<ParamField path="client" type="Client" optional>
Client object (used to deploy a pipeline to Embedchain platform)
Client object (used to deploy an app to Embedchain platform)
</ParamField>
<ParamField path="logger" type="logging.Logger">
Logger object
@@ -36,7 +36,7 @@ Create a RAG pipeline object on Embedchain. This is the main entrypoint for a de
## Usage
You can create an embedchain pipeline instance using the following methods:
You can create an app instance using the following methods:
### Default setting
@@ -127,4 +127,4 @@ app = App.from_config(config_path="config.json")
}
```
</CodeGroup>
</CodeGroup>
+111
View File
@@ -0,0 +1,111 @@
---
title: '🔍 search'
---
`.search()` enables you to uncover the most pertinent context by performing a semantic search across your data sources based on a given query. Refer to the function signature below:
### Parameters
<ParamField path="query" type="str">
Question
</ParamField>
<ParamField path="num_documents" type="int" optional>
Number of relevant documents to fetch. Defaults to `3`
</ParamField>
<ParamField path="where" type="dict" optional>
Key value pair for metadata filtering.
</ParamField>
<ParamField path="raw_filter" type="dict" optional>
Pass raw filter query based on your vector database.
Currently, `raw_filter` param is only supported for Pinecone vector database.
</ParamField>
### Returns
<ResponseField name="answer" type="dict">
Return list of dictionaries that contain the relevant chunk and their source information.
</ResponseField>
## Usage
### Basic
Refer to the following example on how to use the search api:
```python Code example
from embedchain import App
app = App()
app.add("https://www.forbes.com/profile/elon-musk")
context = app.search("What is the net worth of Elon?", num_documents=2)
print(context)
```
### Advanced
#### Metadata filtering using `where` params
Here is an advanced example of `search()` API with metadata filtering on pinecone database:
```python
import os
from embedchain import App
os.environ["PINECONE_API_KEY"] = "xxx"
config = {
"vectordb": {
"provider": "pinecone",
"config": {
"metric": "dotproduct",
"vector_dimension": 1536,
"index_name": "ec-test",
"serverless_config": {"cloud": "aws", "region": "us-west-2"},
},
}
}
app = App.from_config(config=config)
app.add("https://www.forbes.com/profile/bill-gates", metadata={"type": "forbes", "person": "gates"})
app.add("https://en.wikipedia.org/wiki/Bill_Gates", metadata={"type": "wiki", "person": "gates"})
results = app.search("What is the net worth of Bill Gates?", where={"person": "gates"})
print("Num of search results: ", len(results))
```
#### Metadata filtering using `raw_filter` params
Following is an example of metadata filtering by passing the raw filter query that pinecone vector database follows:
```python
import os
from embedchain import App
os.environ["PINECONE_API_KEY"] = "xxx"
config = {
"vectordb": {
"provider": "pinecone",
"config": {
"metric": "dotproduct",
"vector_dimension": 1536,
"index_name": "ec-test",
"serverless_config": {"cloud": "aws", "region": "us-west-2"},
},
}
}
app = App.from_config(config=config)
app.add("https://www.forbes.com/profile/bill-gates", metadata={"year": 2022, "person": "gates"})
app.add("https://en.wikipedia.org/wiki/Bill_Gates", metadata={"year": 2024, "person": "gates"})
print("Filter with person: gates and year > 2023")
raw_filter = {"$and": [{"person": "gates"}, {"year": {"$gt": 2023}}]}
results = app.search("What is the net worth of Bill Gates?", raw_filter=raw_filter)
print("Num of search results: ", len(results))
```
-19
View File
@@ -1,19 +0,0 @@
---
title: 🗑 delete
---
`delete_session_chat_history()` method allows you to delete all previous messages in a chat history.
## Usage
```python
from embedchain import App
app = App()
app.add("https://www.forbes.com/profile/elon-musk")
app.chat("What is the net worth of Elon Musk?")
app.delete_session_chat_history()
```
-31
View File
@@ -1,31 +0,0 @@
---
title: 🚀 deploy
---
Using the `deploy()` method, Embedchain allows developers to easily launch their LLM-powered applications on the [Embedchain Platform](https://app.embedchain.ai). This platform facilitates seamless access to your data's context via a free and user-friendly REST API. Once your pipeline is deployed, you can update your data sources at any time.
The `deploy()` method not only deploys your pipeline but also efficiently manages LLMs, vector databases, embedding models, and data syncing, enabling you to focus on querying, chatting, or searching without the hassle of infrastructure management.
## Usage
```python
from embedchain import App
# Initialize app
app = App()
# Add data source
app.add("https://www.forbes.com/profile/elon-musk")
# Deploy your pipeline to Embedchain Platform
app.deploy()
# 🔑 Enter your Embedchain API key. You can find the API key at https://app.embedchain.ai/settings/keys/
# ec-xxxxxx
# 🛠️ Creating pipeline on the platform...
# 🎉🎉🎉 Pipeline created successfully! View your pipeline: https://app.embedchain.ai/pipelines/xxxxx
# 🛠️ Adding data to your pipeline...
# ✅ Data of type: web_page, value: https://www.forbes.com/profile/elon-musk added successfully.
```
-57
View File
@@ -1,57 +0,0 @@
---
title: '🔍 search'
---
`.search()` enables you to uncover the most pertinent context by performing a semantic search across your data sources based on a given query. Refer to the function signature below:
### Parameters
<ParamField path="query" type="str">
Question
</ParamField>
<ParamField path="num_documents" type="int" optional>
Number of relevant documents to fetch. Defaults to `3`
</ParamField>
### Returns
<ResponseField name="answer" type="dict">
Return list of dictionaries that contain the relevant chunk and their source information.
</ResponseField>
## Usage
Refer to the following example on how to use the search api:
```python Code example
from embedchain import App
# Initialize app
app = App()
# Add data source
app.add("https://www.forbes.com/profile/elon-musk")
# Get relevant context using semantic search
context = app.search("What is the net worth of Elon?", num_documents=2)
print(context)
# Context:
# [
# {
# 'context': 'Elon Musk PROFILEElon MuskCEO, Tesla$221.9BReal Time Net Worth ...',
# 'metadata': {
# 'source': 'https://www.forbes.com/profile/elon-musk',
# 'document_id': 'some_document_id',
# 'score': 0.404,
# }
# },
# {
# 'context': 'company, which is now called X.Wealth HistoryHOVER TO REVEAL NET WORTH ...',
# 'metadata': {
# 'source': 'https://www.forbes.com/profile/elon-musk',
# 'document_id': 'some_document_id',
# 'score': 0.435,
# }
# }
# ]
```
+1 -1
View File
@@ -8,7 +8,7 @@ We believe in building a vibrant and supportive community around embedchain. The
<Card title="Twitter" icon="twitter" href="https://twitter.com/embedchain">
Follow us on Twitter
</Card>
<Card title="Slack" icon="slack" href="https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw" color="#4A154B">
<Card title="Slack" icon="slack" href="https://embedchain.ai/slack" color="#4A154B">
Join our slack community
</Card>
<Card title="Discord" icon="discord" href="https://discord.gg/6PzXDgEjG5" color="#7289DA">
+4 -3
View File
@@ -7,11 +7,12 @@ When we say "custom", we mean that you can customize the loader and chunker to y
```python
from embedchain import App
import your_loader
import your_chunker
from my_module import CustomLoader
from my_module import CustomChunker
app = App()
loader = your_loader()
chunker = your_chunker()
loader = CustomLoader()
chunker = CustomChunker()
app.add("source", data_type="custom", loader=loader, chunker=chunker)
```
+1 -1
View File
@@ -22,7 +22,7 @@ Following is an example of how to use the dropbox loader:
```python
import os
from embedchain import Pipeline as App
from embedchain import App
os.environ["DROPBOX_ACCESS_TOKEN"] = "sl.xxx"
os.environ["OPENAI_API_KEY"] = "sk-xxx"
@@ -19,7 +19,7 @@ The first time you use the loader, you will be prompted to enter your Google acc
```python
from embedchain import Pipeline as App
from embedchain import App
app = App()
+1 -8
View File
@@ -4,13 +4,6 @@ title: '📰 PDF'
You can load any pdf file from your local file system or through a URL.
## Setup
Install the following packages for loading youtube videos which help in transcription.
```bash
pip install pytube youtube-transcript-api
```
## Usage
### Load from a local file
@@ -29,7 +22,7 @@ app = App()
app.add('https://arxiv.org/pdf/1706.03762.pdf', data_type='pdf_file')
app.query("What is the paper 'attention is all you need' about?", citations=True)
# Answer: The paper "Attention Is All You Need" proposes a new network architecture called the Transformer, which is based solely on attention mechanisms. It suggests that complex recurrent or convolutional neural networks can be replaced with a simpler architecture that connects the encoder and decoder through attention. The paper discusses how this approach can improve sequence transduction models, such as neural machine translation.
# Contexts:
# Contexts:
# [
# (
# 'Provided proper attribution is ...',
@@ -2,15 +2,17 @@
title: '📽️ Youtube Channel'
---
To add all the videos from a youtube channel to your app, use the data_type as `youtube_channel`.
## Setup
<Note>
Make sure you have all the required packages installed before using this data type. You can install them by running the following command in your terminal.
```bash
pip install -u "embedchain[youtube]"
pip install -U "embedchain[youtube]"
```
</Note>
## Usage
To add all the videos from a youtube channel to your app, use the data_type as `youtube_channel`.
```python
from embedchain import App
@@ -2,6 +2,16 @@
title: '📺 Youtube Video'
---
## Setup
Make sure you have all the required packages installed before using this data type. You can install them by running the following command in your terminal.
```bash
pip install -U "embedchain[youtube]"
```
## Usage
To add any youtube video to your app, use the data_type as `youtube_video`. Eg:
```python
+21 -1
View File
@@ -40,7 +40,27 @@ app.query("What is OpenAI?")
embedder:
provider: openai
config:
model: 'text-embedding-ada-002'
model: 'text-embedding-3-small'
```
</CodeGroup>
* OpenAI announced two new embedding models: `text-embedding-3-small` and `text-embedding-3-large`. Embedchain supports both these models. Below you can find YAML config for both:
<CodeGroup>
```yaml text-embedding-3-small.yaml
embedder:
provider: openai
config:
model: 'text-embedding-3-small'
```
```yaml text-embedding-3-large.yaml
embedder:
provider: openai
config:
model: 'text-embedding-3-large'
```
</CodeGroup>
+17 -16
View File
@@ -84,7 +84,7 @@ Once you have created your dataset, you can run evaluation on the dataset by pic
For example, you can run evaluation on context relevancy metric using the following code:
```python
from embedchain.eval.metrics import ContextRelevance
from embedchain.evaluation.metrics import ContextRelevance
metric = ContextRelevance()
score = metric.evaluate(dataset)
print(score)
@@ -112,20 +112,21 @@ context_relevance_score = num_relevant_sentences_in_context / num_of_sentences_i
You can run the context relevancy evaluation with the following simple code:
```python
from embedchain.eval.metrics import ContextRelevance
from embedchain.evaluation.metrics import ContextRelevance
metric = ContextRelevance()
score = metric.evaluate(dataset) # 'dataset' is definted in the create dataset section
print(score)
# 0.27975528364849833
```
In the above example, we used sensible defaults for the evaluation. However, you can also configure the evaluation metric as per your needs using the `ContextRelevanceConfig` class.
Here is a more advanced example of how to pass a custom evaluation config for evaluating on context relevance metric:
```python
from embedchain.config.eval.base import ContextRelevanceConfig
from embedchain.eval.metrics import ContextRelevance
from embedchain.config.evaluation.base import ContextRelevanceConfig
from embedchain.evaluation.metrics import ContextRelevance
eval_config = ContextRelevanceConfig(model="gpt-4", api_key="sk-xxx", language="en")
metric = ContextRelevance(config=eval_config)
@@ -144,7 +145,7 @@ metric.evaluate(dataset)
The language of the dataset being evaluated. We need this to determine the understand the context provided in the dataset. Defaults to `en`.
</ParamField>
<ParamField path="prompt" type="str" optional>
The prompt to extract the relevant sentences from the context. Defaults to `CONTEXT_RELEVANCY_PROMPT`, which can be found at `embedchain.config.eval.base` path.
The prompt to extract the relevant sentences from the context. Defaults to `CONTEXT_RELEVANCY_PROMPT`, which can be found at `embedchain.config.evaluation.base` path.
</ParamField>
@@ -161,7 +162,7 @@ answer_relevancy_score = mean(cosine_similarity(generated_questions, original_qu
You can run the answer relevancy evaluation with the following simple code:
```python
from embedchain.eval.metrics import AnswerRelevance
from embedchain.evaluation.metrics import AnswerRelevance
metric = AnswerRelevance()
score = metric.evaluate(dataset)
@@ -172,8 +173,8 @@ print(score)
In the above example, we used sensible defaults for the evaluation. However, you can also configure the evaluation metric as per your needs using the `AnswerRelevanceConfig` class. Here is a more advanced example where you can provide your own evaluation config:
```python
from embedchain.config.eval.base import AnswerRelevanceConfig
from embedchain.eval.metrics import AnswerRelevance
from embedchain.config.evaluation.base import AnswerRelevanceConfig
from embedchain.evaluation.metrics import AnswerRelevance
eval_config = AnswerRelevanceConfig(
model='gpt-4',
@@ -200,7 +201,7 @@ score = metric.evaluate(dataset)
The number of questions to generate for each answer. We use the generated questions to compare the similarity with the original question to determine the score. Defaults to `1`.
</ParamField>
<ParamField path="prompt" type="str" optional>
The prompt to extract the `num_gen_questions` number of questions from the provided answer. Defaults to `ANSWER_RELEVANCY_PROMPT`, which can be found at `embedchain.config.eval.base` path.
The prompt to extract the `num_gen_questions` number of questions from the provided answer. Defaults to `ANSWER_RELEVANCY_PROMPT`, which can be found at `embedchain.config.evaluation.base` path.
</ParamField>
## Groundedness <a id="groundedness"></a>
@@ -214,7 +215,7 @@ groundedness_score = (sum of all verdicts) / (total # of claims)
You can run the groundedness evaluation with the following simple code:
```python
from embedchain.eval.metrics import Groundedness
from embedchain.evaluation.metrics import Groundedness
metric = Groundedness()
score = metric.evaluate(dataset) # dataset from above
print(score)
@@ -224,8 +225,8 @@ print(score)
In the above example, we used sensible defaults for the evaluation. However, you can also configure the evaluation metric as per your needs using the `GroundednessConfig` class. Here is a more advanced example where you can configure the evaluation config:
```python
from embedchain.config.eval.base import GroundednessConfig
from embedchain.eval.metrics import Groundedness
from embedchain.config.evaluation.base import GroundednessConfig
from embedchain.evaluation.metrics import Groundedness
eval_config = GroundednessConfig(model='gpt-4', api_key="sk-xxx")
metric = Groundedness(config=eval_config)
@@ -242,15 +243,15 @@ score = metric.evaluate(dataset)
The openai api key to use for the evaluation. Defaults to `None`. If not provided, we will use the `OPENAI_API_KEY` environment variable.
</ParamField>
<ParamField path="answer_claims_prompt" type="str" optional>
The prompt to extract the claims from the provided answer. Defaults to `GROUNDEDNESS_ANSWER_CLAIMS_PROMPT`, which can be found at `embedchain.config.eval.base` path.
The prompt to extract the claims from the provided answer. Defaults to `GROUNDEDNESS_ANSWER_CLAIMS_PROMPT`, which can be found at `embedchain.config.evaluation.base` path.
</ParamField>
<ParamField path="claims_inference_prompt" type="str" optional>
The prompt to get verdicts on the claims from the answer from the given context. Defaults to `GROUNDEDNESS_CLAIMS_INFERENCE_PROMPT`, which can be found at `embedchain.config.eval.base` path.
The prompt to get verdicts on the claims from the answer from the given context. Defaults to `GROUNDEDNESS_CLAIMS_INFERENCE_PROMPT`, which can be found at `embedchain.config.evaluation.base` path.
</ParamField>
## Custom <a id="custom_metric"></a>
You can also create your own evaluation metric by extending the `BaseMetric` class. You can find the source code for the existing metrics at `embedchain.eval.metrics` path.
You can also create your own evaluation metric by extending the `BaseMetric` class. You can find the source code for the existing metrics at `embedchain.evaluation.metrics` path.
<Note>
You must provide the `name` of your custom metric in the `__init__` method of your class. This name will be used to identify your metric in the evaluation report.
@@ -260,7 +261,7 @@ You must provide the `name` of your custom metric in the `__init__` method of yo
from typing import Optional
from embedchain.config.base_config import BaseConfig
from embedchain.eval.metrics import BaseMetric
from embedchain.evaluation.metrics import BaseMetric
from embedchain.utils.eval import EvalData
class MyCustomMetric(BaseMetric):
+84 -1
View File
@@ -20,6 +20,8 @@ Embedchain comes with built-in support for various popular large language models
<Card title="Hugging Face" href="#hugging-face"></Card>
<Card title="Llama2" href="#llama2"></Card>
<Card title="Vertex AI" href="#vertex-ai"></Card>
<Card title="Mistral AI" href="#mistral-ai"></Card>
<Card title="AWS Bedrock" href="#aws-bedrock"></Card>
</CardGroup>
## OpenAI
@@ -250,7 +252,7 @@ app = App.from_config(config_path="config.yaml")
llm:
provider: azure_openai
config:
model: gpt-35-turbo
model: gpt-3.5-turbo
deployment_name: your_llm_deployment_name
temperature: 0.5
max_tokens: 1000
@@ -620,5 +622,86 @@ llm:
```
</CodeGroup>
## Mistral AI
Obtain the Mistral AI api key from their [console](https://console.mistral.ai/).
<CodeGroup>
```python main.py
os.environ["MISTRAL_API_KEY"] = "xxx"
app = App.from_config(config_path="config.yaml")
app.add("https://www.forbes.com/profile/elon-musk")
response = app.query("what is the net worth of Elon Musk?")
# As of January 16, 2024, Elon Musk's net worth is $225.4 billion.
response = app.chat("which companies does elon own?")
# Elon Musk owns Tesla, SpaceX, Boring Company, Twitter, and X.
response = app.chat("what question did I ask you already?")
# You have asked me several times already which companies Elon Musk owns, specifically Tesla, SpaceX, Boring Company, Twitter, and X.
```
```yaml config.yaml
llm:
provider: mistralai
config:
model: mistral-tiny
temperature: 0.5
max_tokens: 1000
top_p: 1
embedder:
provider: mistralai
config:
model: mistral-embed
```
</CodeGroup>
## AWS Bedrock
### Setup
- Before using the AWS Bedrock LLM, make sure you have the appropriate model access from [Bedrock Console](https://us-east-1.console.aws.amazon.com/bedrock/home?region=us-east-1#/modelaccess).
- You will also need to authenticate the `boto3` client by using a method in the [AWS documentation](https://boto3.amazonaws.com/v1/documentation/api/latest/guide/credentials.html#configuring-credentials)
- You can optionally export an `AWS_REGION`
### Usage
<CodeGroup>
```python main.py
import os
from embedchain import App
os.environ["AWS_ACCESS_KEY_ID"] = "xxx"
os.environ["AWS_SECRET_ACCESS_KEY"] = "xxx"
os.environ["AWS_REGION"] = "us-west-2"
app = App.from_config(config_path="config.yaml")
```
```yaml config.yaml
llm:
provider: aws_bedrock
config:
model: amazon.titan-text-express-v1
# check notes below for model_kwargs
model_kwargs:
temperature: 0.5
topP: 1
maxTokenCount: 1000
```
</CodeGroup>
<br />
<Note>
The model arguments are different for each providers. Please refer to the [AWS Bedrock Documentation](https://us-east-1.console.aws.amazon.com/bedrock/home?region=us-east-1#/providers) to find the appropriate arguments for your model.
</Note>
<br/ >
<Snippet file="missing-llm-tip.mdx" />
+30 -4
View File
@@ -167,7 +167,7 @@ Install pinecone related dependencies using the following command:
pip install --upgrade 'embedchain[pinecone]'
```
In order to use Pinecone as vector database, set the environment variables `PINECONE_API_KEY` and `PINECONE_ENV` which you can find on [Pinecone dashboard](https://app.pinecone.io/).
In order to use Pinecone as vector database, set the environment variable `PINECONE_API_KEY` which you can find on [Pinecone dashboard](https://app.pinecone.io/).
<CodeGroup>
@@ -175,20 +175,46 @@ In order to use Pinecone as vector database, set the environment variables `PINE
from embedchain import App
# load pinecone configuration from yaml file
app = App.from_config(config_path="config.yaml")
app = App.from_config(config_path="pod_config.yaml")
# or
app = App.from_config(config_path="serverless_config.yaml")
```
```yaml config.yaml
```yaml pod_config.yaml
vectordb:
provider: pinecone
config:
metric: cosine
vector_dimension: 1536
collection_name: my-pinecone-index
index_name: my-pinecone-index
pod_config:
environment: gcp-starter
metadata_config:
indexed:
- "url"
- "hash"
```
```yaml serverless_config.yaml
vectordb:
provider: pinecone
config:
metric: cosine
vector_dimension: 1536
index_name: my-pinecone-index
serverless_config:
cloud: aws
region: us-west-2
```
</CodeGroup>
<br />
<Note>
You can find more information about Pinecone configuration [here](https://docs.pinecone.io/docs/manage-indexes#create-a-pod-based-index).
You can also optionally provide `index_name` as a config param in yaml file to specify the index name. If not provided, the index name will be `{collection_name}-{vector_dimension}`.
</Note>
## Qdrant
In order to use Qdrant as a vector database, set the environment variables `QDRANT_URL` and `QDRANT_API_KEY` which you can find on [Qdrant Dashboard](https://cloud.qdrant.io/).
+1 -22
View File
@@ -7,29 +7,8 @@ description: 'Deploy your RAG application to embedchain.ai platform'
Embedchain enables developers to deploy their LLM-powered apps in production using the [Embedchain platform](https://app.embedchain.ai). The platform offers free access to context on your data through its REST API. Once the pipeline is deployed, you can update your data sources anytime after deployment.
See the example below on how to use the deploy your app (for free):
Deployment to Embedchain Platform is currently available on an invitation-only basis. To request access, please submit your information via the provided [Google Form](https://forms.gle/vigN11h7b4Ywat668). We will review your request and respond promptly.
```python
from embedchain import App
# Initialize app
app = App()
# Add data source
app.add("https://www.forbes.com/profile/elon-musk")
# Deploy your pipeline to Embedchain Platform
app.deploy()
# 🔑 Enter your Embedchain API key. You can find the API key at https://app.embedchain.ai/settings/keys/
# ec-xxxxxx
# 🛠️ Creating pipeline on the platform...
# 🎉🎉🎉 Pipeline created successfully! View your pipeline: https://app.embedchain.ai/pipelines/xxxxx
# 🛠️ Adding data to your pipeline...
# ✅ Data of type: web_page, value: https://www.forbes.com/profile/elon-musk added successfully.
```
## Seeking help?
+86
View File
@@ -0,0 +1,86 @@
---
title: 'Railway.app'
description: 'Deploy your RAG application to railway.app'
---
It's easy to host your Embedchain-powered apps and APIs on railway.
Follow the instructions given below to deploy your first application quickly:
## Step-1: Create RAG app
```bash Install embedchain
pip install embedchain
```
<Tip>
**Create a full stack app using Embedchain CLI**
To use your hosted embedchain RAG app, you can easily set up a FastAPI server that can be used anywhere.
To easily set up a FastAPI server, check out [Get started with Full stack](https://docs.embedchain.ai/get-started/full-stack) page.
Hosting this server on railway is super easy!
</Tip>
## Step-2: Set up your project
### With Docker
You can create a `Dockerfile` in the root of the project, with all the instructions. However, this method is sometimes slower in deployment.
### Without Docker
By default, Railway uses Python 3.7. Embedchain requires the python version to be >3.9 in order to install.
To fix this, create a `.python-version` file in the root directory of your project and specify the correct version
```bash .python-version
3.10
```
You also need to create a `requirements.txt` file to specify the requirements.
```bash requirements.txt
python-dotenv
embedchain
fastapi==0.108.0
uvicorn==0.25.0
embedchain
beautifulsoup4
sentence-transformers
```
## Step-3: Deploy to Railway 🚀
1. Go to https://railway.app and create an account.
2. Create a project by clicking on the "Start a new project" button
### With Github
Select `Empty Project` or `Deploy from Github Repo`.
You should be all set!
### Without Github
You can also use the railway CLI to deploy your apps from the terminal, if you don't want to connect a git repository.
To do this, just run this command in your terminal
```bash Install and set up railway CLI
npm i -g @railway/cli
railway login
railway link [projectID]
```
Finally, run `railway up` to deploy your app.
```bash Deploy
railway up
```
## Seeking help?
If you run into issues with deployment, please feel free to reach out to us via any of the following methods:
<Snippet file="get-help.mdx" />
+1 -1
View File
@@ -20,7 +20,7 @@ Embedchain community has been super active in creating demos on top of Embedchai
- [Create Instant ChatBot 🤖 using embedchain](https://databutton.com/v/h3e680h9) by Avra, ([Tweet](https://twitter.com/Avra_b/status/1674704745154641920/))
- [JOBO 🤖 — The AI-driven sidekick to craft your resume](https://try-jobo.com/) by Enrico Willemse, ([LinkedIn Post](https://www.linkedin.com/posts/enrico-willemse_jobai-gptfun-embedchain-activity-7090340080879374336-ueLB/))
- [Explore Your Knowledge Base: Interactive chats over various forms of documents](https://chatdocs.dkedar.com/) by Kedar Dabhadkar, ([LinkedIn Post](https://www.linkedin.com/posts/dkedar7_machinelearning-llmops-activity-7092524836639424513-2O3L/))
- [Chatbot trained on 1000+ videos of Ester hicks the co-author behind the famous book Secret](https://ask-abraham.thoughtseed.repl.co) by Mohan Kumar
- [Chatbot trained on 1000+ videos of Ester hicks the co-author behind the famous book Secret](https://askabraham.tokenofme.io/) by Mohan Kumar
## Templates
+1
View File
@@ -9,6 +9,7 @@ After successfully setting up and testing your RAG app locally, the next step is
<Card title="Fly.io" href="/deployment/fly_io"></Card>
<Card title="Modal.com" href="/deployment/modal_com"></Card>
<Card title="Render.com" href="/deployment/render_com"></Card>
<Card title="Railway.app" href="/deployment/railway"></Card>
<Card title="Streamlit.io" href="/deployment/streamlit_io"></Card>
<Card title="Gradio.app" href="/deployment/gradio_app"></Card>
<Card title="Huggingface.co" href="/deployment/huggingface_spaces"></Card>
+19
View File
@@ -8,6 +8,9 @@ Get started with full-stack RAG applications using Embedchain's easy-to-use CLI
Choose your setup method:
* [Without docker](#without-docker)
* [With Docker](#with-docker)
### Without Docker
Ensure these are installed:
@@ -21,6 +24,14 @@ Install Docker from [Docker's official website](https://docs.docker.com/engine/i
## Quick Start Guide
### Install the package
Before proceeding, make sure you have the Embedchain package installed.
```bash
pip install embedchain -U
```
### Setting Up
For the purpose of the demo, you have to set `OPENAI_API_KEY` to start with but you can choose any llm by changing the configuration easily.
@@ -60,3 +71,11 @@ Open http://localhost:3000 to view the chat UI.
Check out the Embedchain admin panel to see the document chunks for your RAG application.
![full stack chunks](/images/fullstack-chunks.png)
### API Server
If you want to access the API server, you can do so at http://localhost:8000/docs.
![API Server](/images/fullstack-api-server.png)
You can customize the UI and code as per your requirements.
+2 -2
View File
@@ -47,7 +47,7 @@ app.query("What is the net worth of Elon Musk today?")
llm:
provider: huggingface
config:
model: 'mistralai/Mistral-7B-v0.1'
model: 'mistralai/Mistral-7B-Instruct-v0.2'
top_p: 0.5
embedder:
provider: huggingface
@@ -80,4 +80,4 @@ Now that you have created your first app, you can follow any of the links:
* [Introduction](/get-started/introduction)
* [Customization](/components/introduction)
* [Use cases](/use-cases/introduction)
* [Deployment](/get-started/deployment)
* [Deployment](/get-started/deployment)
Binary file not shown.

After

Width:  |  Height:  |  Size: 262 KiB

+13 -11
View File
@@ -142,6 +142,7 @@
"deployment/fly_io",
"deployment/modal_com",
"deployment/render_com",
"deployment/railway",
"deployment/streamlit_io",
"deployment/gradio_app",
"deployment/huggingface_spaces",
@@ -199,18 +200,19 @@
{
"group": "API Reference",
"pages": [
"api-reference/pipeline/overview",
"api-reference/app/overview",
{
"group": "Pipeline methods",
"group": "App methods",
"pages": [
"api-reference/pipeline/add",
"api-reference/pipeline/query",
"api-reference/pipeline/chat",
"api-reference/pipeline/search",
"api-reference/pipeline/deploy",
"api-reference/pipeline/reset",
"api-reference/pipeline/delete",
"api-reference/pipeline/evaluate"
"api-reference/app/add",
"api-reference/app/query",
"api-reference/app/chat",
"api-reference/app/search",
"api-reference/app/get",
"api-reference/app/evaluate",
"api-reference/app/deploy",
"api-reference/app/reset",
"api-reference/app/delete"
]
},
"api-reference/store/openai-assistant",
@@ -238,7 +240,7 @@
"footerSocials": {
"website": "https://embedchain.ai",
"github": "https://github.com/embedchain/embedchain",
"slack": "https://join.slack.com/t/embedchain/shared_invite/zt-22uwz3c46-Zg7cIh5rOBteT_xe1jwLDw",
"slack": "https://embedchain.ai/slack",
"discord": "https://discord.gg/6PzXDgEjG5",
"twitter": "https://twitter.com/embedchain",
"linkedin": "https://www.linkedin.com/company/embedchain"
-3
View File
@@ -1,3 +0,0 @@
---
title: 'FAQs'
---
-3
View File
@@ -1,3 +0,0 @@
---
title: 'Overview'
---
-3
View File
@@ -1,3 +0,0 @@
---
title: 'Quickstart'
---
-3
View File
@@ -1,3 +0,0 @@
---
title: 'Roadmap'
---
-3
View File
@@ -1,3 +0,0 @@
---
title: 'Security'
---
+4 -28
View File
@@ -20,15 +20,15 @@ from embedchain.constants import SQLITE_PATH
from embedchain.embedchain import EmbedChain
from embedchain.embedder.base import BaseEmbedder
from embedchain.embedder.openai import OpenAIEmbedder
from embedchain.eval.base import BaseMetric
from embedchain.eval.metrics import (AnswerRelevance, ContextRelevance,
Groundedness)
from embedchain.evaluation.base import BaseMetric
from embedchain.evaluation.metrics import (AnswerRelevance, ContextRelevance,
Groundedness)
from embedchain.factory import EmbedderFactory, LlmFactory, VectorDBFactory
from embedchain.helpers.json_serializable import register_deserializable
from embedchain.llm.base import BaseLlm
from embedchain.llm.openai import OpenAILlm
from embedchain.telemetry.posthog import AnonymousTelemetry
from embedchain.utils.eval import EvalData, EvalMetric
from embedchain.utils.evaluation import EvalData, EvalMetric
from embedchain.utils.misc import validate_config
from embedchain.vectordb.base import BaseVectorDB
from embedchain.vectordb.chroma import ChromaDB
@@ -250,30 +250,6 @@ class App(EmbedChain):
r.raise_for_status()
return r.json()
def search(self, query, num_documents=3):
"""
Search for similar documents related to the query in the vector database.
"""
# Send anonymous telemetry
self.telemetry.capture(event_name="search", properties=self._telemetry_props)
# TODO: Search will call the endpoint rather than fetching the data from the db itself when deploy=True.
if self.id is None:
where = {"app_id": self.local_id}
context = self.db.query(
query,
n_results=num_documents,
where=where,
citations=True,
)
result = []
for c in context:
result.append({"context": c[0], "metadata": c[1]})
return result
else:
# Make API call to the backend to get the results
NotImplementedError("Search is not implemented yet for the prod mode.")
def _upload_file_to_presigned_url(self, presigned_url, file_path):
try:
with open(file_path, "rb") as file:
+9 -6
View File
@@ -27,7 +27,7 @@ class BaseChunker(JSONSerializable):
chunk_ids = []
id_map = {}
min_chunk_size = config.min_chunk_size if config is not None else 1
logging.info(f"[INFO] Skipping chunks smaller than {min_chunk_size} characters")
logging.info(f"Skipping chunks smaller than {min_chunk_size} characters")
data_result = loader.load_data(src)
data_records = data_result["data"]
doc_id = data_result["doc_id"]
@@ -39,11 +39,14 @@ class BaseChunker(JSONSerializable):
for data in data_records:
content = data["content"]
meta_data = data["meta_data"]
metadata = data["meta_data"]
# add data type to meta data to allow query using data type
meta_data["data_type"] = self.data_type.value
meta_data["doc_id"] = doc_id
url = meta_data["url"]
metadata["data_type"] = self.data_type.value
metadata["doc_id"] = doc_id
# TODO: Currently defaulting to the src as the url. This is done intentianally since some
# of the data types like 'gmail' loader doesn't have the url in the meta data.
url = metadata.get("url", src)
chunks = self.get_chunks(content)
for chunk in chunks:
@@ -53,7 +56,7 @@ class BaseChunker(JSONSerializable):
id_map[chunk_id] = True
chunk_ids.append(chunk_id)
documents.append(chunk)
metadatas.append(meta_data)
metadatas.append(metadata)
return {
"documents": documents,
"ids": chunk_ids,
+1 -1
View File
@@ -292,7 +292,7 @@ def dev(debug, host, port):
run_dev_modal_com()
elif template == "render.com":
run_dev_render_com(debug, host, port)
elif template == "streamlit.io" or template == "hf/streamlit.app":
elif template == "streamlit.io" or template == "hf/streamlit.io":
run_dev_streamlit_io()
elif template == "gradio.app" or template == "hf/gradio.app":
run_dev_gradio()
+6 -1
View File
@@ -6,7 +6,11 @@ from embedchain.helpers.json_serializable import register_deserializable
@register_deserializable
class BaseEmbedderConfig:
def __init__(
self, model: Optional[str] = None, deployment_name: Optional[str] = None, api_key: Optional[str] = None
self,
model: Optional[str] = None,
deployment_name: Optional[str] = None,
vector_dimension: Optional[int] = None,
api_key: Optional[str] = None,
):
"""
Initialize a new instance of an embedder config class.
@@ -18,4 +22,5 @@ class BaseEmbedderConfig:
"""
self.model = model
self.deployment_name = deployment_name
self.vector_dimension = vector_dimension
self.api_key = api_key
+2 -1
View File
@@ -24,7 +24,8 @@ DEFAULT_PROMPT_WITH_HISTORY = """
$context
History: $history
History:
$history
Query: $query
+19 -3
View File
@@ -1,3 +1,4 @@
import os
from typing import Optional
from embedchain.config.vectordb.base import BaseVectorDbConfig
@@ -8,13 +9,28 @@ from embedchain.helpers.json_serializable import register_deserializable
class PineconeDBConfig(BaseVectorDbConfig):
def __init__(
self,
collection_name: Optional[str] = None,
dir: Optional[str] = None,
index_name: Optional[str] = None,
api_key: Optional[str] = None,
vector_dimension: int = 1536,
metric: Optional[str] = "cosine",
pod_config: Optional[dict[str, any]] = None,
serverless_config: Optional[dict[str, any]] = None,
**extra_params: dict[str, any],
):
self.metric = metric
self.api_key = api_key
self.index_name = index_name
self.vector_dimension = vector_dimension
self.extra_params = extra_params
super().__init__(collection_name=collection_name, dir=dir)
if pod_config is None and serverless_config is None:
# If no config is provided, use the default pod spec config
pod_environment = os.environ.get("PINECONE_ENV", "gcp-starter")
self.pod_config = {"environment": pod_environment, "metadata_config": {"indexed": ["*"]}}
else:
self.pod_config = pod_config
self.serverless_config = serverless_config
if self.pod_config and self.serverless_config:
raise ValueError("Only one of pod_config or serverless_config can be provided.")
super().__init__(collection_name=self.index_name, dir=None)
+2 -2
View File
@@ -2,12 +2,12 @@ from dotenv import load_dotenv
from fastapi import FastAPI, responses
from pydantic import BaseModel
from embedchain import Pipeline
from embedchain import App
load_dotenv(".env")
app = FastAPI(title="Embedchain FastAPI App")
embedchain_app = Pipeline()
embedchain_app = App()
class SourceModel(BaseModel):
+2 -2
View File
@@ -2,7 +2,7 @@ from dotenv import load_dotenv
from fastapi import Body, FastAPI, responses
from modal import Image, Secret, Stub, asgi_app
from embedchain import Pipeline
from embedchain import App
load_dotenv(".env")
@@ -18,7 +18,7 @@ stub = Stub(
)
web_app = FastAPI()
embedchain_app = Pipeline(name="embedchain-modal-app")
embedchain_app = App(name="embedchain-modal-app")
@web_app.post("/add")
+2 -2
View File
@@ -1,10 +1,10 @@
from fastapi import FastAPI, responses
from pydantic import BaseModel
from embedchain import Pipeline
from embedchain import App
app = FastAPI(title="Embedchain FastAPI App")
embedchain_app = Pipeline()
embedchain_app = App()
class SourceModel(BaseModel):
+86 -12
View File
@@ -369,7 +369,7 @@ class EmbedChain(JSONSerializable):
metadatas = embeddings_data["metadatas"]
ids = embeddings_data["ids"]
new_doc_id = embeddings_data["doc_id"]
embeddings = embeddings_data.get("embeddings")
if existing_doc_id and existing_doc_id == new_doc_id:
print("Doc content has not changed. Skipping creating chunks and embeddings")
return [], [], [], 0
@@ -429,22 +429,36 @@ class EmbedChain(JSONSerializable):
if dry_run:
return list(documents), metadatas, ids, 0
# Count before, to calculate a delta in the end.
chunks_before_addition = self.db.count()
self.db.add(
embeddings=embeddings,
documents=documents,
metadatas=metadatas,
ids=ids,
**kwargs,
)
count_new_chunks = self.db.count() - chunks_before_addition
# Filter out empty documents and ensure they meet the API requirements
valid_documents = [doc for doc in documents if doc and isinstance(doc, str)]
documents = valid_documents
# Chunk documents into batches of 2048 and handle each batch
# helps wigth large loads of embeddings that hit OpenAI limits
document_batches = [documents[i:i+2048] for i in range(0, len(documents), 2048)]
for batch in document_batches:
try:
# Add only valid batches
if batch:
self.db.add(documents=batch, metadatas=metadatas, ids=ids, **kwargs)
except Exception as e:
print(f"Failed to add batch due to a bad request: {e}")
# Handle the error, e.g., by logging, retrying, or skipping
pass
count_new_chunks = self.db.count() - chunks_before_addition
print(f"Successfully saved {src} ({chunker.data_type}). New chunks count: {count_new_chunks}")
return list(documents), metadatas, ids, count_new_chunks
@staticmethod
def _format_result(results):
return [
@@ -479,7 +493,9 @@ class EmbedChain(JSONSerializable):
:return: List of contents of the document that matched your query
:rtype: list[str]
"""
print("Query passed in config:", config)
query_config = config or self.llm.config
print("Final config:", query_config)
if where is not None:
where = where
else:
@@ -490,6 +506,7 @@ class EmbedChain(JSONSerializable):
if self.config.id is not None:
where.update({"app_id": self.config.id})
print('Number documents', query_config)
contexts = self.db.query(
input_query=input_query,
n_results=query_config.number_documents,
@@ -640,6 +657,41 @@ class EmbedChain(JSONSerializable):
else:
return answer
def search(self, query, num_documents=3, where=None, raw_filter=None):
"""
Search for similar documents related to the query in the vector database.
Args:
query (str): The query to use.
num_documents (int, optional): Number of similar documents to fetch. Defaults to 3.
where (dict[str, any], optional): Filter criteria for the search.
raw_filter (dict[str, any], optional): Advanced raw filter criteria for the search.
Raises:
ValueError: If both `raw_filter` and `where` are used simultaneously.
Returns:
list[dict]: A list of dictionaries, each containing the 'context' and 'metadata' of a document.
"""
# Send anonymous telemetry
self.telemetry.capture(event_name="search", properties=self._telemetry_props)
if raw_filter and where:
raise ValueError("You can't use both `raw_filter` and `where` together.")
filter_type = "raw_filter" if raw_filter else "where"
filter_criteria = raw_filter if raw_filter else where
params = {
"input_query": query,
"n_results": num_documents,
"citations": True,
"app_id": self.config.id,
filter_type: filter_criteria,
}
return [{"context": c[0], "metadata": c[1]} for c in self.db.query(**params)]
def set_collection_name(self, name: str):
"""
Set the name of the collection. A collection is an isolated space for vectors.
@@ -667,9 +719,19 @@ class EmbedChain(JSONSerializable):
# Send anonymous telemetry
self.telemetry.capture(event_name="reset", properties=self._telemetry_props)
def get_history(self, num_rounds: int = 10, display_format: bool = True, session_id: Optional[str] = "default"):
def get_history(
self,
num_rounds: int = 10,
display_format: bool = True,
session_id: Optional[str] = "default",
fetch_all: bool = False,
):
history = self.llm.memory.get(
app_id=self.config.id, session_id=session_id, num_rounds=num_rounds, display_format=display_format
app_id=self.config.id,
session_id=session_id,
num_rounds=num_rounds,
display_format=display_format,
fetch_all=fetch_all,
)
return history
@@ -680,3 +742,15 @@ class EmbedChain(JSONSerializable):
def delete_all_chat_history(self, app_id: str):
self.llm.memory.delete(app_id=app_id)
self.llm.update_history(app_id=app_id)
def delete(self, source_id: str):
"""
Deletes the data from the database.
:param source_hash: The hash of the source.
:type source_hash: str
"""
self.db.delete(where={"hash": source_id})
logging.info(f"Successfully deleted {source_id}")
# Send anonymous telemetry
if self.config.collect_metrics:
self.telemetry.capture(event_name="delete", properties=self._telemetry_props)
+12 -5
View File
@@ -1,4 +1,4 @@
from typing import Optional
from typing import Optional, Union
import google.generativeai as genai
from chromadb import EmbeddingFunction, Embeddings
@@ -13,12 +13,19 @@ class GoogleAIEmbeddingFunction(EmbeddingFunction):
super().__init__()
self.config = config or GoogleAIEmbedderConfig()
def __call__(self, input_: str) -> Embeddings:
def __call__(self, input: Union[list[str], str]) -> Embeddings:
model = self.config.model
title = self.config.title
task_type = self.config.task_type
embeddings = genai.embed_content(model=model, content=input_, task_type=task_type, title=title)
return embeddings["embedding"]
if isinstance(input, str):
input_ = [input]
else:
input_ = input
data = genai.embed_content(model=model, content=input_, task_type=task_type, title=title)
embeddings = data["embedding"]
if isinstance(input_, str):
embeddings = [embeddings]
return embeddings
class GoogleAIEmbedder(BaseEmbedder):
@@ -27,5 +34,5 @@ class GoogleAIEmbedder(BaseEmbedder):
embedding_fn = GoogleAIEmbeddingFunction(config=config)
self.set_embedding_fn(embedding_fn=embedding_fn)
vector_dimension = VectorDimensions.GOOGLE_AI.value
vector_dimension = self.config.vector_dimension or VectorDimensions.GOOGLE_AI.value
self.set_vector_dimension(vector_dimension=vector_dimension)
+1 -1
View File
@@ -16,5 +16,5 @@ class GPT4AllEmbedder(BaseEmbedder):
embedding_fn = BaseEmbedder._langchain_default_concept(embeddings)
self.set_embedding_fn(embedding_fn=embedding_fn)
vector_dimension = VectorDimensions.GPT4ALL.value
vector_dimension = self.config.vector_dimension or VectorDimensions.GPT4ALL.value
self.set_vector_dimension(vector_dimension=vector_dimension)
+1 -1
View File
@@ -15,5 +15,5 @@ class HuggingFaceEmbedder(BaseEmbedder):
embedding_fn = BaseEmbedder._langchain_default_concept(embeddings)
self.set_embedding_fn(embedding_fn=embedding_fn)
vector_dimension = VectorDimensions.HUGGING_FACE.value
vector_dimension = self.config.vector_dimension or VectorDimensions.HUGGING_FACE.value
self.set_vector_dimension(vector_dimension=vector_dimension)
+46
View File
@@ -0,0 +1,46 @@
import os
from typing import Optional, Union
from chromadb import EmbeddingFunction, Embeddings
from embedchain.config import BaseEmbedderConfig
from embedchain.embedder.base import BaseEmbedder
from embedchain.models import VectorDimensions
class MistralAIEmbeddingFunction(EmbeddingFunction):
def __init__(self, config: BaseEmbedderConfig) -> None:
super().__init__()
try:
from langchain_mistralai import MistralAIEmbeddings
except ModuleNotFoundError:
raise ModuleNotFoundError(
"The required dependencies for MistralAI are not installed."
'Please install with `pip install --upgrade "embedchain[mistralai]"`'
) from None
self.config = config
api_key = self.config.api_key or os.getenv("MISTRAL_API_KEY")
self.client = MistralAIEmbeddings(mistral_api_key=api_key)
self.client.model = self.config.model
def __call__(self, input: Union[list[str], str]) -> Embeddings:
if isinstance(input, str):
input_ = [input]
else:
input_ = input
response = self.client.embed_documents(input_)
return response
class MistralAIEmbedder(BaseEmbedder):
def __init__(self, config: Optional[BaseEmbedderConfig] = None):
super().__init__(config)
if self.config.model is None:
self.config.model = "mistral-embed"
embedding_fn = MistralAIEmbeddingFunction(config=self.config)
self.set_embedding_fn(embedding_fn=embedding_fn)
vector_dimension = self.config.vector_dimension or VectorDimensions.MISTRAL_AI.value
self.set_vector_dimension(vector_dimension=vector_dimension)
+2 -1
View File
@@ -32,4 +32,5 @@ class OpenAIEmbedder(BaseEmbedder):
model_name=self.config.model,
)
self.set_embedding_fn(embedding_fn=embedding_fn)
self.set_vector_dimension(vector_dimension=VectorDimensions.OPENAI.value)
vector_dimension = self.config.vector_dimension or VectorDimensions.OPENAI.value
self.set_vector_dimension(vector_dimension=vector_dimension)
+1 -1
View File
@@ -15,5 +15,5 @@ class VertexAIEmbedder(BaseEmbedder):
embedding_fn = BaseEmbedder._langchain_default_concept(embeddings)
self.set_embedding_fn(embedding_fn=embedding_fn)
vector_dimension = VectorDimensions.VERTEX_AI.value
vector_dimension = self.config.vector_dimension or VectorDimensions.VERTEX_AI.value
self.set_vector_dimension(vector_dimension=vector_dimension)
@@ -1,6 +1,6 @@
from abc import ABC, abstractmethod
from embedchain.utils.eval import EvalData
from embedchain.utils.evaluation import EvalData
class BaseMetric(ABC):
@@ -8,9 +8,9 @@ import numpy as np
from openai import OpenAI
from tqdm import tqdm
from embedchain.config.eval.base import AnswerRelevanceConfig
from embedchain.eval.base import BaseMetric
from embedchain.utils.eval import EvalData, EvalMetric
from embedchain.config.evaluation.base import AnswerRelevanceConfig
from embedchain.evaluation.base import BaseMetric
from embedchain.utils.evaluation import EvalData, EvalMetric
class AnswerRelevance(BaseMetric):
@@ -8,9 +8,9 @@ import pysbd
from openai import OpenAI
from tqdm import tqdm
from embedchain.config.eval.base import ContextRelevanceConfig
from embedchain.eval.base import BaseMetric
from embedchain.utils.eval import EvalData, EvalMetric
from embedchain.config.evaluation.base import ContextRelevanceConfig
from embedchain.evaluation.base import BaseMetric
from embedchain.utils.evaluation import EvalData, EvalMetric
class ContextRelevance(BaseMetric):
@@ -8,9 +8,9 @@ import numpy as np
from openai import OpenAI
from tqdm import tqdm
from embedchain.config.eval.base import GroundednessConfig
from embedchain.eval.base import BaseMetric
from embedchain.utils.eval import EvalData, EvalMetric
from embedchain.config.evaluation.base import GroundednessConfig
from embedchain.evaluation.base import BaseMetric
from embedchain.utils.evaluation import EvalData, EvalMetric
class Groundedness(BaseMetric):
@@ -21,7 +21,7 @@ class Groundedness(BaseMetric):
def __init__(self, config: Optional[GroundednessConfig] = None):
super().__init__(name=EvalMetric.GROUNDEDNESS.value)
self.config = config or GroundednessConfig()
api_key = self.config.api_key or os.environ["OPENAI_API_KEY"]
api_key = self.config.api_key or os.getenv("OPENAI_API_KEY")
if not api_key:
raise ValueError("Please set the OPENAI_API_KEY environment variable or pass the `api_key` in config.")
self.client = OpenAI(api_key=api_key)
+4
View File
@@ -21,6 +21,8 @@ class LlmFactory:
"openai": "embedchain.llm.openai.OpenAILlm",
"vertexai": "embedchain.llm.vertex_ai.VertexAILlm",
"google": "embedchain.llm.google.GoogleLlm",
"aws_bedrock": "embedchain.llm.aws_bedrock.AWSBedrockLlm",
"mistralai": "embedchain.llm.mistralai.MistralAILlm",
}
provider_to_config_class = {
"embedchain": "embedchain.config.llm.base.BaseLlmConfig",
@@ -50,12 +52,14 @@ class EmbedderFactory:
"openai": "embedchain.embedder.openai.OpenAIEmbedder",
"vertexai": "embedchain.embedder.vertexai.VertexAIEmbedder",
"google": "embedchain.embedder.google.GoogleAIEmbedder",
"mistralai": "embedchain.embedder.mistralai.MistralAIEmbedder",
}
provider_to_config_class = {
"azure_openai": "embedchain.config.embedder.base.BaseEmbedderConfig",
"openai": "embedchain.config.embedder.base.BaseEmbedderConfig",
"gpt4all": "embedchain.config.embedder.base.BaseEmbedderConfig",
"google": "embedchain.config.embedder.google.GoogleAIEmbedderConfig",
"huggingface": "embedchain.config.embedder.base.BaseEmbedderConfig",
}
@classmethod
+49
View File
@@ -0,0 +1,49 @@
import os
from typing import Optional
from langchain.llms import Bedrock
from embedchain.config import BaseLlmConfig
from embedchain.helpers.json_serializable import register_deserializable
from embedchain.llm.base import BaseLlm
@register_deserializable
class AWSBedrockLlm(BaseLlm):
def __init__(self, config: Optional[BaseLlmConfig] = None):
super().__init__(config)
def get_llm_model_answer(self, prompt) -> str:
response = self._get_answer(prompt, self.config)
return response
def _get_answer(self, prompt: str, config: BaseLlmConfig) -> str:
try:
import boto3
except ModuleNotFoundError:
raise ModuleNotFoundError(
"The required dependencies for AWSBedrock are not installed."
'Please install with `pip install --upgrade "embedchain[aws-bedrock]"`'
) from None
self.boto_client = boto3.client("bedrock-runtime", "us-west-2" or os.environ.get("AWS_REGION"))
kwargs = {
"model_id": config.model or "amazon.titan-text-express-v1",
"client": self.boto_client,
"model_kwargs": config.model_kwargs
or {
"temperature": config.temperature,
},
}
if config.stream:
from langchain.callbacks.streaming_stdout import \
StreamingStdOutCallbackHandler
callbacks = [StreamingStdOutCallbackHandler()]
llm = Bedrock(**kwargs, streaming=config.stream, callbacks=callbacks)
else:
llm = Bedrock(**kwargs)
return llm(prompt)
+11 -7
View File
@@ -5,9 +5,7 @@ from typing import Any, Optional
from langchain.schema import BaseMessage as LCBaseMessage
from embedchain.config import BaseLlmConfig
from embedchain.config.llm.base import (DEFAULT_PROMPT,
DEFAULT_PROMPT_WITH_HISTORY_TEMPLATE,
DOCS_SITE_PROMPT_TEMPLATE)
from embedchain.config.llm.base import DEFAULT_PROMPT, DEFAULT_PROMPT_WITH_HISTORY_TEMPLATE, DOCS_SITE_PROMPT_TEMPLATE
from embedchain.helpers.json_serializable import JSONSerializable
from embedchain.memory.base import ChatHistory
from embedchain.memory.message import ChatMessage
@@ -65,6 +63,14 @@ class BaseLlm(JSONSerializable):
self.memory.add(app_id=app_id, chat_message=chat_message, session_id=session_id)
self.update_history(app_id=app_id, session_id=session_id)
def _format_history(self) -> str:
"""Format history to be used in prompt
:return: Formatted history
:rtype: str
"""
return "\n".join(self.history)
def generate_prompt(self, input_query: str, contexts: list[str], **kwargs: dict[str, Any]) -> str:
"""
Generates a prompt based on the given query and context, ready to be
@@ -84,10 +90,8 @@ class BaseLlm(JSONSerializable):
prompt_contains_history = self.config._validate_prompt_history(self.config.prompt)
if prompt_contains_history:
# Prompt contains history
# If there is no history yet, we insert `- no history -`
prompt = self.config.prompt.substitute(
context=context_string, query=input_query, history=self.history or "- no history -"
context=context_string, query=input_query, history=self._format_history() or "No history"
)
elif self.history and not prompt_contains_history:
# History is present, but not included in the prompt.
@@ -98,7 +102,7 @@ class BaseLlm(JSONSerializable):
):
# swap in the template with history
prompt = DEFAULT_PROMPT_WITH_HISTORY_TEMPLATE.substitute(
context=context_string, query=input_query, history=self.history
context=context_string, query=input_query, history=self._format_history()
)
else:
# If we can't swap in the default, we still proceed but tell users that the history is ignored.
+52
View File
@@ -0,0 +1,52 @@
import os
from typing import Optional
from embedchain.config import BaseLlmConfig
from embedchain.helpers.json_serializable import register_deserializable
from embedchain.llm.base import BaseLlm
@register_deserializable
class MistralAILlm(BaseLlm):
def __init__(self, config: Optional[BaseLlmConfig] = None):
super().__init__(config)
if not self.config.api_key and "MISTRAL_API_KEY" not in os.environ:
raise ValueError("Please set the MISTRAL_API_KEY environment variable or pass it in the config.")
def get_llm_model_answer(self, prompt):
return MistralAILlm._get_answer(prompt=prompt, config=self.config)
@staticmethod
def _get_answer(prompt: str, config: BaseLlmConfig):
try:
from langchain_core.messages import HumanMessage, SystemMessage
from langchain_mistralai.chat_models import ChatMistralAI
except ModuleNotFoundError:
raise ModuleNotFoundError(
"The required dependencies for MistralAI are not installed."
'Please install with `pip install --upgrade "embedchain[mistralai]"`'
) from None
api_key = config.api_key or os.getenv("MISTRAL_API_KEY")
client = ChatMistralAI(mistral_api_key=api_key)
messages = []
if config.system_prompt:
messages.append(SystemMessage(content=config.system_prompt))
messages.append(HumanMessage(content=prompt))
kwargs = {
"model": config.model or "mistral-tiny",
"temperature": config.temperature,
"max_tokens": config.max_tokens,
"top_p": config.top_p,
}
# TODO: Add support for streaming
if config.stream:
answer = ""
for chunk in client.stream(**kwargs, input=messages):
answer += chunk.content
return answer
else:
response = client.invoke(**kwargs, input=messages)
answer = response.content
return answer
+6 -8
View File
@@ -35,21 +35,19 @@ class OpenAILlm(BaseLlm):
if config.top_p:
kwargs["model_kwargs"]["top_p"] = config.top_p
if config.stream:
from langchain.callbacks.streaming_stdout import \
StreamingStdOutCallbackHandler
from langchain.callbacks.streaming_stdout import StreamingStdOutCallbackHandler
callbacks = config.callbacks if config.callbacks else [StreamingStdOutCallbackHandler()]
chat = ChatOpenAI(**kwargs, streaming=config.stream, callbacks=callbacks, api_key=api_key)
llm = ChatOpenAI(**kwargs, streaming=config.stream, callbacks=callbacks, api_key=api_key)
else:
chat = ChatOpenAI(**kwargs, api_key=api_key)
llm = ChatOpenAI(**kwargs, api_key=api_key)
if self.functions is not None:
from langchain.chains.openai_functions import \
create_openai_fn_runnable
from langchain.chains.openai_functions import create_openai_fn_runnable
from langchain.prompts import ChatPromptTemplate
structured_prompt = ChatPromptTemplate.from_messages(messages)
runnable = create_openai_fn_runnable(functions=self.functions, prompt=structured_prompt, llm=chat)
runnable = create_openai_fn_runnable(functions=self.functions, prompt=structured_prompt, llm=llm)
fn_res = runnable.invoke(
{
"input": prompt,
@@ -57,4 +55,4 @@ class OpenAILlm(BaseLlm):
)
messages.append(AIMessage(content=json.dumps(fn_res)))
return chat(messages).content
return llm(messages).content
+3 -3
View File
@@ -14,19 +14,19 @@ class DiscourseLoader(BaseLoader):
super().__init__()
if not config:
raise ValueError(
"DiscourseLoader requires a config. Check the documentation for the correct format - `https://docs.embedchain.ai/data-sources/discourse`" # noqa: E501
"DiscourseLoader requires a config. Check the documentation for the correct format - `https://docs.embedchain.ai/components/data-sources/discourse`" # noqa: E501
)
self.domain = config.get("domain")
if not self.domain:
raise ValueError(
"DiscourseLoader requires a domain. Check the documentation for the correct format - `https://docs.embedchain.ai/data-sources/discourse`" # noqa: E501
"DiscourseLoader requires a domain. Check the documentation for the correct format - `https://docs.embedchain.ai/components/data-sources/discourse`" # noqa: E501
)
def _check_query(self, query):
if not query or not isinstance(query, str):
raise ValueError(
"DiscourseLoader requires a query. Check the documentation for the correct format - `https://docs.embedchain.ai/data-sources/discourse`" # noqa: E501
"DiscourseLoader requires a query. Check the documentation for the correct format - `https://docs.embedchain.ai/components/data-sources/discourse`" # noqa: E501
)
def _load_post(self, post_id):
+3 -1
View File
@@ -36,7 +36,9 @@ class JSONReader:
return ["\n".join(useful_lines)]
VALID_URL_PATTERN = "^https:\/\/[0-9A-Za-z]+(\.[0-9A-Za-z]+)*\/[0-9A-Za-z_\/]*\.json$"
VALID_URL_PATTERN = (
"^https?://(?:www\.)?(?:\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}|[a-zA-Z0-9.-]+)(?::\d+)?/(?:[^/\s]+/)*[^/\s]+\.json$"
)
class JSONLoader(BaseLoader):
+4 -1
View File
@@ -15,7 +15,10 @@ from embedchain.utils.misc import clean_string
class PdfFileLoader(BaseLoader):
def load_data(self, url):
"""Load data from a PDF file."""
loader = PyPDFLoader(url)
headers = {
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/98.0.4758.102 Safari/537.36", # noqa:E501
}
loader = PyPDFLoader(url, headers=headers)
data = []
all_content = []
pages = loader.load_and_split()
+4 -1
View File
@@ -31,10 +31,13 @@ class SitemapLoader(BaseLoader):
def load_data(self, sitemap_source):
output = []
web_page_loader = WebPageLoader()
headers = {
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/98.0.4758.102 Safari/537.36", # noqa:E501
}
if urlparse(sitemap_source).scheme in ("http", "https"):
try:
response = requests.get(sitemap_source)
response = requests.get(sitemap_source, headers=headers)
response.raise_for_status()
soup = BeautifulSoup(response.text, "xml")
except requests.RequestException as e:
+4 -1
View File
@@ -22,7 +22,10 @@ class WebPageLoader(BaseLoader):
def load_data(self, url):
"""Load data from a web page using a shared requests' session."""
response = self._session.get(url, timeout=30)
headers = {
"User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/98.0.4758.102 Safari/537.36", # noqa:E501
}
response = self._session.get(url, headers=headers, timeout=30)
response.raise_for_status()
data = response.content
content = self._get_clean_content(data, url)
+2 -2
View File
@@ -92,12 +92,12 @@ class ChatHistory:
"""
if fetch_all:
additional_query = "ORDER BY created_at DESC"
additional_query = "ORDER BY created_at ASC"
params = (app_id,)
else:
additional_query = """
AND session_id=?
ORDER BY created_at DESC
ORDER BY created_at ASC
LIMIT ?
"""
params = (app_id, session_id, num_rounds)
+1
View File
@@ -8,3 +8,4 @@ class VectorDimensions(Enum):
VERTEX_AI = 768
HUGGING_FACE = 384
GOOGLE_AI = 768
MISTRAL_AI = 1024
+1 -1
View File
@@ -88,7 +88,7 @@ class OpenAIAssistant:
if Path(source).is_file():
return source
data_type = data_type or detect_datatype(source)
formatter = DataFormatter(data_type=DataType(data_type), config=AddConfig(), kwargs={})
formatter = DataFormatter(data_type=DataType(data_type), config=AddConfig())
data = formatter.loader.load_data(source)["data"]
return self._save_temp_data(data=data[0]["content"].encode(), source=source)
+32 -5
View File
@@ -201,10 +201,16 @@ def detect_datatype(source: Any) -> DataType:
formatted_source = format_source(str(source), 30)
if url:
from langchain.document_loaders.youtube import \
ALLOWED_NETLOCK as YOUTUBE_ALLOWED_NETLOCS
YOUTUBE_ALLOWED_NETLOCKS = {
"www.youtube.com",
"m.youtube.com",
"youtu.be",
"youtube.com",
"vid.plus",
"www.youtube-nocookie.com",
}
if url.netloc in YOUTUBE_ALLOWED_NETLOCS:
if url.netloc in YOUTUBE_ALLOWED_NETLOCKS:
logging.debug(f"Source of `{formatted_source}` detected as `youtube_video`.")
return DataType.YOUTUBE_VIDEO
@@ -400,6 +406,8 @@ def validate_config(config_data):
"llama2",
"vertexai",
"google",
"aws_bedrock",
"mistralai",
),
Optional("config"): {
Optional("model"): str,
@@ -416,6 +424,7 @@ def validate_config(config_data):
Optional("query_type"): str,
Optional("api_key"): str,
Optional("endpoint"): str,
Optional("model_kwargs"): dict,
},
},
Optional("vectordb"): {
@@ -425,23 +434,41 @@ def validate_config(config_data):
Optional("config"): object, # TODO: add particular config schema for each provider
},
Optional("embedder"): {
Optional("provider"): Or("openai", "gpt4all", "huggingface", "vertexai", "azure_openai", "google"),
Optional("provider"): Or(
"openai",
"gpt4all",
"huggingface",
"vertexai",
"azure_openai",
"google",
"mistralai",
),
Optional("config"): {
Optional("model"): Optional(str),
Optional("deployment_name"): Optional(str),
Optional("api_key"): str,
Optional("title"): str,
Optional("task_type"): str,
Optional("vector_dimension"): int,
},
},
Optional("embedding_model"): {
Optional("provider"): Or("openai", "gpt4all", "huggingface", "vertexai", "azure_openai", "google"),
Optional("provider"): Or(
"openai",
"gpt4all",
"huggingface",
"vertexai",
"azure_openai",
"google",
"mistralai",
),
Optional("config"): {
Optional("model"): str,
Optional("deployment_name"): str,
Optional("api_key"): str,
Optional("title"): str,
Optional("task_type"): str,
Optional("vector_dimension"): int,
},
},
Optional("chunker"): {
+5
View File
@@ -75,3 +75,8 @@ class BaseVectorDB(JSONSerializable):
:type name: str
"""
raise NotImplementedError
def delete(self):
"""Delete from database."""
raise NotImplementedError
+16 -8
View File
@@ -79,6 +79,8 @@ class ChromaDB(BaseVectorDB):
def _generate_where_clause(where: dict[str, any]) -> dict[str, any]:
# If only one filter is supplied, return it as is
# (no need to wrap in $and based on chroma docs)
if where is None:
return {}
if len(where.keys()) <= 1:
return where
where_filters = []
@@ -129,17 +131,13 @@ class ChromaDB(BaseVectorDB):
def add(
self,
embeddings: list[list[float]],
documents: list[str],
metadatas: list[object],
ids: list[str],
**kwargs: Optional[dict[str, Any]],
) -> Any:
"""
Add vectors to chroma database
:param embeddings: list of embeddings to add
:type embeddings: list[list[str]]
:param documents: Documents
:type documents: list[str]
:param metadatas: Metadatas
@@ -184,9 +182,10 @@ class ChromaDB(BaseVectorDB):
self,
input_query: list[str],
n_results: int,
where: dict[str, any],
where: Optional[dict[str, any]] = None,
raw_filter: Optional[dict[str, any]] = None,
citations: bool = False,
**kwargs: Optional[dict[str, Any]],
**kwargs: Optional[dict[str, any]],
) -> Union[list[tuple[str, dict]], list[str]]:
"""
Query contents from vector database based on vector similarity
@@ -197,6 +196,8 @@ class ChromaDB(BaseVectorDB):
:type n_results: int
:param where: to filter data
:type where: dict[str, Any]
:param raw_filter: Raw filter to apply
:type raw_filter: dict[str, Any]
:param citations: we use citations boolean param to return context along with the answer.
:type citations: bool, default is False.
:raises InvalidDimensionException: Dimensions do not match.
@@ -204,14 +205,21 @@ class ChromaDB(BaseVectorDB):
along with url of the source and doc_id (if citations flag is true)
:rtype: list[str], if citations=False, otherwise list[tuple[str, str, str]]
"""
if where and raw_filter:
raise ValueError("Both `where` and `raw_filter` cannot be used together.")
where_clause = {}
if raw_filter:
where_clause = raw_filter
if where:
where_clause = self._generate_where_clause(where)
try:
result = self.collection.query(
query_texts=[
input_query,
],
n_results=n_results,
where=self._generate_where_clause(where),
**kwargs,
where=where_clause,
)
except InvalidDimensionException as e:
raise InvalidDimensionException(
+28 -11
View File
@@ -99,18 +99,27 @@ class ElasticsearchDB(BaseVectorDB):
query = {"bool": {"must": [{"ids": {"values": ids}}]}}
else:
query = {"bool": {"must": []}}
if "app_id" in where:
app_id = where["app_id"]
query["bool"]["must"].append({"term": {"metadata.app_id": app_id}})
response = self.client.search(index=self._get_index(), query=query, _source=False, size=limit)
if where:
for key, value in where.items():
query["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
response = self.client.search(index=self._get_index(), query=query, _source=True, size=limit)
docs = response["hits"]["hits"]
ids = [doc["_id"] for doc in docs]
return {"ids": set(ids)}
doc_ids = [doc["_source"]["metadata"]["doc_id"] for doc in docs]
# Result is modified for compatibility with other vector databases
# TODO: Add method in vector database to return result in a standard format
result = {"ids": ids, "metadatas": []}
for doc_id in doc_ids:
result["metadatas"].append({"doc_id": doc_id})
return result
def add(
self,
embeddings: list[list[float]],
documents: list[str],
metadatas: list[object],
ids: list[str],
@@ -118,8 +127,6 @@ class ElasticsearchDB(BaseVectorDB):
) -> Any:
"""
add data in vector database
:param embeddings: list of embeddings to add
:type embeddings: list[list[str]]
:param documents: list of texts to add
:type documents: list[str]
:param metadatas: list of metadata associated with docs
@@ -189,9 +196,11 @@ class ElasticsearchDB(BaseVectorDB):
},
}
}
if "app_id" in where:
app_id = where["app_id"]
query["script_score"]["query"] = {"match": {"metadata.app_id": app_id}}
if where:
for key, value in where.items():
query["script_score"]["query"]["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
_source = ["text", "metadata"]
response = self.client.search(index=self._get_index(), query=query, _source=_source, size=n_results)
docs = response["hits"]["hits"]
@@ -247,3 +256,11 @@ class ElasticsearchDB(BaseVectorDB):
# NOTE: The method is preferred to an attribute, because if collection name changes,
# it's always up-to-date.
return f"{self.config.collection_name}_{self.embedder.vector_dimension}".lower()
def delete(self, where):
"""Delete documents from the database."""
query = {"query": {"bool": {"must": []}}}
for key, value in where.items():
query["query"]["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
self.client.delete_by_query(index=self._get_index(), body=query)
self.client.indices.refresh(index=self._get_index())
+14 -25
View File
@@ -96,9 +96,9 @@ class OpenSearchDB(BaseVectorDB):
else:
query["query"] = {"bool": {"must": []}}
if "app_id" in where:
app_id = where["app_id"]
query["query"]["bool"]["must"].append({"term": {"metadata.app_id.keyword": app_id}})
if where:
for key, value in where.items():
query["query"]["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
# OpenSearch syntax is different from Elasticsearch
response = self.client.search(index=self._get_index(), body=query, _source=True, size=limit)
@@ -114,22 +114,10 @@ class OpenSearchDB(BaseVectorDB):
result["metadatas"].append({"doc_id": doc_id})
return result
def add(
self,
embeddings: list[list[str]],
documents: list[str],
metadatas: list[object],
ids: list[str],
**kwargs: Optional[dict[str, any]],
):
"""Add data in vector database.
def add(self, documents: list[str], metadatas: list[object], ids: list[str], **kwargs: Optional[dict[str, any]]):
"""Adds documents to the opensearch index"""
Args:
embeddings (list[list[str]]): list of embeddings to add.
documents (list[str]): list of texts to add.
metadatas (list[object]): list of metadata associated with docs.
ids (list[str]): IDs of docs.
"""
embeddings = self.embedder.embedding_fn(documents)
for batch_start in tqdm(range(0, len(documents), self.BATCH_SIZE), desc="Inserting batches in opensearch"):
batch_end = batch_start + self.BATCH_SIZE
batch_documents = documents[batch_start:batch_end]
@@ -188,9 +176,11 @@ class OpenSearchDB(BaseVectorDB):
)
pre_filter = {"match_all": {}} # default
if "app_id" in where:
app_id = where["app_id"]
pre_filter = {"bool": {"must": [{"term": {"metadata.app_id.keyword": app_id}}]}}
if len(where) > 0:
pre_filter = {"bool": {"must": []}}
for key, value in where.items():
pre_filter["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
docs = docsearch.similarity_search_with_score(
input_query,
search_type="script_scoring",
@@ -248,10 +238,9 @@ class OpenSearchDB(BaseVectorDB):
def delete(self, where):
"""Deletes a document from the OpenSearch index"""
if "doc_id" not in where:
raise ValueError("doc_id is required to delete a document")
query = {"query": {"bool": {"must": [{"term": {"metadata.doc_id": where["doc_id"]}}]}}}
query = {"query": {"bool": {"must": []}}}
for key, value in where.items():
query["query"]["bool"]["must"].append({"term": {f"metadata.{key}.keyword": value}})
self.client.delete_by_query(index=self._get_index(), body=query)
def _get_index(self) -> str:
+85 -52
View File
@@ -41,7 +41,7 @@ class PineconeDB(BaseVectorDB):
"Please make sure the type is right and that you are passing an instance."
)
self.config = config
self.client = self._setup_pinecone_index()
self._setup_pinecone_index()
# Call parent init here because embedder is needed
super().__init__(config=self.config)
@@ -52,20 +52,30 @@ class PineconeDB(BaseVectorDB):
if not self.embedder:
raise ValueError("Embedder not set. Please set an embedder with `set_embedder` before initialization.")
# Loads the Pinecone index or creates it if not present.
def _setup_pinecone_index(self):
pinecone.init(
api_key=os.environ.get("PINECONE_API_KEY"),
environment=os.environ.get("PINECONE_ENV"),
**self.config.extra_params,
)
self.index_name = self._get_index_name()
indexes = pinecone.list_indexes()
if indexes is None or self.index_name not in indexes:
pinecone.create_index(
name=self.index_name, metric=self.config.metric, dimension=self.config.vector_dimension
"""
Loads the Pinecone index or creates it if not present.
"""
api_key = self.config.api_key or os.environ.get("PINECONE_API_KEY")
if not api_key:
raise ValueError("Please set the PINECONE_API_KEY environment variable or pass it in config.")
self.client = pinecone.Pinecone(api_key=api_key, **self.config.extra_params)
indexes = self.client.list_indexes().names()
if indexes is None or self.config.index_name not in indexes:
if self.config.pod_config:
spec = pinecone.PodSpec(**self.config.pod_config)
elif self.config.serverless_config:
spec = pinecone.ServerlessSpec(**self.config.serverless_config)
else:
raise ValueError("No pod_config or serverless_config found.")
self.client.create_index(
name=self.config.index_name,
metric=self.config.metric,
dimension=self.config.vector_dimension,
spec=spec,
)
return pinecone.Index(self.index_name)
self.pinecone_index = self.client.Index(self.config.index_name)
def get(self, ids: Optional[list[str]] = None, where: Optional[dict[str, any]] = None, limit: Optional[int] = None):
"""
@@ -79,16 +89,19 @@ class PineconeDB(BaseVectorDB):
:rtype: Set[str]
"""
existing_ids = list()
metadatas = []
if ids is not None:
for i in range(0, len(ids), 1000):
result = self.client.fetch(ids=ids[i : i + 1000])
batch_existing_ids = list(result.get("vectors").keys())
result = self.pinecone_index.fetch(ids=ids[i : i + 1000])
vectors = result.get("vectors")
batch_existing_ids = list(vectors.keys())
existing_ids.extend(batch_existing_ids)
return {"ids": existing_ids}
metadatas.extend([vectors.get(ids).get("metadata") for ids in batch_existing_ids])
return {"ids": existing_ids, "metadatas": metadatas}
def add(
self,
embeddings: list[list[float]],
documents: list[str],
metadatas: list[object],
ids: list[str],
@@ -104,7 +117,6 @@ class PineconeDB(BaseVectorDB):
:type ids: list[str]
"""
docs = []
print("Adding documents to Pinecone...")
embeddings = self.embedder.embedding_fn(documents)
for id, text, metadata, embedding in zip(ids, documents, metadatas, embeddings):
docs.append(
@@ -115,43 +127,51 @@ class PineconeDB(BaseVectorDB):
}
)
for chunk in chunks(docs, self.BATCH_SIZE, desc="Adding chunks in batches..."):
self.client.upsert(chunk, **kwargs)
for chunk in chunks(docs, self.BATCH_SIZE, desc="Adding chunks in batches"):
self.pinecone_index.upsert(chunk, **kwargs)
def query(
self,
input_query: list[str],
n_results: int,
where: dict[str, any],
where: Optional[dict[str, any]] = None,
raw_filter: Optional[dict[str, any]] = None,
citations: bool = False,
app_id: Optional[str] = None,
**kwargs: Optional[dict[str, any]],
) -> Union[list[tuple[str, dict]], list[str]]:
"""
query contents from vector database based on vector similarity
:param input_query: list of query string
:type input_query: list[str]
:param n_results: no of similar documents to fetch from database
:type n_results: int
:param where: Optional. to filter data
:type where: dict[str, any]
:param citations: we use citations boolean param to return context along with the answer.
:type citations: bool, default is False.
:return: The content of the document that matched your query,
along with url of the source and doc_id (if citations flag is true)
:rtype: list[str], if citations=False, otherwise list[tuple[str, str, str]]
Query contents from vector database based on vector similarity.
Args:
input_query (list[str]): List of query strings.
n_results (int): Number of similar documents to fetch from the database.
where (dict[str, any], optional): Filter criteria for the search.
raw_filter (dict[str, any], optional): Advanced raw filter criteria for the search.
citations (bool, optional): Flag to return context along with metadata. Defaults to False.
app_id (str, optional): Application ID to be passed to Pinecone.
Returns:
Union[list[tuple[str, dict]], list[str]]: List of document contexts, optionally with metadata.
"""
query_filter = raw_filter if raw_filter is not None else self._generate_filter(where)
if app_id:
query_filter["app_id"] = {"$eq": app_id}
query_vector = self.embedder.embedding_fn([input_query])[0]
data = self.client.query(vector=query_vector, filter=where, top_k=n_results, include_metadata=True, **kwargs)
contexts = []
for doc in data["matches"]:
metadata = doc["metadata"]
context = metadata["text"]
if citations:
metadata["score"] = doc["score"]
contexts.append(tuple((context, metadata)))
else:
contexts.append(context)
return contexts
data = self.pinecone_index.query(
vector=query_vector,
filter=query_filter,
top_k=n_results,
include_metadata=True,
**kwargs,
)
return [
(metadata.get("text"), {**metadata, "score": doc.get("score")}) if citations else metadata.get("text")
for doc in data.get("matches", [])
for metadata in [doc.get("metadata", {})]
]
def set_collection_name(self, name: str):
"""
@@ -171,7 +191,8 @@ class PineconeDB(BaseVectorDB):
:return: number of documents
:rtype: int
"""
return self.client.describe_index_stats()["total_vector_count"]
data = self.pinecone_index.describe_index_stats()
return data["total_vector_count"]
def _get_or_create_db(self):
"""Called during initialization"""
@@ -182,14 +203,26 @@ class PineconeDB(BaseVectorDB):
Resets the database. Deletes all embeddings irreversibly.
"""
# Delete all data from the database
pinecone.delete_index(self.index_name)
self.client.delete_index(self.config.index_name)
self._setup_pinecone_index()
# Pinecone only allows alphanumeric characters and "-" in the index name
def _get_index_name(self) -> str:
"""Get the Pinecone index for a collection
@staticmethod
def _generate_filter(where: dict):
query = {}
for k, v in where.items():
query[k] = {"$eq": v}
return query
:return: Pinecone index
:rtype: str
def delete(self, where: dict):
"""Delete from database.
:param ids: list of ids to delete
:type ids: list[str]
"""
return f"{self.config.collection_name}-{self.config.vector_dimension}".lower().replace("_", "-")
# Deleting with filters is not supported for `starter` index type.
# Follow `https://docs.pinecone.io/docs/metadata-filtering#deleting-vectors-by-metadata-filter` for more details
db_filter = self._generate_filter(where)
try:
self.pinecone_index.delete(filter=db_filter)
except Exception as e:
print(f"Failed to delete from Pinecone: {e}")
return
+42 -20
View File
@@ -11,6 +11,8 @@ try:
except ImportError:
raise ImportError("Qdrant requires extra dependencies. Install with `pip install embedchain[qdrant]`") from None
from tqdm import tqdm
from embedchain.config.vectordb.qdrant import QdrantDBConfig
from embedchain.vectordb.base import BaseVectorDB
@@ -48,7 +50,6 @@ class QdrantDB(BaseVectorDB):
raise ValueError("Embedder not set. Please set an embedder with `set_embedder` before initialization.")
self.collection_name = self._get_or_create_collection()
self.metadata_keys = {"data_type", "doc_id", "url", "hash", "app_id", "text"}
all_collections = self.client.get_collections()
collection_names = [collection.name for collection in all_collections.collections]
if self.collection_name not in collection_names:
@@ -82,21 +83,23 @@ class QdrantDB(BaseVectorDB):
:return: All the existing IDs
:rtype: Set[str]
"""
if ids is None or len(ids) == 0:
return {"ids": []}
keys = set(where.keys() if where is not None else set())
qdrant_must_filters = [
models.FieldCondition(
key="identifier",
match=models.MatchAny(
any=ids,
),
qdrant_must_filters = []
if ids:
qdrant_must_filters.append(
models.FieldCondition(
key="identifier",
match=models.MatchAny(
any=ids,
),
)
)
]
if len(keys.intersection(self.metadata_keys)) != 0:
for key in keys.intersection(self.metadata_keys):
if len(keys) > 0:
for key in keys:
qdrant_must_filters.append(
models.FieldCondition(
key="metadata.{}".format(key),
@@ -108,6 +111,7 @@ class QdrantDB(BaseVectorDB):
offset = 0
existing_ids = []
metadatas = []
while offset is not None:
response = self.client.scroll(
collection_name=self.collection_name,
@@ -118,19 +122,17 @@ class QdrantDB(BaseVectorDB):
offset = response[1]
for doc in response[0]:
existing_ids.append(doc.payload["identifier"])
return {"ids": existing_ids}
metadatas.append(doc.payload["metadata"])
return {"ids": existing_ids, "metadatas": metadatas}
def add(
self,
embeddings: list[list[float]],
documents: list[str],
metadatas: list[object],
ids: list[str],
**kwargs: Optional[dict[str, any]],
):
"""add data in vector database
:param embeddings: list of embeddings for the corresponding documents to be added
:type documents: list[list[float]]
:param documents: list of texts to add
:type documents: list[str]
:param metadatas: list of metadata associated with docs
@@ -146,7 +148,8 @@ class QdrantDB(BaseVectorDB):
metadata["text"] = document
qdrant_ids.append(str(uuid.uuid4()))
payloads.append({"identifier": id, "text": document, "metadata": copy.deepcopy(metadata)})
for i in range(0, len(qdrant_ids), self.BATCH_SIZE):
for i in tqdm(range(0, len(qdrant_ids), self.BATCH_SIZE), desc="Adding data in batches"):
self.client.upsert(
collection_name=self.collection_name,
points=Batch(
@@ -183,16 +186,17 @@ class QdrantDB(BaseVectorDB):
keys = set(where.keys() if where is not None else set())
qdrant_must_filters = []
if len(keys.intersection(self.metadata_keys)) != 0:
for key in keys.intersection(self.metadata_keys):
if len(keys) > 0:
for key in keys:
qdrant_must_filters.append(
models.FieldCondition(
key="payload.metadata.{}".format(key),
key="metadata.{}".format(key),
match=models.MatchValue(
value=where.get(key),
),
)
)
results = self.client.search(
collection_name=self.collection_name,
query_filter=models.Filter(must=qdrant_must_filters),
@@ -231,3 +235,21 @@ class QdrantDB(BaseVectorDB):
raise TypeError("Collection name must be a string")
self.config.collection_name = name
self.collection_name = self._get_or_create_collection()
@staticmethod
def _generate_query(where: dict):
must_fields = []
for key, value in where.items():
must_fields.append(
models.FieldCondition(
key=f"metadata.{key}",
match=models.MatchValue(
value=value,
),
)
)
return models.Filter(must=must_fields)
def delete(self, where: dict):
db_filter = self._generate_query(where)
self.client.delete(collection_name=self.collection_name, points_selector=db_filter)
+92 -41
View File
@@ -1,6 +1,6 @@
import copy
import os
from typing import Any, Optional, Union
from typing import Optional, Union
try:
import weaviate
@@ -45,6 +45,9 @@ class WeaviateDB(BaseVectorDB):
auth_client_secret=weaviate.AuthApiKey(api_key=os.environ.get("WEAVIATE_API_KEY")),
**self.config.extra_params,
)
# Since weaviate uses graphQL, we need to keep track of metadata keys added in the vectordb.
# This is needed to filter data while querying.
self.metadata_keys = {"data_type", "doc_id", "url", "hash", "app_id"}
# Call parent init here because embedder is needed
super().__init__(config=self.config)
@@ -58,7 +61,6 @@ class WeaviateDB(BaseVectorDB):
raise ValueError("Embedder not set. Please set an embedder with `set_embedder` before initialization.")
self.index_name = self._get_index_name()
self.metadata_keys = {"data_type", "doc_id", "url", "hash", "app_id"}
if not self.client.schema.exists(self.index_name):
# id is a reserved field in Weaviate, hence we had to change the name of the id field to identifier
# The none vectorizer is crucial as we have our own custom embedding function
@@ -127,41 +129,67 @@ class WeaviateDB(BaseVectorDB):
:return: ids
:rtype: Set[str]
"""
weaviate_where_operands = []
if ids is None or len(ids) == 0:
return {"ids": []}
if ids:
for doc_id in ids:
weaviate_where_operands.append({"path": ["identifier"], "operator": "Equal", "valueText": doc_id})
keys = set(where.keys() if where is not None else set())
if len(keys) > 0:
for key in keys:
weaviate_where_operands.append(
{
"path": ["metadata", self.index_name + "_metadata", key],
"operator": "Equal",
"valueText": where.get(key),
}
)
if len(weaviate_where_operands) == 1:
weaviate_where_clause = weaviate_where_operands[0]
else:
weaviate_where_clause = {"operator": "And", "operands": weaviate_where_operands}
existing_ids = []
metadatas = []
cursor = None
offset = 0
has_iterated_once = False
query_metadata_keys = self.metadata_keys.union(keys)
while cursor is not None or not has_iterated_once:
has_iterated_once = True
results = self._query_with_cursor(
self.client.query.get(self.index_name, ["identifier"])
results = self._query_with_offset(
self.client.query.get(
self.index_name,
[
"identifier",
weaviate.LinkTo("metadata", self.index_name + "_metadata", list(query_metadata_keys)),
],
)
.with_where(weaviate_where_clause)
.with_additional(["id"])
.with_limit(self.BATCH_SIZE),
cursor,
.with_limit(limit or self.BATCH_SIZE),
offset,
)
fetched_results = results["data"]["Get"].get(self.index_name, [])
if len(fetched_results) == 0:
if not fetched_results:
break
for result in fetched_results:
existing_ids.append(result["identifier"])
metadatas.append(result["metadata"][0])
cursor = result["_additional"]["id"]
offset += 1
return {"ids": existing_ids}
if limit is not None and len(existing_ids) >= limit:
break
def add(
self,
embeddings: list[list[float]],
documents: list[str],
metadatas: list[object],
ids: list[str],
**kwargs: Optional[dict[str, any]],
):
return {"ids": existing_ids, "metadatas": metadatas}
def add(self, documents: list[str], metadatas: list[object], ids: list[str], **kwargs: Optional[dict[str, any]]):
"""add data in vector database
:param embeddings: list of embeddings for the corresponding documents to be added
:type documents: list[list[float]]
:param documents: list of texts to add
:type documents: list[str]
:param metadatas: list of metadata associated with docs
@@ -191,12 +219,7 @@ class WeaviateDB(BaseVectorDB):
)
def query(
self,
input_query: list[str],
n_results: int,
where: dict[str, any],
citations: bool = False,
**kwargs: Optional[dict[str, Any]],
self, input_query: list[str], n_results: int, where: dict[str, any], citations: bool = False
) -> Union[list[tuple[str, dict]], list[str]]:
"""
query contents from vector database based on vector similarity
@@ -215,21 +238,20 @@ class WeaviateDB(BaseVectorDB):
query_vector = self.embedder.embedding_fn([input_query])[0]
keys = set(where.keys() if where is not None else set())
data_fields = ["text"]
query_metadata_keys = self.metadata_keys.union(keys)
if citations:
data_fields.append(weaviate.LinkTo("metadata", self.index_name + "_metadata", list(self.metadata_keys)))
data_fields.append(weaviate.LinkTo("metadata", self.index_name + "_metadata", list(query_metadata_keys)))
if len(keys.intersection(self.metadata_keys)) != 0:
if len(keys) > 0:
weaviate_where_operands = []
for key in keys:
if key in self.metadata_keys:
weaviate_where_operands.append(
{
"path": ["metadata", self.index_name + "_metadata", key],
"operator": "Equal",
"valueText": where.get(key),
}
)
weaviate_where_operands.append(
{
"path": ["metadata", self.index_name + "_metadata", key],
"operator": "Equal",
"valueText": where.get(key),
}
)
if len(weaviate_where_operands) == 1:
weaviate_where_clause = weaviate_where_operands[0]
else:
@@ -252,6 +274,9 @@ class WeaviateDB(BaseVectorDB):
.do()
)
if results["data"]["Get"].get(self.index_name) is None:
return []
docs = results["data"]["Get"].get(self.index_name)
contexts = []
for doc in docs:
@@ -303,11 +328,37 @@ class WeaviateDB(BaseVectorDB):
:return: Weaviate index
:rtype: str
"""
return f"{self.config.collection_name}_{self.embedder.vector_dimension}".capitalize()
return f"{self.config.collection_name}_{self.embedder.vector_dimension}".capitalize().replace("-", "_")
@staticmethod
def _query_with_cursor(query, cursor):
if cursor is not None:
query.with_after(cursor)
def _query_with_offset(query, offset):
if offset:
query.with_offset(offset)
results = query.do()
return results
def _generate_query(self, where: dict):
weaviate_where_operands = []
for key, value in where.items():
weaviate_where_operands.append(
{
"path": ["metadata", self.index_name + "_metadata", key],
"operator": "Equal",
"valueText": value,
}
)
if len(weaviate_where_operands) == 1:
weaviate_where_clause = weaviate_where_operands[0]
else:
weaviate_where_clause = {"operator": "And", "operands": weaviate_where_operands}
return weaviate_where_clause
def delete(self, where: dict):
"""Delete from database.
:param where: to filter data
:type where: dict[str, any]
"""
query = self._generate_query(where)
self.client.batch.delete_objects(self.index_name, where=query)
+38 -26
View File
@@ -69,6 +69,7 @@ class ZillizVectorDB(BaseVectorDB):
FieldSchema(name="id", dtype=DataType.VARCHAR, is_primary=True, max_length=512),
FieldSchema(name="text", dtype=DataType.VARCHAR, max_length=2048),
FieldSchema(name="embeddings", dtype=DataType.FLOAT_VECTOR, dim=self.embedder.vector_dimension),
FieldSchema(name="metadata", dtype=DataType.JSON),
]
schema = CollectionSchema(fields, enable_dynamic_field=True)
@@ -94,21 +95,29 @@ class ZillizVectorDB(BaseVectorDB):
:return: Existing documents.
:rtype: Set[str]
"""
if ids is None or len(ids) == 0 or self.collection.num_entities == 0:
return {"ids": []}
data_ids = []
metadatas = []
if self.collection.num_entities == 0 or self.collection.is_empty:
return {"ids": data_ids, "metadatas": metadatas}
if not self.collection.is_empty:
filter_ = f"id in {ids}"
results = self.client.query(
collection_name=self.config.collection_name, filter=filter_, output_fields=["id"]
)
results = [res["id"] for res in results]
filter_ = ""
if ids:
filter_ = f'id in "{ids}"'
return {"ids": set(results)}
if where:
if filter_:
filter_ += " and "
filter_ = f"{self._generate_zilliz_filter(where)}"
results = self.client.query(collection_name=self.config.collection_name, filter=filter_, output_fields=["*"])
for res in results:
data_ids.append(res.get("id"))
metadatas.append(res.get("metadata", {}))
return {"ids": data_ids, "metadatas": metadatas}
def add(
self,
embeddings: list[list[float]],
documents: list[str],
metadatas: list[object],
ids: list[str],
@@ -118,7 +127,7 @@ class ZillizVectorDB(BaseVectorDB):
embeddings = self.embedder.embedding_fn(documents)
for id, doc, metadata, embedding in zip(ids, documents, metadatas, embeddings):
data = {**metadata, "id": id, "text": doc, "embeddings": embedding}
data = {"id": id, "text": doc, "embeddings": embedding, "metadata": metadata}
self.client.insert(collection_name=self.config.collection_name, data=data, **kwargs)
self.collection.load()
@@ -129,7 +138,7 @@ class ZillizVectorDB(BaseVectorDB):
self,
input_query: list[str],
n_results: int,
where: dict[str, any],
where: dict[str, Any],
citations: bool = False,
**kwargs: Optional[dict[str, Any]],
) -> Union[list[tuple[str, dict]], list[str]]:
@@ -141,7 +150,7 @@ class ZillizVectorDB(BaseVectorDB):
:param n_results: no of similar documents to fetch from database
:type n_results: int
:param where: to filter data
:type where: str
:type where: dict[str, Any]
:raises InvalidDimensionException: Dimensions do not match.
:param citations: we use citations boolean param to return context along with the answer.
:type citations: bool, default is False.
@@ -153,16 +162,15 @@ class ZillizVectorDB(BaseVectorDB):
if self.collection.is_empty:
return []
if not isinstance(where, str):
where = None
output_fields = ["*"]
input_query_vector = self.embedder.embedding_fn([input_query])
query_vector = input_query_vector[0]
query_filter = self._generate_zilliz_filter(where)
query_result = self.client.search(
collection_name=self.config.collection_name,
data=[query_vector],
filter=query_filter,
limit=n_results,
output_fields=output_fields,
**kwargs,
@@ -174,12 +182,10 @@ class ZillizVectorDB(BaseVectorDB):
score = query["distance"]
context = data["text"]
if "embeddings" in data:
data.pop("embeddings")
if citations:
data["score"] = score
contexts.append(tuple((context, data)))
metadata = data.get("metadata", {})
metadata["score"] = score
contexts.append(tuple((context, metadata)))
else:
contexts.append(context)
return contexts
@@ -217,7 +223,13 @@ class ZillizVectorDB(BaseVectorDB):
raise TypeError("Collection name must be a string")
self.config.collection_name = name
def delete(self, keys: Union[list, str, int]):
def _generate_zilliz_filter(self, where: dict[str, str]):
operands = []
for key, value in where.items():
operands.append(f'(metadata["{key}"] == "{value}")')
return " and ".join(operands)
def delete(self, where: dict[str, Any]):
"""
Delete the embeddings from DB. Zilliz only support deleting with keys.
@@ -225,7 +237,7 @@ class ZillizVectorDB(BaseVectorDB):
:param keys: Primary keys of the table entries to delete.
:type keys: Union[list, str, int]
"""
self.client.delete(
collection_name=self.config.collection_name,
pks=keys,
)
data = self.get(where=where)
keys = data.get("ids", [])
if keys:
self.client.delete(collection_name=self.config.collection_name, pks=keys)
+9 -7
View File
@@ -72,13 +72,15 @@
"outputs": [],
"source": [
"app = App.from_config(config={\n",
" \"provider\": \"chroma\",\n",
" \"config\": {\n",
" \"collection_name\": \"my-collection\",\n",
" \"host\": \"your-chromadb-url.com\",\n",
" \"port\": 5200,\n",
" \"allow_reset\": True\n",
" }\n",
" \"vectordb\": {\n",
" \"provider\": \"chroma\",\n",
" \"config\": {\n",
" \"collection_name\": \"my-collection\",\n",
" \"host\": \"your-chromadb-url.com\",\n",
" \"port\": 5200,\n",
" \"allow_reset\": True\n",
" }\n",
" }\n",
"})"
]
},
Generated
+203 -73
View File
@@ -383,6 +383,47 @@ files = [
{file = "blinker-1.6.3.tar.gz", hash = "sha256:152090d27c1c5c722ee7e48504b02d76502811ce02e1523553b4cf8c8b3d3a8d"},
]
[[package]]
name = "boto3"
version = "1.34.22"
description = "The AWS SDK for Python"
optional = true
python-versions = ">= 3.8"
files = [
{file = "boto3-1.34.22-py3-none-any.whl", hash = "sha256:5909cd1393143576265c692e908a9ae495492c04a0ffd4bae8578adc2e44729e"},
{file = "boto3-1.34.22.tar.gz", hash = "sha256:a98c0b86f6044ff8314cc2361e1ef574d674318313ab5606ccb4a6651c7a3f8c"},
]
[package.dependencies]
botocore = ">=1.34.22,<1.35.0"
jmespath = ">=0.7.1,<2.0.0"
s3transfer = ">=0.10.0,<0.11.0"
[package.extras]
crt = ["botocore[crt] (>=1.21.0,<2.0a0)"]
[[package]]
name = "botocore"
version = "1.34.22"
description = "Low-level, data-driven core of boto 3."
optional = true
python-versions = ">= 3.8"
files = [
{file = "botocore-1.34.22-py3-none-any.whl", hash = "sha256:e5f7775975b9213507fbcf846a96b7a2aec2a44fc12a44585197b014a4ab0889"},
{file = "botocore-1.34.22.tar.gz", hash = "sha256:c47ba4286c576150d1b6ca6df69a87b5deff3d23bd84da8bcf8431ebac3c40ba"},
]
[package.dependencies]
jmespath = ">=0.7.1,<2.0.0"
python-dateutil = ">=2.1,<3.0.0"
urllib3 = [
{version = ">=1.25.4,<1.27", markers = "python_version < \"3.10\""},
{version = ">=1.25.4,<2.1", markers = "python_version >= \"3.10\""},
]
[package.extras]
crt = ["awscrt (==0.19.19)"]
[[package]]
name = "brotli"
version = "1.1.0"
@@ -1219,25 +1260,6 @@ files = [
{file = "distro-1.8.0.tar.gz", hash = "sha256:02e111d1dc6a50abb8eed6bf31c3e48ed8b0830d1ea2a1b78c61765c2513fdd8"},
]
[[package]]
name = "dnspython"
version = "2.4.2"
description = "DNS toolkit"
optional = true
python-versions = ">=3.8,<4.0"
files = [
{file = "dnspython-2.4.2-py3-none-any.whl", hash = "sha256:57c6fbaaeaaf39c891292012060beb141791735dbb4004798328fc2c467402d8"},
{file = "dnspython-2.4.2.tar.gz", hash = "sha256:8dcfae8c7460a2f84b4072e26f1c9f4101ca20c071649cb7c34e8b6a93d58984"},
]
[package.extras]
dnssec = ["cryptography (>=2.6,<42.0)"]
doh = ["h2 (>=4.1.0)", "httpcore (>=0.17.3)", "httpx (>=0.24.1)"]
doq = ["aioquic (>=0.9.20)"]
idna = ["idna (>=2.1,<4.0)"]
trio = ["trio (>=0.14,<0.23)"]
wmi = ["wmi (>=1.5.1,<2.0.0)"]
[[package]]
name = "docx2txt"
version = "0.8"
@@ -2487,24 +2509,24 @@ files = [
[[package]]
name = "httpcore"
version = "0.18.0"
version = "1.0.2"
description = "A minimal low-level HTTP client."
optional = false
python-versions = ">=3.8"
files = [
{file = "httpcore-0.18.0-py3-none-any.whl", hash = "sha256:adc5398ee0a476567bf87467063ee63584a8bce86078bf748e48754f60202ced"},
{file = "httpcore-0.18.0.tar.gz", hash = "sha256:13b5e5cd1dca1a6636a6aaea212b19f4f85cd88c366a2b82304181b769aab3c9"},
{file = "httpcore-1.0.2-py3-none-any.whl", hash = "sha256:096cc05bca73b8e459a1fc3dcf585148f63e534eae4339559c9b8a8d6399acc7"},
{file = "httpcore-1.0.2.tar.gz", hash = "sha256:9fc092e4799b26174648e54b74ed5f683132a464e95643b226e00c2ed2fa6535"},
]
[package.dependencies]
anyio = ">=3.0,<5.0"
certifi = "*"
h11 = ">=0.13,<0.15"
sniffio = "==1.*"
[package.extras]
asyncio = ["anyio (>=4.0,<5.0)"]
http2 = ["h2 (>=3,<5)"]
socks = ["socksio (==1.*)"]
trio = ["trio (>=0.22.0,<0.23.0)"]
[[package]]
name = "httplib2"
@@ -2569,21 +2591,22 @@ test = ["Cython (>=0.29.24,<0.30.0)"]
[[package]]
name = "httpx"
version = "0.25.0"
version = "0.25.2"
description = "The next generation HTTP client."
optional = false
python-versions = ">=3.8"
files = [
{file = "httpx-0.25.0-py3-none-any.whl", hash = "sha256:181ea7f8ba3a82578be86ef4171554dd45fec26a02556a744db029a0a27b7100"},
{file = "httpx-0.25.0.tar.gz", hash = "sha256:47ecda285389cb32bb2691cc6e069e3ab0205956f681c5b2ad2325719751d875"},
{file = "httpx-0.25.2-py3-none-any.whl", hash = "sha256:a05d3d052d9b2dfce0e3896636467f8a5342fb2b902c819428e1ac65413ca118"},
{file = "httpx-0.25.2.tar.gz", hash = "sha256:8b8fcaa0c8ea7b05edd69a094e63a2094c4efcb48129fb757361bc423c0ad9e8"},
]
[package.dependencies]
anyio = "*"
brotli = {version = "*", optional = true, markers = "platform_python_implementation == \"CPython\" and extra == \"brotli\""}
brotlicffi = {version = "*", optional = true, markers = "platform_python_implementation != \"CPython\" and extra == \"brotli\""}
certifi = "*"
h2 = {version = ">=3,<5", optional = true, markers = "extra == \"http2\""}
httpcore = ">=0.18.0,<0.19.0"
httpcore = "==1.*"
idna = "*"
sniffio = "*"
socksio = {version = "==1.*", optional = true, markers = "extra == \"socks\""}
@@ -2809,6 +2832,17 @@ MarkupSafe = ">=2.0"
[package.extras]
i18n = ["Babel (>=2.7)"]
[[package]]
name = "jmespath"
version = "1.0.1"
description = "JSON Matching Expressions"
optional = true
python-versions = ">=3.7"
files = [
{file = "jmespath-1.0.1-py3-none-any.whl", hash = "sha256:02e2e4cc71b5bcab88332eebf907519190dd9e6e82107fa7f83b1003a6252980"},
{file = "jmespath-1.0.1.tar.gz", hash = "sha256:90261b206d6defd58fdd5e85f478bf633a2901798906be2ad389150c5c60edbe"},
]
[[package]]
name = "joblib"
version = "1.3.2"
@@ -3024,6 +3058,45 @@ openai = ["openai (<2)", "tiktoken (>=0.3.2,<0.6.0)"]
qdrant = ["qdrant-client (>=1.3.1,<2.0.0)"]
text-helpers = ["chardet (>=5.1.0,<6.0.0)"]
[[package]]
name = "langchain-core"
version = "0.1.12"
description = "Building applications with LLMs through composability"
optional = true
python-versions = ">=3.8.1,<4.0"
files = [
{file = "langchain_core-0.1.12-py3-none-any.whl", hash = "sha256:d11c6262f7a9deff7de8fdf14498b8a951020dfed3a80f2358ab731ad04abef0"},
{file = "langchain_core-0.1.12.tar.gz", hash = "sha256:f18e9300e9a07589b3e280e51befbc5a4513f535949406e55eb7a2dc40c3ce66"},
]
[package.dependencies]
anyio = ">=3,<5"
jsonpatch = ">=1.33,<2.0"
langsmith = ">=0.0.63,<0.1.0"
packaging = ">=23.2,<24.0"
pydantic = ">=1,<3"
PyYAML = ">=5.3"
requests = ">=2,<3"
tenacity = ">=8.1.0,<9.0.0"
[package.extras]
extended-testing = ["jinja2 (>=3,<4)"]
[[package]]
name = "langchain-mistralai"
version = "0.0.3"
description = "An integration package connecting Mistral and LangChain"
optional = true
python-versions = ">=3.8.1,<4.0"
files = [
{file = "langchain_mistralai-0.0.3-py3-none-any.whl", hash = "sha256:ebb8ba3d7978b5ee16f7e09512ffa434e00bc9863f1537f1a5f5203882d99619"},
{file = "langchain_mistralai-0.0.3.tar.gz", hash = "sha256:2e45ee0118df8e4b5577ce8c4f89743059801e473f40a8b7c89cb99dd715f423"},
]
[package.dependencies]
langchain-core = ">=0.1,<0.2"
mistralai = ">=0.0.11,<0.0.12"
[[package]]
name = "langdetect"
version = "1.0.9"
@@ -3112,24 +3185,6 @@ files = [
{file = "lit-17.0.2.tar.gz", hash = "sha256:d6a551eab550f81023c82a260cd484d63970d2be9fd7588111208e7d2ff62212"},
]
[[package]]
name = "loguru"
version = "0.7.2"
description = "Python logging made (stupidly) simple"
optional = true
python-versions = ">=3.5"
files = [
{file = "loguru-0.7.2-py3-none-any.whl", hash = "sha256:003d71e3d3ed35f0f8984898359d65b79e5b21943f78af86aa5491210429b8eb"},
{file = "loguru-0.7.2.tar.gz", hash = "sha256:e671a53522515f34fd406340ee968cb9ecafbc4b36c679da03c18fd8d0bd51ac"},
]
[package.dependencies]
colorama = {version = ">=0.3.4", markers = "sys_platform == \"win32\""}
win32-setctime = {version = ">=1.0.0", markers = "sys_platform == \"win32\""}
[package.extras]
dev = ["Sphinx (==7.2.5)", "colorama (==0.4.5)", "colorama (==0.4.6)", "exceptiongroup (==1.1.3)", "freezegun (==1.1.0)", "freezegun (==1.2.2)", "mypy (==v0.910)", "mypy (==v0.971)", "mypy (==v1.4.1)", "mypy (==v1.5.1)", "pre-commit (==3.4.0)", "pytest (==6.1.2)", "pytest (==7.4.0)", "pytest-cov (==2.12.1)", "pytest-cov (==4.1.0)", "pytest-mypy-plugins (==1.9.3)", "pytest-mypy-plugins (==3.0.0)", "sphinx-autobuild (==2021.3.14)", "sphinx-rtd-theme (==1.3.0)", "tox (==3.27.1)", "tox (==4.11.0)"]
[[package]]
name = "lxml"
version = "4.9.3"
@@ -3458,6 +3513,22 @@ files = [
certifi = "*"
urllib3 = "*"
[[package]]
name = "mistralai"
version = "0.0.11"
description = ""
optional = true
python-versions = ">=3.8,<4.0"
files = [
{file = "mistralai-0.0.11-py3-none-any.whl", hash = "sha256:fb2a240a3985420c4e7db48eb5077d6d6dbc5e83cac0dd948c20342fb48087ee"},
{file = "mistralai-0.0.11.tar.gz", hash = "sha256:383072715531198305dab829ab3749b64933bbc2549354f3c9ebc43c17b912cf"},
]
[package.dependencies]
httpx = ">=0.25.2,<0.26.0"
orjson = ">=3.9.10,<4.0.0"
pydantic = ">=2.5.2,<3.0.0"
[[package]]
name = "mock"
version = "5.1.0"
@@ -4155,9 +4226,9 @@ files = [
[package.dependencies]
numpy = [
{version = ">=1.21.0", markers = "python_version == \"3.9\" and platform_system == \"Darwin\" and platform_machine == \"arm64\""},
{version = ">=1.19.3", markers = "platform_system == \"Linux\" and platform_machine == \"aarch64\" and python_version >= \"3.8\" and python_version < \"3.10\" or python_version > \"3.9\" and python_version < \"3.10\" or python_version >= \"3.9\" and platform_system != \"Darwin\" and python_version < \"3.10\" or python_version >= \"3.9\" and platform_machine != \"arm64\" and python_version < \"3.10\""},
{version = ">=1.21.4", markers = "python_version >= \"3.10\" and platform_system == \"Darwin\" and python_version < \"3.11\""},
{version = ">=1.21.2", markers = "platform_system != \"Darwin\" and python_version >= \"3.10\" and python_version < \"3.11\""},
{version = ">=1.19.3", markers = "platform_system == \"Linux\" and platform_machine == \"aarch64\" and python_version >= \"3.8\" and python_version < \"3.10\" or python_version > \"3.9\" and python_version < \"3.10\" or python_version >= \"3.9\" and platform_system != \"Darwin\" and python_version < \"3.10\" or python_version >= \"3.9\" and platform_machine != \"arm64\" and python_version < \"3.10\""},
{version = ">=1.23.5", markers = "python_version >= \"3.11\""},
]
@@ -4294,6 +4365,65 @@ files = [
{file = "opentelemetry_semantic_conventions-0.42b0.tar.gz", hash = "sha256:44ae67a0a3252a05072877857e5cc1242c98d4cf12870159f1a94bec800d38ec"},
]
[[package]]
name = "orjson"
version = "3.9.12"
description = "Fast, correct Python JSON library supporting dataclasses, datetimes, and numpy"
optional = true
python-versions = ">=3.8"
files = [
{file = "orjson-3.9.12-cp310-cp310-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:6b4e2bed7d00753c438e83b613923afdd067564ff7ed696bfe3a7b073a236e07"},
{file = "orjson-3.9.12-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:bd1b8ec63f0bf54a50b498eedeccdca23bd7b658f81c524d18e410c203189365"},
{file = "orjson-3.9.12-cp310-cp310-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:ab8add018a53665042a5ae68200f1ad14c7953fa12110d12d41166f111724656"},
{file = "orjson-3.9.12-cp310-cp310-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:12756a108875526b76e505afe6d6ba34960ac6b8c5ec2f35faf73ef161e97e07"},
{file = "orjson-3.9.12-cp310-cp310-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:890e7519c0c70296253660455f77e3a194554a3c45e42aa193cdebc76a02d82b"},
{file = "orjson-3.9.12-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:d664880d7f016efbae97c725b243b33c2cbb4851ddc77f683fd1eec4a7894146"},
{file = "orjson-3.9.12-cp310-cp310-musllinux_1_1_aarch64.whl", hash = "sha256:cfdaede0fa5b500314ec7b1249c7e30e871504a57004acd116be6acdda3b8ab3"},
{file = "orjson-3.9.12-cp310-cp310-musllinux_1_1_x86_64.whl", hash = "sha256:6492ff5953011e1ba9ed1bf086835fd574bd0a3cbe252db8e15ed72a30479081"},
{file = "orjson-3.9.12-cp310-none-win32.whl", hash = "sha256:29bf08e2eadb2c480fdc2e2daae58f2f013dff5d3b506edd1e02963b9ce9f8a9"},
{file = "orjson-3.9.12-cp310-none-win_amd64.whl", hash = "sha256:0fc156fba60d6b50743337ba09f052d8afc8b64595112996d22f5fce01ab57da"},
{file = "orjson-3.9.12-cp311-cp311-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:2849f88a0a12b8d94579b67486cbd8f3a49e36a4cb3d3f0ab352c596078c730c"},
{file = "orjson-3.9.12-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:3186b18754befa660b31c649a108a915493ea69b4fc33f624ed854ad3563ac65"},
{file = "orjson-3.9.12-cp311-cp311-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:cbbf313c9fb9d4f6cf9c22ced4b6682230457741daeb3d7060c5d06c2e73884a"},
{file = "orjson-3.9.12-cp311-cp311-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:99e8cd005b3926c3db9b63d264bd05e1bf4451787cc79a048f27f5190a9a0311"},
{file = "orjson-3.9.12-cp311-cp311-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:59feb148392d9155f3bfed0a2a3209268e000c2c3c834fb8fe1a6af9392efcbf"},
{file = "orjson-3.9.12-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:a4ae815a172a1f073b05b9e04273e3b23e608a0858c4e76f606d2d75fcabde0c"},
{file = "orjson-3.9.12-cp311-cp311-musllinux_1_1_aarch64.whl", hash = "sha256:ed398f9a9d5a1bf55b6e362ffc80ac846af2122d14a8243a1e6510a4eabcb71e"},
{file = "orjson-3.9.12-cp311-cp311-musllinux_1_1_x86_64.whl", hash = "sha256:d3cfb76600c5a1e6be91326b8f3b83035a370e727854a96d801c1ea08b708073"},
{file = "orjson-3.9.12-cp311-none-win32.whl", hash = "sha256:a2b6f5252c92bcab3b742ddb3ac195c0fa74bed4319acd74f5d54d79ef4715dc"},
{file = "orjson-3.9.12-cp311-none-win_amd64.whl", hash = "sha256:c95488e4aa1d078ff5776b58f66bd29d628fa59adcb2047f4efd3ecb2bd41a71"},
{file = "orjson-3.9.12-cp312-cp312-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:d6ce2062c4af43b92b0221ed4f445632c6bf4213f8a7da5396a122931377acd9"},
{file = "orjson-3.9.12-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:950951799967558c214cd6cceb7ceceed6f81d2c3c4135ee4a2c9c69f58aa225"},
{file = "orjson-3.9.12-cp312-cp312-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:2dfaf71499d6fd4153f5c86eebb68e3ec1bf95851b030a4b55c7637a37bbdee4"},
{file = "orjson-3.9.12-cp312-cp312-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:659a8d7279e46c97661839035a1a218b61957316bf0202674e944ac5cfe7ed83"},
{file = "orjson-3.9.12-cp312-cp312-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:af17fa87bccad0b7f6fd8ac8f9cbc9ee656b4552783b10b97a071337616db3e4"},
{file = "orjson-3.9.12-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:cd52dec9eddf4c8c74392f3fd52fa137b5f2e2bed1d9ae958d879de5f7d7cded"},
{file = "orjson-3.9.12-cp312-cp312-musllinux_1_1_aarch64.whl", hash = "sha256:640e2b5d8e36b970202cfd0799d11a9a4ab46cf9212332cd642101ec952df7c8"},
{file = "orjson-3.9.12-cp312-cp312-musllinux_1_1_x86_64.whl", hash = "sha256:daa438bd8024e03bcea2c5a92cd719a663a58e223fba967296b6ab9992259dbf"},
{file = "orjson-3.9.12-cp312-none-win_amd64.whl", hash = "sha256:1bb8f657c39ecdb924d02e809f992c9aafeb1ad70127d53fb573a6a6ab59d549"},
{file = "orjson-3.9.12-cp38-cp38-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:f4098c7674901402c86ba6045a551a2ee345f9f7ed54eeffc7d86d155c8427e5"},
{file = "orjson-3.9.12-cp38-cp38-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:5586a533998267458fad3a457d6f3cdbddbcce696c916599fa8e2a10a89b24d3"},
{file = "orjson-3.9.12-cp38-cp38-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:54071b7398cd3f90e4bb61df46705ee96cb5e33e53fc0b2f47dbd9b000e238e1"},
{file = "orjson-3.9.12-cp38-cp38-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:67426651faa671b40443ea6f03065f9c8e22272b62fa23238b3efdacd301df31"},
{file = "orjson-3.9.12-cp38-cp38-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:4a0cd56e8ee56b203abae7d482ac0d233dbfb436bb2e2d5cbcb539fe1200a312"},
{file = "orjson-3.9.12-cp38-cp38-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:a84a0c3d4841a42e2571b1c1ead20a83e2792644c5827a606c50fc8af7ca4bee"},
{file = "orjson-3.9.12-cp38-cp38-musllinux_1_1_aarch64.whl", hash = "sha256:09d60450cda3fa6c8ed17770c3a88473a16460cd0ff2ba74ef0df663b6fd3bb8"},
{file = "orjson-3.9.12-cp38-cp38-musllinux_1_1_x86_64.whl", hash = "sha256:bc82a4db9934a78ade211cf2e07161e4f068a461c1796465d10069cb50b32a80"},
{file = "orjson-3.9.12-cp38-none-win32.whl", hash = "sha256:61563d5d3b0019804d782137a4f32c72dc44c84e7d078b89d2d2a1adbaa47b52"},
{file = "orjson-3.9.12-cp38-none-win_amd64.whl", hash = "sha256:410f24309fbbaa2fab776e3212a81b96a1ec6037259359a32ea79fbccfcf76aa"},
{file = "orjson-3.9.12-cp39-cp39-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:e773f251258dd82795fd5daeac081d00b97bacf1548e44e71245543374874bcf"},
{file = "orjson-3.9.12-cp39-cp39-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:b159baecfda51c840a619948c25817d37733a4d9877fea96590ef8606468b362"},
{file = "orjson-3.9.12-cp39-cp39-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:975e72e81a249174840d5a8df977d067b0183ef1560a32998be340f7e195c730"},
{file = "orjson-3.9.12-cp39-cp39-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:06e42e899dde61eb1851a9fad7f1a21b8e4be063438399b63c07839b57668f6c"},
{file = "orjson-3.9.12-cp39-cp39-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:5c157e999e5694475a5515942aebeed6e43f7a1ed52267c1c93dcfde7d78d421"},
{file = "orjson-3.9.12-cp39-cp39-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:dde1bc7c035f2d03aa49dc8642d9c6c9b1a81f2470e02055e76ed8853cfae0c3"},
{file = "orjson-3.9.12-cp39-cp39-musllinux_1_1_aarch64.whl", hash = "sha256:b0e9d73cdbdad76a53a48f563447e0e1ce34bcecef4614eb4b146383e6e7d8c9"},
{file = "orjson-3.9.12-cp39-cp39-musllinux_1_1_x86_64.whl", hash = "sha256:96e44b21fe407b8ed48afbb3721f3c8c8ce17e345fbe232bd4651ace7317782d"},
{file = "orjson-3.9.12-cp39-none-win32.whl", hash = "sha256:cbd0f3555205bf2a60f8812133f2452d498dbefa14423ba90fe89f32276f7abf"},
{file = "orjson-3.9.12-cp39-none-win_amd64.whl", hash = "sha256:03ea7ee7e992532c2f4a06edd7ee1553f0644790553a118e003e3c405add41fa"},
{file = "orjson-3.9.12.tar.gz", hash = "sha256:da908d23a3b3243632b523344403b128722a5f45e278a8343c2bb67538dff0e4"},
]
[[package]]
name = "overrides"
version = "7.4.0"
@@ -4513,25 +4643,20 @@ tests = ["check-manifest", "coverage", "defusedxml", "markdown2", "olefile", "pa
[[package]]
name = "pinecone-client"
version = "2.2.4"
version = "3.0.1"
description = "Pinecone client and SDK"
optional = true
python-versions = ">=3.8"
python-versions = ">=3.8,<3.13"
files = [
{file = "pinecone-client-2.2.4.tar.gz", hash = "sha256:2c1cc1d6648b2be66e944db2ffa59166a37b9164d1135ad525d9cd8b1e298168"},
{file = "pinecone_client-2.2.4-py3-none-any.whl", hash = "sha256:5bf496c01c2f82f4e5c2dc977cc5062ecd7168b8ed90743b09afcc8c7eb242ec"},
{file = "pinecone_client-3.0.1-py3-none-any.whl", hash = "sha256:c9bb21c23a9088c6198c839be5538ed3f733d152d5fbeaafcc020c1b70b62c2d"},
{file = "pinecone_client-3.0.1.tar.gz", hash = "sha256:626a0055852c88f1462fc2e132f21d2b078f9a0a74c70b17fe07df3081c6615f"},
]
[package.dependencies]
dnspython = ">=2.0.0"
loguru = ">=0.5.0"
numpy = ">=1.22.0"
python-dateutil = ">=2.5.3"
pyyaml = ">=5.4"
requests = ">=2.19.0"
certifi = ">=2019.11.17"
tqdm = ">=4.64.1"
typing-extensions = ">=3.7.4"
urllib3 = ">=1.21.1"
urllib3 = ">=1.26.0"
[package.extras]
grpc = ["googleapis-common-protos (>=1.53.0)", "grpc-gateway-protoc-gen-openapiv2 (==0.1.0)", "grpcio (>=1.44.0)", "lz4 (>=3.1.3)", "protobuf (>=3.20.0,<3.21.0)"]
@@ -5976,6 +6101,23 @@ files = [
{file = "ruff-0.1.11.tar.gz", hash = "sha256:f9d4d88cb6eeb4dfe20f9f0519bd2eaba8119bde87c3d5065c541dbae2b5a2cb"},
]
[[package]]
name = "s3transfer"
version = "0.10.0"
description = "An Amazon S3 Transfer Manager"
optional = true
python-versions = ">= 3.8"
files = [
{file = "s3transfer-0.10.0-py3-none-any.whl", hash = "sha256:3cdb40f5cfa6966e812209d0994f2a4709b561c88e90cf00c2696d2df4e56b2e"},
{file = "s3transfer-0.10.0.tar.gz", hash = "sha256:d0c8bbf672d5eebbe4e57945e23b972d963f07d82f661cabf678a5c88831595b"},
]
[package.dependencies]
botocore = ">=1.33.2,<2.0a.0"
[package.extras]
crt = ["botocore[crt] (>=1.33.2,<2.0a.0)"]
[[package]]
name = "safetensors"
version = "0.4.0"
@@ -7835,20 +7977,6 @@ files = [
[package.extras]
test = ["pytest (>=6.0.0)", "setuptools (>=65)"]
[[package]]
name = "win32-setctime"
version = "1.1.0"
description = "A small Python utility to set file creation time on Windows"
optional = true
python-versions = ">=3.5"
files = [
{file = "win32_setctime-1.1.0-py3-none-any.whl", hash = "sha256:231db239e959c2fe7eb1d7dc129f11172354f98361c4fa2d6d2d7e278baa8aad"},
{file = "win32_setctime-1.1.0.tar.gz", hash = "sha256:15cf5750465118d6929ae4de4eb46e8edae9a5634350c01ba582df868e932cb2"},
]
[package.extras]
dev = ["black (>=19.3b0)", "pytest (>=4.6.2)"]
[[package]]
name = "wrapt"
version = "1.15.0"
@@ -8098,6 +8226,7 @@ docs = ["furo", "jaraco.packaging (>=9.3)", "jaraco.tidelift (>=1.4)", "rst.link
testing = ["big-O", "jaraco.functools", "jaraco.itertools", "more-itertools", "pytest (>=6)", "pytest-black (>=0.3.7)", "pytest-checkdocs (>=2.4)", "pytest-cov", "pytest-enabler (>=2.2)", "pytest-ignore-flaky", "pytest-mypy (>=0.9.1)", "pytest-ruff"]
[extras]
aws-bedrock = ["boto3"]
cohere = ["cohere"]
dataloaders = ["docx2txt", "duckduckgo-search", "pytube", "sentence-transformers", "unstructured", "youtube-transcript-api"]
discord = ["discord"]
@@ -8110,6 +8239,7 @@ googledrive = ["google-api-python-client", "google-auth-httplib2", "google-auth-
huggingface-hub = ["huggingface_hub"]
llama2 = ["replicate"]
milvus = ["pymilvus"]
mistralai = ["langchain-mistralai"]
modal = ["modal"]
mysql = ["mysql-connector-python"]
opensearch = ["opensearch-py"]
@@ -8130,4 +8260,4 @@ youtube = ["youtube-transcript-api", "yt_dlp"]
[metadata]
lock-version = "2.0"
python-versions = ">=3.9,<3.12"
content-hash = "02bd85e14374a9dc9b59523b8fb4baea7068251976ba7f87722cac94a9974ccc"
content-hash = "a16addd3362ae70c79b15677c6815f708677f11f636093f4e1f5084ba44b5a36"
+7 -3
View File
@@ -1,7 +1,7 @@
[tool.poetry]
name = "embedchain"
version = "0.1.63"
description = "Data platform for LLMs - Load, index, retrieve and sync any unstructured data"
version = "0.1.77"
description = "Simplest open source retrieval(RAG) framework"
authors = [
"Taranjeet Singh <taranjeet@embedchain.ai>",
"Deshraj Yadav <deshraj@embedchain.ai>",
@@ -123,7 +123,7 @@ cohere = { version = "^4.27", optional = true }
together = { version = "^0.2.8", optional = true }
weaviate-client = { version = "^3.24.1", optional = true }
docx2txt = { version = "^0.8", optional = true }
pinecone-client = { version = "^2.2.4", optional = true }
pinecone-client = { version = "^3.0.0", optional = true }
qdrant-client = { version = "1.6.3", optional = true }
unstructured = {extras = ["local-inference", "all-docs"], version = "^0.10.18", optional = true}
huggingface_hub = { version = "^0.17.3", optional = true }
@@ -149,6 +149,8 @@ google-auth-oauthlib = { version = "^1.2.0", optional = true }
google-auth = { version = "^2.25.2", optional = true }
google-auth-httplib2 = { version = "^0.2.0", optional = true }
google-api-core = { version = "^2.15.0", optional = true }
boto3 = { version = "^1.34.20", optional = true }
langchain-mistralai = { version = "^0.0.3", optional = true }
[tool.poetry.group.dev.dependencies]
black = "^23.3.0"
@@ -214,6 +216,8 @@ rss_feed = [
google = ["google-generativeai"]
modal = ["modal"]
dropbox = ["dropbox"]
aws_bedrock = ["boto3"]
mistralai = ["langchain-mistralai"]
[tool.poetry.group.docs.dependencies]
@@ -0,0 +1,223 @@
import numpy as np
import pytest
from embedchain.config.evaluation.base import AnswerRelevanceConfig
from embedchain.evaluation.metrics import AnswerRelevance
from embedchain.utils.evaluation import EvalData, EvalMetric
@pytest.fixture
def mock_data():
return [
EvalData(
contexts=[
"This is a test context 1.",
],
question="This is a test question 1.",
answer="This is a test answer 1.",
),
EvalData(
contexts=[
"This is a test context 2-1.",
"This is a test context 2-2.",
],
question="This is a test question 2.",
answer="This is a test answer 2.",
),
]
@pytest.fixture
def mock_answer_relevance_metric(monkeypatch):
monkeypatch.setenv("OPENAI_API_KEY", "test_api_key")
metric = AnswerRelevance()
return metric
def test_answer_relevance_init(monkeypatch):
monkeypatch.setenv("OPENAI_API_KEY", "test_api_key")
metric = AnswerRelevance()
assert metric.name == EvalMetric.ANSWER_RELEVANCY.value
assert metric.config.model == "gpt-4"
assert metric.config.embedder == "text-embedding-ada-002"
assert metric.config.api_key is None
assert metric.config.num_gen_questions == 1
monkeypatch.delenv("OPENAI_API_KEY")
def test_answer_relevance_init_with_config():
metric = AnswerRelevance(config=AnswerRelevanceConfig(api_key="test_api_key"))
assert metric.name == EvalMetric.ANSWER_RELEVANCY.value
assert metric.config.model == "gpt-4"
assert metric.config.embedder == "text-embedding-ada-002"
assert metric.config.api_key == "test_api_key"
assert metric.config.num_gen_questions == 1
def test_answer_relevance_init_without_api_key(monkeypatch):
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
with pytest.raises(ValueError):
AnswerRelevance()
def test_generate_prompt(mock_answer_relevance_metric, mock_data):
prompt = mock_answer_relevance_metric._generate_prompt(mock_data[0])
assert "This is a test answer 1." in prompt
prompt = mock_answer_relevance_metric._generate_prompt(mock_data[1])
assert "This is a test answer 2." in prompt
def test_generate_questions(mock_answer_relevance_metric, mock_data, monkeypatch):
monkeypatch.setattr(
mock_answer_relevance_metric.client.chat.completions,
"create",
lambda model, messages: type(
"obj",
(object,),
{
"choices": [
type(
"obj",
(object,),
{"message": type("obj", (object,), {"content": "This is a test question response.\n"})},
)
]
},
)(),
)
prompt = mock_answer_relevance_metric._generate_prompt(mock_data[0])
questions = mock_answer_relevance_metric._generate_questions(prompt)
assert len(questions) == 1
monkeypatch.setattr(
mock_answer_relevance_metric.client.chat.completions,
"create",
lambda model, messages: type(
"obj",
(object,),
{
"choices": [
type("obj", (object,), {"message": type("obj", (object,), {"content": "question 1?\nquestion2?"})})
]
},
)(),
)
prompt = mock_answer_relevance_metric._generate_prompt(mock_data[1])
questions = mock_answer_relevance_metric._generate_questions(prompt)
assert len(questions) == 2
def test_generate_embedding(mock_answer_relevance_metric, mock_data, monkeypatch):
monkeypatch.setattr(
mock_answer_relevance_metric.client.embeddings,
"create",
lambda input, model: type("obj", (object,), {"data": [type("obj", (object,), {"embedding": [1, 2, 3]})]})(),
)
embedding = mock_answer_relevance_metric._generate_embedding("This is a test question.")
assert len(embedding) == 3
def test_compute_similarity(mock_answer_relevance_metric, mock_data):
original = np.array([1, 2, 3])
generated = np.array([[1, 2, 3], [1, 2, 3]])
similarity = mock_answer_relevance_metric._compute_similarity(original, generated)
assert len(similarity) == 2
assert similarity[0] == 1.0
assert similarity[1] == 1.0
def test_compute_score(mock_answer_relevance_metric, mock_data, monkeypatch):
monkeypatch.setattr(
mock_answer_relevance_metric.client.chat.completions,
"create",
lambda model, messages: type(
"obj",
(object,),
{
"choices": [
type(
"obj",
(object,),
{"message": type("obj", (object,), {"content": "This is a test question response.\n"})},
)
]
},
)(),
)
monkeypatch.setattr(
mock_answer_relevance_metric.client.embeddings,
"create",
lambda input, model: type("obj", (object,), {"data": [type("obj", (object,), {"embedding": [1, 2, 3]})]})(),
)
score = mock_answer_relevance_metric._compute_score(mock_data[0])
assert score == 1.0
monkeypatch.setattr(
mock_answer_relevance_metric.client.chat.completions,
"create",
lambda model, messages: type(
"obj",
(object,),
{
"choices": [
type("obj", (object,), {"message": type("obj", (object,), {"content": "question 1?\nquestion2?"})})
]
},
)(),
)
monkeypatch.setattr(
mock_answer_relevance_metric.client.embeddings,
"create",
lambda input, model: type("obj", (object,), {"data": [type("obj", (object,), {"embedding": [1, 2, 3]})]})(),
)
score = mock_answer_relevance_metric._compute_score(mock_data[1])
assert score == 1.0
def test_evaluate(mock_answer_relevance_metric, mock_data, monkeypatch):
monkeypatch.setattr(
mock_answer_relevance_metric.client.chat.completions,
"create",
lambda model, messages: type(
"obj",
(object,),
{
"choices": [
type(
"obj",
(object,),
{"message": type("obj", (object,), {"content": "This is a test question response.\n"})},
)
]
},
)(),
)
monkeypatch.setattr(
mock_answer_relevance_metric.client.embeddings,
"create",
lambda input, model: type("obj", (object,), {"data": [type("obj", (object,), {"embedding": [1, 2, 3]})]})(),
)
score = mock_answer_relevance_metric.evaluate(mock_data)
assert score == 1.0
monkeypatch.setattr(
mock_answer_relevance_metric.client.chat.completions,
"create",
lambda model, messages: type(
"obj",
(object,),
{
"choices": [
type("obj", (object,), {"message": type("obj", (object,), {"content": "question 1?\nquestion2?"})})
]
},
)(),
)
monkeypatch.setattr(
mock_answer_relevance_metric.client.embeddings,
"create",
lambda input, model: type("obj", (object,), {"data": [type("obj", (object,), {"embedding": [1, 2, 3]})]})(),
)
score = mock_answer_relevance_metric.evaluate(mock_data)
assert score == 1.0
@@ -0,0 +1,100 @@
import pytest
from embedchain.config.evaluation.base import ContextRelevanceConfig
from embedchain.evaluation.metrics import ContextRelevance
from embedchain.utils.evaluation import EvalData, EvalMetric
@pytest.fixture
def mock_data():
return [
EvalData(
contexts=[
"This is a test context 1.",
],
question="This is a test question 1.",
answer="This is a test answer 1.",
),
EvalData(
contexts=[
"This is a test context 2-1.",
"This is a test context 2-2.",
],
question="This is a test question 2.",
answer="This is a test answer 2.",
),
]
@pytest.fixture
def mock_context_relevance_metric(monkeypatch):
monkeypatch.setenv("OPENAI_API_KEY", "test_api_key")
metric = ContextRelevance()
return metric
def test_context_relevance_init(monkeypatch):
monkeypatch.setenv("OPENAI_API_KEY", "test_api_key")
metric = ContextRelevance()
assert metric.name == EvalMetric.CONTEXT_RELEVANCY.value
assert metric.config.model == "gpt-4"
assert metric.config.api_key is None
assert metric.config.language == "en"
monkeypatch.delenv("OPENAI_API_KEY")
def test_context_relevance_init_with_config():
metric = ContextRelevance(config=ContextRelevanceConfig(api_key="test_api_key"))
assert metric.name == EvalMetric.CONTEXT_RELEVANCY.value
assert metric.config.model == "gpt-4"
assert metric.config.api_key == "test_api_key"
assert metric.config.language == "en"
def test_context_relevance_init_without_api_key(monkeypatch):
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
with pytest.raises(ValueError):
ContextRelevance()
def test_sentence_segmenter(mock_context_relevance_metric):
text = "This is a test sentence. This is another sentence."
assert mock_context_relevance_metric._sentence_segmenter(text) == [
"This is a test sentence. ",
"This is another sentence.",
]
def test_compute_score(mock_context_relevance_metric, mock_data, monkeypatch):
monkeypatch.setattr(
mock_context_relevance_metric.client.chat.completions,
"create",
lambda model, messages: type(
"obj",
(object,),
{
"choices": [
type("obj", (object,), {"message": type("obj", (object,), {"content": "This is a test reponse."})})
]
},
)(),
)
assert mock_context_relevance_metric._compute_score(mock_data[0]) == 1.0
assert mock_context_relevance_metric._compute_score(mock_data[1]) == 0.5
def test_evaluate(mock_context_relevance_metric, mock_data, monkeypatch):
monkeypatch.setattr(
mock_context_relevance_metric.client.chat.completions,
"create",
lambda model, messages: type(
"obj",
(object,),
{
"choices": [
type("obj", (object,), {"message": type("obj", (object,), {"content": "This is a test reponse."})})
]
},
)(),
)
assert mock_context_relevance_metric.evaluate(mock_data) == 0.75
@@ -0,0 +1,152 @@
import numpy as np
import pytest
from embedchain.config.evaluation.base import GroundednessConfig
from embedchain.evaluation.metrics import Groundedness
from embedchain.utils.evaluation import EvalData, EvalMetric
@pytest.fixture
def mock_data():
return [
EvalData(
contexts=[
"This is a test context 1.",
],
question="This is a test question 1.",
answer="This is a test answer 1.",
),
EvalData(
contexts=[
"This is a test context 2-1.",
"This is a test context 2-2.",
],
question="This is a test question 2.",
answer="This is a test answer 2.",
),
]
@pytest.fixture
def mock_groundedness_metric(monkeypatch):
monkeypatch.setenv("OPENAI_API_KEY", "test_api_key")
metric = Groundedness()
return metric
def test_groundedness_init(monkeypatch):
monkeypatch.setenv("OPENAI_API_KEY", "test_api_key")
metric = Groundedness()
assert metric.name == EvalMetric.GROUNDEDNESS.value
assert metric.config.model == "gpt-4"
assert metric.config.api_key is None
monkeypatch.delenv("OPENAI_API_KEY")
def test_groundedness_init_with_config():
metric = Groundedness(config=GroundednessConfig(api_key="test_api_key"))
assert metric.name == EvalMetric.GROUNDEDNESS.value
assert metric.config.model == "gpt-4"
assert metric.config.api_key == "test_api_key"
def test_groundedness_init_without_api_key(monkeypatch):
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
with pytest.raises(ValueError):
Groundedness()
def test_generate_answer_claim_prompt(mock_groundedness_metric, mock_data):
prompt = mock_groundedness_metric._generate_answer_claim_prompt(data=mock_data[0])
assert "This is a test question 1." in prompt
assert "This is a test answer 1." in prompt
def test_get_claim_statements(mock_groundedness_metric, mock_data, monkeypatch):
monkeypatch.setattr(
mock_groundedness_metric.client.chat.completions,
"create",
lambda *args, **kwargs: type(
"obj",
(object,),
{
"choices": [
type(
"obj",
(object,),
{
"message": type(
"obj",
(object,),
{
"content": """This is a test answer 1.
This is a test answer 2.
This is a test answer 3."""
},
)
},
)
]
},
)(),
)
prompt = mock_groundedness_metric._generate_answer_claim_prompt(data=mock_data[0])
claim_statements = mock_groundedness_metric._get_claim_statements(prompt=prompt)
assert len(claim_statements) == 3
assert "This is a test answer 1." in claim_statements
def test_generate_claim_inference_prompt(mock_groundedness_metric, mock_data):
prompt = mock_groundedness_metric._generate_answer_claim_prompt(data=mock_data[0])
claim_statements = [
"This is a test claim 1.",
"This is a test claim 2.",
]
prompt = mock_groundedness_metric._generate_claim_inference_prompt(
data=mock_data[0], claim_statements=claim_statements
)
assert "This is a test context 1." in prompt
assert "This is a test claim 1." in prompt
def test_get_claim_verdict_scores(mock_groundedness_metric, mock_data, monkeypatch):
monkeypatch.setattr(
mock_groundedness_metric.client.chat.completions,
"create",
lambda *args, **kwargs: type(
"obj",
(object,),
{"choices": [type("obj", (object,), {"message": type("obj", (object,), {"content": "1\n0\n-1"})})]},
)(),
)
prompt = mock_groundedness_metric._generate_answer_claim_prompt(data=mock_data[0])
claim_statements = mock_groundedness_metric._get_claim_statements(prompt=prompt)
prompt = mock_groundedness_metric._generate_claim_inference_prompt(
data=mock_data[0], claim_statements=claim_statements
)
claim_verdict_scores = mock_groundedness_metric._get_claim_verdict_scores(prompt=prompt)
assert len(claim_verdict_scores) == 3
assert claim_verdict_scores[0] == 1
assert claim_verdict_scores[1] == 0
def test_compute_score(mock_groundedness_metric, mock_data, monkeypatch):
monkeypatch.setattr(
mock_groundedness_metric,
"_get_claim_statements",
lambda *args, **kwargs: np.array(
[
"This is a test claim 1.",
"This is a test claim 2.",
]
),
)
monkeypatch.setattr(mock_groundedness_metric, "_get_claim_verdict_scores", lambda *args, **kwargs: np.array([1, 0]))
score = mock_groundedness_metric._compute_score(data=mock_data[0])
assert score == 0.5
def test_evaluate(mock_groundedness_metric, mock_data, monkeypatch):
monkeypatch.setattr(mock_groundedness_metric, "_compute_score", lambda *args, **kwargs: 0.5)
score = mock_groundedness_metric.evaluate(dataset=mock_data)
assert score == 0.5
+56
View File
@@ -0,0 +1,56 @@
import pytest
from langchain.callbacks.streaming_stdout import StreamingStdOutCallbackHandler
from embedchain.config import BaseLlmConfig
from embedchain.llm.aws_bedrock import AWSBedrockLlm
@pytest.fixture
def config(monkeypatch):
monkeypatch.setenv("AWS_ACCESS_KEY_ID", "test_access_key_id")
monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test_secret_access_key")
monkeypatch.setenv("OPENAI_API_KEY", "test_api_key")
config = BaseLlmConfig(
model="amazon.titan-text-express-v1",
model_kwargs={
"temperature": 0.5,
"topP": 1,
"maxTokenCount": 1000,
},
)
yield config
monkeypatch.delenv("AWS_ACCESS_KEY_ID")
monkeypatch.delenv("AWS_SECRET_ACCESS_KEY")
monkeypatch.delenv("OPENAI_API_KEY")
def test_get_llm_model_answer(config, mocker):
mocked_get_answer = mocker.patch("embedchain.llm.aws_bedrock.AWSBedrockLlm._get_answer", return_value="Test answer")
llm = AWSBedrockLlm(config)
answer = llm.get_llm_model_answer("Test query")
assert answer == "Test answer"
mocked_get_answer.assert_called_once_with("Test query", config)
def test_get_llm_model_answer_empty_prompt(config, mocker):
mocked_get_answer = mocker.patch("embedchain.llm.aws_bedrock.AWSBedrockLlm._get_answer", return_value="Test answer")
llm = AWSBedrockLlm(config)
answer = llm.get_llm_model_answer("")
assert answer == "Test answer"
mocked_get_answer.assert_called_once_with("", config)
def test_get_llm_model_answer_with_streaming(config, mocker):
config.stream = True
mocked_bedrock_chat = mocker.patch("embedchain.llm.aws_bedrock.Bedrock")
llm = AWSBedrockLlm(config)
llm.get_llm_model_answer("Test query")
mocked_bedrock_chat.assert_called_once()
callbacks = [callback[1]["callbacks"] for callback in mocked_bedrock_chat.call_args_list]
assert any(isinstance(callback[0], StreamingStdOutCallbackHandler) for callback in callbacks)
+60
View File
@@ -0,0 +1,60 @@
import pytest
from embedchain.config import BaseLlmConfig
from embedchain.llm.mistralai import MistralAILlm
@pytest.fixture
def mistralai_llm_config(monkeypatch):
monkeypatch.setenv("MISTRAL_API_KEY", "fake_api_key")
yield BaseLlmConfig(model="mistral-tiny", max_tokens=100, temperature=0.7, top_p=0.5, stream=False)
monkeypatch.delenv("MISTRAL_API_KEY", raising=False)
def test_mistralai_llm_init_missing_api_key(monkeypatch):
monkeypatch.delenv("MISTRAL_API_KEY", raising=False)
with pytest.raises(ValueError, match="Please set the MISTRAL_API_KEY environment variable."):
MistralAILlm()
def test_mistralai_llm_init(monkeypatch):
monkeypatch.setenv("MISTRAL_API_KEY", "fake_api_key")
llm = MistralAILlm()
assert llm is not None
def test_get_llm_model_answer(monkeypatch, mistralai_llm_config):
def mock_get_answer(prompt, config):
return "Generated Text"
monkeypatch.setattr(MistralAILlm, "_get_answer", mock_get_answer)
llm = MistralAILlm(config=mistralai_llm_config)
result = llm.get_llm_model_answer("test prompt")
assert result == "Generated Text"
def test_get_llm_model_answer_with_system_prompt(monkeypatch, mistralai_llm_config):
mistralai_llm_config.system_prompt = "Test system prompt"
monkeypatch.setattr(MistralAILlm, "_get_answer", lambda prompt, config: "Generated Text")
llm = MistralAILlm(config=mistralai_llm_config)
result = llm.get_llm_model_answer("test prompt")
assert result == "Generated Text"
def test_get_llm_model_answer_empty_prompt(monkeypatch, mistralai_llm_config):
monkeypatch.setattr(MistralAILlm, "_get_answer", lambda prompt, config: "Generated Text")
llm = MistralAILlm(config=mistralai_llm_config)
result = llm.get_llm_model_answer("")
assert result == "Generated Text"
def test_get_llm_model_answer_without_system_prompt(monkeypatch, mistralai_llm_config):
mistralai_llm_config.system_prompt = None
monkeypatch.setattr(MistralAILlm, "_get_answer", lambda prompt, config: "Generated Text")
llm = MistralAILlm(config=mistralai_llm_config)
result = llm.get_llm_model_answer("test prompt")
assert result == "Generated Text"
+1 -1
View File
@@ -35,7 +35,7 @@ class TestFactories:
("gpt4all", {}, embedchain.embedder.gpt4all.GPT4AllEmbedder),
(
"huggingface",
{"model": "sentence-transformers/all-mpnet-base-v2"},
{"model": "sentence-transformers/all-mpnet-base-v2", "vector_dimension": 768},
embedchain.embedder.huggingface.HuggingFaceEmbedder,
),
("vertexai", {"model": "textembedding-gecko"}, embedchain.embedder.vertexai.VertexAIEmbedder),
+2 -3
View File
@@ -28,14 +28,13 @@ class TestEsDB(unittest.TestCase):
# Assert that the Elasticsearch client is stored in the ElasticsearchDB class.
self.assertEqual(self.db.client, mock_client.return_value)
# Create some dummy data.
embeddings = [[1, 2, 3], [4, 5, 6]]
# Create some dummy data
documents = ["This is a document.", "This is another document."]
metadatas = [{"url": "url_1", "doc_id": "doc_id_1"}, {"url": "url_2", "doc_id": "doc_id_2"}]
ids = ["doc_1", "doc_2"]
# Add the data to the database.
self.db.add(embeddings, documents, metadatas, ids)
self.db.add(documents, metadatas, ids)
search_response = {
"hits": {

Some files were not shown because too many files have changed in this diff Show More