Compare commits
43 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 2985b667b0 | |||
| 31bb0e7f0f | |||
| 8f28264aec | |||
| ec4fb11aa5 | |||
| b210723de1 | |||
| 433f99dd78 | |||
| e75c05112e | |||
| d2a5b50ff8 | |||
| 120690afd4 | |||
| 344dbeee42 | |||
| 3fe3b0320a | |||
| 75896b647f | |||
| 446d0975aa | |||
| b7d365119c | |||
| 2d9fbd4e49 | |||
| 22e14b5e65 | |||
| 1a654beea4 | |||
| f50f8a444a | |||
| 3cc3a0058d | |||
| ae473b5e3c | |||
| efb7e31565 | |||
| 069d265338 | |||
| 751a3a4bd1 | |||
| cb0499407e | |||
| 9afc6878c8 | |||
| 0b5b12575a | |||
| d79d30bf0c | |||
| 59600e2a5b | |||
| e572b5a3dc | |||
| 5b46daaee4 | |||
| 2784bae772 | |||
| 325e11f0de | |||
| 7444f59e3c | |||
| affe319460 | |||
| 862ff6cca6 | |||
| c020e65a50 | |||
| f582c1fe25 | |||
| 785929c502 | |||
| 68ec6615b1 | |||
| e2cca61cd3 | |||
| 69e83adae0 | |||
| 9e24aee40d | |||
| 3cff5e9898 |
@@ -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">
|
||||
|
||||
@@ -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,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">
|
||||
|
||||
@@ -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)
|
||||
```
|
||||
@@ -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>
|
||||
@@ -0,0 +1,41 @@
|
||||
---
|
||||
title: '📝 evaluate'
|
||||
---
|
||||
|
||||
`evaluate()` method is used to evaluate the performance of a RAG app. You can find the signature below:
|
||||
|
||||
### Parameters
|
||||
|
||||
<ParamField path="question" type="Union[str, list[str]]">
|
||||
A question or a list of questions to evaluate your app on.
|
||||
</ParamField>
|
||||
<ParamField path="metrics" type="Optional[list[Union[BaseMetric, str]]]" optional>
|
||||
The metrics to evaluate your app on. Defaults to all metrics: `["context_relevancy", "answer_relevancy", "groundedness"]`
|
||||
</ParamField>
|
||||
<ParamField path="num_workers" type="int" optional>
|
||||
Specify the number of threads to use for parallel processing.
|
||||
</ParamField>
|
||||
|
||||
### Returns
|
||||
|
||||
<ResponseField name="metrics" type="dict">
|
||||
Returns the metrics you have chosen to evaluate your app on as a dictionary.
|
||||
</ResponseField>
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from embedchain import App
|
||||
|
||||
app = App()
|
||||
|
||||
# add data source
|
||||
app.add("https://www.forbes.com/profile/elon-musk")
|
||||
|
||||
# run evaluation
|
||||
app.evaluate("what is the net worth of Elon Musk?")
|
||||
# {'answer_relevancy': 0.958019958036268, 'context_relevancy': 0.12903225806451613}
|
||||
|
||||
# or
|
||||
# app.evaluate(["what is the net worth of Elon Musk?", "which companies does Elon Musk own?"])
|
||||
```
|
||||
@@ -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>
|
||||
@@ -1,19 +0,0 @@
|
||||
---
|
||||
title: 🗑 delete
|
||||
---
|
||||
|
||||
`delete_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_chat_history()
|
||||
```
|
||||
@@ -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">
|
||||
|
||||
@@ -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)
|
||||
```
|
||||
|
||||
@@ -10,6 +10,16 @@ Ensure your app has the following settings activated:
|
||||
|
||||
- In the Permissions section, enable `files.content.read` and `files.metadata.read`.
|
||||
|
||||
## Usage
|
||||
|
||||
Install the `dropbox` pypi package:
|
||||
|
||||
```bash
|
||||
pip install dropbox
|
||||
```
|
||||
|
||||
Following is an example of how to use the dropbox loader:
|
||||
|
||||
```python
|
||||
import os
|
||||
from embedchain import Pipeline as App
|
||||
|
||||
@@ -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]"
|
||||
```
|
||||
</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
|
||||
|
||||
@@ -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>
|
||||
|
||||
@@ -0,0 +1,275 @@
|
||||
---
|
||||
title: 🔬 Evaluation
|
||||
---
|
||||
|
||||
## Overview
|
||||
|
||||
We provide out-of-the-box evaluation metrics for your RAG application. You can use them to evaluate your RAG applications and compare against different settings of your production RAG application.
|
||||
|
||||
Currently, we provide support for following evaluation metrics:
|
||||
|
||||
<CardGroup cols={3}>
|
||||
<Card title="Context Relevancy" href="#context_relevancy"></Card>
|
||||
<Card title="Answer Relevancy" href="#answer_relevancy"></Card>
|
||||
<Card title="Groundedness" href="#groundedness"></Card>
|
||||
<Card title="Custom Metric" href="#custom_metric"></Card>
|
||||
</CardGroup>
|
||||
|
||||
## Quickstart
|
||||
|
||||
Here is a basic example of running evaluation:
|
||||
|
||||
```python example.py
|
||||
from embedchain import App
|
||||
|
||||
app = App()
|
||||
|
||||
# Add data sources
|
||||
app.add("https://www.forbes.com/profile/elon-musk")
|
||||
|
||||
# Run evaluation
|
||||
app.evaluate(["What is the net worth of Elon Musk?", "How many companies Elon Musk owns?"])
|
||||
# {'answer_relevancy': 0.9987286412340826, 'groundedness': 1.0, 'context_relevancy': 0.3571428571428571}
|
||||
```
|
||||
|
||||
Under the hood, Embedchain does the following:
|
||||
|
||||
1. Runs semantic search in the vector database and fetches context
|
||||
2. LLM call with question, context to fetch the answer
|
||||
3. Run evaluation on following metrics: `context relevancy`, `groundedness`, and `answer relevancy` and return result
|
||||
|
||||
## Advanced Usage
|
||||
|
||||
We use OpenAI's `gpt-4` model as default LLM model for automatic evaluation. Hence, we require you to set `OPENAI_API_KEY` as an environment variable.
|
||||
|
||||
### Step-1: Create dataset
|
||||
|
||||
In order to evaluate your RAG application, you have to setup a dataset. A data point in the dataset consists of `questions`, `contexts`, `answer`. Here is an example of how to create a dataset for evaluation:
|
||||
|
||||
```python
|
||||
from embedchain.utils.eval import EvalData
|
||||
|
||||
data = [
|
||||
{
|
||||
"question": "What is the net worth of Elon Musk?",
|
||||
"contexts": [
|
||||
"Elon Musk PROFILEElon MuskCEO, ...",
|
||||
"a Twitter poll on whether the journalists' ...",
|
||||
"2016 and run by Jared Birchall.[335]...",
|
||||
],
|
||||
"answer": "As of the information provided, Elon Musk's net worth is $241.6 billion.",
|
||||
},
|
||||
{
|
||||
"question": "which companies does Elon Musk own?",
|
||||
"contexts": [
|
||||
"of December 2023[update], ...",
|
||||
"ThielCofounderView ProfileTeslaHolds ...",
|
||||
"Elon Musk PROFILEElon MuskCEO, ...",
|
||||
],
|
||||
"answer": "Elon Musk owns several companies, including Tesla, SpaceX, Neuralink, and The Boring Company.",
|
||||
},
|
||||
]
|
||||
|
||||
dataset = []
|
||||
|
||||
for d in data:
|
||||
eval_data = EvalData(question=d["question"], contexts=d["contexts"], answer=d["answer"])
|
||||
dataset.append(eval_data)
|
||||
```
|
||||
|
||||
### Step-2: Run evaluation
|
||||
|
||||
Once you have created your dataset, you can run evaluation on the dataset by picking the metric you want to run evaluation on.
|
||||
|
||||
For example, you can run evaluation on context relevancy metric using the following code:
|
||||
|
||||
```python
|
||||
from embedchain.evaluation.metrics import ContextRelevance
|
||||
metric = ContextRelevance()
|
||||
score = metric.evaluate(dataset)
|
||||
print(score)
|
||||
```
|
||||
|
||||
You can choose a different metric or write your own to run evaluation on. You can check the following links:
|
||||
|
||||
- [Context Relevancy](#context_relevancy)
|
||||
- [Answer relenvancy](#answer_relevancy)
|
||||
- [Groundedness](#groundedness)
|
||||
- [Build your own metric](#custom_metric)
|
||||
|
||||
## Metrics
|
||||
|
||||
### Context Relevancy <a id="context_relevancy"></a>
|
||||
|
||||
Context relevancy is a metric to determine "how relevant the context is to the question". We use OpenAI's `gpt-4` model to determine the relevancy of the context. We achieve this by prompting the model with the question and the context and asking it to return relevant sentences from the context. We then use the following formula to determine the score:
|
||||
|
||||
```
|
||||
context_relevance_score = num_relevant_sentences_in_context / num_of_sentences_in_context
|
||||
```
|
||||
|
||||
#### Examples
|
||||
|
||||
You can run the context relevancy evaluation with the following simple code:
|
||||
|
||||
```python
|
||||
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.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)
|
||||
metric.evaluate(dataset)
|
||||
```
|
||||
|
||||
#### `ContextRelevanceConfig`
|
||||
|
||||
<ParamField path="model" type="str" optional>
|
||||
The model to use for the evaluation. Defaults to `gpt-4`. We only support openai's models for now.
|
||||
</ParamField>
|
||||
<ParamField path="api_key" type="str" optional>
|
||||
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="language" type="str" optional>
|
||||
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.evaluation.base` path.
|
||||
</ParamField>
|
||||
|
||||
|
||||
### Answer Relevancy <a id="answer_relevancy"></a>
|
||||
|
||||
Answer relevancy is a metric to determine how relevant the answer is to the question. We prompt the model with the answer and asking it to generate questions from the answer. We then use the cosine similarity between the generated questions and the original question to determine the score.
|
||||
|
||||
```
|
||||
answer_relevancy_score = mean(cosine_similarity(generated_questions, original_question))
|
||||
```
|
||||
|
||||
#### Examples
|
||||
|
||||
You can run the answer relevancy evaluation with the following simple code:
|
||||
|
||||
```python
|
||||
from embedchain.evaluation.metrics import AnswerRelevance
|
||||
|
||||
metric = AnswerRelevance()
|
||||
score = metric.evaluate(dataset)
|
||||
print(score)
|
||||
# 0.9505334177461916
|
||||
```
|
||||
|
||||
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.evaluation.base import AnswerRelevanceConfig
|
||||
from embedchain.evaluation.metrics import AnswerRelevance
|
||||
|
||||
eval_config = AnswerRelevanceConfig(
|
||||
model='gpt-4',
|
||||
embedder="text-embedding-ada-002",
|
||||
api_key="sk-xxx",
|
||||
num_gen_questions=2
|
||||
)
|
||||
metric = AnswerRelevance(config=eval_config)
|
||||
score = metric.evaluate(dataset)
|
||||
```
|
||||
|
||||
#### `AnswerRelevanceConfig`
|
||||
|
||||
<ParamField path="model" type="str" optional>
|
||||
The model to use for the evaluation. Defaults to `gpt-4`. We only support openai's models for now.
|
||||
</ParamField>
|
||||
<ParamField path="embedder" type="str" optional>
|
||||
The embedder to use for embedding the text. Defaults to `text-embedding-ada-002`. We only support openai's embedders for now.
|
||||
</ParamField>
|
||||
<ParamField path="api_key" type="str" optional>
|
||||
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="num_gen_questions" type="int" optional>
|
||||
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.evaluation.base` path.
|
||||
</ParamField>
|
||||
|
||||
## Groundedness <a id="groundedness"></a>
|
||||
|
||||
Groundedness is a metric to determine how grounded the answer is to the context. We use OpenAI's `gpt-4` model to determine the groundedness of the answer. We achieve this by prompting the model with the answer and asking it to generate claims from the answer. We then again prompt the model with the context and the generated claims to determine the verdict on the claims. We then use the following formula to determine the score:
|
||||
|
||||
```
|
||||
groundedness_score = (sum of all verdicts) / (total # of claims)
|
||||
```
|
||||
|
||||
You can run the groundedness evaluation with the following simple code:
|
||||
|
||||
```python
|
||||
from embedchain.evaluation.metrics import Groundedness
|
||||
metric = Groundedness()
|
||||
score = metric.evaluate(dataset) # dataset from above
|
||||
print(score)
|
||||
# 1.0
|
||||
```
|
||||
|
||||
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.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)
|
||||
score = metric.evaluate(dataset)
|
||||
```
|
||||
|
||||
|
||||
#### `GroundednessConfig`
|
||||
|
||||
<ParamField path="model" type="str" optional>
|
||||
The model to use for the evaluation. Defaults to `gpt-4`. We only support openai's models for now.
|
||||
</ParamField>
|
||||
<ParamField path="api_key" type="str" optional>
|
||||
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.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.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.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.
|
||||
</Note>
|
||||
|
||||
```python
|
||||
from typing import Optional
|
||||
|
||||
from embedchain.config.base_config import BaseConfig
|
||||
from embedchain.evaluation.metrics import BaseMetric
|
||||
from embedchain.utils.eval import EvalData
|
||||
|
||||
class MyCustomMetric(BaseMetric):
|
||||
def __init__(self, config: Optional[BaseConfig] = None):
|
||||
super().__init__(name="my_custom_metric")
|
||||
|
||||
def evaluate(self, dataset: list[EvalData]):
|
||||
score = 0.0
|
||||
# write your evaluation logic here
|
||||
return score
|
||||
```
|
||||
@@ -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,84 @@ 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 `AWS_ACCESS_KEY_ID` and `AWS_SECRET_ACCESS_KEY` to authenticate the API with AWS. You can find these in your [AWS Console](https://us-east-1.console.aws.amazon.com/iam/home?region=us-east-1#/users).
|
||||
|
||||
|
||||
### 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"
|
||||
|
||||
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" />
|
||||
|
||||
@@ -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
|
||||
pod_config:
|
||||
environment: gcp-starter
|
||||
metadata_config:
|
||||
indexed:
|
||||
- "url"
|
||||
- "hash"
|
||||
```
|
||||
|
||||
```yaml serverless_config.yaml
|
||||
vectordb:
|
||||
provider: pinecone
|
||||
config:
|
||||
metric: cosine
|
||||
vector_dimension: 1536
|
||||
collection_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/).
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
+14
-11
@@ -131,7 +131,8 @@
|
||||
},
|
||||
"components/llms",
|
||||
"components/vector-databases",
|
||||
"components/embedding-models"
|
||||
"components/embedding-models",
|
||||
"components/evaluation"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -198,17 +199,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/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",
|
||||
@@ -236,7 +239,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"
|
||||
|
||||
+108
-2
@@ -1,13 +1,15 @@
|
||||
import ast
|
||||
import concurrent.futures
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sqlite3
|
||||
import uuid
|
||||
from typing import Any, Optional
|
||||
from typing import Any, Optional, Union
|
||||
|
||||
import requests
|
||||
import yaml
|
||||
from tqdm import tqdm
|
||||
|
||||
from embedchain.cache import (Config, ExactMatchEvaluation,
|
||||
SearchDistanceEvaluation, cache,
|
||||
@@ -18,11 +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.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.evaluation import EvalData, EvalMetric
|
||||
from embedchain.utils.misc import validate_config
|
||||
from embedchain.vectordb.base import BaseVectorDB
|
||||
from embedchain.vectordb.chroma import ChromaDB
|
||||
@@ -393,7 +399,7 @@ class App(EmbedChain):
|
||||
|
||||
if config_path:
|
||||
file_extension = os.path.splitext(config_path)[1]
|
||||
with open(config_path, "r") as file:
|
||||
with open(config_path, "r", encoding="UTF-8") as file:
|
||||
if file_extension in [".yaml", ".yml"]:
|
||||
config_data = yaml.safe_load(file)
|
||||
elif file_extension == ".json":
|
||||
@@ -455,3 +461,103 @@ class App(EmbedChain):
|
||||
chunker=chunker_config_data,
|
||||
cache_config=cache_config,
|
||||
)
|
||||
|
||||
def _eval(self, dataset: list[EvalData], metric: Union[BaseMetric, str]):
|
||||
"""
|
||||
Evaluate the app on a dataset for a given metric.
|
||||
"""
|
||||
metric_str = metric.name if isinstance(metric, BaseMetric) else metric
|
||||
eval_class_map = {
|
||||
EvalMetric.CONTEXT_RELEVANCY.value: ContextRelevance,
|
||||
EvalMetric.ANSWER_RELEVANCY.value: AnswerRelevance,
|
||||
EvalMetric.GROUNDEDNESS.value: Groundedness,
|
||||
}
|
||||
|
||||
if metric_str in eval_class_map:
|
||||
return eval_class_map[metric_str]().evaluate(dataset)
|
||||
|
||||
# Handle the case for custom metrics
|
||||
if isinstance(metric, BaseMetric):
|
||||
return metric.evaluate(dataset)
|
||||
else:
|
||||
raise ValueError(f"Invalid metric: {metric}")
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
questions: Union[str, list[str]],
|
||||
metrics: Optional[list[Union[BaseMetric, str]]] = None,
|
||||
num_workers: int = 4,
|
||||
):
|
||||
"""
|
||||
Evaluate the app on a question.
|
||||
|
||||
param: questions: A question or a list of questions to evaluate.
|
||||
type: questions: Union[str, list[str]]
|
||||
param: metrics: A list of metrics to evaluate. Defaults to all metrics.
|
||||
type: metrics: Optional[list[Union[BaseMetric, str]]]
|
||||
param: num_workers: Number of workers to use for parallel processing.
|
||||
type: num_workers: int
|
||||
return: A dictionary containing the evaluation results.
|
||||
rtype: dict
|
||||
"""
|
||||
if "OPENAI_API_KEY" not in os.environ:
|
||||
raise ValueError("Please set the OPENAI_API_KEY environment variable with permission to use `gpt4` model.")
|
||||
|
||||
queries, answers, contexts = [], [], []
|
||||
if isinstance(questions, list):
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=num_workers) as executor:
|
||||
future_to_data = {executor.submit(self.query, q, citations=True): q for q in questions}
|
||||
for future in tqdm(
|
||||
concurrent.futures.as_completed(future_to_data),
|
||||
total=len(future_to_data),
|
||||
desc="Getting answer and contexts for questions",
|
||||
):
|
||||
question = future_to_data[future]
|
||||
queries.append(question)
|
||||
answer, context = future.result()
|
||||
answers.append(answer)
|
||||
contexts.append(list(map(lambda x: x[0], context)))
|
||||
else:
|
||||
answer, context = self.query(questions, citations=True)
|
||||
queries = [questions]
|
||||
answers = [answer]
|
||||
contexts = [list(map(lambda x: x[0], context))]
|
||||
|
||||
metrics = metrics or [
|
||||
EvalMetric.CONTEXT_RELEVANCY.value,
|
||||
EvalMetric.ANSWER_RELEVANCY.value,
|
||||
EvalMetric.GROUNDEDNESS.value,
|
||||
]
|
||||
|
||||
logging.info(f"Collecting data from {len(queries)} questions for evaluation...")
|
||||
dataset = []
|
||||
for q, a, c in zip(queries, answers, contexts):
|
||||
dataset.append(EvalData(question=q, answer=a, contexts=c))
|
||||
|
||||
logging.info(f"Evaluating {len(dataset)} data points...")
|
||||
result = {}
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=num_workers) as executor:
|
||||
future_to_metric = {executor.submit(self._eval, dataset, metric): metric for metric in metrics}
|
||||
for future in tqdm(
|
||||
concurrent.futures.as_completed(future_to_metric),
|
||||
total=len(future_to_metric),
|
||||
desc="Evaluating metrics",
|
||||
):
|
||||
metric = future_to_metric[future]
|
||||
if isinstance(metric, BaseMetric):
|
||||
result[metric.name] = future.result()
|
||||
else:
|
||||
result[metric] = future.result()
|
||||
|
||||
if self.config.collect_metrics:
|
||||
telemetry_props = self._telemetry_props
|
||||
metrics_names = []
|
||||
for metric in metrics:
|
||||
if isinstance(metric, BaseMetric):
|
||||
metrics_names.append(metric.name)
|
||||
else:
|
||||
metrics_names.append(metric)
|
||||
telemetry_props["metrics"] = metrics_names
|
||||
self.telemetry.capture(event_name="evaluate", properties=telemetry_props)
|
||||
|
||||
return result
|
||||
|
||||
@@ -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
@@ -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,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
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
from .base import (AnswerRelevanceConfig, ContextRelevanceConfig, # noqa: F401
|
||||
GroundednessConfig)
|
||||
@@ -0,0 +1,92 @@
|
||||
from typing import Optional
|
||||
|
||||
from embedchain.config.base_config import BaseConfig
|
||||
|
||||
ANSWER_RELEVANCY_PROMPT = """
|
||||
Please provide $num_gen_questions questions from the provided answer.
|
||||
You must provide the complete question, if are not able to provide the complete question, return empty string ("").
|
||||
Please only provide one question per line without numbers or bullets to distinguish them.
|
||||
You must only provide the questions and no other text.
|
||||
|
||||
$answer
|
||||
""" # noqa:E501
|
||||
|
||||
|
||||
CONTEXT_RELEVANCY_PROMPT = """
|
||||
Please extract relevant sentences from the provided context that is required to answer the given question.
|
||||
If no relevant sentences are found, or if you believe the question cannot be answered from the given context, return the empty string ("").
|
||||
While extracting candidate sentences you're not allowed to make any changes to sentences from given context or make up any sentences.
|
||||
You must only provide sentences from the given context and nothing else.
|
||||
|
||||
Context: $context
|
||||
Question: $question
|
||||
""" # noqa:E501
|
||||
|
||||
GROUNDEDNESS_ANSWER_CLAIMS_PROMPT = """
|
||||
Please provide one or more statements from each sentence of the provided answer.
|
||||
You must provide the symantically equivalent statements for each sentence of the answer.
|
||||
You must provide the complete statement, if are not able to provide the complete statement, return empty string ("").
|
||||
Please only provide one statement per line WITHOUT numbers or bullets.
|
||||
If the question provided is not being answered in the provided answer, return empty string ("").
|
||||
You must only provide the statements and no other text.
|
||||
|
||||
$question
|
||||
$answer
|
||||
""" # noqa:E501
|
||||
|
||||
GROUNDEDNESS_CLAIMS_INFERENCE_PROMPT = """
|
||||
Given the context and the provided claim statements, please provide a verdict for each claim statement whether it can be completely infered from the given context or not.
|
||||
Use only "1" (yes), "0" (no) and "-1" (null) for "yes", "no" or "null" respectively.
|
||||
You must provide one verdict per line, ONLY WITH "1", "0" or "-1" as per your verdict to the given statement and nothing else.
|
||||
You must provide the verdicts in the same order as the claim statements.
|
||||
|
||||
Contexts:
|
||||
$context
|
||||
|
||||
Claim statements:
|
||||
$claim_statements
|
||||
""" # noqa:E501
|
||||
|
||||
|
||||
class GroundednessConfig(BaseConfig):
|
||||
def __init__(
|
||||
self,
|
||||
model: str = "gpt-4",
|
||||
api_key: Optional[str] = None,
|
||||
answer_claims_prompt: str = GROUNDEDNESS_ANSWER_CLAIMS_PROMPT,
|
||||
claims_inference_prompt: str = GROUNDEDNESS_CLAIMS_INFERENCE_PROMPT,
|
||||
):
|
||||
self.model = model
|
||||
self.api_key = api_key
|
||||
self.answer_claims_prompt = answer_claims_prompt
|
||||
self.claims_inference_prompt = claims_inference_prompt
|
||||
|
||||
|
||||
class AnswerRelevanceConfig(BaseConfig):
|
||||
def __init__(
|
||||
self,
|
||||
model: str = "gpt-4",
|
||||
embedder: str = "text-embedding-ada-002",
|
||||
api_key: Optional[str] = None,
|
||||
num_gen_questions: int = 1,
|
||||
prompt: str = ANSWER_RELEVANCY_PROMPT,
|
||||
):
|
||||
self.model = model
|
||||
self.embedder = embedder
|
||||
self.api_key = api_key
|
||||
self.num_gen_questions = num_gen_questions
|
||||
self.prompt = prompt
|
||||
|
||||
|
||||
class ContextRelevanceConfig(BaseConfig):
|
||||
def __init__(
|
||||
self,
|
||||
model: str = "gpt-4",
|
||||
api_key: Optional[str] = None,
|
||||
language: str = "en",
|
||||
prompt: str = CONTEXT_RELEVANCY_PROMPT,
|
||||
):
|
||||
self.model = model
|
||||
self.api_key = api_key
|
||||
self.language = language
|
||||
self.prompt = prompt
|
||||
@@ -145,7 +145,7 @@ class BaseLlmConfig(BaseConfig):
|
||||
self.endpoint = endpoint
|
||||
self.model_kwargs = model_kwargs
|
||||
|
||||
if type(prompt) is str:
|
||||
if isinstance(prompt, str):
|
||||
prompt = Template(prompt)
|
||||
|
||||
if self.validate_prompt(prompt):
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from embedchain.config.vectordb.base import BaseVectorDbConfig
|
||||
@@ -9,12 +10,29 @@ class PineconeDBConfig(BaseVectorDbConfig):
|
||||
def __init__(
|
||||
self,
|
||||
collection_name: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
index_name: Optional[str] = None,
|
||||
dir: 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.vector_dimension = vector_dimension
|
||||
self.extra_params = extra_params
|
||||
super().__init__(collection_name=collection_name, dir=dir)
|
||||
self.index_name = index_name or f"{collection_name}-{vector_dimension}".lower().replace("_", "-")
|
||||
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=collection_name, dir=None)
|
||||
|
||||
+30
-19
@@ -7,12 +7,9 @@ from typing import Any, Optional, Union
|
||||
from dotenv import load_dotenv
|
||||
from langchain.docstore.document import Document
|
||||
|
||||
from embedchain.cache import (
|
||||
adapt,
|
||||
get_gptcache_session,
|
||||
gptcache_data_convert,
|
||||
gptcache_update_cache_callback,
|
||||
)
|
||||
from embedchain.cache import (adapt, get_gptcache_session,
|
||||
gptcache_data_convert,
|
||||
gptcache_update_cache_callback)
|
||||
from embedchain.chunkers.base_chunker import BaseChunker
|
||||
from embedchain.config import AddConfig, BaseLlmConfig, ChunkerConfig
|
||||
from embedchain.config.base_app_config import BaseAppConfig
|
||||
@@ -22,7 +19,8 @@ from embedchain.embedder.base import BaseEmbedder
|
||||
from embedchain.helpers.json_serializable import JSONSerializable
|
||||
from embedchain.llm.base import BaseLlm
|
||||
from embedchain.loaders.base_loader import BaseLoader
|
||||
from embedchain.models.data_type import DataType, DirectDataType, IndirectDataType, SpecialDataType
|
||||
from embedchain.models.data_type import (DataType, DirectDataType,
|
||||
IndirectDataType, SpecialDataType)
|
||||
from embedchain.telemetry.posthog import AnonymousTelemetry
|
||||
from embedchain.utils.misc import detect_datatype, is_valid_json_string
|
||||
from embedchain.vectordb.base import BaseVectorDB
|
||||
@@ -371,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
|
||||
@@ -435,13 +433,7 @@ class EmbedChain(JSONSerializable):
|
||||
# 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,
|
||||
)
|
||||
self.db.add(documents=documents, metadatas=metadatas, ids=ids, **kwargs)
|
||||
count_new_chunks = self.db.count() - chunks_before_addition
|
||||
|
||||
print(f"Successfully saved {src} ({chunker.data_type}). New chunks count: {count_new_chunks}")
|
||||
@@ -665,13 +657,32 @@ class EmbedChain(JSONSerializable):
|
||||
self.db.reset()
|
||||
self.cursor.execute("DELETE FROM data_sources WHERE pipeline_id = ?", (self.config.id,))
|
||||
self.connection.commit()
|
||||
self.delete_chat_history()
|
||||
self.delete_all_chat_history(app_id=self.config.id)
|
||||
# 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):
|
||||
return self.llm.memory.get(app_id=self.config.id, num_rounds=num_rounds, display_format=display_format)
|
||||
def get_history(self, num_rounds: int = 10, display_format: bool = True, session_id: Optional[str] = "default"):
|
||||
history = self.llm.memory.get(
|
||||
app_id=self.config.id, session_id=session_id, num_rounds=num_rounds, display_format=display_format
|
||||
)
|
||||
return history
|
||||
|
||||
def delete_chat_history(self, session_id: str = "default"):
|
||||
def delete_session_chat_history(self, session_id: str = "default"):
|
||||
self.llm.memory.delete(app_id=self.config.id, session_id=session_id)
|
||||
self.llm.update_history(app_id=self.config.id)
|
||||
|
||||
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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from embedchain.utils.evaluation import EvalData
|
||||
|
||||
|
||||
class BaseMetric(ABC):
|
||||
"""Base class for a metric.
|
||||
|
||||
This class provides a common interface for all metrics.
|
||||
"""
|
||||
|
||||
def __init__(self, name: str = "base_metric"):
|
||||
"""
|
||||
Initialize the BaseMetric.
|
||||
"""
|
||||
self.name = name
|
||||
|
||||
@abstractmethod
|
||||
def evaluate(self, dataset: list[EvalData]):
|
||||
"""
|
||||
Abstract method to evaluate the dataset.
|
||||
|
||||
This method should be implemented by subclasses to perform the actual
|
||||
evaluation on the dataset.
|
||||
|
||||
:param dataset: dataset to evaluate
|
||||
:type dataset: list[EvalData]
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
@@ -0,0 +1,3 @@
|
||||
from .answer_relevancy import AnswerRelevance # noqa: F401
|
||||
from .context_relevancy import ContextRelevance # noqa: F401
|
||||
from .groundedness import Groundedness # noqa: F401
|
||||
@@ -0,0 +1,93 @@
|
||||
import concurrent.futures
|
||||
import logging
|
||||
import os
|
||||
from string import Template
|
||||
from typing import Optional
|
||||
|
||||
import numpy as np
|
||||
from openai import OpenAI
|
||||
from tqdm import tqdm
|
||||
|
||||
from embedchain.config.evaluation.base import AnswerRelevanceConfig
|
||||
from embedchain.evaluation.base import BaseMetric
|
||||
from embedchain.utils.evaluation import EvalData, EvalMetric
|
||||
|
||||
|
||||
class AnswerRelevance(BaseMetric):
|
||||
"""
|
||||
Metric for evaluating the relevance of answers.
|
||||
"""
|
||||
|
||||
def __init__(self, config: Optional[AnswerRelevanceConfig] = AnswerRelevanceConfig()):
|
||||
super().__init__(name=EvalMetric.ANSWER_RELEVANCY.value)
|
||||
self.config = config
|
||||
api_key = self.config.api_key or os.getenv("OPENAI_API_KEY")
|
||||
if not api_key:
|
||||
raise ValueError("API key not found. Set 'OPENAI_API_KEY' or pass it in the config.")
|
||||
self.client = OpenAI(api_key=api_key)
|
||||
|
||||
def _generate_prompt(self, data: EvalData) -> str:
|
||||
"""
|
||||
Generates a prompt based on the provided data.
|
||||
"""
|
||||
return Template(self.config.prompt).substitute(
|
||||
num_gen_questions=self.config.num_gen_questions, answer=data.answer
|
||||
)
|
||||
|
||||
def _generate_questions(self, prompt: str) -> list[str]:
|
||||
"""
|
||||
Generates questions from the prompt.
|
||||
"""
|
||||
response = self.client.chat.completions.create(
|
||||
model=self.config.model,
|
||||
messages=[{"role": "user", "content": prompt}],
|
||||
)
|
||||
return response.choices[0].message.content.strip().split("\n")
|
||||
|
||||
def _generate_embedding(self, question: str) -> np.ndarray:
|
||||
"""
|
||||
Generates the embedding for a question.
|
||||
"""
|
||||
response = self.client.embeddings.create(
|
||||
input=question,
|
||||
model=self.config.embedder,
|
||||
)
|
||||
return np.array(response.data[0].embedding)
|
||||
|
||||
def _compute_similarity(self, original: np.ndarray, generated: np.ndarray) -> float:
|
||||
"""
|
||||
Computes the cosine similarity between two embeddings.
|
||||
"""
|
||||
original = original.reshape(1, -1)
|
||||
norm = np.linalg.norm(original) * np.linalg.norm(generated, axis=1)
|
||||
return np.dot(generated, original.T).flatten() / norm
|
||||
|
||||
def _compute_score(self, data: EvalData) -> float:
|
||||
"""
|
||||
Computes the relevance score for a given data item.
|
||||
"""
|
||||
prompt = self._generate_prompt(data)
|
||||
generated_questions = self._generate_questions(prompt)
|
||||
original_embedding = self._generate_embedding(data.question)
|
||||
generated_embeddings = np.array([self._generate_embedding(q) for q in generated_questions])
|
||||
similarities = self._compute_similarity(original_embedding, generated_embeddings)
|
||||
return np.mean(similarities)
|
||||
|
||||
def evaluate(self, dataset: list[EvalData]) -> float:
|
||||
"""
|
||||
Evaluates the dataset and returns the average answer relevance score.
|
||||
"""
|
||||
results = []
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor() as executor:
|
||||
future_to_data = {executor.submit(self._compute_score, data): data for data in dataset}
|
||||
for future in tqdm(
|
||||
concurrent.futures.as_completed(future_to_data), total=len(dataset), desc="Evaluating Answer Relevancy"
|
||||
):
|
||||
data = future_to_data[future]
|
||||
try:
|
||||
results.append(future.result())
|
||||
except Exception as e:
|
||||
logging.error(f"Error evaluating answer relevancy for {data}: {e}")
|
||||
|
||||
return np.mean(results) if results else 0.0
|
||||
@@ -0,0 +1,69 @@
|
||||
import concurrent.futures
|
||||
import os
|
||||
from string import Template
|
||||
from typing import Optional
|
||||
|
||||
import numpy as np
|
||||
import pysbd
|
||||
from openai import OpenAI
|
||||
from tqdm import tqdm
|
||||
|
||||
from embedchain.config.evaluation.base import ContextRelevanceConfig
|
||||
from embedchain.evaluation.base import BaseMetric
|
||||
from embedchain.utils.evaluation import EvalData, EvalMetric
|
||||
|
||||
|
||||
class ContextRelevance(BaseMetric):
|
||||
"""
|
||||
Metric for evaluating the relevance of context in a dataset.
|
||||
"""
|
||||
|
||||
def __init__(self, config: Optional[ContextRelevanceConfig] = ContextRelevanceConfig()):
|
||||
super().__init__(name=EvalMetric.CONTEXT_RELEVANCY.value)
|
||||
self.config = config
|
||||
api_key = self.config.api_key or os.getenv("OPENAI_API_KEY")
|
||||
if not api_key:
|
||||
raise ValueError("API key not found. Set 'OPENAI_API_KEY' or pass it in the config.")
|
||||
self.client = OpenAI(api_key=api_key)
|
||||
self._sbd = pysbd.Segmenter(language=self.config.language, clean=False)
|
||||
|
||||
def _sentence_segmenter(self, text: str) -> list[str]:
|
||||
"""
|
||||
Segments the given text into sentences.
|
||||
"""
|
||||
return self._sbd.segment(text)
|
||||
|
||||
def _compute_score(self, data: EvalData) -> float:
|
||||
"""
|
||||
Computes the context relevance score for a given data item.
|
||||
"""
|
||||
original_context = "\n".join(data.contexts)
|
||||
prompt = Template(self.config.prompt).substitute(context=original_context, question=data.question)
|
||||
response = self.client.chat.completions.create(
|
||||
model=self.config.model, messages=[{"role": "user", "content": prompt}]
|
||||
)
|
||||
useful_context = response.choices[0].message.content.strip()
|
||||
useful_context_sentences = self._sentence_segmenter(useful_context)
|
||||
original_context_sentences = self._sentence_segmenter(original_context)
|
||||
|
||||
if not original_context_sentences:
|
||||
return 0.0
|
||||
return len(useful_context_sentences) / len(original_context_sentences)
|
||||
|
||||
def evaluate(self, dataset: list[EvalData]) -> float:
|
||||
"""
|
||||
Evaluates the dataset and returns the average context relevance score.
|
||||
"""
|
||||
scores = []
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor() as executor:
|
||||
futures = [executor.submit(self._compute_score, data) for data in dataset]
|
||||
for future in tqdm(
|
||||
concurrent.futures.as_completed(futures), total=len(dataset), desc="Evaluating Context Relevancy"
|
||||
):
|
||||
try:
|
||||
scores.append(future.result())
|
||||
except Exception as e:
|
||||
print(f"Error during evaluation: {e}")
|
||||
|
||||
return np.mean(scores) if scores else 0.0
|
||||
@@ -0,0 +1,102 @@
|
||||
import concurrent.futures
|
||||
import logging
|
||||
import os
|
||||
from string import Template
|
||||
from typing import Optional
|
||||
|
||||
import numpy as np
|
||||
from openai import OpenAI
|
||||
from tqdm import tqdm
|
||||
|
||||
from embedchain.config.evaluation.base import GroundednessConfig
|
||||
from embedchain.evaluation.base import BaseMetric
|
||||
from embedchain.utils.evaluation import EvalData, EvalMetric
|
||||
|
||||
|
||||
class Groundedness(BaseMetric):
|
||||
"""
|
||||
Metric for groundedness of answer from the given contexts.
|
||||
"""
|
||||
|
||||
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.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)
|
||||
|
||||
def _generate_answer_claim_prompt(self, data: EvalData) -> str:
|
||||
"""
|
||||
Generate the prompt for the given data.
|
||||
"""
|
||||
prompt = Template(self.config.answer_claims_prompt).substitute(question=data.question, answer=data.answer)
|
||||
return prompt
|
||||
|
||||
def _get_claim_statements(self, prompt: str) -> np.ndarray:
|
||||
"""
|
||||
Get claim statements from the answer.
|
||||
"""
|
||||
response = self.client.chat.completions.create(
|
||||
model=self.config.model,
|
||||
messages=[{"role": "user", "content": f"{prompt}"}],
|
||||
)
|
||||
result = response.choices[0].message.content.strip()
|
||||
claim_statements = np.array([statement for statement in result.split("\n") if statement])
|
||||
return claim_statements
|
||||
|
||||
def _generate_claim_inference_prompt(self, data: EvalData, claim_statements: list[str]) -> str:
|
||||
"""
|
||||
Generate the claim inference prompt for the given data and claim statements.
|
||||
"""
|
||||
prompt = Template(self.config.claims_inference_prompt).substitute(
|
||||
context="\n".join(data.contexts), claim_statements="\n".join(claim_statements)
|
||||
)
|
||||
return prompt
|
||||
|
||||
def _get_claim_verdict_scores(self, prompt: str) -> np.ndarray:
|
||||
"""
|
||||
Get verdicts for claim statements.
|
||||
"""
|
||||
response = self.client.chat.completions.create(
|
||||
model=self.config.model,
|
||||
messages=[{"role": "user", "content": f"{prompt}"}],
|
||||
)
|
||||
result = response.choices[0].message.content.strip()
|
||||
claim_verdicts = result.split("\n")
|
||||
verdict_score_map = {"1": 1, "0": 0, "-1": np.nan}
|
||||
verdict_scores = np.array([verdict_score_map[verdict] for verdict in claim_verdicts])
|
||||
return verdict_scores
|
||||
|
||||
def _compute_score(self, data: EvalData) -> float:
|
||||
"""
|
||||
Compute the groundedness score for a single data point.
|
||||
"""
|
||||
answer_claims_prompt = self._generate_answer_claim_prompt(data)
|
||||
claim_statements = self._get_claim_statements(answer_claims_prompt)
|
||||
|
||||
claim_inference_prompt = self._generate_claim_inference_prompt(data, claim_statements)
|
||||
verdict_scores = self._get_claim_verdict_scores(claim_inference_prompt)
|
||||
return np.sum(verdict_scores) / claim_statements.size
|
||||
|
||||
def evaluate(self, dataset: list[EvalData]):
|
||||
"""
|
||||
Evaluate the dataset and returns the average groundedness score.
|
||||
"""
|
||||
results = []
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor() as executor:
|
||||
future_to_data = {executor.submit(self._compute_score, data): data for data in dataset}
|
||||
for future in tqdm(
|
||||
concurrent.futures.as_completed(future_to_data),
|
||||
total=len(future_to_data),
|
||||
desc="Evaluating Groundedness",
|
||||
):
|
||||
data = future_to_data[future]
|
||||
try:
|
||||
score = future.result()
|
||||
results.append(score)
|
||||
except Exception as e:
|
||||
logging.error(f"Error while evaluating groundedness for data point {data}: {e}")
|
||||
|
||||
return np.mean(results) if results else 0.0
|
||||
@@ -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,6 +52,7 @@ 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",
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
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")
|
||||
|
||||
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)
|
||||
@@ -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
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
+54
-15
@@ -53,7 +53,7 @@ class ChatHistory:
|
||||
logging.info(f"Added chat memory to db with id: {memory_id}")
|
||||
return memory_id
|
||||
|
||||
def delete(self, app_id: str, session_id: str):
|
||||
def delete(self, app_id: str, session_id: Optional[str] = None):
|
||||
"""
|
||||
Delete all chat history for a given app_id and session_id.
|
||||
This is useful for deleting chat history for a given user.
|
||||
@@ -63,25 +63,50 @@ class ChatHistory:
|
||||
|
||||
:return: None
|
||||
"""
|
||||
DELETE_CHAT_HISTORY_QUERY = "DELETE FROM ec_chat_history WHERE app_id=? AND session_id=?"
|
||||
self.cursor.execute(DELETE_CHAT_HISTORY_QUERY, (app_id, session_id))
|
||||
if session_id:
|
||||
DELETE_CHAT_HISTORY_QUERY = "DELETE FROM ec_chat_history WHERE app_id=? AND session_id=?"
|
||||
params = (app_id, session_id)
|
||||
else:
|
||||
DELETE_CHAT_HISTORY_QUERY = "DELETE FROM ec_chat_history WHERE app_id=?"
|
||||
params = (app_id,)
|
||||
|
||||
self.cursor.execute(DELETE_CHAT_HISTORY_QUERY, params)
|
||||
self.connection.commit()
|
||||
|
||||
def get(self, app_id, session_id, num_rounds=10, display_format=False) -> list[ChatMessage]:
|
||||
def get(
|
||||
self, app_id, session_id: str = "default", num_rounds=10, fetch_all: bool = False, display_format=False
|
||||
) -> list[ChatMessage]:
|
||||
"""
|
||||
Get the most recent num_rounds rounds of conversations
|
||||
between human and AI, for a given app_id.
|
||||
Get the chat history for a given app_id.
|
||||
|
||||
param: app_id - The app_id to get chat history
|
||||
param: session_id (optional) - The session_id to get chat history. Defaults to "default"
|
||||
param: num_rounds (optional) - The number of rounds to get chat history. Defaults to 10
|
||||
param: fetch_all (optional) - Whether to fetch all chat history or not. Defaults to False
|
||||
param: display_format (optional) - Whether to return the chat history in display format. Defaults to False
|
||||
"""
|
||||
|
||||
QUERY = """
|
||||
base_query = """
|
||||
SELECT * FROM ec_chat_history
|
||||
WHERE app_id=? AND session_id=?
|
||||
ORDER BY created_at DESC
|
||||
LIMIT ?
|
||||
WHERE app_id=?
|
||||
"""
|
||||
|
||||
if fetch_all:
|
||||
additional_query = "ORDER BY created_at DESC"
|
||||
params = (app_id,)
|
||||
else:
|
||||
additional_query = """
|
||||
AND session_id=?
|
||||
ORDER BY created_at DESC
|
||||
LIMIT ?
|
||||
"""
|
||||
params = (app_id, session_id, num_rounds)
|
||||
|
||||
QUERY = base_query + additional_query
|
||||
|
||||
self.cursor.execute(
|
||||
QUERY,
|
||||
(app_id, session_id, num_rounds),
|
||||
params,
|
||||
)
|
||||
|
||||
results = self.cursor.fetchall()
|
||||
@@ -91,7 +116,15 @@ class ChatHistory:
|
||||
metadata = self._deserialize_json(metadata=metadata)
|
||||
# Return list of dict if display_format is True
|
||||
if display_format:
|
||||
history.append({"human": question, "ai": answer, "metadata": metadata, "timestamp": timestamp})
|
||||
history.append(
|
||||
{
|
||||
"session_id": session_id,
|
||||
"human": question,
|
||||
"ai": answer,
|
||||
"metadata": metadata,
|
||||
"timestamp": timestamp,
|
||||
}
|
||||
)
|
||||
else:
|
||||
memory = ChatMessage()
|
||||
memory.add_user_message(question, metadata=metadata)
|
||||
@@ -99,7 +132,7 @@ class ChatHistory:
|
||||
history.append(memory)
|
||||
return history
|
||||
|
||||
def count(self, app_id: str, session_id: str):
|
||||
def count(self, app_id: str, session_id: Optional[str] = None):
|
||||
"""
|
||||
Count the number of chat messages for a given app_id and session_id.
|
||||
|
||||
@@ -108,8 +141,14 @@ class ChatHistory:
|
||||
|
||||
:return: The number of chat messages for a given app_id and session_id
|
||||
"""
|
||||
QUERY = "SELECT COUNT(*) FROM ec_chat_history WHERE app_id=? AND session_id=?"
|
||||
self.cursor.execute(QUERY, (app_id, session_id))
|
||||
if session_id:
|
||||
QUERY = "SELECT COUNT(*) FROM ec_chat_history WHERE app_id=? AND session_id=?"
|
||||
params = (app_id, session_id)
|
||||
else:
|
||||
QUERY = "SELECT COUNT(*) FROM ec_chat_history WHERE app_id=?"
|
||||
params = (app_id,)
|
||||
|
||||
self.cursor.execute(QUERY, params)
|
||||
count = self.cursor.fetchone()[0]
|
||||
return count
|
||||
|
||||
|
||||
@@ -8,3 +8,4 @@ class VectorDimensions(Enum):
|
||||
VERTEX_AI = 768
|
||||
HUGGING_FACE = 384
|
||||
GOOGLE_AI = 768
|
||||
MISTRAL_AI = 1024
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
from enum import Enum
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class EvalMetric(Enum):
|
||||
CONTEXT_RELEVANCY = "context_relevancy"
|
||||
ANSWER_RELEVANCY = "answer_relevancy"
|
||||
GROUNDEDNESS = "groundedness"
|
||||
|
||||
|
||||
class EvalData(BaseModel):
|
||||
question: str
|
||||
contexts: list[str]
|
||||
answer: str
|
||||
ground_truth: Optional[str] = None # Not used as of now
|
||||
@@ -201,9 +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
|
||||
|
||||
@@ -399,6 +406,8 @@ def validate_config(config_data):
|
||||
"llama2",
|
||||
"vertexai",
|
||||
"google",
|
||||
"aws_bedrock",
|
||||
"mistralai",
|
||||
),
|
||||
Optional("config"): {
|
||||
Optional("model"): str,
|
||||
@@ -415,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"): {
|
||||
@@ -424,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"): {
|
||||
|
||||
@@ -75,3 +75,8 @@ class BaseVectorDB(JSONSerializable):
|
||||
:type name: str
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def delete(self):
|
||||
"""Delete from database."""
|
||||
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -129,17 +129,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
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional, Union
|
||||
|
||||
@@ -41,7 +42,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 +53,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 +90,23 @@ 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])
|
||||
|
||||
if where is not None:
|
||||
logging.warning("Filtering is not supported by Pinecone")
|
||||
|
||||
return {"ids": existing_ids, "metadatas": metadatas}
|
||||
|
||||
def add(
|
||||
self,
|
||||
embeddings: list[list[float]],
|
||||
documents: list[str],
|
||||
metadatas: list[object],
|
||||
ids: list[str],
|
||||
@@ -115,8 +133,8 @@ 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,
|
||||
@@ -141,13 +159,20 @@ class PineconeDB(BaseVectorDB):
|
||||
:rtype: list[str], if citations=False, otherwise list[tuple[str, str, str]]
|
||||
"""
|
||||
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)
|
||||
query_filter = self._generate_filter(where)
|
||||
data = self.pinecone_index.query(
|
||||
vector=query_vector,
|
||||
filter=query_filter,
|
||||
top_k=n_results,
|
||||
include_metadata=True,
|
||||
**kwargs,
|
||||
)
|
||||
contexts = []
|
||||
for doc in data["matches"]:
|
||||
metadata = doc["metadata"]
|
||||
context = metadata["text"]
|
||||
for doc in data.get("matches", []):
|
||||
metadata = doc.get("metadata", {})
|
||||
context = metadata.get("text")
|
||||
if citations:
|
||||
metadata["score"] = doc["score"]
|
||||
metadata["score"] = doc.get("score")
|
||||
contexts.append(tuple((context, metadata)))
|
||||
else:
|
||||
contexts.append(context)
|
||||
@@ -171,7 +196,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 +208,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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
@@ -303,11 +325,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)
|
||||
|
||||
@@ -6,15 +6,8 @@ from embedchain.helpers.json_serializable import register_deserializable
|
||||
from embedchain.vectordb.base import BaseVectorDB
|
||||
|
||||
try:
|
||||
from pymilvus import (
|
||||
Collection,
|
||||
CollectionSchema,
|
||||
DataType,
|
||||
FieldSchema,
|
||||
MilvusClient,
|
||||
connections,
|
||||
utility,
|
||||
)
|
||||
from pymilvus import (Collection, CollectionSchema, DataType, FieldSchema,
|
||||
MilvusClient, connections, utility)
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"Zilliz requires extra dependencies. Install with `pip install --upgrade embedchain[milvus]`"
|
||||
@@ -76,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)
|
||||
@@ -101,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],
|
||||
@@ -125,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()
|
||||
@@ -136,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]]:
|
||||
@@ -148,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.
|
||||
@@ -160,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,
|
||||
@@ -181,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
|
||||
@@ -224,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.
|
||||
|
||||
@@ -232,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)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import logging
|
||||
import os
|
||||
|
||||
import aiofiles
|
||||
import yaml
|
||||
from database import Base, SessionLocal, engine
|
||||
from fastapi import Depends, FastAPI, HTTPException, UploadFile
|
||||
@@ -74,8 +75,8 @@ async def create_app_using_default_config(app_id: str, config: UploadFile = None
|
||||
yaml.safe_load(contents)
|
||||
# TODO: validate the config yaml file here
|
||||
yaml_path = f"configs/{app_id}.yaml"
|
||||
with open(yaml_path, "w") as file:
|
||||
file.write(str(contents, "utf-8"))
|
||||
async with aiofiles.open(yaml_path, mode="w") as file_out:
|
||||
await file_out.write(str(contents, "utf-8"))
|
||||
except yaml.YAMLError as exc:
|
||||
raise HTTPException(detail=f"Error parsing YAML: {exc}", status_code=400)
|
||||
|
||||
|
||||
@@ -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
+232
-117
@@ -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)"]
|
||||
@@ -4677,32 +4802,6 @@ files = [
|
||||
{file = "protobuf-4.21.12.tar.gz", hash = "sha256:7cd532c4566d0e6feafecc1059d04c7915aec8e182d1cf7adee8b24ef1e2e6ab"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "psutil"
|
||||
version = "5.9.5"
|
||||
description = "Cross-platform lib for process and system monitoring in Python."
|
||||
optional = true
|
||||
python-versions = ">=2.7, !=3.0.*, !=3.1.*, !=3.2.*, !=3.3.*"
|
||||
files = [
|
||||
{file = "psutil-5.9.5-cp27-cp27m-macosx_10_9_x86_64.whl", hash = "sha256:be8929ce4313f9f8146caad4272f6abb8bf99fc6cf59344a3167ecd74f4f203f"},
|
||||
{file = "psutil-5.9.5-cp27-cp27m-manylinux2010_i686.whl", hash = "sha256:ab8ed1a1d77c95453db1ae00a3f9c50227ebd955437bcf2a574ba8adbf6a74d5"},
|
||||
{file = "psutil-5.9.5-cp27-cp27m-manylinux2010_x86_64.whl", hash = "sha256:4aef137f3345082a3d3232187aeb4ac4ef959ba3d7c10c33dd73763fbc063da4"},
|
||||
{file = "psutil-5.9.5-cp27-cp27mu-manylinux2010_i686.whl", hash = "sha256:ea8518d152174e1249c4f2a1c89e3e6065941df2fa13a1ab45327716a23c2b48"},
|
||||
{file = "psutil-5.9.5-cp27-cp27mu-manylinux2010_x86_64.whl", hash = "sha256:acf2aef9391710afded549ff602b5887d7a2349831ae4c26be7c807c0a39fac4"},
|
||||
{file = "psutil-5.9.5-cp27-none-win32.whl", hash = "sha256:5b9b8cb93f507e8dbaf22af6a2fd0ccbe8244bf30b1baad6b3954e935157ae3f"},
|
||||
{file = "psutil-5.9.5-cp27-none-win_amd64.whl", hash = "sha256:8c5f7c5a052d1d567db4ddd231a9d27a74e8e4a9c3f44b1032762bd7b9fdcd42"},
|
||||
{file = "psutil-5.9.5-cp36-abi3-macosx_10_9_x86_64.whl", hash = "sha256:3c6f686f4225553615612f6d9bc21f1c0e305f75d7d8454f9b46e901778e7217"},
|
||||
{file = "psutil-5.9.5-cp36-abi3-manylinux_2_12_i686.manylinux2010_i686.manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:7a7dd9997128a0d928ed4fb2c2d57e5102bb6089027939f3b722f3a210f9a8da"},
|
||||
{file = "psutil-5.9.5-cp36-abi3-manylinux_2_12_x86_64.manylinux2010_x86_64.manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:89518112647f1276b03ca97b65cc7f64ca587b1eb0278383017c2a0dcc26cbe4"},
|
||||
{file = "psutil-5.9.5-cp36-abi3-win32.whl", hash = "sha256:104a5cc0e31baa2bcf67900be36acde157756b9c44017b86b2c049f11957887d"},
|
||||
{file = "psutil-5.9.5-cp36-abi3-win_amd64.whl", hash = "sha256:b258c0c1c9d145a1d5ceffab1134441c4c5113b2417fafff7315a917a026c3c9"},
|
||||
{file = "psutil-5.9.5-cp38-abi3-macosx_11_0_arm64.whl", hash = "sha256:c607bb3b57dc779d55e1554846352b4e358c10fff3abf3514a7a6601beebdb30"},
|
||||
{file = "psutil-5.9.5.tar.gz", hash = "sha256:5410638e4df39c54d957fc51ce03048acd8e6d60abc0f5107af51e5fb566eb3c"},
|
||||
]
|
||||
|
||||
[package.extras]
|
||||
test = ["enum34", "ipaddress", "mock", "pywin32", "wmi"]
|
||||
|
||||
[[package]]
|
||||
name = "psycopg"
|
||||
version = "3.1.12"
|
||||
@@ -5293,6 +5392,16 @@ files = [
|
||||
{file = "pyreadline3-3.4.1.tar.gz", hash = "sha256:6f3d1f7b8a31ba32b73917cefc1f28cc660562f39aea8646d30bd6eff21f7bae"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pysbd"
|
||||
version = "0.3.4"
|
||||
description = "pysbd (Python Sentence Boundary Disambiguation) is a rule-based sentence boundary detection that works out-of-the-box across many languages."
|
||||
optional = false
|
||||
python-versions = ">=3"
|
||||
files = [
|
||||
{file = "pysbd-0.3.4-py3-none-any.whl", hash = "sha256:cd838939b7b0b185fcf86b0baf6636667dfb6e474743beeff878e9f42e022953"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "pytesseract"
|
||||
version = "0.3.10"
|
||||
@@ -5968,29 +6077,47 @@ pyasn1 = ">=0.1.3"
|
||||
|
||||
[[package]]
|
||||
name = "ruff"
|
||||
version = "0.0.220"
|
||||
description = "An extremely fast Python linter, written in Rust."
|
||||
version = "0.1.11"
|
||||
description = "An extremely fast Python linter and code formatter, written in Rust."
|
||||
optional = false
|
||||
python-versions = ">=3.7"
|
||||
files = [
|
||||
{file = "ruff-0.0.220-py3-none-macosx_10_7_x86_64.whl", hash = "sha256:152e6697aca6aea991cdd37922c34a3e4db4828822c4663122326e6051e0f68a"},
|
||||
{file = "ruff-0.0.220-py3-none-macosx_10_9_x86_64.macosx_11_0_arm64.macosx_10_9_universal2.whl", hash = "sha256:127887a00d53beb7c0c78a8b4bbdda2f14f07db7b3571feb6855cb32862cb88d"},
|
||||
{file = "ruff-0.0.220-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:91235ff448786f8f3b856c104fd6c4fe11e835b0db75da5fdf337e1ed5d454da"},
|
||||
{file = "ruff-0.0.220-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:2f0a104afc32012048627317ae8b0940e3f11a717905aed3fc26931a873e3b29"},
|
||||
{file = "ruff-0.0.220-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:eee1deddf1671860e056a78938176600108857a527c078038627b284a554723c"},
|
||||
{file = "ruff-0.0.220-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:071082d09c953924eccfd88ffd0d71119ddd6fc7767f3c31549a1cd0651ba586"},
|
||||
{file = "ruff-0.0.220-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:a5ddfc945a9076c9779b52c1f7296cf8d8e6919e619c4522617bc37b60eddd2e"},
|
||||
{file = "ruff-0.0.220-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:15f387fd156430353fb61d2b609f1c38d2e9096e2fce31149da5cf08b73f04a8"},
|
||||
{file = "ruff-0.0.220-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:b5688983f21ac64bbcca8d84f4107733cc2d62c1354ea1a6b85eb9ead32328cc"},
|
||||
{file = "ruff-0.0.220-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:e2b0c9dbff13649ded5ee92d6a47d720e8471461e0a4eba01bf3474f851cb2f0"},
|
||||
{file = "ruff-0.0.220-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:9b540fa9f90f46f656b34fb73b738613562974599903a1f0d40bdd1a8180bfab"},
|
||||
{file = "ruff-0.0.220-py3-none-musllinux_1_2_i686.whl", hash = "sha256:a061c17c2b0f81193fca5e53829b6c0569c5c7d393cc4fc1c192ce0a64d3b9ca"},
|
||||
{file = "ruff-0.0.220-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:42677089abd7db6f8aefa3dbe7a82fea4e2a43a08bc7bf6d3f58ec3d76c63712"},
|
||||
{file = "ruff-0.0.220-py3-none-win32.whl", hash = "sha256:f8821cfc204b38140afe870bcd4cc6c836bbd2f820b92df66b8fe8b8d71a3772"},
|
||||
{file = "ruff-0.0.220-py3-none-win_amd64.whl", hash = "sha256:8a1d678a224afd7149afbe497c97c3ccdc6c42632ee84fb0e3f68d190c1ccec1"},
|
||||
{file = "ruff-0.0.220.tar.gz", hash = "sha256:621f7f063c0d13570b709fb9904a329ddb9a614fdafc786718afd43e97440c34"},
|
||||
{file = "ruff-0.1.11-py3-none-macosx_10_12_x86_64.macosx_11_0_arm64.macosx_10_12_universal2.whl", hash = "sha256:a7f772696b4cdc0a3b2e527fc3c7ccc41cdcb98f5c80fdd4f2b8c50eb1458196"},
|
||||
{file = "ruff-0.1.11-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:934832f6ed9b34a7d5feea58972635c2039c7a3b434fe5ba2ce015064cb6e955"},
|
||||
{file = "ruff-0.1.11-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:ea0d3e950e394c4b332bcdd112aa566010a9f9c95814844a7468325290aabfd9"},
|
||||
{file = "ruff-0.1.11-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:9bd4025b9c5b429a48280785a2b71d479798a69f5c2919e7d274c5f4b32c3607"},
|
||||
{file = "ruff-0.1.11-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:e1ad00662305dcb1e987f5ec214d31f7d6a062cae3e74c1cbccef15afd96611d"},
|
||||
{file = "ruff-0.1.11-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:4b077ce83f47dd6bea1991af08b140e8b8339f0ba8cb9b7a484c30ebab18a23f"},
|
||||
{file = "ruff-0.1.11-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:c4a88efecec23c37b11076fe676e15c6cdb1271a38f2b415e381e87fe4517f18"},
|
||||
{file = "ruff-0.1.11-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:5b25093dad3b055667730a9b491129c42d45e11cdb7043b702e97125bcec48a1"},
|
||||
{file = "ruff-0.1.11-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:231d8fb11b2cc7c0366a326a66dafc6ad449d7fcdbc268497ee47e1334f66f77"},
|
||||
{file = "ruff-0.1.11-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:09c415716884950080921dd6237767e52e227e397e2008e2bed410117679975b"},
|
||||
{file = "ruff-0.1.11-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:0f58948c6d212a6b8d41cd59e349751018797ce1727f961c2fa755ad6208ba45"},
|
||||
{file = "ruff-0.1.11-py3-none-musllinux_1_2_i686.whl", hash = "sha256:190a566c8f766c37074d99640cd9ca3da11d8deae2deae7c9505e68a4a30f740"},
|
||||
{file = "ruff-0.1.11-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:6464289bd67b2344d2a5d9158d5eb81025258f169e69a46b741b396ffb0cda95"},
|
||||
{file = "ruff-0.1.11-py3-none-win32.whl", hash = "sha256:9b8f397902f92bc2e70fb6bebfa2139008dc72ae5177e66c383fa5426cb0bf2c"},
|
||||
{file = "ruff-0.1.11-py3-none-win_amd64.whl", hash = "sha256:eb85ee287b11f901037a6683b2374bb0ec82928c5cbc984f575d0437979c521a"},
|
||||
{file = "ruff-0.1.11-py3-none-win_arm64.whl", hash = "sha256:97ce4d752f964ba559c7023a86e5f8e97f026d511e48013987623915431c7ea9"},
|
||||
{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"
|
||||
@@ -7850,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"
|
||||
@@ -8113,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"]
|
||||
@@ -8125,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"]
|
||||
@@ -8145,4 +8260,4 @@ youtube = ["youtube-transcript-api", "yt_dlp"]
|
||||
[metadata]
|
||||
lock-version = "2.0"
|
||||
python-versions = ">=3.9,<3.12"
|
||||
content-hash = "8def3cb3aa4737793eaacd9358c092e0331f001044f5cacca513fc47faf44b06"
|
||||
content-hash = "a16addd3362ae70c79b15677c6815f708677f11f636093f4e1f5084ba44b5a36"
|
||||
|
||||
+10
-5
@@ -1,7 +1,7 @@
|
||||
[tool.poetry]
|
||||
name = "embedchain"
|
||||
version = "0.1.59"
|
||||
description = "Data platform for LLMs - Load, index, retrieve and sync any unstructured data"
|
||||
version = "0.1.72"
|
||||
description = "Simplest open source retrieval(RAG) framework"
|
||||
authors = [
|
||||
"Taranjeet Singh <taranjeet@embedchain.ai>",
|
||||
"Deshraj Yadav <deshraj@embedchain.ai>",
|
||||
@@ -22,7 +22,7 @@ build-backend = "poetry.core.masonry.api"
|
||||
requires = ["poetry-core"]
|
||||
|
||||
[tool.ruff]
|
||||
select = ["E", "F"]
|
||||
select = ["ASYNC", "E", "F"]
|
||||
ignore = []
|
||||
fixable = ["ALL"]
|
||||
unfixable = []
|
||||
@@ -102,6 +102,7 @@ rich = "^13.7.0"
|
||||
beautifulsoup4 = "^4.12.2"
|
||||
pypdf = "^3.11.0"
|
||||
gptcache = "^0.1.43"
|
||||
pysbd = "^0.3.4"
|
||||
tiktoken = { version = "^0.4.0", optional = true }
|
||||
youtube-transcript-api = { version = "^0.6.1", optional = true }
|
||||
pytube = { version = "^15.0.0", optional = true }
|
||||
@@ -122,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 }
|
||||
@@ -148,11 +149,13 @@ 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"
|
||||
pre-commit = "^3.2.2"
|
||||
ruff = "^0.0.220"
|
||||
ruff = "^0.1.11"
|
||||
pytest = "^7.3.1"
|
||||
pytest-mock = "^3.10.0"
|
||||
pytest-env = "^0.8.1"
|
||||
@@ -213,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
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
@@ -44,6 +44,10 @@ def test_get(chat_memory_instance):
|
||||
|
||||
assert len(recent_memories) == 5
|
||||
|
||||
all_memories = chat_memory_instance.get(app_id, fetch_all=True)
|
||||
|
||||
assert len(all_memories) == 6
|
||||
|
||||
|
||||
def test_delete_chat_history(chat_memory_instance):
|
||||
app_id = "test_app"
|
||||
@@ -59,9 +63,26 @@ def test_delete_chat_history(chat_memory_instance):
|
||||
|
||||
chat_memory_instance.add(app_id, session_id, chat_message)
|
||||
|
||||
session_id_2 = "test_session_2"
|
||||
|
||||
for i in range(1, 6):
|
||||
human_message = f"Question {i}"
|
||||
ai_message = f"Answer {i}"
|
||||
|
||||
chat_message = ChatMessage()
|
||||
chat_message.add_user_message(human_message)
|
||||
chat_message.add_ai_message(ai_message)
|
||||
|
||||
chat_memory_instance.add(app_id, session_id_2, chat_message)
|
||||
|
||||
chat_memory_instance.delete(app_id, session_id)
|
||||
|
||||
assert chat_memory_instance.count(app_id, session_id) == 0
|
||||
assert chat_memory_instance.count(app_id) == 5
|
||||
|
||||
chat_memory_instance.delete(app_id)
|
||||
|
||||
assert chat_memory_instance.count(app_id) == 0
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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": {
|
||||
|
||||
+206
-87
@@ -1,106 +1,225 @@
|
||||
from unittest import mock
|
||||
from unittest.mock import patch
|
||||
import pytest
|
||||
|
||||
from embedchain import App
|
||||
from embedchain.config import AppConfig
|
||||
from embedchain.embedder.base import BaseEmbedder
|
||||
from embedchain.config.vectordb.pinecone import PineconeDBConfig
|
||||
from embedchain.vectordb.pinecone import PineconeDB
|
||||
|
||||
|
||||
class TestPinecone:
|
||||
@patch("embedchain.vectordb.pinecone.pinecone")
|
||||
def test_init(self, pinecone_mock):
|
||||
"""Test that the PineconeDB can be initialized."""
|
||||
# Create a PineconeDB instance
|
||||
PineconeDB()
|
||||
@pytest.fixture
|
||||
def pinecone_pod_config():
|
||||
return PineconeDBConfig(
|
||||
collection_name="test_collection",
|
||||
api_key="test_api_key",
|
||||
vector_dimension=3,
|
||||
pod_config={"environment": "test_environment", "metadata_config": {"indexed": ["*"]}},
|
||||
)
|
||||
|
||||
# Assert that the Pinecone client was initialized
|
||||
pinecone_mock.init.assert_called_once()
|
||||
pinecone_mock.list_indexes.assert_called_once()
|
||||
pinecone_mock.Index.assert_called_once()
|
||||
|
||||
@patch("embedchain.vectordb.pinecone.pinecone")
|
||||
def test_set_embedder(self, pinecone_mock):
|
||||
"""Test that the embedder can be set."""
|
||||
@pytest.fixture
|
||||
def pinecone_serverless_config():
|
||||
return PineconeDBConfig(
|
||||
collection_name="test_collection",
|
||||
api_key="test_api_key",
|
||||
vector_dimension=3,
|
||||
serverless_config={
|
||||
"cloud": "test_cloud",
|
||||
"region": "test_region",
|
||||
},
|
||||
)
|
||||
|
||||
# Set the embedder
|
||||
embedder = BaseEmbedder()
|
||||
|
||||
# Create a PineconeDB instance
|
||||
def test_pinecone_init_without_config(monkeypatch):
|
||||
monkeypatch.setenv("PINECONE_API_KEY", "test_api_key")
|
||||
monkeypatch.setattr("embedchain.vectordb.pinecone.PineconeDB._setup_pinecone_index", lambda x: x)
|
||||
monkeypatch.setattr("embedchain.vectordb.pinecone.PineconeDB._get_or_create_db", lambda x: x)
|
||||
pinecone_db = PineconeDB()
|
||||
|
||||
assert isinstance(pinecone_db, PineconeDB)
|
||||
assert isinstance(pinecone_db.config, PineconeDBConfig)
|
||||
assert pinecone_db.config.pod_config == {"environment": "gcp-starter", "metadata_config": {"indexed": ["*"]}}
|
||||
monkeypatch.delenv("PINECONE_API_KEY")
|
||||
|
||||
|
||||
def test_pinecone_init_with_config(pinecone_pod_config, pinecone_serverless_config, monkeypatch):
|
||||
monkeypatch.setattr("embedchain.vectordb.pinecone.PineconeDB._setup_pinecone_index", lambda x: x)
|
||||
monkeypatch.setattr("embedchain.vectordb.pinecone.PineconeDB._get_or_create_db", lambda x: x)
|
||||
pinecone_db = PineconeDB(config=pinecone_pod_config)
|
||||
|
||||
assert isinstance(pinecone_db, PineconeDB)
|
||||
assert isinstance(pinecone_db.config, PineconeDBConfig)
|
||||
|
||||
assert pinecone_db.config.pod_config == pinecone_pod_config.pod_config
|
||||
|
||||
pinecone_db = PineconeDB(config=pinecone_pod_config)
|
||||
|
||||
assert isinstance(pinecone_db, PineconeDB)
|
||||
assert isinstance(pinecone_db.config, PineconeDBConfig)
|
||||
|
||||
assert pinecone_db.config.serverless_config == pinecone_pod_config.serverless_config
|
||||
|
||||
|
||||
class MockListIndexes:
|
||||
def names(self):
|
||||
return ["test_collection"]
|
||||
|
||||
|
||||
class MockPineconeIndex:
|
||||
db = []
|
||||
|
||||
def __init__(*args, **kwargs):
|
||||
pass
|
||||
|
||||
def upsert(self, chunk, **kwargs):
|
||||
self.db.extend([c for c in chunk])
|
||||
return
|
||||
|
||||
def delete(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
def query(self, *args, **kwargs):
|
||||
return {
|
||||
"matches": [
|
||||
{
|
||||
"metadata": {
|
||||
"key": "value",
|
||||
"text": "text_1",
|
||||
},
|
||||
"score": 0.1,
|
||||
},
|
||||
{
|
||||
"metadata": {
|
||||
"key": "value",
|
||||
"text": "text_2",
|
||||
},
|
||||
"score": 0.2,
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
def fetch(self, *args, **kwargs):
|
||||
return {
|
||||
"vectors": {
|
||||
"key_1": {
|
||||
"metadata": {
|
||||
"source": "1",
|
||||
}
|
||||
},
|
||||
"key_2": {
|
||||
"metadata": {
|
||||
"source": "2",
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
def describe_index_stats(self, *args, **kwargs):
|
||||
return {"total_vector_count": len(self.db)}
|
||||
|
||||
|
||||
class MockPineconeClient:
|
||||
def __init__(*args, **kwargs):
|
||||
pass
|
||||
|
||||
def list_indexes(self):
|
||||
return MockListIndexes()
|
||||
|
||||
def create_index(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
def Index(self, *args, **kwargs):
|
||||
return MockPineconeIndex()
|
||||
|
||||
def delete_index(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
|
||||
class MockPinecone:
|
||||
def __init__(*args, **kwargs):
|
||||
pass
|
||||
|
||||
def Pinecone(*args, **kwargs):
|
||||
return MockPineconeClient()
|
||||
|
||||
def PodSpec(*args, **kwargs):
|
||||
pass
|
||||
|
||||
def ServerlessSpec(*args, **kwargs):
|
||||
pass
|
||||
|
||||
|
||||
class MockEmbedder:
|
||||
def embedding_fn(self, documents):
|
||||
return [[1, 1, 1] for d in documents]
|
||||
|
||||
|
||||
def test_setup_pinecone_index(pinecone_pod_config, pinecone_serverless_config, monkeypatch):
|
||||
monkeypatch.setattr("embedchain.vectordb.pinecone.pinecone", MockPinecone)
|
||||
monkeypatch.setenv("PINECONE_API_KEY", "test_api_key")
|
||||
pinecone_db = PineconeDB(config=pinecone_pod_config)
|
||||
pinecone_db._setup_pinecone_index()
|
||||
|
||||
assert pinecone_db.client is not None
|
||||
assert pinecone_db.config.index_name == "test-collection-3"
|
||||
assert pinecone_db.client.list_indexes().names() == ["test_collection"]
|
||||
assert pinecone_db.pinecone_index is not None
|
||||
|
||||
pinecone_db = PineconeDB(config=pinecone_serverless_config)
|
||||
pinecone_db._setup_pinecone_index()
|
||||
|
||||
assert pinecone_db.client is not None
|
||||
assert pinecone_db.config.index_name == "test-collection-3"
|
||||
assert pinecone_db.client.list_indexes().names() == ["test_collection"]
|
||||
assert pinecone_db.pinecone_index is not None
|
||||
|
||||
|
||||
def test_get(monkeypatch):
|
||||
def mock_pinecone_db():
|
||||
monkeypatch.setenv("PINECONE_API_KEY", "test_api_key")
|
||||
monkeypatch.setattr("embedchain.vectordb.pinecone.PineconeDB._setup_pinecone_index", lambda x: x)
|
||||
monkeypatch.setattr("embedchain.vectordb.pinecone.PineconeDB._get_or_create_db", lambda x: x)
|
||||
db = PineconeDB()
|
||||
app_config = AppConfig(collect_metrics=False)
|
||||
App(config=app_config, db=db, embedding_model=embedder)
|
||||
db.pinecone_index = MockPineconeIndex()
|
||||
return db
|
||||
|
||||
# Assert that the embedder was set
|
||||
assert db.embedder == embedder
|
||||
pinecone_mock.init.assert_called_once()
|
||||
pinecone_db = mock_pinecone_db()
|
||||
ids = pinecone_db.get(["key_1", "key_2"])
|
||||
assert ids == {"ids": ["key_1", "key_2"], "metadatas": [{"source": "1"}, {"source": "2"}]}
|
||||
|
||||
@patch("embedchain.vectordb.pinecone.pinecone")
|
||||
def test_add_documents(self, pinecone_mock):
|
||||
"""Test that documents can be added to the database."""
|
||||
pinecone_client_mock = pinecone_mock.Index.return_value
|
||||
|
||||
embedding_function = mock.Mock()
|
||||
base_embedder = BaseEmbedder()
|
||||
base_embedder.set_embedding_fn(embedding_function)
|
||||
vectors = [[0, 0, 0], [1, 1, 1]]
|
||||
embedding_function.return_value = vectors
|
||||
# Create a PineconeDb instance
|
||||
def test_add(monkeypatch):
|
||||
def mock_pinecone_db():
|
||||
monkeypatch.setenv("PINECONE_API_KEY", "test_api_key")
|
||||
monkeypatch.setattr("embedchain.vectordb.pinecone.PineconeDB._setup_pinecone_index", lambda x: x)
|
||||
monkeypatch.setattr("embedchain.vectordb.pinecone.PineconeDB._get_or_create_db", lambda x: x)
|
||||
db = PineconeDB()
|
||||
app_config = AppConfig(collect_metrics=False)
|
||||
App(config=app_config, db=db, embedding_model=base_embedder)
|
||||
db.pinecone_index = MockPineconeIndex()
|
||||
db._set_embedder(MockEmbedder())
|
||||
return db
|
||||
|
||||
# Add some documents to the database
|
||||
documents = ["This is a document.", "This is another document."]
|
||||
metadatas = [{}, {}]
|
||||
ids = ["doc1", "doc2"]
|
||||
db.add(vectors, documents, metadatas, ids)
|
||||
pinecone_db = mock_pinecone_db()
|
||||
pinecone_db.add(["text_1", "text_2"], [{"key_1": "value_1"}, {"key_2": "value_2"}], ["key_1", "key_2"])
|
||||
assert pinecone_db.count() == 2
|
||||
|
||||
expected_pinecone_upsert_args = [
|
||||
{"id": "doc1", "values": [0, 0, 0], "metadata": {"text": "This is a document."}},
|
||||
{"id": "doc2", "values": [1, 1, 1], "metadata": {"text": "This is another document."}},
|
||||
]
|
||||
# Assert that the Pinecone client was called to upsert the documents
|
||||
pinecone_client_mock.upsert.assert_called_once_with(tuple(expected_pinecone_upsert_args))
|
||||
pinecone_db.add(["text_3", "text_4"], [{"key_3": "value_3"}, {"key_4": "value_4"}], ["key_3", "key_4"])
|
||||
assert pinecone_db.count() == 4
|
||||
|
||||
@patch("embedchain.vectordb.pinecone.pinecone")
|
||||
def test_query_documents(self, pinecone_mock):
|
||||
"""Test that documents can be queried from the database."""
|
||||
pinecone_client_mock = pinecone_mock.Index.return_value
|
||||
|
||||
embedding_function = mock.Mock()
|
||||
base_embedder = BaseEmbedder()
|
||||
base_embedder.set_embedding_fn(embedding_function)
|
||||
vectors = [[0, 0, 0]]
|
||||
embedding_function.return_value = vectors
|
||||
# Create a PineconeDB instance
|
||||
def test_query(monkeypatch):
|
||||
def mock_pinecone_db():
|
||||
monkeypatch.setenv("PINECONE_API_KEY", "test_api_key")
|
||||
monkeypatch.setattr("embedchain.vectordb.pinecone.PineconeDB._setup_pinecone_index", lambda x: x)
|
||||
monkeypatch.setattr("embedchain.vectordb.pinecone.PineconeDB._get_or_create_db", lambda x: x)
|
||||
db = PineconeDB()
|
||||
app_config = AppConfig(collect_metrics=False)
|
||||
App(config=app_config, db=db, embedding_model=base_embedder)
|
||||
db.pinecone_index = MockPineconeIndex()
|
||||
db._set_embedder(MockEmbedder())
|
||||
return db
|
||||
|
||||
# Query the database for documents that are similar to "document"
|
||||
input_query = ["document"]
|
||||
n_results = 1
|
||||
db.query(input_query, n_results, where={})
|
||||
|
||||
# Assert that the Pinecone client was called to query the database
|
||||
pinecone_client_mock.query.assert_called_once_with(
|
||||
vector=db.embedder.embedding_fn(input_query)[0], top_k=n_results, filter={}, include_metadata=True
|
||||
)
|
||||
|
||||
@patch("embedchain.vectordb.pinecone.pinecone")
|
||||
def test_reset(self, pinecone_mock):
|
||||
"""Test that the database can be reset."""
|
||||
# Create a PineconeDb instance
|
||||
db = PineconeDB()
|
||||
app_config = AppConfig(collect_metrics=False)
|
||||
App(config=app_config, db=db, embedding_model=BaseEmbedder())
|
||||
|
||||
# Reset the database
|
||||
db.reset()
|
||||
|
||||
# Assert that the Pinecone client was called to delete the index
|
||||
pinecone_mock.delete_index.assert_called_once_with(db.index_name)
|
||||
|
||||
# Assert that the index is recreated
|
||||
pinecone_mock.Index.assert_called_with(db.index_name)
|
||||
pinecone_db = mock_pinecone_db()
|
||||
# without citations
|
||||
results = pinecone_db.query(["text_1", "text_2"], n_results=2, where={})
|
||||
assert results == ["text_1", "text_2"]
|
||||
# with citations
|
||||
results = pinecone_db.query(["text_1", "text_2"], n_results=2, where={}, citations=True)
|
||||
assert results == [
|
||||
("text_1", {"key": "value", "text": "text_1", "score": 0.1}),
|
||||
("text_2", {"key": "value", "text": "text_2", "score": 0.2}),
|
||||
]
|
||||
|
||||
@@ -56,9 +56,9 @@ class TestQdrantDB(unittest.TestCase):
|
||||
App(config=app_config, db=db, embedding_model=embedder)
|
||||
|
||||
resp = db.get(ids=[], where={})
|
||||
self.assertEqual(resp, {"ids": []})
|
||||
self.assertEqual(resp, {"ids": [], "metadatas": []})
|
||||
resp2 = db.get(ids=["123", "456"], where={"url": "https://ai.ai"})
|
||||
self.assertEqual(resp2, {"ids": []})
|
||||
self.assertEqual(resp2, {"ids": [], "metadatas": []})
|
||||
|
||||
@patch("embedchain.vectordb.qdrant.QdrantClient")
|
||||
@patch.object(uuid, "uuid4", side_effect=TEST_UUIDS)
|
||||
@@ -75,11 +75,10 @@ class TestQdrantDB(unittest.TestCase):
|
||||
app_config = AppConfig(collect_metrics=False)
|
||||
App(config=app_config, db=db, embedding_model=embedder)
|
||||
|
||||
embeddings = [[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]]
|
||||
documents = ["This is a test document.", "This is another test document."]
|
||||
metadatas = [{}, {}]
|
||||
ids = ["123", "456"]
|
||||
db.add(embeddings, documents, metadatas, ids)
|
||||
db.add(documents, metadatas, ids)
|
||||
qdrant_client_mock.return_value.upsert.assert_called_once_with(
|
||||
collection_name="embedchain-store-1526",
|
||||
points=Batch(
|
||||
@@ -96,7 +95,7 @@ class TestQdrantDB(unittest.TestCase):
|
||||
"metadata": {"text": "This is another test document."},
|
||||
},
|
||||
],
|
||||
vectors=embeddings,
|
||||
vectors=[[1, 2, 3], [4, 5, 6]],
|
||||
),
|
||||
)
|
||||
|
||||
@@ -120,7 +119,7 @@ class TestQdrantDB(unittest.TestCase):
|
||||
query_filter=models.Filter(
|
||||
must=[
|
||||
models.FieldCondition(
|
||||
key="payload.metadata.doc_id",
|
||||
key="metadata.doc_id",
|
||||
match=models.MatchValue(
|
||||
value="123",
|
||||
),
|
||||
|
||||
@@ -29,7 +29,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
weaviate_client_schema_mock.exists.return_value = False
|
||||
# Set the embedder
|
||||
embedder = BaseEmbedder()
|
||||
embedder.set_vector_dimension(1526)
|
||||
embedder.set_vector_dimension(1536)
|
||||
embedder.set_embedding_fn(mock_embedding_fn)
|
||||
|
||||
# Create a Weaviate instance
|
||||
@@ -40,7 +40,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
expected_class_obj = {
|
||||
"classes": [
|
||||
{
|
||||
"class": "Embedchain_store_1526",
|
||||
"class": "Embedchain_store_1536",
|
||||
"vectorizer": "none",
|
||||
"properties": [
|
||||
{
|
||||
@@ -53,12 +53,12 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
},
|
||||
{
|
||||
"name": "metadata",
|
||||
"dataType": ["Embedchain_store_1526_metadata"],
|
||||
"dataType": ["Embedchain_store_1536_metadata"],
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"class": "Embedchain_store_1526_metadata",
|
||||
"class": "Embedchain_store_1536_metadata",
|
||||
"vectorizer": "none",
|
||||
"properties": [
|
||||
{
|
||||
@@ -88,7 +88,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
|
||||
# Assert that the Weaviate client was initialized
|
||||
weaviate_mock.Client.assert_called_once()
|
||||
self.assertEqual(db.index_name, "Embedchain_store_1526")
|
||||
self.assertEqual(db.index_name, "Embedchain_store_1536")
|
||||
weaviate_client_schema_mock.create.assert_called_once_with(expected_class_obj)
|
||||
|
||||
@patch("embedchain.vectordb.weaviate.weaviate")
|
||||
@@ -97,7 +97,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
weaviate_client_mock = weaviate_mock.Client.return_value
|
||||
|
||||
embedder = BaseEmbedder()
|
||||
embedder.set_vector_dimension(1526)
|
||||
embedder.set_vector_dimension(1536)
|
||||
embedder.set_embedding_fn(mock_embedding_fn)
|
||||
|
||||
# Create a Weaviate instance
|
||||
@@ -117,7 +117,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
|
||||
# Set the embedder
|
||||
embedder = BaseEmbedder()
|
||||
embedder.set_vector_dimension(1526)
|
||||
embedder.set_vector_dimension(1536)
|
||||
embedder.set_embedding_fn(mock_embedding_fn)
|
||||
|
||||
# Create a Weaviate instance
|
||||
@@ -126,30 +126,21 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
App(config=app_config, db=db, embedding_model=embedder)
|
||||
db.BATCH_SIZE = 1
|
||||
|
||||
embeddings = [[1, 2, 3], [4, 5, 6]]
|
||||
documents = ["This is a test document.", "This is another test document."]
|
||||
metadatas = [None, None]
|
||||
ids = ["123", "456"]
|
||||
db.add(embeddings, documents, metadatas, ids)
|
||||
documents = ["This is test document"]
|
||||
metadatas = [None]
|
||||
ids = ["id_1"]
|
||||
db.add(documents, metadatas, ids)
|
||||
|
||||
# Check if the document was added to the database.
|
||||
weaviate_client_batch_mock.configure.assert_called_once_with(batch_size=1, timeout_retries=3)
|
||||
weaviate_client_batch_enter_mock.add_data_object.assert_any_call(
|
||||
data_object={"text": documents[0]}, class_name="Embedchain_store_1526_metadata", vector=embeddings[0]
|
||||
)
|
||||
weaviate_client_batch_enter_mock.add_data_object.assert_any_call(
|
||||
data_object={"text": documents[1]}, class_name="Embedchain_store_1526_metadata", vector=embeddings[1]
|
||||
data_object={"text": documents[0]}, class_name="Embedchain_store_1536_metadata", vector=[1, 2, 3]
|
||||
)
|
||||
|
||||
weaviate_client_batch_enter_mock.add_data_object.assert_any_call(
|
||||
data_object={"identifier": ids[0], "text": documents[0]},
|
||||
class_name="Embedchain_store_1526",
|
||||
vector=embeddings[0],
|
||||
)
|
||||
weaviate_client_batch_enter_mock.add_data_object.assert_any_call(
|
||||
data_object={"identifier": ids[1], "text": documents[1]},
|
||||
class_name="Embedchain_store_1526",
|
||||
vector=embeddings[1],
|
||||
data_object={"text": documents[0]},
|
||||
class_name="Embedchain_store_1536_metadata",
|
||||
vector=[1, 2, 3],
|
||||
)
|
||||
|
||||
@patch("embedchain.vectordb.weaviate.weaviate")
|
||||
@@ -161,7 +152,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
|
||||
# Set the embedder
|
||||
embedder = BaseEmbedder()
|
||||
embedder.set_vector_dimension(1526)
|
||||
embedder.set_vector_dimension(1536)
|
||||
embedder.set_embedding_fn(mock_embedding_fn)
|
||||
|
||||
# Create a Weaviate instance
|
||||
@@ -172,7 +163,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
# Query for the document.
|
||||
db.query(input_query=["This is a test document."], n_results=1, where={})
|
||||
|
||||
weaviate_client_query_mock.get.assert_called_once_with("Embedchain_store_1526", ["text"])
|
||||
weaviate_client_query_mock.get.assert_called_once_with("Embedchain_store_1536", ["text"])
|
||||
weaviate_client_query_get_mock.with_near_vector.assert_called_once_with({"vector": [1, 2, 3]})
|
||||
|
||||
@patch("embedchain.vectordb.weaviate.weaviate")
|
||||
@@ -185,7 +176,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
|
||||
# Set the embedder
|
||||
embedder = BaseEmbedder()
|
||||
embedder.set_vector_dimension(1526)
|
||||
embedder.set_vector_dimension(1536)
|
||||
embedder.set_embedding_fn(mock_embedding_fn)
|
||||
|
||||
# Create a Weaviate instance
|
||||
@@ -196,9 +187,9 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
# Query for the document.
|
||||
db.query(input_query=["This is a test document."], n_results=1, where={"doc_id": "123"})
|
||||
|
||||
weaviate_client_query_mock.get.assert_called_once_with("Embedchain_store_1526", ["text"])
|
||||
weaviate_client_query_mock.get.assert_called_once_with("Embedchain_store_1536", ["text"])
|
||||
weaviate_client_query_get_mock.with_where.assert_called_once_with(
|
||||
{"operator": "Equal", "path": ["metadata", "Embedchain_store_1526_metadata", "doc_id"], "valueText": "123"}
|
||||
{"operator": "Equal", "path": ["metadata", "Embedchain_store_1536_metadata", "doc_id"], "valueText": "123"}
|
||||
)
|
||||
weaviate_client_query_get_where_mock.with_near_vector.assert_called_once_with({"vector": [1, 2, 3]})
|
||||
|
||||
@@ -210,7 +201,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
|
||||
# Set the embedder
|
||||
embedder = BaseEmbedder()
|
||||
embedder.set_vector_dimension(1526)
|
||||
embedder.set_vector_dimension(1536)
|
||||
embedder.set_embedding_fn(mock_embedding_fn)
|
||||
|
||||
# Create a Weaviate instance
|
||||
@@ -222,7 +213,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
db.reset()
|
||||
|
||||
weaviate_client_batch_mock.delete_objects.assert_called_once_with(
|
||||
"Embedchain_store_1526", where={"path": ["identifier"], "operator": "Like", "valueText": ".*"}
|
||||
"Embedchain_store_1536", where={"path": ["identifier"], "operator": "Like", "valueText": ".*"}
|
||||
)
|
||||
|
||||
@patch("embedchain.vectordb.weaviate.weaviate")
|
||||
@@ -233,7 +224,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
|
||||
# Set the embedder
|
||||
embedder = BaseEmbedder()
|
||||
embedder.set_vector_dimension(1526)
|
||||
embedder.set_vector_dimension(1536)
|
||||
embedder.set_embedding_fn(mock_embedding_fn)
|
||||
|
||||
# Create a Weaviate instance
|
||||
@@ -244,4 +235,4 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
# Reset the database.
|
||||
db.count()
|
||||
|
||||
weaviate_client_query.aggregate.assert_called_once_with("Embedchain_store_1526")
|
||||
weaviate_client_query.aggregate.assert_called_once_with("Embedchain_store_1536")
|
||||
|
||||
@@ -130,7 +130,11 @@ class TestZillizDBCollection:
|
||||
[
|
||||
{
|
||||
"distance": 0.0,
|
||||
"entity": {"text": "result_doc", "url": "url_1", "doc_id": "doc_id_1", "embeddings": [1, 2, 3]},
|
||||
"entity": {
|
||||
"text": "result_doc",
|
||||
"embeddings": [1, 2, 3],
|
||||
"metadata": {"url": "url_1", "doc_id": "doc_id_1"},
|
||||
},
|
||||
}
|
||||
]
|
||||
]
|
||||
@@ -141,6 +145,7 @@ class TestZillizDBCollection:
|
||||
mock_search.assert_called_with(
|
||||
collection_name=mock_config.collection_name,
|
||||
data=["query_vector"],
|
||||
filter="",
|
||||
limit=1,
|
||||
output_fields=["*"],
|
||||
)
|
||||
@@ -155,10 +160,9 @@ class TestZillizDBCollection:
|
||||
mock_search.assert_called_with(
|
||||
collection_name=mock_config.collection_name,
|
||||
data=["query_vector"],
|
||||
filter="",
|
||||
limit=1,
|
||||
output_fields=["*"],
|
||||
)
|
||||
|
||||
assert query_result_with_citations == [
|
||||
("result_doc", {"text": "result_doc", "url": "url_1", "doc_id": "doc_id_1", "score": 0.0})
|
||||
]
|
||||
assert query_result_with_citations == [("result_doc", {"url": "url_1", "doc_id": "doc_id_1", "score": 0.0})]
|
||||
|
||||
Reference in New Issue
Block a user