Compare commits
13 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| ebbf90f4aa | |||
| 4f119692f1 | |||
| bbe56107fb | |||
| bd654e7aac | |||
| 33500a7ce2 | |||
| 4880557d51 | |||
| ea09b5f7f0 | |||
| 5258fd91ea | |||
| b305d674de | |||
| 7c24601d0f | |||
| 50c0285cb2 | |||
| 0a78198bb5 | |||
| edaeb78ccf |
@@ -11,7 +11,7 @@ install:
|
||||
|
||||
install_all:
|
||||
poetry install --all-extras
|
||||
poetry run pip install pinecone-text pinecone-client langchain-anthropic "unstructured[local-inference, all-docs]" ollama deepgram-sdk==3.2.7 langchain-huggingface
|
||||
poetry run pip install pinecone-text pinecone-client langchain-anthropic "unstructured[local-inference, all-docs]" ollama langchain_together==0.1.3 langchain_cohere==0.1.5 deepgram-sdk==3.2.7 langchain-huggingface psutil
|
||||
|
||||
install_es:
|
||||
poetry install --extras elasticsearch
|
||||
|
||||
@@ -30,6 +30,7 @@ llm:
|
||||
response_format:
|
||||
type: json_object
|
||||
api_version: 2024-02-01
|
||||
http_client_proxies: http://testproxy.mem0.net:8000
|
||||
prompt: |
|
||||
Use the following pieces of context to answer the query at the end.
|
||||
If you don't know the answer, just say that you don't know, don't try to make up an answer.
|
||||
@@ -89,7 +90,8 @@ cache:
|
||||
"system_prompt": "Act as William Shakespeare. Answer the following questions in the style of William Shakespeare.",
|
||||
"api_key": "sk-xxx",
|
||||
"model_kwargs": {"response_format": {"type": "json_object"}},
|
||||
"api_version": "2024-02-01"
|
||||
"api_version": "2024-02-01",
|
||||
"http_client_proxies": "http://testproxy.mem0.net:8000",
|
||||
}
|
||||
},
|
||||
"vectordb": {
|
||||
@@ -150,7 +152,8 @@ config = {
|
||||
"Act as William Shakespeare. Answer the following questions in the style of William Shakespeare."
|
||||
),
|
||||
'api_key': 'sk-xxx',
|
||||
"model_kwargs": {"response_format": {"type": "json_object"}}
|
||||
"model_kwargs": {"response_format": {"type": "json_object"}},
|
||||
"http_client_proxies": "http://testproxy.mem0.net:8000",
|
||||
}
|
||||
},
|
||||
'vectordb': {
|
||||
@@ -206,17 +209,21 @@ Alright, let's dive into what each key means in the yaml config above:
|
||||
- `top_p` (Float): Controls the diversity of word selection. A higher value (closer to 1) makes word selection more diverse.
|
||||
- `stream` (Boolean): Controls if the response is streamed back to the user (set to false).
|
||||
- `online` (Boolean): Controls whether to use internet to get more context for answering query (set to false).
|
||||
- `token_usage` (Boolean): Controls whether to use token usage for the querying models (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.
|
||||
- `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.
|
||||
- `http_client_proxies` (Dict | String): The proxy server settings used to create `self.http_client` using `httpx.Client(proxies=http_client_proxies)`
|
||||
- `http_async_client_proxies` (Dict | String): The proxy server settings for async calls used to create `self.http_async_client` using `httpx.AsyncClient(proxies=http_async_client_proxies)`
|
||||
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`:
|
||||
- `collection_name` (String): The initial collection name for the vectordb, set to 'full-stack-app'.
|
||||
- `dir` (String): The directory for the local database, set to 'db'.
|
||||
- `allow_reset` (Boolean): Indicates whether resetting the vectordb is allowed, set to true.
|
||||
- `batch_size` (Integer): The batch size for docs insertion in vectordb, defaults to `100`
|
||||
<Note>We recommend you to checkout vectordb specific config [here](https://docs.embedchain.ai/components/vector-databases)</Note>
|
||||
4. `embedder` Section:
|
||||
- `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).
|
||||
@@ -228,6 +235,7 @@ Alright, let's dive into what each key means in the yaml config above:
|
||||
- `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.
|
||||
- `model_kwargs` (Dict): Used to pass extra arguments to embedders.
|
||||
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.
|
||||
@@ -241,6 +249,9 @@ Alright, let's dive into what each key means in the yaml config above:
|
||||
- `config` (Optional): The config for initializing the cache. If not provided, sensible default values are used as mentioned below.
|
||||
- `similarity_threshold` (Float): The threshold for similarity evaluation. Defaults to `0.8`.
|
||||
- `auto_flush` (Integer): The number of queries after which the cache is flushed. Defaults to `20`.
|
||||
7. `memory` Section: (Optional)
|
||||
- `api_key` (String): The API key of mem0.
|
||||
- `top_k` (Integer): The number of top-k results to return. Defaults to `10`.
|
||||
<Note>
|
||||
If you provide a cache section, the app will automatically configure and use a cache to store the results of the language model. This is useful if you want to speed up the response time and save inference cost of your app.
|
||||
</Note>
|
||||
|
||||
@@ -144,3 +144,28 @@ 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)
|
||||
```
|
||||
|
||||
### With Mem0 to store chat history
|
||||
|
||||
Mem0 is a cutting-edge long-term memory for LLMs to enable personalization for the GenAI stack. It enables LLMs to remember past interactions and provide more personalized responses.
|
||||
|
||||
Follow these steps to use Mem0 to enable memory for personalization in your apps:
|
||||
- Install the [`mem0`](https://docs.mem0.ai/) package using `pip install memzero`.
|
||||
- Get the api_key from [Mem0 Platform](https://app.mem0.ai/).
|
||||
- Provide api_key in config under `memory`, refer [Configurations](docs/api-reference/advanced/configuration.mdx).
|
||||
|
||||
```python with mem0
|
||||
from embedchain import App
|
||||
|
||||
config = {
|
||||
"memory": {
|
||||
"api_key": "m0-xxx",
|
||||
"top_k": 5
|
||||
}
|
||||
}
|
||||
|
||||
app = App.from_config(config=config)
|
||||
app.add("https://www.forbes.com/profile/elon-musk")
|
||||
|
||||
app.chat("What is the net worth of Elon Musk?")
|
||||
```
|
||||
@@ -192,6 +192,8 @@ embedder:
|
||||
provider: huggingface
|
||||
config:
|
||||
model: 'sentence-transformers/all-mpnet-base-v2'
|
||||
model_kwargs:
|
||||
trust_remote_code: True # Only use if you trust your embedder
|
||||
```
|
||||
|
||||
</CodeGroup>
|
||||
|
||||
@@ -840,6 +840,62 @@ answer = app.query("What is the net worth of Elon Musk today?")
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
## Token Usage
|
||||
|
||||
You can get the cost of the query by setting `token_usage` to `True` in the config file. This will return the token details: `prompt_tokens`, `completion_tokens`, `total_tokens`, `total_cost`, `cost_currency`.
|
||||
The list of paid LLMs that support token usage are:
|
||||
- OpenAI
|
||||
- Vertex AI
|
||||
- Anthropic
|
||||
- Cohere
|
||||
- Together
|
||||
- Groq
|
||||
- Mistral AI
|
||||
- NVIDIA AI
|
||||
|
||||
Here is an example of how to use token usage:
|
||||
<CodeGroup>
|
||||
|
||||
```python main.py
|
||||
os.environ["OPENAI_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?")
|
||||
# {'answer': 'Elon Musk's net worth is $209.9 billion as of 6/9/24.',
|
||||
# 'usage': {'prompt_tokens': 1228,
|
||||
# 'completion_tokens': 21,
|
||||
# 'total_tokens': 1249,
|
||||
# 'total_cost': 0.001884,
|
||||
# 'cost_currency': 'USD'}
|
||||
# }
|
||||
|
||||
|
||||
response = app.chat("Which companies did Elon Musk found?")
|
||||
# {'answer': 'Elon Musk founded six companies, including Tesla, which is an electric car maker, SpaceX, a rocket producer, and the Boring Company, a tunneling startup.',
|
||||
# 'usage': {'prompt_tokens': 1616,
|
||||
# 'completion_tokens': 34,
|
||||
# 'total_tokens': 1650,
|
||||
# 'total_cost': 0.002492,
|
||||
# 'cost_currency': 'USD'}
|
||||
# }
|
||||
```
|
||||
|
||||
```yaml config.yaml
|
||||
llm:
|
||||
provider: openai
|
||||
config:
|
||||
model: gpt-3.5-turbo
|
||||
temperature: 0.5
|
||||
max_tokens: 1000
|
||||
token_usage: true
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
If a model is missing and you'd like to add it to `model_prices_and_context_window.json`, please feel free to open a PR.
|
||||
|
||||
<br/ >
|
||||
|
||||
<Snippet file="missing-llm-tip.mdx" />
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 329 KiB |
@@ -0,0 +1,52 @@
|
||||
---
|
||||
title: "🧊 Helicone"
|
||||
description: "Implement Helicone, the open-source LLM observability platform, with Embedchain. Monitor, debug, and optimize your AI applications effortlessly."
|
||||
"twitter:title": "Helicone LLM Observability for Embedchain"
|
||||
---
|
||||
|
||||
Get started with [Helicone](https://www.helicone.ai/), the open-source LLM observability platform for developers to monitor, debug, and optimize their applications.
|
||||
|
||||
To use Helicone, you need to do the following steps.
|
||||
|
||||
## Integration Steps
|
||||
|
||||
<Steps>
|
||||
<Step title="Create an account + Generate an API Key">
|
||||
Log into [Helicone](https://www.helicone.ai) or create an account. Once you have an account, you
|
||||
can generate an [API key](https://helicone.ai/developer).
|
||||
|
||||
<Note>
|
||||
Make sure to generate a [write only API key](helicone-headers/helicone-auth).
|
||||
</Note>
|
||||
|
||||
</Step>
|
||||
<Step title="Set base_url in the your code">
|
||||
You can configure your base_url and OpenAI API key in your codebase
|
||||
<CodeGroup>
|
||||
|
||||
```python main.py
|
||||
import os
|
||||
from embedchain import App
|
||||
|
||||
# Modify the base path and add a Helicone URL
|
||||
os.environ["OPENAI_API_BASE"] = "https://oai.helicone.ai/{YOUR_HELICONE_API_KEY}/v1"
|
||||
# Add your OpenAI API Key
|
||||
os.environ["OPENAI_API_KEY"] = "{YOUR_OPENAI_API_KEY}"
|
||||
|
||||
app = App()
|
||||
|
||||
# Add data to your app
|
||||
app.add("https://en.wikipedia.org/wiki/Elon_Musk")
|
||||
|
||||
# Query your app
|
||||
print(app.query("How many companies did Elon found? Which companies?"))
|
||||
```
|
||||
|
||||
</CodeGroup>
|
||||
</Step>
|
||||
<Step title="Now you can see all passing requests through Embedchain in Helicone">
|
||||
<img src="/images/helicone-embedchain.png" alt="Embedchain requests" />
|
||||
</Step>
|
||||
</Steps>
|
||||
|
||||
Check out [Helicone](https://www.helicone.ai) to see more use cases!
|
||||
+14
-21
@@ -19,9 +19,7 @@
|
||||
"modeToggle": {
|
||||
"default": "dark"
|
||||
},
|
||||
"openapi": [
|
||||
"/rest-api.json"
|
||||
],
|
||||
"openapi": ["/rest-api.json"],
|
||||
"metadata": {
|
||||
"og:image": "/images/og.png",
|
||||
"twitter:site": "@embedchain"
|
||||
@@ -70,7 +68,8 @@
|
||||
"integration/langsmith",
|
||||
"integration/chainlit",
|
||||
"integration/streamlit-mistral",
|
||||
"integration/openlit"
|
||||
"integration/openlit",
|
||||
"integration/helicone"
|
||||
]
|
||||
}
|
||||
]
|
||||
@@ -132,13 +131,13 @@
|
||||
{
|
||||
"group": "🗄️ Vector databases",
|
||||
"pages": [
|
||||
"components/vector-databases/chromadb",
|
||||
"components/vector-databases/elasticsearch",
|
||||
"components/vector-databases/pinecone",
|
||||
"components/vector-databases/opensearch",
|
||||
"components/vector-databases/qdrant",
|
||||
"components/vector-databases/weaviate",
|
||||
"components/vector-databases/zilliz"
|
||||
"components/vector-databases/chromadb",
|
||||
"components/vector-databases/elasticsearch",
|
||||
"components/vector-databases/pinecone",
|
||||
"components/vector-databases/opensearch",
|
||||
"components/vector-databases/qdrant",
|
||||
"components/vector-databases/weaviate",
|
||||
"components/vector-databases/zilliz"
|
||||
]
|
||||
},
|
||||
"components/llms",
|
||||
@@ -161,9 +160,7 @@
|
||||
},
|
||||
{
|
||||
"group": "Community",
|
||||
"pages": [
|
||||
"community/connect-with-us"
|
||||
]
|
||||
"pages": ["community/connect-with-us"]
|
||||
},
|
||||
{
|
||||
"group": "Examples",
|
||||
@@ -203,9 +200,7 @@
|
||||
},
|
||||
{
|
||||
"group": "Showcase",
|
||||
"pages": [
|
||||
"examples/showcase"
|
||||
]
|
||||
"pages": ["examples/showcase"]
|
||||
},
|
||||
{
|
||||
"group": "API Reference",
|
||||
@@ -241,9 +236,7 @@
|
||||
},
|
||||
{
|
||||
"group": "Product",
|
||||
"pages": [
|
||||
"product/release-notes"
|
||||
]
|
||||
"pages": ["product/release-notes"]
|
||||
}
|
||||
],
|
||||
"footerSocials": {
|
||||
@@ -281,4 +274,4 @@
|
||||
"destination": "/get-started/introduction"
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
+21
-6
@@ -9,19 +9,24 @@ import requests
|
||||
import yaml
|
||||
from tqdm import tqdm
|
||||
|
||||
from embedchain.cache import (Config, ExactMatchEvaluation,
|
||||
SearchDistanceEvaluation, cache,
|
||||
gptcache_data_manager, gptcache_pre_function)
|
||||
from mem0 import Mem0
|
||||
from embedchain.cache import (
|
||||
Config,
|
||||
ExactMatchEvaluation,
|
||||
SearchDistanceEvaluation,
|
||||
cache,
|
||||
gptcache_data_manager,
|
||||
gptcache_pre_function,
|
||||
)
|
||||
from embedchain.client import Client
|
||||
from embedchain.config import AppConfig, CacheConfig, ChunkerConfig
|
||||
from embedchain.config import AppConfig, CacheConfig, ChunkerConfig, Mem0Config
|
||||
from embedchain.core.db.database import get_session, init_db, setup_engine
|
||||
from embedchain.core.db.models import DataSource
|
||||
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.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
|
||||
@@ -55,6 +60,7 @@ class App(EmbedChain):
|
||||
auto_deploy: bool = False,
|
||||
chunker: ChunkerConfig = None,
|
||||
cache_config: CacheConfig = None,
|
||||
memory_config: Mem0Config = None,
|
||||
log_level: int = logging.WARN,
|
||||
):
|
||||
"""
|
||||
@@ -95,6 +101,7 @@ class App(EmbedChain):
|
||||
self.id = None
|
||||
self.chunker = ChunkerConfig(**chunker) if chunker else None
|
||||
self.cache_config = cache_config
|
||||
self.memory_config = memory_config
|
||||
|
||||
self.config = config or AppConfig()
|
||||
self.name = self.config.name
|
||||
@@ -123,6 +130,11 @@ class App(EmbedChain):
|
||||
if self.cache_config is not None:
|
||||
self._init_cache()
|
||||
|
||||
# If memory_config is provided, initializing the memory ...
|
||||
self.mem0_client = None
|
||||
if self.memory_config is not None:
|
||||
self.mem0_client = Mem0(api_key=self.memory_config.api_key)
|
||||
|
||||
# Send anonymous telemetry
|
||||
self._telemetry_props = {"class": self.__class__.__name__}
|
||||
self.telemetry = AnonymousTelemetry(enabled=self.config.collect_metrics)
|
||||
@@ -365,11 +377,13 @@ class App(EmbedChain):
|
||||
app_config_data = config_data.get("app", {}).get("config", {})
|
||||
vector_db_config_data = config_data.get("vectordb", {})
|
||||
embedding_model_config_data = config_data.get("embedding_model", config_data.get("embedder", {}))
|
||||
memory_config_data = config_data.get("memory", {})
|
||||
llm_config_data = config_data.get("llm", {})
|
||||
chunker_config_data = config_data.get("chunker", {})
|
||||
cache_config_data = config_data.get("cache", None)
|
||||
|
||||
app_config = AppConfig(**app_config_data)
|
||||
memory_config = Mem0Config(**memory_config_data) if memory_config_data else None
|
||||
|
||||
vector_db_provider = vector_db_config_data.get("provider", "chroma")
|
||||
vector_db = VectorDBFactory.create(vector_db_provider, vector_db_config_data.get("config", {}))
|
||||
@@ -403,6 +417,7 @@ class App(EmbedChain):
|
||||
auto_deploy=auto_deploy,
|
||||
chunker=chunker_config_data,
|
||||
cache_config=cache_config,
|
||||
memory_config=memory_config,
|
||||
)
|
||||
|
||||
def _eval(self, dataset: list[EvalData], metric: Union[BaseMetric, str]):
|
||||
|
||||
@@ -12,3 +12,4 @@ from .vectordb.chroma import ChromaDbConfig
|
||||
from .vectordb.elasticsearch import ElasticsearchDBConfig
|
||||
from .vectordb.opensearch import OpenSearchDBConfig
|
||||
from .vectordb.zilliz import ZillizDBConfig
|
||||
from .mem0_config import Mem0Config
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Optional
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
|
||||
@@ -13,6 +13,7 @@ class BaseEmbedderConfig:
|
||||
endpoint: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
model_kwargs: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
"""
|
||||
Initialize a new instance of an embedder config class.
|
||||
@@ -29,6 +30,8 @@ class BaseEmbedderConfig:
|
||||
:type api_key: Optional[str], optional
|
||||
:param api_base: huggingface api base, defaults to None
|
||||
:type api_base: Optional[str], optional
|
||||
:param model_kwargs: key-value arguments for the embedding model, defaults a dict inside init.
|
||||
:type model_kwargs: Optional[Dict[str, Any]], defaults a dict inside init.
|
||||
"""
|
||||
self.model = model
|
||||
self.deployment_name = deployment_name
|
||||
@@ -36,3 +39,4 @@ class BaseEmbedderConfig:
|
||||
self.endpoint = endpoint
|
||||
self.api_key = api_key
|
||||
self.api_base = api_base
|
||||
self.model_kwargs = model_kwargs or {}
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from string import Template
|
||||
from typing import Any, Mapping, Optional
|
||||
from typing import Any, Mapping, Optional, Dict, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from embedchain.config.base_config import BaseConfig
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
@@ -47,6 +50,35 @@ Query: $query
|
||||
Answer:
|
||||
""" # noqa:E501
|
||||
|
||||
DEFAULT_PROMPT_WITH_MEM0_MEMORY = """
|
||||
You are a Q&A expert system. Your responses must always be rooted in the context provided for each query. You are also provided with the conversation history and memories with the user. Make sure to use relevant context from conversation history and memories as needed.
|
||||
|
||||
Here are some guidelines to follow:
|
||||
|
||||
1. Refrain from explicitly mentioning the context provided in your response.
|
||||
2. Take into consideration the conversation history and memories provided.
|
||||
3. The context should silently guide your answers without being directly acknowledged.
|
||||
4. Do not use phrases such as 'According to the context provided', 'Based on the context, ...' etc.
|
||||
|
||||
Context information:
|
||||
----------------------
|
||||
$context
|
||||
----------------------
|
||||
|
||||
Conversation history:
|
||||
----------------------
|
||||
$history
|
||||
----------------------
|
||||
|
||||
Memories/Preferences:
|
||||
----------------------
|
||||
$memories
|
||||
----------------------
|
||||
|
||||
Query: $query
|
||||
Answer:
|
||||
""" # noqa:E501
|
||||
|
||||
DOCS_SITE_DEFAULT_PROMPT = """
|
||||
You are an expert AI assistant for developer support product. Your responses must always be rooted in the context provided for each query. Wherever possible, give complete code snippet. Dont make up any code snippet on your own.
|
||||
|
||||
@@ -67,6 +99,7 @@ Answer:
|
||||
|
||||
DEFAULT_PROMPT_TEMPLATE = Template(DEFAULT_PROMPT)
|
||||
DEFAULT_PROMPT_WITH_HISTORY_TEMPLATE = Template(DEFAULT_PROMPT_WITH_HISTORY)
|
||||
DEFAULT_PROMPT_WITH_MEM0_MEMORY_TEMPLATE = Template(DEFAULT_PROMPT_WITH_MEM0_MEMORY)
|
||||
DOCS_SITE_PROMPT_TEMPLATE = Template(DOCS_SITE_DEFAULT_PROMPT)
|
||||
query_re = re.compile(r"\$\{*query\}*")
|
||||
context_re = re.compile(r"\$\{*context\}*")
|
||||
@@ -90,6 +123,7 @@ class BaseLlmConfig(BaseConfig):
|
||||
top_p: float = 1,
|
||||
stream: bool = False,
|
||||
online: bool = False,
|
||||
token_usage: bool = False,
|
||||
deployment_name: Optional[str] = None,
|
||||
system_prompt: Optional[str] = None,
|
||||
where: dict[str, Any] = None,
|
||||
@@ -99,8 +133,8 @@ class BaseLlmConfig(BaseConfig):
|
||||
base_url: Optional[str] = None,
|
||||
endpoint: Optional[str] = None,
|
||||
model_kwargs: Optional[dict[str, Any]] = None,
|
||||
http_client: Optional[Any] = None,
|
||||
http_async_client: Optional[Any] = None,
|
||||
http_client_proxies: Optional[Union[Dict, str]] = None,
|
||||
http_async_client_proxies: Optional[Union[Dict, str]] = None,
|
||||
local: Optional[bool] = False,
|
||||
default_headers: Optional[Mapping[str, str]] = None,
|
||||
api_version: Optional[str] = None,
|
||||
@@ -133,6 +167,8 @@ class BaseLlmConfig(BaseConfig):
|
||||
:type stream: bool, optional
|
||||
:param online: Controls whether to use internet for answering query, defaults to False
|
||||
:type online: bool, optional
|
||||
:param token_usage: Controls whether to return token usage in response, defaults to False
|
||||
:type token_usage: bool, optional
|
||||
:param deployment_name: t.b.a., defaults to None
|
||||
:type deployment_name: Optional[str], optional
|
||||
:param system_prompt: System prompt string, defaults to None
|
||||
@@ -149,6 +185,11 @@ class BaseLlmConfig(BaseConfig):
|
||||
:type callbacks: Optional[list], optional
|
||||
:param query_type: The type of query to use, defaults to None
|
||||
:type query_type: Optional[str], optional
|
||||
:param http_client_proxies: The proxy server settings used to create self.http_client, defaults to None
|
||||
:type http_client_proxies: Optional[Dict | str], optional
|
||||
:param http_async_client_proxies: The proxy server settings for async calls used to create
|
||||
self.http_async_client, defaults to None
|
||||
:type http_async_client_proxies: Optional[Dict | str], optional
|
||||
:param local: If True, the model will be run locally, defaults to False (for huggingface provider)
|
||||
:type local: Optional[bool], optional
|
||||
:param default_headers: Set additional HTTP headers to be sent with requests to OpenAI
|
||||
@@ -173,6 +214,8 @@ class BaseLlmConfig(BaseConfig):
|
||||
self.max_tokens = max_tokens
|
||||
self.model = model
|
||||
self.top_p = top_p
|
||||
self.online = online
|
||||
self.token_usage = token_usage
|
||||
self.deployment_name = deployment_name
|
||||
self.system_prompt = system_prompt
|
||||
self.query_type = query_type
|
||||
@@ -181,13 +224,19 @@ class BaseLlmConfig(BaseConfig):
|
||||
self.base_url = base_url
|
||||
self.endpoint = endpoint
|
||||
self.model_kwargs = model_kwargs
|
||||
self.http_client = http_client
|
||||
self.http_async_client = http_async_client
|
||||
self.http_client = httpx.Client(proxies=http_client_proxies) if http_client_proxies else None
|
||||
self.http_async_client = (
|
||||
httpx.AsyncClient(proxies=http_async_client_proxies) if http_async_client_proxies else None
|
||||
)
|
||||
self.local = local
|
||||
self.default_headers = default_headers
|
||||
self.online = online
|
||||
self.api_version = api_version
|
||||
|
||||
if token_usage:
|
||||
f = open("model_prices_and_context_window.json")
|
||||
self.model_pricing_map = json.load(f)
|
||||
|
||||
if isinstance(prompt, str):
|
||||
prompt = Template(prompt)
|
||||
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
from typing import Any, Optional
|
||||
|
||||
from embedchain.config.base_config import BaseConfig
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
|
||||
|
||||
@register_deserializable
|
||||
class Mem0Config(BaseConfig):
|
||||
def __init__(self, api_key: str, top_k: Optional[int] = 10):
|
||||
self.api_key = api_key
|
||||
self.top_k = top_k
|
||||
|
||||
@staticmethod
|
||||
def from_config(config: Optional[dict[str, Any]]):
|
||||
if config is None:
|
||||
return Mem0Config()
|
||||
else:
|
||||
return Mem0Config(
|
||||
api_key=config.get("api_key", ""),
|
||||
init_config=config.get("top_k", 10),
|
||||
)
|
||||
@@ -12,6 +12,7 @@ class ChromaDbConfig(BaseVectorDbConfig):
|
||||
dir: Optional[str] = None,
|
||||
host: Optional[str] = None,
|
||||
port: Optional[str] = None,
|
||||
batch_size: Optional[int] = 100,
|
||||
allow_reset=False,
|
||||
chroma_settings: Optional[dict] = None,
|
||||
):
|
||||
@@ -26,6 +27,8 @@ class ChromaDbConfig(BaseVectorDbConfig):
|
||||
:type host: Optional[str], optional
|
||||
:param port: Database connection remote port. Use this if you run Embedchain as a client, defaults to None
|
||||
:type port: Optional[str], optional
|
||||
:param batch_size: Number of items to insert in one batch, defaults to 100
|
||||
:type batch_size: Optional[int], optional
|
||||
:param allow_reset: Resets the database. defaults to False
|
||||
:type allow_reset: bool
|
||||
:param chroma_settings: Chroma settings dict, defaults to None
|
||||
@@ -34,4 +37,5 @@ class ChromaDbConfig(BaseVectorDbConfig):
|
||||
|
||||
self.chroma_settings = chroma_settings
|
||||
self.allow_reset = allow_reset
|
||||
self.batch_size = batch_size
|
||||
super().__init__(collection_name=collection_name, dir=dir, host=host, port=port)
|
||||
|
||||
@@ -13,6 +13,7 @@ class ElasticsearchDBConfig(BaseVectorDbConfig):
|
||||
dir: Optional[str] = None,
|
||||
es_url: Union[str, list[str]] = None,
|
||||
cloud_id: Optional[str] = None,
|
||||
batch_size: Optional[int] = 100,
|
||||
**ES_EXTRA_PARAMS: dict[str, any],
|
||||
):
|
||||
"""
|
||||
@@ -24,6 +25,10 @@ class ElasticsearchDBConfig(BaseVectorDbConfig):
|
||||
:type dir: Optional[str], optional
|
||||
:param es_url: elasticsearch url or list of nodes url to be used for connection, defaults to None
|
||||
:type es_url: Union[str, list[str]], optional
|
||||
:param cloud_id: cloud id of the elasticsearch cluster, defaults to None
|
||||
:type cloud_id: Optional[str], optional
|
||||
:param batch_size: Number of items to insert in one batch, defaults to 100
|
||||
:type batch_size: Optional[int], optional
|
||||
:param ES_EXTRA_PARAMS: extra params dict that can be passed to elasticsearch.
|
||||
:type ES_EXTRA_PARAMS: dict[str, Any], optional
|
||||
"""
|
||||
@@ -46,4 +51,6 @@ class ElasticsearchDBConfig(BaseVectorDbConfig):
|
||||
and not self.ES_EXTRA_PARAMS.get("bearer_auth")
|
||||
):
|
||||
self.ES_EXTRA_PARAMS["api_key"] = os.environ.get("ELASTICSEARCH_API_KEY")
|
||||
|
||||
self.batch_size = batch_size
|
||||
super().__init__(collection_name=collection_name, dir=dir)
|
||||
|
||||
@@ -13,6 +13,7 @@ class OpenSearchDBConfig(BaseVectorDbConfig):
|
||||
vector_dimension: int = 1536,
|
||||
collection_name: Optional[str] = None,
|
||||
dir: Optional[str] = None,
|
||||
batch_size: Optional[int] = 100,
|
||||
**extra_params: dict[str, any],
|
||||
):
|
||||
"""
|
||||
@@ -28,10 +29,13 @@ class OpenSearchDBConfig(BaseVectorDbConfig):
|
||||
:type vector_dimension: int, optional
|
||||
:param dir: Path to the database directory, where the database is stored, defaults to None
|
||||
:type dir: Optional[str], optional
|
||||
:param batch_size: Number of items to insert in one batch, defaults to 100
|
||||
:type batch_size: Optional[int], optional
|
||||
"""
|
||||
self.opensearch_url = opensearch_url
|
||||
self.http_auth = http_auth
|
||||
self.vector_dimension = vector_dimension
|
||||
self.extra_params = extra_params
|
||||
self.batch_size = batch_size
|
||||
|
||||
super().__init__(collection_name=collection_name, dir=dir)
|
||||
|
||||
@@ -17,6 +17,7 @@ class PineconeDBConfig(BaseVectorDbConfig):
|
||||
serverless_config: Optional[dict[str, any]] = None,
|
||||
hybrid_search: bool = False,
|
||||
bm25_encoder: any = None,
|
||||
batch_size: Optional[int] = 100,
|
||||
**extra_params: dict[str, any],
|
||||
):
|
||||
self.metric = metric
|
||||
@@ -26,6 +27,7 @@ class PineconeDBConfig(BaseVectorDbConfig):
|
||||
self.extra_params = extra_params
|
||||
self.hybrid_search = hybrid_search
|
||||
self.bm25_encoder = bm25_encoder
|
||||
self.batch_size = batch_size
|
||||
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")
|
||||
|
||||
@@ -18,6 +18,7 @@ class QdrantDBConfig(BaseVectorDbConfig):
|
||||
hnsw_config: Optional[dict[str, any]] = None,
|
||||
quantization_config: Optional[dict[str, any]] = None,
|
||||
on_disk: Optional[bool] = None,
|
||||
batch_size: Optional[int] = 10,
|
||||
**extra_params: dict[str, any],
|
||||
):
|
||||
"""
|
||||
@@ -36,9 +37,12 @@ class QdrantDBConfig(BaseVectorDbConfig):
|
||||
This setting saves RAM by (slightly) increasing the response time.
|
||||
Note: those payload values that are involved in filtering and are indexed - remain in RAM.
|
||||
:type on_disk: bool, optional, defaults to None
|
||||
:param batch_size: Number of items to insert in one batch, defaults to 10
|
||||
:type batch_size: Optional[int], optional
|
||||
"""
|
||||
self.hnsw_config = hnsw_config
|
||||
self.quantization_config = quantization_config
|
||||
self.on_disk = on_disk
|
||||
self.batch_size = batch_size
|
||||
self.extra_params = extra_params
|
||||
super().__init__(collection_name=collection_name, dir=dir)
|
||||
|
||||
@@ -10,7 +10,9 @@ class WeaviateDBConfig(BaseVectorDbConfig):
|
||||
self,
|
||||
collection_name: Optional[str] = None,
|
||||
dir: Optional[str] = None,
|
||||
batch_size: Optional[int] = 100,
|
||||
**extra_params: dict[str, any],
|
||||
):
|
||||
self.batch_size = batch_size
|
||||
self.extra_params = extra_params
|
||||
super().__init__(collection_name=collection_name, dir=dir)
|
||||
|
||||
+66
-19
@@ -6,9 +6,7 @@ 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
|
||||
@@ -18,8 +16,7 @@ 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.utils.misc import detect_datatype, is_valid_json_string
|
||||
from embedchain.vectordb.base import BaseVectorDB
|
||||
|
||||
@@ -55,6 +52,8 @@ class EmbedChain(JSONSerializable):
|
||||
"""
|
||||
self.config = config
|
||||
self.cache_config = None
|
||||
self.memory_config = None
|
||||
self.mem0_client = None
|
||||
# Llm
|
||||
self.llm = llm
|
||||
# Database has support for config assignment for backwards compatibility
|
||||
@@ -478,7 +477,7 @@ class EmbedChain(JSONSerializable):
|
||||
where: Optional[dict] = None,
|
||||
citations: bool = False,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> Union[tuple[str, list[tuple[str, dict]]], str]:
|
||||
) -> Union[tuple[str, list[tuple[str, dict]]], str, dict[str, Any]]:
|
||||
"""
|
||||
Queries the vector database based on the given input query.
|
||||
Gets relevant doc based on the query and then passes it to an
|
||||
@@ -501,7 +500,9 @@ class EmbedChain(JSONSerializable):
|
||||
:type kwargs: dict[str, Any]
|
||||
:return: The answer to the query, with citations if the citation flag is True
|
||||
or the dry run result
|
||||
:rtype: str, if citations is False, otherwise tuple[str, list[tuple[str,str,str]]]
|
||||
:rtype: str, if citations is False and token_usage is False, otherwise if citations is true then
|
||||
tuple[str, list[tuple[str,str,str]]] and if token_usage is true then
|
||||
tuple[str, list[tuple[str,str,str]], dict[str, Any]]
|
||||
"""
|
||||
contexts = self._retrieve_from_database(
|
||||
input_query=input_query, config=config, where=where, citations=citations, **kwargs
|
||||
@@ -524,17 +525,29 @@ class EmbedChain(JSONSerializable):
|
||||
dry_run=dry_run,
|
||||
)
|
||||
else:
|
||||
answer = self.llm.query(
|
||||
input_query=input_query, contexts=contexts_data_for_llm_query, config=config, dry_run=dry_run
|
||||
)
|
||||
if self.llm.config.token_usage:
|
||||
answer, token_info = self.llm.query(
|
||||
input_query=input_query, contexts=contexts_data_for_llm_query, config=config, dry_run=dry_run
|
||||
)
|
||||
else:
|
||||
answer = self.llm.query(
|
||||
input_query=input_query, contexts=contexts_data_for_llm_query, config=config, dry_run=dry_run
|
||||
)
|
||||
|
||||
# Send anonymous telemetry
|
||||
self.telemetry.capture(event_name="query", properties=self._telemetry_props)
|
||||
|
||||
if citations:
|
||||
if self.llm.config.token_usage:
|
||||
return {"answer": answer, "contexts": contexts, "usage": token_info}
|
||||
return answer, contexts
|
||||
else:
|
||||
return answer
|
||||
if self.llm.config.token_usage:
|
||||
return {"answer": answer, "usage": token_info}
|
||||
|
||||
logger.warning(
|
||||
"Starting from v0.1.125 the return type of query method will be changed to tuple containing `answer`."
|
||||
)
|
||||
return answer
|
||||
|
||||
def chat(
|
||||
self,
|
||||
@@ -545,7 +558,7 @@ class EmbedChain(JSONSerializable):
|
||||
where: Optional[dict[str, str]] = None,
|
||||
citations: bool = False,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> Union[tuple[str, list[tuple[str, dict]]], str]:
|
||||
) -> Union[tuple[str, list[tuple[str, dict]]], str, dict[str, Any]]:
|
||||
"""
|
||||
Queries the vector database on the given input query.
|
||||
Gets relevant doc based on the query and then passes it to an
|
||||
@@ -572,7 +585,9 @@ class EmbedChain(JSONSerializable):
|
||||
:type kwargs: dict[str, Any]
|
||||
:return: The answer to the query, with citations if the citation flag is True
|
||||
or the dry run result
|
||||
:rtype: str, if citations is False, otherwise tuple[str, list[tuple[str,str,str]]]
|
||||
:rtype: str, if citations is False and token_usage is False, otherwise if citations is true then
|
||||
tuple[str, list[tuple[str,str,str]]] and if token_usage is true then
|
||||
tuple[str, list[tuple[str,str,str]], dict[str, Any]]
|
||||
"""
|
||||
contexts = self._retrieve_from_database(
|
||||
input_query=input_query, config=config, where=where, citations=citations, **kwargs
|
||||
@@ -582,6 +597,12 @@ class EmbedChain(JSONSerializable):
|
||||
else:
|
||||
contexts_data_for_llm_query = contexts
|
||||
|
||||
memories = None
|
||||
if self.mem0_client:
|
||||
memories = self.mem0_client.search(
|
||||
query=input_query, agent_id=self.config.id, session_id=session_id, limit=self.memory_config.top_k
|
||||
)
|
||||
|
||||
# Update the history beforehand so that we can handle multiple chat sessions in the same python session
|
||||
self.llm.update_history(app_id=self.config.id, session_id=session_id)
|
||||
|
||||
@@ -600,9 +621,28 @@ class EmbedChain(JSONSerializable):
|
||||
)
|
||||
else:
|
||||
logger.debug("Cache disabled. Running chat without cache.")
|
||||
answer = self.llm.chat(
|
||||
input_query=input_query, contexts=contexts_data_for_llm_query, config=config, dry_run=dry_run
|
||||
)
|
||||
if self.llm.config.token_usage:
|
||||
answer, token_info = self.llm.query(
|
||||
input_query=input_query,
|
||||
contexts=contexts_data_for_llm_query,
|
||||
config=config,
|
||||
dry_run=dry_run,
|
||||
memories=memories,
|
||||
)
|
||||
else:
|
||||
answer = self.llm.query(
|
||||
input_query=input_query,
|
||||
contexts=contexts_data_for_llm_query,
|
||||
config=config,
|
||||
dry_run=dry_run,
|
||||
memories=memories,
|
||||
)
|
||||
|
||||
# Add to Mem0 memory if enabled
|
||||
# TODO: Might need to prepend with some text like:
|
||||
# "Remember user preferences from following user query: {input_query}"
|
||||
if self.mem0_client:
|
||||
self.mem0_client.add(data=input_query, agent_id=self.config.id, session_id=session_id)
|
||||
|
||||
# add conversation in memory
|
||||
self.llm.add_history(self.config.id, input_query, answer, session_id=session_id)
|
||||
@@ -611,9 +651,16 @@ class EmbedChain(JSONSerializable):
|
||||
self.telemetry.capture(event_name="chat", properties=self._telemetry_props)
|
||||
|
||||
if citations:
|
||||
if self.llm.config.token_usage:
|
||||
return {"answer": answer, "contexts": contexts, "usage": token_info}
|
||||
return answer, contexts
|
||||
else:
|
||||
return answer
|
||||
if self.llm.config.token_usage:
|
||||
return {"answer": answer, "usage": token_info}
|
||||
|
||||
logger.warning(
|
||||
"Starting from v0.1.125 the return type of query method will be changed to tuple containing `answer`."
|
||||
)
|
||||
return answer
|
||||
|
||||
def search(self, query, num_documents=3, where=None, raw_filter=None, namespace=None):
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
from typing import Optional
|
||||
|
||||
from langchain_community.embeddings import AzureOpenAIEmbeddings
|
||||
|
||||
from embedchain.config import BaseEmbedderConfig
|
||||
from embedchain.embedder.base import BaseEmbedder
|
||||
from embedchain.models import VectorDimensions
|
||||
|
||||
|
||||
class AzureOpenAIEmbedder(BaseEmbedder):
|
||||
def __init__(self, config: Optional[BaseEmbedderConfig] = None):
|
||||
super().__init__(config=config)
|
||||
|
||||
if self.config.model is None:
|
||||
self.config.model = "text-embedding-ada-002"
|
||||
|
||||
embeddings = AzureOpenAIEmbeddings(deployment=self.config.deployment_name)
|
||||
embedding_fn = BaseEmbedder._langchain_default_concept(embeddings)
|
||||
|
||||
self.set_embedding_fn(embedding_fn=embedding_fn)
|
||||
vector_dimension = self.config.vector_dimension or VectorDimensions.OPENAI.value
|
||||
self.set_vector_dimension(vector_dimension=vector_dimension)
|
||||
@@ -9,10 +9,10 @@ class GPT4AllEmbedder(BaseEmbedder):
|
||||
def __init__(self, config: Optional[BaseEmbedderConfig] = None):
|
||||
super().__init__(config=config)
|
||||
|
||||
from langchain.embeddings import \
|
||||
GPT4AllEmbeddings as LangchainGPT4AllEmbeddings
|
||||
from langchain_community.embeddings import GPT4AllEmbeddings as LangchainGPT4AllEmbeddings
|
||||
|
||||
embeddings = LangchainGPT4AllEmbeddings()
|
||||
model_name = self.config.model or "all-MiniLM-L6-v2-f16.gguf"
|
||||
embeddings = LangchainGPT4AllEmbeddings(model_name=model_name)
|
||||
embedding_fn = BaseEmbedder._langchain_default_concept(embeddings)
|
||||
self.set_embedding_fn(embedding_fn=embedding_fn)
|
||||
|
||||
|
||||
@@ -31,7 +31,8 @@ class HuggingFaceEmbedder(BaseEmbedder):
|
||||
huggingfacehub_api_token=self.config.api_key or os.getenv("HUGGINGFACE_ACCESS_TOKEN"),
|
||||
)
|
||||
else:
|
||||
embeddings = HuggingFaceEmbeddings(model_name=self.config.model)
|
||||
embeddings = HuggingFaceEmbeddings(model_name=self.config.model, model_kwargs=self.config.model_kwargs)
|
||||
|
||||
embedding_fn = BaseEmbedder._langchain_default_concept(embeddings)
|
||||
self.set_embedding_fn(embedding_fn=embedding_fn)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ import os
|
||||
from typing import Optional
|
||||
|
||||
from chromadb.utils.embedding_functions import OpenAIEmbeddingFunction
|
||||
from langchain_openai.embeddings import AzureOpenAIEmbeddings
|
||||
|
||||
|
||||
from embedchain.config import BaseEmbedderConfig
|
||||
from embedchain.embedder.base import BaseEmbedder
|
||||
@@ -19,20 +19,14 @@ class OpenAIEmbedder(BaseEmbedder):
|
||||
api_key = self.config.api_key or os.environ["OPENAI_API_KEY"]
|
||||
api_base = self.config.api_base or os.environ.get("OPENAI_API_BASE")
|
||||
|
||||
if self.config.deployment_name:
|
||||
embeddings = AzureOpenAIEmbeddings(deployment=self.config.deployment_name)
|
||||
embedding_fn = BaseEmbedder._langchain_default_concept(embeddings)
|
||||
else:
|
||||
if api_key is None and os.getenv("OPENAI_ORGANIZATION") is None:
|
||||
raise ValueError(
|
||||
"OPENAI_API_KEY or OPENAI_ORGANIZATION environment variables not provided"
|
||||
) # noqa:E501
|
||||
embedding_fn = OpenAIEmbeddingFunction(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
organization_id=os.getenv("OPENAI_ORGANIZATION"),
|
||||
model_name=self.config.model,
|
||||
)
|
||||
if api_key is None and os.getenv("OPENAI_ORGANIZATION") is None:
|
||||
raise ValueError("OPENAI_API_KEY or OPENAI_ORGANIZATION environment variables not provided") # noqa:E501
|
||||
embedding_fn = OpenAIEmbeddingFunction(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
organization_id=os.getenv("OPENAI_ORGANIZATION"),
|
||||
model_name=self.config.model,
|
||||
)
|
||||
self.set_embedding_fn(embedding_fn=embedding_fn)
|
||||
vector_dimension = self.config.vector_dimension or VectorDimensions.OPENAI.value
|
||||
self.set_vector_dimension(vector_dimension=vector_dimension)
|
||||
|
||||
@@ -50,7 +50,7 @@ class LlmFactory:
|
||||
|
||||
class EmbedderFactory:
|
||||
provider_to_class = {
|
||||
"azure_openai": "embedchain.embedder.openai.OpenAIEmbedder",
|
||||
"azure_openai": "embedchain.embedder.azure_openai.AzureOpenAIEmbedder",
|
||||
"gpt4all": "embedchain.embedder.gpt4all.GPT4AllEmbedder",
|
||||
"huggingface": "embedchain.embedder.huggingface.HuggingFaceEmbedder",
|
||||
"openai": "embedchain.embedder.openai.OpenAIEmbedder",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
|
||||
try:
|
||||
from langchain_anthropic import ChatAnthropic
|
||||
@@ -21,8 +21,27 @@ class AnthropicLlm(BaseLlm):
|
||||
if not self.config.api_key and "ANTHROPIC_API_KEY" not in os.environ:
|
||||
raise ValueError("Please set the ANTHROPIC_API_KEY environment variable or pass it in the config.")
|
||||
|
||||
def get_llm_model_answer(self, prompt):
|
||||
return AnthropicLlm._get_answer(prompt=prompt, config=self.config)
|
||||
def get_llm_model_answer(self, prompt) -> tuple[str, Optional[dict[str, Any]]]:
|
||||
if self.config.token_usage:
|
||||
response, token_info = self._get_answer(prompt, self.config)
|
||||
model_name = "anthropic/" + self.config.model
|
||||
if model_name not in self.config.model_pricing_map:
|
||||
raise ValueError(
|
||||
f"Model {model_name} not found in `model_prices_and_context_window.json`. \
|
||||
You can disable token usage by setting `token_usage` to False."
|
||||
)
|
||||
total_cost = (
|
||||
self.config.model_pricing_map[model_name]["input_cost_per_token"] * token_info["input_tokens"]
|
||||
) + self.config.model_pricing_map[model_name]["output_cost_per_token"] * token_info["output_tokens"]
|
||||
response_token_info = {
|
||||
"prompt_tokens": token_info["input_tokens"],
|
||||
"completion_tokens": token_info["output_tokens"],
|
||||
"total_tokens": token_info["input_tokens"] + token_info["output_tokens"],
|
||||
"total_cost": round(total_cost, 10),
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
return response, response_token_info
|
||||
return self._get_answer(prompt, self.config)
|
||||
|
||||
@staticmethod
|
||||
def _get_answer(prompt: str, config: BaseLlmConfig) -> str:
|
||||
@@ -34,4 +53,7 @@ class AnthropicLlm(BaseLlm):
|
||||
|
||||
messages = BaseLlm._get_messages(prompt, system_prompt=config.system_prompt)
|
||||
|
||||
return chat(messages).content
|
||||
chat_response = chat.invoke(messages)
|
||||
if config.token_usage:
|
||||
return chat_response.content, chat_response.response_metadata["token_usage"]
|
||||
return chat_response.content
|
||||
|
||||
+46
-13
@@ -5,9 +5,12 @@ from typing import Any, Optional
|
||||
from langchain.schema import BaseMessage as LCBaseMessage
|
||||
|
||||
from embedchain.config import BaseLlmConfig
|
||||
from embedchain.config.llm.base import (DEFAULT_PROMPT,
|
||||
DEFAULT_PROMPT_WITH_HISTORY_TEMPLATE,
|
||||
DOCS_SITE_PROMPT_TEMPLATE)
|
||||
from embedchain.config.llm.base import (
|
||||
DEFAULT_PROMPT,
|
||||
DEFAULT_PROMPT_WITH_HISTORY_TEMPLATE,
|
||||
DEFAULT_PROMPT_WITH_MEM0_MEMORY_TEMPLATE,
|
||||
DOCS_SITE_PROMPT_TEMPLATE,
|
||||
)
|
||||
from embedchain.helpers.json_serializable import JSONSerializable
|
||||
from embedchain.memory.base import ChatHistory
|
||||
from embedchain.memory.message import ChatMessage
|
||||
@@ -74,6 +77,16 @@ class BaseLlm(JSONSerializable):
|
||||
"""
|
||||
return "\n".join(self.history)
|
||||
|
||||
def _format_memories(self, memories: list[dict]) -> str:
|
||||
"""Format memories to be used in prompt
|
||||
|
||||
:param memories: Memories to format
|
||||
:type memories: list[dict]
|
||||
:return: Formatted memories
|
||||
:rtype: str
|
||||
"""
|
||||
return "\n".join([memory["text"] for memory in memories])
|
||||
|
||||
def generate_prompt(self, input_query: str, contexts: list[str], **kwargs: dict[str, Any]) -> str:
|
||||
"""
|
||||
Generates a prompt based on the given query and context, ready to be
|
||||
@@ -88,6 +101,7 @@ class BaseLlm(JSONSerializable):
|
||||
"""
|
||||
context_string = " | ".join(contexts)
|
||||
web_search_result = kwargs.get("web_search_result", "")
|
||||
memories = kwargs.get("memories", None)
|
||||
if web_search_result:
|
||||
context_string = self._append_search_and_context(context_string, web_search_result)
|
||||
|
||||
@@ -103,10 +117,19 @@ class BaseLlm(JSONSerializable):
|
||||
not self.config._validate_prompt_history(self.config.prompt)
|
||||
and self.config.prompt.template == DEFAULT_PROMPT
|
||||
):
|
||||
# swap in the template with history
|
||||
prompt = DEFAULT_PROMPT_WITH_HISTORY_TEMPLATE.substitute(
|
||||
context=context_string, query=input_query, history=self._format_history()
|
||||
)
|
||||
if memories:
|
||||
# swap in the template with Mem0 memory template
|
||||
prompt = DEFAULT_PROMPT_WITH_MEM0_MEMORY_TEMPLATE.substitute(
|
||||
context=context_string,
|
||||
query=input_query,
|
||||
history=self._format_history(),
|
||||
memories=self._format_memories(memories),
|
||||
)
|
||||
else:
|
||||
# swap in the template with history
|
||||
prompt = DEFAULT_PROMPT_WITH_HISTORY_TEMPLATE.substitute(
|
||||
context=context_string, query=input_query, history=self._format_history()
|
||||
)
|
||||
else:
|
||||
# If we can't swap in the default, we still proceed but tell users that the history is ignored.
|
||||
logger.warning(
|
||||
@@ -164,7 +187,7 @@ class BaseLlm(JSONSerializable):
|
||||
return search.run(input_query)
|
||||
|
||||
@staticmethod
|
||||
def _stream_response(answer: Any) -> Generator[Any, Any, None]:
|
||||
def _stream_response(answer: Any, token_info: Optional[dict[str, Any]] = None) -> Generator[Any, Any, None]:
|
||||
"""Generator to be used as streaming response
|
||||
|
||||
:param answer: Answer chunk from llm
|
||||
@@ -177,8 +200,10 @@ class BaseLlm(JSONSerializable):
|
||||
streamed_answer = streamed_answer + chunk
|
||||
yield chunk
|
||||
logger.info(f"Answer: {streamed_answer}")
|
||||
if token_info:
|
||||
logger.info(f"Token Info: {token_info}")
|
||||
|
||||
def query(self, input_query: str, contexts: list[str], config: BaseLlmConfig = None, dry_run=False):
|
||||
def query(self, input_query: str, contexts: list[str], config: BaseLlmConfig = None, dry_run=False, memories=None):
|
||||
"""
|
||||
Queries the vector database based on the given input query.
|
||||
Gets relevant doc based on the query and then passes it to an
|
||||
@@ -214,16 +239,24 @@ class BaseLlm(JSONSerializable):
|
||||
k = {}
|
||||
if self.config.online:
|
||||
k["web_search_result"] = self.access_search_and_get_results(input_query)
|
||||
k["memories"] = memories
|
||||
prompt = self.generate_prompt(input_query, contexts, **k)
|
||||
logger.info(f"Prompt: {prompt}")
|
||||
if dry_run:
|
||||
return prompt
|
||||
|
||||
answer = self.get_answer_from_llm(prompt)
|
||||
if self.config.token_usage:
|
||||
answer, token_info = self.get_answer_from_llm(prompt)
|
||||
else:
|
||||
answer = self.get_answer_from_llm(prompt)
|
||||
if isinstance(answer, str):
|
||||
logger.info(f"Answer: {answer}")
|
||||
if self.config.token_usage:
|
||||
return answer, token_info
|
||||
return answer
|
||||
else:
|
||||
if self.config.token_usage:
|
||||
return self._stream_response(answer, token_info)
|
||||
return self._stream_response(answer)
|
||||
finally:
|
||||
if config:
|
||||
@@ -276,13 +309,13 @@ class BaseLlm(JSONSerializable):
|
||||
if dry_run:
|
||||
return prompt
|
||||
|
||||
answer = self.get_answer_from_llm(prompt)
|
||||
answer, token_info = self.get_answer_from_llm(prompt)
|
||||
if isinstance(answer, str):
|
||||
logger.info(f"Answer: {answer}")
|
||||
return answer
|
||||
return answer, token_info
|
||||
else:
|
||||
# this is a streamed response and needs to be handled differently.
|
||||
return self._stream_response(answer)
|
||||
return self._stream_response(answer, token_info)
|
||||
finally:
|
||||
if config:
|
||||
# Restore previous config
|
||||
|
||||
+37
-14
@@ -1,8 +1,8 @@
|
||||
import importlib
|
||||
import os
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
|
||||
from langchain_community.llms.cohere import Cohere
|
||||
from langchain_cohere import ChatCohere
|
||||
|
||||
from embedchain.config import BaseLlmConfig
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
@@ -17,27 +17,50 @@ class CohereLlm(BaseLlm):
|
||||
except ModuleNotFoundError:
|
||||
raise ModuleNotFoundError(
|
||||
"The required dependencies for Cohere are not installed."
|
||||
'Please install with `pip install --upgrade "embedchain[cohere]"`'
|
||||
"Please install with `pip install langchain_cohere==1.16.0`"
|
||||
) from None
|
||||
|
||||
super().__init__(config=config)
|
||||
if not self.config.api_key and "COHERE_API_KEY" not in os.environ:
|
||||
raise ValueError("Please set the COHERE_API_KEY environment variable or pass it in the config.")
|
||||
|
||||
def get_llm_model_answer(self, prompt):
|
||||
def get_llm_model_answer(self, prompt) -> tuple[str, Optional[dict[str, Any]]]:
|
||||
if self.config.system_prompt:
|
||||
raise ValueError("CohereLlm does not support `system_prompt`")
|
||||
return CohereLlm._get_answer(prompt=prompt, config=self.config)
|
||||
|
||||
if self.config.token_usage:
|
||||
response, token_info = self._get_answer(prompt, self.config)
|
||||
model_name = "cohere/" + self.config.model
|
||||
if model_name not in self.config.model_pricing_map:
|
||||
raise ValueError(
|
||||
f"Model {model_name} not found in `model_prices_and_context_window.json`. \
|
||||
You can disable token usage by setting `token_usage` to False."
|
||||
)
|
||||
total_cost = (
|
||||
self.config.model_pricing_map[model_name]["input_cost_per_token"] * token_info["input_tokens"]
|
||||
) + self.config.model_pricing_map[model_name]["output_cost_per_token"] * token_info["output_tokens"]
|
||||
response_token_info = {
|
||||
"prompt_tokens": token_info["input_tokens"],
|
||||
"completion_tokens": token_info["output_tokens"],
|
||||
"total_tokens": token_info["input_tokens"] + token_info["output_tokens"],
|
||||
"total_cost": round(total_cost, 10),
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
return response, response_token_info
|
||||
return self._get_answer(prompt, self.config)
|
||||
|
||||
@staticmethod
|
||||
def _get_answer(prompt: str, config: BaseLlmConfig) -> str:
|
||||
api_key = config.api_key or os.getenv("COHERE_API_KEY")
|
||||
llm = Cohere(
|
||||
cohere_api_key=api_key,
|
||||
model=config.model,
|
||||
max_tokens=config.max_tokens,
|
||||
temperature=config.temperature,
|
||||
p=config.top_p,
|
||||
)
|
||||
api_key = config.api_key or os.environ["COHERE_API_KEY"]
|
||||
kwargs = {
|
||||
"model_name": config.model or "command-r",
|
||||
"temperature": config.temperature,
|
||||
"max_tokens": config.max_tokens,
|
||||
"together_api_key": api_key,
|
||||
}
|
||||
|
||||
return llm.invoke(prompt)
|
||||
chat = ChatCohere(**kwargs)
|
||||
chat_response = chat.invoke(prompt)
|
||||
if config.token_usage:
|
||||
return chat_response.content, chat_response.response_metadata["token_count"]
|
||||
return chat_response.content
|
||||
|
||||
+27
-5
@@ -1,5 +1,5 @@
|
||||
import os
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
|
||||
from langchain.callbacks.streaming_stdout import StreamingStdOutCallbackHandler
|
||||
from langchain.schema import HumanMessage, SystemMessage
|
||||
@@ -22,9 +22,27 @@ class GroqLlm(BaseLlm):
|
||||
if not self.config.api_key and "GROQ_API_KEY" not in os.environ:
|
||||
raise ValueError("Please set the GROQ_API_KEY environment variable or pass it in the config.")
|
||||
|
||||
def get_llm_model_answer(self, prompt) -> str:
|
||||
response = self._get_answer(prompt, self.config)
|
||||
return response
|
||||
def get_llm_model_answer(self, prompt) -> tuple[str, Optional[dict[str, Any]]]:
|
||||
if self.config.token_usage:
|
||||
response, token_info = self._get_answer(prompt, self.config)
|
||||
model_name = "groq/" + self.config.model
|
||||
if model_name not in self.config.model_pricing_map:
|
||||
raise ValueError(
|
||||
f"Model {model_name} not found in `model_prices_and_context_window.json`. \
|
||||
You can disable token usage by setting `token_usage` to False."
|
||||
)
|
||||
total_cost = (
|
||||
self.config.model_pricing_map[model_name]["input_cost_per_token"] * token_info["prompt_tokens"]
|
||||
) + self.config.model_pricing_map[model_name]["output_cost_per_token"] * token_info["completion_tokens"]
|
||||
response_token_info = {
|
||||
"prompt_tokens": token_info["prompt_tokens"],
|
||||
"completion_tokens": token_info["completion_tokens"],
|
||||
"total_tokens": token_info["prompt_tokens"] + token_info["completion_tokens"],
|
||||
"total_cost": round(total_cost, 10),
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
return response, response_token_info
|
||||
return self._get_answer(prompt, self.config)
|
||||
|
||||
def _get_answer(self, prompt: str, config: BaseLlmConfig) -> str:
|
||||
messages = []
|
||||
@@ -42,4 +60,8 @@ class GroqLlm(BaseLlm):
|
||||
chat = ChatGroq(**kwargs, streaming=config.stream, callbacks=callbacks, api_key=api_key)
|
||||
else:
|
||||
chat = ChatGroq(**kwargs)
|
||||
return chat.invoke(messages).content
|
||||
|
||||
chat_response = chat.invoke(prompt)
|
||||
if self.config.token_usage:
|
||||
return chat_response.content, chat_response.response_metadata["token_usage"]
|
||||
return chat_response.content
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import os
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
|
||||
from embedchain.config import BaseLlmConfig
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
@@ -13,8 +13,27 @@ class MistralAILlm(BaseLlm):
|
||||
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)
|
||||
def get_llm_model_answer(self, prompt) -> tuple[str, Optional[dict[str, Any]]]:
|
||||
if self.config.token_usage:
|
||||
response, token_info = self._get_answer(prompt, self.config)
|
||||
model_name = "mistralai/" + self.config.model
|
||||
if model_name not in self.config.model_pricing_map:
|
||||
raise ValueError(
|
||||
f"Model {model_name} not found in `model_prices_and_context_window.json`. \
|
||||
You can disable token usage by setting `token_usage` to False."
|
||||
)
|
||||
total_cost = (
|
||||
self.config.model_pricing_map[model_name]["input_cost_per_token"] * token_info["prompt_tokens"]
|
||||
) + self.config.model_pricing_map[model_name]["output_cost_per_token"] * token_info["completion_tokens"]
|
||||
response_token_info = {
|
||||
"prompt_tokens": token_info["prompt_tokens"],
|
||||
"completion_tokens": token_info["completion_tokens"],
|
||||
"total_tokens": token_info["prompt_tokens"] + token_info["completion_tokens"],
|
||||
"total_cost": round(total_cost, 10),
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
return response, response_token_info
|
||||
return self._get_answer(prompt, self.config)
|
||||
|
||||
@staticmethod
|
||||
def _get_answer(prompt: str, config: BaseLlmConfig):
|
||||
@@ -47,6 +66,7 @@ class MistralAILlm(BaseLlm):
|
||||
answer += chunk.content
|
||||
return answer
|
||||
else:
|
||||
response = client.invoke(**kwargs, input=messages)
|
||||
answer = response.content
|
||||
return answer
|
||||
chat_response = client.invoke(**kwargs, input=messages)
|
||||
if config.token_usage:
|
||||
return chat_response.content, chat_response.response_metadata["token_usage"]
|
||||
return chat_response.content
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import os
|
||||
from collections.abc import Iterable
|
||||
from typing import Optional, Union
|
||||
from typing import Any, Optional, Union
|
||||
|
||||
from langchain.callbacks.manager import CallbackManager
|
||||
from langchain.callbacks.stdout import StdOutCallbackHandler
|
||||
@@ -25,8 +25,27 @@ class NvidiaLlm(BaseLlm):
|
||||
if not self.config.api_key and "NVIDIA_API_KEY" not in os.environ:
|
||||
raise ValueError("Please set the NVIDIA_API_KEY environment variable or pass it in the config.")
|
||||
|
||||
def get_llm_model_answer(self, prompt):
|
||||
return self._get_answer(prompt=prompt, config=self.config)
|
||||
def get_llm_model_answer(self, prompt) -> tuple[str, Optional[dict[str, Any]]]:
|
||||
if self.config.token_usage:
|
||||
response, token_info = self._get_answer(prompt, self.config)
|
||||
model_name = "nvidia/" + self.config.model
|
||||
if model_name not in self.config.model_pricing_map:
|
||||
raise ValueError(
|
||||
f"Model {model_name} not found in `model_prices_and_context_window.json`. \
|
||||
You can disable token usage by setting `token_usage` to False."
|
||||
)
|
||||
total_cost = (
|
||||
self.config.model_pricing_map[model_name]["input_cost_per_token"] * token_info["input_tokens"]
|
||||
) + self.config.model_pricing_map[model_name]["output_cost_per_token"] * token_info["output_tokens"]
|
||||
response_token_info = {
|
||||
"prompt_tokens": token_info["input_tokens"],
|
||||
"completion_tokens": token_info["output_tokens"],
|
||||
"total_tokens": token_info["input_tokens"] + token_info["output_tokens"],
|
||||
"total_cost": round(total_cost, 10),
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
return response, response_token_info
|
||||
return self._get_answer(prompt, self.config)
|
||||
|
||||
@staticmethod
|
||||
def _get_answer(prompt: str, config: BaseLlmConfig) -> Union[str, Iterable]:
|
||||
@@ -43,4 +62,7 @@ class NvidiaLlm(BaseLlm):
|
||||
if labels:
|
||||
params["labels"] = labels
|
||||
llm = ChatNVIDIA(**params, callback_manager=CallbackManager(callback_manager))
|
||||
return llm.invoke(prompt).content if labels is None else llm.invoke(prompt, labels=labels).content
|
||||
chat_response = llm.invoke(prompt) if labels is None else llm.invoke(prompt, labels=labels)
|
||||
if config.token_usage:
|
||||
return chat_response.content, chat_response.response_metadata["token_usage"]
|
||||
return chat_response.content
|
||||
|
||||
@@ -23,9 +23,28 @@ class OpenAILlm(BaseLlm):
|
||||
self.tools = tools
|
||||
super().__init__(config=config)
|
||||
|
||||
def get_llm_model_answer(self, prompt) -> str:
|
||||
response = self._get_answer(prompt, self.config)
|
||||
return response
|
||||
def get_llm_model_answer(self, prompt) -> tuple[str, Optional[dict[str, Any]]]:
|
||||
if self.config.token_usage:
|
||||
response, token_info = self._get_answer(prompt, self.config)
|
||||
model_name = "openai/" + self.config.model
|
||||
if model_name not in self.config.model_pricing_map:
|
||||
raise ValueError(
|
||||
f"Model {model_name} not found in `model_prices_and_context_window.json`. \
|
||||
You can disable token usage by setting `token_usage` to False."
|
||||
)
|
||||
total_cost = (
|
||||
self.config.model_pricing_map[model_name]["input_cost_per_token"] * token_info["prompt_tokens"]
|
||||
) + self.config.model_pricing_map[model_name]["output_cost_per_token"] * token_info["completion_tokens"]
|
||||
response_token_info = {
|
||||
"prompt_tokens": token_info["prompt_tokens"],
|
||||
"completion_tokens": token_info["completion_tokens"],
|
||||
"total_tokens": token_info["prompt_tokens"] + token_info["completion_tokens"],
|
||||
"total_cost": round(total_cost, 10),
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
return response, response_token_info
|
||||
|
||||
return self._get_answer(prompt, self.config)
|
||||
|
||||
def _get_answer(self, prompt: str, config: BaseLlmConfig) -> str:
|
||||
messages = []
|
||||
@@ -56,11 +75,20 @@ class OpenAILlm(BaseLlm):
|
||||
http_async_client=config.http_async_client,
|
||||
)
|
||||
else:
|
||||
chat = ChatOpenAI(**kwargs, api_key=api_key, base_url=base_url)
|
||||
chat = ChatOpenAI(
|
||||
**kwargs,
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
http_client=config.http_client,
|
||||
http_async_client=config.http_async_client,
|
||||
)
|
||||
if self.tools:
|
||||
return self._query_function_call(chat, self.tools, messages)
|
||||
|
||||
return chat.invoke(messages).content
|
||||
chat_response = chat.invoke(messages)
|
||||
if self.config.token_usage:
|
||||
return chat_response.content, chat_response.response_metadata["token_usage"]
|
||||
return chat_response.content
|
||||
|
||||
def _query_function_call(
|
||||
self,
|
||||
@@ -69,8 +97,7 @@ class OpenAILlm(BaseLlm):
|
||||
messages: list[BaseMessage],
|
||||
) -> str:
|
||||
from langchain.output_parsers.openai_tools import JsonOutputToolsParser
|
||||
from langchain_core.utils.function_calling import \
|
||||
convert_to_openai_tool
|
||||
from langchain_core.utils.function_calling import convert_to_openai_tool
|
||||
|
||||
openai_tools = [convert_to_openai_tool(tools)]
|
||||
chat = chat.bind(tools=openai_tools).pipe(JsonOutputToolsParser())
|
||||
|
||||
+41
-13
@@ -1,8 +1,13 @@
|
||||
import importlib
|
||||
import os
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
|
||||
from langchain_community.llms import Together
|
||||
try:
|
||||
from langchain_together import ChatTogether
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"Please install the langchain_together package by running `pip install langchain_together==0.1.3`."
|
||||
)
|
||||
|
||||
from embedchain.config import BaseLlmConfig
|
||||
from embedchain.helpers.json_serializable import register_deserializable
|
||||
@@ -24,20 +29,43 @@ class TogetherLlm(BaseLlm):
|
||||
if not self.config.api_key and "TOGETHER_API_KEY" not in os.environ:
|
||||
raise ValueError("Please set the TOGETHER_API_KEY environment variable or pass it in the config.")
|
||||
|
||||
def get_llm_model_answer(self, prompt):
|
||||
def get_llm_model_answer(self, prompt) -> tuple[str, Optional[dict[str, Any]]]:
|
||||
if self.config.system_prompt:
|
||||
raise ValueError("TogetherLlm does not support `system_prompt`")
|
||||
return TogetherLlm._get_answer(prompt=prompt, config=self.config)
|
||||
|
||||
if self.config.token_usage:
|
||||
response, token_info = self._get_answer(prompt, self.config)
|
||||
model_name = "together/" + self.config.model
|
||||
if model_name not in self.config.model_pricing_map:
|
||||
raise ValueError(
|
||||
f"Model {model_name} not found in `model_prices_and_context_window.json`. \
|
||||
You can disable token usage by setting `token_usage` to False."
|
||||
)
|
||||
total_cost = (
|
||||
self.config.model_pricing_map[model_name]["input_cost_per_token"] * token_info["prompt_tokens"]
|
||||
) + self.config.model_pricing_map[model_name]["output_cost_per_token"] * token_info["completion_tokens"]
|
||||
response_token_info = {
|
||||
"prompt_tokens": token_info["prompt_tokens"],
|
||||
"completion_tokens": token_info["completion_tokens"],
|
||||
"total_tokens": token_info["prompt_tokens"] + token_info["completion_tokens"],
|
||||
"total_cost": round(total_cost, 10),
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
return response, response_token_info
|
||||
return self._get_answer(prompt, self.config)
|
||||
|
||||
@staticmethod
|
||||
def _get_answer(prompt: str, config: BaseLlmConfig) -> str:
|
||||
api_key = config.api_key or os.getenv("TOGETHER_API_KEY")
|
||||
llm = Together(
|
||||
together_api_key=api_key,
|
||||
model=config.model,
|
||||
max_tokens=config.max_tokens,
|
||||
temperature=config.temperature,
|
||||
top_p=config.top_p,
|
||||
)
|
||||
api_key = config.api_key or os.environ["TOGETHER_API_KEY"]
|
||||
kwargs = {
|
||||
"model_name": config.model or "mixtral-8x7b-32768",
|
||||
"temperature": config.temperature,
|
||||
"max_tokens": config.max_tokens,
|
||||
"together_api_key": api_key,
|
||||
}
|
||||
|
||||
return llm.invoke(prompt)
|
||||
chat = ChatTogether(**kwargs)
|
||||
chat_response = chat.invoke(prompt)
|
||||
if config.token_usage:
|
||||
return chat_response.content, chat_response.response_metadata["token_usage"]
|
||||
return chat_response.content
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import importlib
|
||||
import logging
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
|
||||
from langchain.callbacks.streaming_stdout import StreamingStdOutCallbackHandler
|
||||
from langchain_google_vertexai import ChatVertexAI
|
||||
@@ -24,16 +24,35 @@ class VertexAILlm(BaseLlm):
|
||||
) from None
|
||||
super().__init__(config=config)
|
||||
|
||||
def get_llm_model_answer(self, prompt):
|
||||
return VertexAILlm._get_answer(prompt=prompt, config=self.config)
|
||||
def get_llm_model_answer(self, prompt) -> tuple[str, Optional[dict[str, Any]]]:
|
||||
if self.config.token_usage:
|
||||
response, token_info = self._get_answer(prompt, self.config)
|
||||
model_name = "vertexai/" + self.config.model
|
||||
if model_name not in self.config.model_pricing_map:
|
||||
raise ValueError(
|
||||
f"Model {model_name} not found in `model_prices_and_context_window.json`. \
|
||||
You can disable token usage by setting `token_usage` to False."
|
||||
)
|
||||
total_cost = (
|
||||
self.config.model_pricing_map[model_name]["input_cost_per_token"] * token_info["prompt_token_count"]
|
||||
) + self.config.model_pricing_map[model_name]["output_cost_per_token"] * token_info[
|
||||
"candidates_token_count"
|
||||
]
|
||||
response_token_info = {
|
||||
"prompt_tokens": token_info["prompt_token_count"],
|
||||
"completion_tokens": token_info["candidates_token_count"],
|
||||
"total_tokens": token_info["prompt_token_count"] + token_info["candidates_token_count"],
|
||||
"total_cost": round(total_cost, 10),
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
return response, response_token_info
|
||||
return self._get_answer(prompt, self.config)
|
||||
|
||||
@staticmethod
|
||||
def _get_answer(prompt: str, config: BaseLlmConfig) -> str:
|
||||
if config.top_p and config.top_p != 1:
|
||||
logger.warning("Config option `top_p` is not supported by this model.")
|
||||
|
||||
messages = BaseLlm._get_messages(prompt, system_prompt=config.system_prompt)
|
||||
|
||||
if config.stream:
|
||||
callbacks = config.callbacks if config.callbacks else [StreamingStdOutCallbackHandler()]
|
||||
llm = ChatVertexAI(
|
||||
@@ -42,4 +61,8 @@ class VertexAILlm(BaseLlm):
|
||||
else:
|
||||
llm = ChatVertexAI(temperature=config.temperature, model=config.model)
|
||||
|
||||
return llm.invoke(messages).content
|
||||
messages = VertexAILlm._get_messages(prompt)
|
||||
chat_response = llm.invoke(messages)
|
||||
if config.token_usage:
|
||||
return chat_response.content, chat_response.response_metadata["usage_metadata"]
|
||||
return chat_response.content
|
||||
|
||||
@@ -428,6 +428,7 @@ def validate_config(config_data):
|
||||
Optional("top_p"): Or(float, int),
|
||||
Optional("stream"): bool,
|
||||
Optional("online"): bool,
|
||||
Optional("token_usage"): bool,
|
||||
Optional("template"): str,
|
||||
Optional("prompt"): str,
|
||||
Optional("system_prompt"): str,
|
||||
@@ -442,6 +443,8 @@ def validate_config(config_data):
|
||||
Optional("base_url"): str,
|
||||
Optional("default_headers"): dict,
|
||||
Optional("api_version"): Or(str, datetime.date),
|
||||
Optional("http_client_proxies"): Or(str, dict),
|
||||
Optional("http_async_client_proxies"): Or(str, dict),
|
||||
},
|
||||
},
|
||||
Optional("vectordb"): {
|
||||
@@ -474,6 +477,7 @@ def validate_config(config_data):
|
||||
Optional("vector_dimension"): int,
|
||||
Optional("base_url"): str,
|
||||
Optional("endpoint"): str,
|
||||
Optional("model_kwargs"): dict,
|
||||
},
|
||||
},
|
||||
Optional("embedding_model"): {
|
||||
@@ -516,6 +520,10 @@ def validate_config(config_data):
|
||||
Optional("auto_flush"): int,
|
||||
},
|
||||
},
|
||||
Optional("memory"): {
|
||||
"api_key": str,
|
||||
Optional("top_k"): int,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -29,8 +29,6 @@ logger = logging.getLogger(__name__)
|
||||
class ChromaDB(BaseVectorDB):
|
||||
"""Vector database using ChromaDB."""
|
||||
|
||||
BATCH_SIZE = 100
|
||||
|
||||
def __init__(self, config: Optional[ChromaDbConfig] = None):
|
||||
"""Initialize a new ChromaDB instance
|
||||
|
||||
@@ -44,6 +42,7 @@ class ChromaDB(BaseVectorDB):
|
||||
|
||||
self.settings = Settings(anonymized_telemetry=False)
|
||||
self.settings.allow_reset = self.config.allow_reset if hasattr(self.config, "allow_reset") else False
|
||||
self.batch_size = self.config.batch_size
|
||||
if self.config.chroma_settings:
|
||||
for key, value in self.config.chroma_settings.items():
|
||||
if hasattr(self.settings, key):
|
||||
@@ -155,12 +154,13 @@ class ChromaDB(BaseVectorDB):
|
||||
" Ids size: {}".format(len(documents), len(metadatas), len(ids))
|
||||
)
|
||||
|
||||
for i in tqdm(range(0, len(documents), self.BATCH_SIZE), desc="Inserting batches in chromadb"):
|
||||
for i in tqdm(range(0, len(documents), self.batch_size), desc="Inserting batches in chromadb"):
|
||||
self.collection.add(
|
||||
documents=documents[i : i + self.BATCH_SIZE],
|
||||
metadatas=metadatas[i : i + self.BATCH_SIZE],
|
||||
ids=ids[i : i + self.BATCH_SIZE],
|
||||
documents=documents[i : i + self.batch_size],
|
||||
metadatas=metadatas[i : i + self.batch_size],
|
||||
ids=ids[i : i + self.batch_size],
|
||||
)
|
||||
self.config
|
||||
|
||||
@staticmethod
|
||||
def _format_result(results: QueryResult) -> list[tuple[Document, float]]:
|
||||
|
||||
@@ -23,8 +23,6 @@ class ElasticsearchDB(BaseVectorDB):
|
||||
Elasticsearch as vector database
|
||||
"""
|
||||
|
||||
BATCH_SIZE = 100
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Optional[ElasticsearchDBConfig] = None,
|
||||
@@ -57,6 +55,7 @@ class ElasticsearchDB(BaseVectorDB):
|
||||
"Something is wrong with your config. Please check again - `https://docs.embedchain.ai/components/vector-databases#elasticsearch`" # noqa: E501
|
||||
)
|
||||
|
||||
self.batch_size = self.config.batch_size
|
||||
# Call parent init here because embedder is needed
|
||||
super().__init__(config=self.config)
|
||||
|
||||
@@ -140,7 +139,9 @@ class ElasticsearchDB(BaseVectorDB):
|
||||
embeddings = self.embedder.embedding_fn(documents)
|
||||
|
||||
for chunk in chunks(
|
||||
list(zip(ids, documents, metadatas, embeddings)), self.BATCH_SIZE, desc="Inserting batches in elasticsearch"
|
||||
list(zip(ids, documents, metadatas, embeddings)),
|
||||
self.batch_size,
|
||||
desc="Inserting batches in elasticsearch",
|
||||
): # noqa: E501
|
||||
ids, docs, metadatas, embeddings = [], [], [], []
|
||||
for id, text, metadata, embedding in chunk:
|
||||
|
||||
@@ -18,8 +18,6 @@ class LanceDB(BaseVectorDB):
|
||||
LanceDB as vector database
|
||||
"""
|
||||
|
||||
BATCH_SIZE = 100
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Optional[LanceDBConfig] = None,
|
||||
|
||||
@@ -28,8 +28,6 @@ class OpenSearchDB(BaseVectorDB):
|
||||
OpenSearch as vector database
|
||||
"""
|
||||
|
||||
BATCH_SIZE = 100
|
||||
|
||||
def __init__(self, config: OpenSearchDBConfig):
|
||||
"""OpenSearch as vector database.
|
||||
|
||||
@@ -39,6 +37,7 @@ class OpenSearchDB(BaseVectorDB):
|
||||
if config is None:
|
||||
raise ValueError("OpenSearchDBConfig is required")
|
||||
self.config = config
|
||||
self.batch_size = self.config.batch_size
|
||||
self.client = OpenSearch(
|
||||
hosts=[self.config.opensearch_url],
|
||||
http_auth=self.config.http_auth,
|
||||
@@ -120,8 +119,8 @@ class OpenSearchDB(BaseVectorDB):
|
||||
"""Adds documents to the opensearch index"""
|
||||
|
||||
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
|
||||
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]
|
||||
batch_embeddings = embeddings[batch_start:batch_end]
|
||||
|
||||
|
||||
@@ -25,8 +25,6 @@ class PineconeDB(BaseVectorDB):
|
||||
Pinecone as vector database
|
||||
"""
|
||||
|
||||
BATCH_SIZE = 100
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Optional[PineconeDBConfig] = None,
|
||||
@@ -50,6 +48,7 @@ class PineconeDB(BaseVectorDB):
|
||||
|
||||
# Setup BM25Encoder if sparse vectors are to be used
|
||||
self.bm25_encoder = None
|
||||
self.batch_size = self.config.batch_size
|
||||
if self.config.hybrid_search:
|
||||
logger.info("Initializing BM25Encoder for sparse vectors..")
|
||||
self.bm25_encoder = self.config.bm25_encoder if self.config.bm25_encoder else BM25Encoder.default()
|
||||
@@ -103,10 +102,9 @@ class PineconeDB(BaseVectorDB):
|
||||
existing_ids = list()
|
||||
metadatas = []
|
||||
|
||||
batch_size = 100
|
||||
if ids is not None:
|
||||
for i in range(0, len(ids), batch_size):
|
||||
result = self.pinecone_index.fetch(ids=ids[i : i + batch_size])
|
||||
for i in range(0, len(ids), self.batch_size):
|
||||
result = self.pinecone_index.fetch(ids=ids[i : i + self.batch_size])
|
||||
vectors = result.get("vectors")
|
||||
batch_existing_ids = list(vectors.keys())
|
||||
existing_ids.extend(batch_existing_ids)
|
||||
@@ -145,7 +143,7 @@ class PineconeDB(BaseVectorDB):
|
||||
},
|
||||
)
|
||||
|
||||
for chunk in chunks(docs, self.BATCH_SIZE, desc="Adding chunks in batches"):
|
||||
for chunk in chunks(docs, self.batch_size, desc="Adding chunks in batches"):
|
||||
self.pinecone_index.upsert(chunk, **kwargs)
|
||||
|
||||
def query(
|
||||
|
||||
@@ -21,8 +21,6 @@ class QdrantDB(BaseVectorDB):
|
||||
Qdrant as vector database
|
||||
"""
|
||||
|
||||
BATCH_SIZE = 10
|
||||
|
||||
def __init__(self, config: QdrantDBConfig = None):
|
||||
"""
|
||||
Qdrant as vector database
|
||||
@@ -37,6 +35,7 @@ class QdrantDB(BaseVectorDB):
|
||||
"Please make sure the type is right and that you are passing an instance."
|
||||
)
|
||||
self.config = config
|
||||
self.batch_size = self.config.batch_size
|
||||
self.client = QdrantClient(url=os.getenv("QDRANT_URL"), api_key=os.getenv("QDRANT_API_KEY"))
|
||||
# Call parent init here because embedder is needed
|
||||
super().__init__(config=self.config)
|
||||
@@ -116,7 +115,7 @@ class QdrantDB(BaseVectorDB):
|
||||
collection_name=self.collection_name,
|
||||
scroll_filter=models.Filter(must=qdrant_must_filters),
|
||||
offset=offset,
|
||||
limit=self.BATCH_SIZE,
|
||||
limit=self.batch_size,
|
||||
)
|
||||
offset = response[1]
|
||||
for doc in response[0]:
|
||||
@@ -148,13 +147,13 @@ class QdrantDB(BaseVectorDB):
|
||||
qdrant_ids.append(id)
|
||||
payloads.append({"identifier": id, "text": document, "metadata": copy.deepcopy(metadata)})
|
||||
|
||||
for i in tqdm(range(0, len(qdrant_ids), self.BATCH_SIZE), desc="Adding data in batches"):
|
||||
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(
|
||||
ids=qdrant_ids[i : i + self.BATCH_SIZE],
|
||||
payloads=payloads[i : i + self.BATCH_SIZE],
|
||||
vectors=embeddings[i : i + self.BATCH_SIZE],
|
||||
ids=qdrant_ids[i : i + self.batch_size],
|
||||
payloads=payloads[i : i + self.batch_size],
|
||||
vectors=embeddings[i : i + self.batch_size],
|
||||
),
|
||||
**kwargs,
|
||||
)
|
||||
@@ -251,4 +250,4 @@ class QdrantDB(BaseVectorDB):
|
||||
|
||||
def delete(self, where: dict):
|
||||
db_filter = self._generate_query(where)
|
||||
self.client.delete(collection_name=self.collection_name, points_selector=db_filter)
|
||||
self.client.delete(collection_name=self.collection_name, points_selector=db_filter)
|
||||
|
||||
@@ -20,8 +20,6 @@ class WeaviateDB(BaseVectorDB):
|
||||
Weaviate as vector database
|
||||
"""
|
||||
|
||||
BATCH_SIZE = 100
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Optional[WeaviateDBConfig] = None,
|
||||
@@ -40,6 +38,7 @@ class WeaviateDB(BaseVectorDB):
|
||||
"Please make sure the type is right and that you are passing an instance."
|
||||
)
|
||||
self.config = config
|
||||
self.batch_size = self.config.batch_size
|
||||
self.client = weaviate.Client(
|
||||
url=os.environ.get("WEAVIATE_ENDPOINT"),
|
||||
auth_client_secret=weaviate.AuthApiKey(api_key=os.environ.get("WEAVIATE_API_KEY")),
|
||||
@@ -169,7 +168,7 @@ class WeaviateDB(BaseVectorDB):
|
||||
)
|
||||
.with_where(weaviate_where_clause)
|
||||
.with_additional(["id"])
|
||||
.with_limit(limit or self.BATCH_SIZE),
|
||||
.with_limit(limit or self.batch_size),
|
||||
offset,
|
||||
)
|
||||
|
||||
@@ -198,7 +197,7 @@ class WeaviateDB(BaseVectorDB):
|
||||
:type ids: list[str]
|
||||
"""
|
||||
embeddings = self.embedder.embedding_fn(documents)
|
||||
self.client.batch.configure(batch_size=self.BATCH_SIZE, timeout_retries=3) # Configure batch
|
||||
self.client.batch.configure(batch_size=self.batch_size, timeout_retries=3) # Configure batch
|
||||
with self.client.batch as batch: # Initialize a batch process
|
||||
for id, text, metadata, embedding in zip(ids, documents, metadatas, embeddings):
|
||||
doc = {"identifier": id, "text": text}
|
||||
|
||||
@@ -0,0 +1,803 @@
|
||||
{
|
||||
"openai/gpt-4": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00003,
|
||||
"output_cost_per_token": 0.00006
|
||||
},
|
||||
"openai/gpt-4o": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000005,
|
||||
"output_cost_per_token": 0.000015
|
||||
},
|
||||
"openai/gpt-4o-2024-05-13": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000005,
|
||||
"output_cost_per_token": 0.000015
|
||||
},
|
||||
"openai/gpt-4-turbo-preview": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00001,
|
||||
"output_cost_per_token": 0.00003
|
||||
},
|
||||
"openai/gpt-4-0314": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00003,
|
||||
"output_cost_per_token": 0.00006
|
||||
},
|
||||
"openai/gpt-4-0613": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00003,
|
||||
"output_cost_per_token": 0.00006
|
||||
},
|
||||
"openai/gpt-4-32k": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00006,
|
||||
"output_cost_per_token": 0.00012
|
||||
},
|
||||
"openai/gpt-4-32k-0314": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00006,
|
||||
"output_cost_per_token": 0.00012
|
||||
},
|
||||
"openai/gpt-4-32k-0613": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00006,
|
||||
"output_cost_per_token": 0.00012
|
||||
},
|
||||
"openai/gpt-4-turbo": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00001,
|
||||
"output_cost_per_token": 0.00003
|
||||
},
|
||||
"openai/gpt-4-turbo-2024-04-09": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00001,
|
||||
"output_cost_per_token": 0.00003
|
||||
},
|
||||
"openai/gpt-4-1106-preview": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00001,
|
||||
"output_cost_per_token": 0.00003
|
||||
},
|
||||
"openai/gpt-4-0125-preview": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00001,
|
||||
"output_cost_per_token": 0.00003
|
||||
},
|
||||
"openai/gpt-3.5-turbo": {
|
||||
"max_tokens": 4097,
|
||||
"max_input_tokens": 16385,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.0000015,
|
||||
"output_cost_per_token": 0.000002
|
||||
},
|
||||
"openai/gpt-3.5-turbo-0301": {
|
||||
"max_tokens": 4097,
|
||||
"max_input_tokens": 4097,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.0000015,
|
||||
"output_cost_per_token": 0.000002
|
||||
},
|
||||
"openai/gpt-3.5-turbo-0613": {
|
||||
"input_cost_per_token": 0.0000015,
|
||||
"output_cost_per_token": 0.000002
|
||||
},
|
||||
"openai/gpt-3.5-turbo-1106": {
|
||||
"max_tokens": 16385,
|
||||
"max_input_tokens": 16385,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.0000010,
|
||||
"output_cost_per_token": 0.0000020
|
||||
},
|
||||
"openai/gpt-3.5-turbo-0125": {
|
||||
"max_tokens": 16385,
|
||||
"max_input_tokens": 16385,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.0000005,
|
||||
"output_cost_per_token": 0.0000015
|
||||
},
|
||||
"openai/gpt-3.5-turbo-16k": {
|
||||
"max_tokens": 16385,
|
||||
"max_input_tokens": 16385,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000003,
|
||||
"output_cost_per_token": 0.000004
|
||||
},
|
||||
"openai/gpt-3.5-turbo-16k-0613": {
|
||||
"max_tokens": 16385,
|
||||
"max_input_tokens": 16385,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000003,
|
||||
"output_cost_per_token": 0.000004
|
||||
},
|
||||
"openai/text-embedding-3-large": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 8191,
|
||||
"output_vector_size": 3072,
|
||||
"input_cost_per_token": 0.00000013,
|
||||
"output_cost_per_token": 0.000000
|
||||
},
|
||||
"openai/text-embedding-3-small": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 8191,
|
||||
"output_vector_size": 1536,
|
||||
"input_cost_per_token": 0.00000002,
|
||||
"output_cost_per_token": 0.000000
|
||||
},
|
||||
"openai/text-embedding-ada-002": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 8191,
|
||||
"output_vector_size": 1536,
|
||||
"input_cost_per_token": 0.0000001,
|
||||
"output_cost_per_token": 0.000000
|
||||
},
|
||||
"openai/text-embedding-ada-002-v2": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 8191,
|
||||
"input_cost_per_token": 0.0000001,
|
||||
"output_cost_per_token": 0.000000
|
||||
},
|
||||
"openai/babbage-002": {
|
||||
"max_tokens": 16384,
|
||||
"max_input_tokens": 16384,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.0000004,
|
||||
"output_cost_per_token": 0.0000004
|
||||
},
|
||||
"openai/davinci-002": {
|
||||
"max_tokens": 16384,
|
||||
"max_input_tokens": 16384,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000002,
|
||||
"output_cost_per_token": 0.000002
|
||||
},
|
||||
"openai/gpt-3.5-turbo-instruct": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.0000015,
|
||||
"output_cost_per_token": 0.000002
|
||||
},
|
||||
"openai/gpt-3.5-turbo-instruct-0914": {
|
||||
"max_tokens": 4097,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 4097,
|
||||
"input_cost_per_token": 0.0000015,
|
||||
"output_cost_per_token": 0.000002
|
||||
},
|
||||
"azure/gpt-4o": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000005,
|
||||
"output_cost_per_token": 0.000015
|
||||
},
|
||||
"azure/gpt-4-turbo-2024-04-09": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00001,
|
||||
"output_cost_per_token": 0.00003
|
||||
},
|
||||
"azure/gpt-4-0125-preview": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00001,
|
||||
"output_cost_per_token": 0.00003
|
||||
},
|
||||
"azure/gpt-4-1106-preview": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00001,
|
||||
"output_cost_per_token": 0.00003
|
||||
},
|
||||
"azure/gpt-4-0613": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00003,
|
||||
"output_cost_per_token": 0.00006
|
||||
},
|
||||
"azure/gpt-4-32k-0613": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00006,
|
||||
"output_cost_per_token": 0.00012
|
||||
},
|
||||
"azure/gpt-4-32k": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00006,
|
||||
"output_cost_per_token": 0.00012
|
||||
},
|
||||
"azure/gpt-4": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00003,
|
||||
"output_cost_per_token": 0.00006
|
||||
},
|
||||
"azure/gpt-4-turbo": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00001,
|
||||
"output_cost_per_token": 0.00003
|
||||
},
|
||||
"azure/gpt-4-turbo-vision-preview": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00001,
|
||||
"output_cost_per_token": 0.00003
|
||||
},
|
||||
"azure/gpt-3.5-turbo-16k-0613": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 16385,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000003,
|
||||
"output_cost_per_token": 0.000004
|
||||
},
|
||||
"azure/gpt-3.5-turbo-1106": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 16384,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.0000015,
|
||||
"output_cost_per_token": 0.000002
|
||||
},
|
||||
"azure/gpt-3.5-turbo-0125": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 16384,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.0000005,
|
||||
"output_cost_per_token": 0.0000015
|
||||
},
|
||||
"azure/gpt-3.5-turbo-16k": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 16385,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000003,
|
||||
"output_cost_per_token": 0.000004
|
||||
},
|
||||
"azure/gpt-3.5-turbo": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4097,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.0000005,
|
||||
"output_cost_per_token": 0.0000015
|
||||
},
|
||||
"azure/gpt-3.5-turbo-instruct-0914": {
|
||||
"max_tokens": 4097,
|
||||
"max_input_tokens": 4097,
|
||||
"input_cost_per_token": 0.0000015,
|
||||
"output_cost_per_token": 0.000002
|
||||
},
|
||||
"azure/gpt-3.5-turbo-instruct": {
|
||||
"max_tokens": 4097,
|
||||
"max_input_tokens": 4097,
|
||||
"input_cost_per_token": 0.0000015,
|
||||
"output_cost_per_token": 0.000002
|
||||
},
|
||||
"azure/text-embedding-ada-002": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 8191,
|
||||
"input_cost_per_token": 0.0000001,
|
||||
"output_cost_per_token": 0.000000
|
||||
},
|
||||
"azure/text-embedding-3-large": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 8191,
|
||||
"input_cost_per_token": 0.00000013,
|
||||
"output_cost_per_token": 0.000000
|
||||
},
|
||||
"azure/text-embedding-3-small": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 8191,
|
||||
"input_cost_per_token": 0.00000002,
|
||||
"output_cost_per_token": 0.000000
|
||||
},
|
||||
"mistralai/mistral-tiny": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.00000025,
|
||||
"output_cost_per_token": 0.00000025
|
||||
},
|
||||
"mistralai/mistral-small": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.000001,
|
||||
"output_cost_per_token": 0.000003
|
||||
},
|
||||
"mistralai/mistral-small-latest": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.000001,
|
||||
"output_cost_per_token": 0.000003
|
||||
},
|
||||
"mistralai/mistral-medium": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.0000027,
|
||||
"output_cost_per_token": 0.0000081
|
||||
},
|
||||
"mistralai/mistral-medium-latest": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.0000027,
|
||||
"output_cost_per_token": 0.0000081
|
||||
},
|
||||
"mistralai/mistral-medium-2312": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.0000027,
|
||||
"output_cost_per_token": 0.0000081
|
||||
},
|
||||
"mistralai/mistral-large-latest": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.000004,
|
||||
"output_cost_per_token": 0.000012
|
||||
},
|
||||
"mistralai/mistral-large-2402": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.000004,
|
||||
"output_cost_per_token": 0.000012
|
||||
},
|
||||
"mistralai/open-mistral-7b": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.00000025,
|
||||
"output_cost_per_token": 0.00000025
|
||||
},
|
||||
"mistralai/open-mixtral-8x7b": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.0000007,
|
||||
"output_cost_per_token": 0.0000007
|
||||
},
|
||||
"mistralai/open-mixtral-8x22b": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 64000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.000002,
|
||||
"output_cost_per_token": 0.000006
|
||||
},
|
||||
"mistralai/codestral-latest": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.000001,
|
||||
"output_cost_per_token": 0.000003
|
||||
},
|
||||
"mistralai/codestral-2405": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.000001,
|
||||
"output_cost_per_token": 0.000003
|
||||
},
|
||||
"mistralai/mistral-embed": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 8192,
|
||||
"input_cost_per_token": 0.0000001,
|
||||
"output_cost_per_token": 0.0
|
||||
},
|
||||
"groq/llama2-70b-4096": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4096,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00000070,
|
||||
"output_cost_per_token": 0.00000080
|
||||
},
|
||||
"groq/llama3-8b-8192": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.00000010,
|
||||
"output_cost_per_token": 0.00000010
|
||||
},
|
||||
"groq/llama3-70b-8192": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.00000064,
|
||||
"output_cost_per_token": 0.00000080
|
||||
},
|
||||
"groq/mixtral-8x7b-32768": {
|
||||
"max_tokens": 32768,
|
||||
"max_input_tokens": 32768,
|
||||
"max_output_tokens": 32768,
|
||||
"input_cost_per_token": 0.00000027,
|
||||
"output_cost_per_token": 0.00000027
|
||||
},
|
||||
"groq/gemma-7b-it": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.00000010,
|
||||
"output_cost_per_token": 0.00000010
|
||||
},
|
||||
"anthropic/claude-instant-1": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 100000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.00000163,
|
||||
"output_cost_per_token": 0.00000551
|
||||
},
|
||||
"anthropic/claude-instant-1.2": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 100000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.000000163,
|
||||
"output_cost_per_token": 0.000000551
|
||||
},
|
||||
"anthropic/claude-2": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 100000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.000008,
|
||||
"output_cost_per_token": 0.000024
|
||||
},
|
||||
"anthropic/claude-2.1": {
|
||||
"max_tokens": 8191,
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 8191,
|
||||
"input_cost_per_token": 0.000008,
|
||||
"output_cost_per_token": 0.000024
|
||||
},
|
||||
"anthropic/claude-3-haiku-20240307": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00000025,
|
||||
"output_cost_per_token": 0.00000125
|
||||
},
|
||||
"anthropic/claude-3-opus-20240229": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000015,
|
||||
"output_cost_per_token": 0.000075
|
||||
},
|
||||
"anthropic/claude-3-sonnet-20240229": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000003,
|
||||
"output_cost_per_token": 0.000015
|
||||
},
|
||||
"vertexai/chat-bison": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/chat-bison@001": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/chat-bison@002": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 8192,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/chat-bison-32k": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/code-bison": {
|
||||
"max_tokens": 1024,
|
||||
"max_input_tokens": 6144,
|
||||
"max_output_tokens": 1024,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/code-bison@001": {
|
||||
"max_tokens": 1024,
|
||||
"max_input_tokens": 6144,
|
||||
"max_output_tokens": 1024,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/code-gecko@001": {
|
||||
"max_tokens": 64,
|
||||
"max_input_tokens": 2048,
|
||||
"max_output_tokens": 64,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/code-gecko@002": {
|
||||
"max_tokens": 64,
|
||||
"max_input_tokens": 2048,
|
||||
"max_output_tokens": 64,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/code-gecko": {
|
||||
"max_tokens": 64,
|
||||
"max_input_tokens": 2048,
|
||||
"max_output_tokens": 64,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/codechat-bison": {
|
||||
"max_tokens": 1024,
|
||||
"max_input_tokens": 6144,
|
||||
"max_output_tokens": 1024,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/codechat-bison@001": {
|
||||
"max_tokens": 1024,
|
||||
"max_input_tokens": 6144,
|
||||
"max_output_tokens": 1024,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/codechat-bison-32k": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.000000125,
|
||||
"output_cost_per_token": 0.000000125
|
||||
},
|
||||
"vertexai/gemini-pro": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 32760,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.00000025,
|
||||
"output_cost_per_token": 0.0000005
|
||||
},
|
||||
"vertexai/gemini-1.0-pro": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 32760,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.00000025,
|
||||
"output_cost_per_token": 0.0000005
|
||||
},
|
||||
"vertexai/gemini-1.0-pro-001": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 32760,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.00000025,
|
||||
"output_cost_per_token": 0.0000005
|
||||
},
|
||||
"vertexai/gemini-1.0-pro-002": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 32760,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.00000025,
|
||||
"output_cost_per_token": 0.0000005
|
||||
},
|
||||
"vertexai/gemini-1.5-pro": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.000000625,
|
||||
"output_cost_per_token": 0.000001875
|
||||
},
|
||||
"vertexai/gemini-1.5-flash-001": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0,
|
||||
"output_cost_per_token": 0
|
||||
},
|
||||
"vertexai/gemini-1.5-flash-preview-0514": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0,
|
||||
"output_cost_per_token": 0
|
||||
},
|
||||
"vertexai/gemini-1.5-pro-001": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.000000625,
|
||||
"output_cost_per_token": 0.000001875
|
||||
},
|
||||
"vertexai/gemini-1.5-pro-preview-0514": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.000000625,
|
||||
"output_cost_per_token": 0.000001875
|
||||
},
|
||||
"vertexai/gemini-1.5-pro-preview-0215": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.000000625,
|
||||
"output_cost_per_token": 0.000001875
|
||||
},
|
||||
"vertexai/gemini-1.5-pro-preview-0409": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0.000000625,
|
||||
"output_cost_per_token": 0.000001875
|
||||
},
|
||||
"vertexai/gemini-experimental": {
|
||||
"max_tokens": 8192,
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 8192,
|
||||
"input_cost_per_token": 0,
|
||||
"output_cost_per_token": 0
|
||||
},
|
||||
"vertexai/gemini-pro-vision": {
|
||||
"max_tokens": 2048,
|
||||
"max_input_tokens": 16384,
|
||||
"max_output_tokens": 2048,
|
||||
"max_images_per_prompt": 16,
|
||||
"max_videos_per_prompt": 1,
|
||||
"max_video_length": 2,
|
||||
"input_cost_per_token": 0.00000025,
|
||||
"output_cost_per_token": 0.0000005
|
||||
},
|
||||
"vertexai/gemini-1.0-pro-vision": {
|
||||
"max_tokens": 2048,
|
||||
"max_input_tokens": 16384,
|
||||
"max_output_tokens": 2048,
|
||||
"max_images_per_prompt": 16,
|
||||
"max_videos_per_prompt": 1,
|
||||
"max_video_length": 2,
|
||||
"input_cost_per_token": 0.00000025,
|
||||
"output_cost_per_token": 0.0000005
|
||||
},
|
||||
"vertexai/gemini-1.0-pro-vision-001": {
|
||||
"max_tokens": 2048,
|
||||
"max_input_tokens": 16384,
|
||||
"max_output_tokens": 2048,
|
||||
"max_images_per_prompt": 16,
|
||||
"max_videos_per_prompt": 1,
|
||||
"max_video_length": 2,
|
||||
"input_cost_per_token": 0.00000025,
|
||||
"output_cost_per_token": 0.0000005
|
||||
},
|
||||
"vertexai/claude-3-sonnet@20240229": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000003,
|
||||
"output_cost_per_token": 0.000015
|
||||
},
|
||||
"vertexai/claude-3-haiku@20240307": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00000025,
|
||||
"output_cost_per_token": 0.00000125
|
||||
},
|
||||
"vertexai/claude-3-opus@20240229": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000015,
|
||||
"output_cost_per_token": 0.000075
|
||||
},
|
||||
"cohere/command-r": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.00000050,
|
||||
"output_cost_per_token": 0.0000015
|
||||
},
|
||||
"cohere/command-light": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4096,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000015,
|
||||
"output_cost_per_token": 0.000015
|
||||
},
|
||||
"cohere/command-r-plus": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000003,
|
||||
"output_cost_per_token": 0.000015
|
||||
},
|
||||
"cohere/command-nightly": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4096,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000015,
|
||||
"output_cost_per_token": 0.000015
|
||||
},
|
||||
"cohere/command": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4096,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000015,
|
||||
"output_cost_per_token": 0.000015
|
||||
},
|
||||
"cohere/command-medium-beta": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4096,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000015,
|
||||
"output_cost_per_token": 0.000015
|
||||
},
|
||||
"cohere/command-xlarge-beta": {
|
||||
"max_tokens": 4096,
|
||||
"max_input_tokens": 4096,
|
||||
"max_output_tokens": 4096,
|
||||
"input_cost_per_token": 0.000015,
|
||||
"output_cost_per_token": 0.000015
|
||||
},
|
||||
"together/together-ai-up-to-3b": {
|
||||
"input_cost_per_token": 0.0000001,
|
||||
"output_cost_per_token": 0.0000001
|
||||
},
|
||||
"together/together-ai-3.1b-7b": {
|
||||
"input_cost_per_token": 0.0000002,
|
||||
"output_cost_per_token": 0.0000002
|
||||
},
|
||||
"together/together-ai-7.1b-20b": {
|
||||
"max_tokens": 1000,
|
||||
"input_cost_per_token": 0.0000004,
|
||||
"output_cost_per_token": 0.0000004
|
||||
},
|
||||
"together/together-ai-20.1b-40b": {
|
||||
"input_cost_per_token": 0.0000008,
|
||||
"output_cost_per_token": 0.0000008
|
||||
},
|
||||
"together/together-ai-40.1b-70b": {
|
||||
"input_cost_per_token": 0.0000009,
|
||||
"output_cost_per_token": 0.0000009
|
||||
},
|
||||
"together/mistralai/Mixtral-8x7B-Instruct-v0.1": {
|
||||
"input_cost_per_token": 0.0000006,
|
||||
"output_cost_per_token": 0.0000006
|
||||
}
|
||||
}
|
||||
Generated
+788
-784
File diff suppressed because it is too large
Load Diff
+8
-7
@@ -1,6 +1,6 @@
|
||||
[tool.poetry]
|
||||
name = "embedchain"
|
||||
version = "0.1.113"
|
||||
version = "0.1.116"
|
||||
description = "Simplest open source retrieval (RAG) framework"
|
||||
authors = [
|
||||
"Taranjeet Singh <taranjeet@embedchain.ai>",
|
||||
@@ -93,7 +93,7 @@ color = true
|
||||
[tool.poetry.dependencies]
|
||||
python = ">=3.9,<=3.13"
|
||||
python-dotenv = "^1.0.0"
|
||||
langchain = "^0.1.4"
|
||||
langchain = ">0.2,<=0.3"
|
||||
requests = "^2.31.0"
|
||||
openai = ">=1.1.1"
|
||||
chromadb = "^0.4.24"
|
||||
@@ -103,6 +103,7 @@ beautifulsoup4 = "^4.12.2"
|
||||
pypdf = "^4.0.1"
|
||||
gptcache = "^0.1.43"
|
||||
pysbd = "^0.3.4"
|
||||
memzero = "^0.0.7"
|
||||
tiktoken = { version = "^0.7.0", optional = true }
|
||||
youtube-transcript-api = { version = "^0.6.1", optional = true }
|
||||
pytube = { version = "^15.0.0", optional = true }
|
||||
@@ -150,12 +151,13 @@ 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 }
|
||||
langchain-mistralai = { version = "^0.1.9", optional = true }
|
||||
langchain-openai = "^0.1.7"
|
||||
langchain-google-vertexai = { version = "^0.0.5", optional = true }
|
||||
langchain-google-vertexai = { version = "^1.0.6", optional = true }
|
||||
sqlalchemy = "^2.0.27"
|
||||
alembic = "^1.13.1"
|
||||
langchain-cohere = "^0.1.4"
|
||||
langchain-community = "^0.2.6"
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
black = "^23.3.0"
|
||||
@@ -183,9 +185,8 @@ slack = ["slack-sdk", "flask"]
|
||||
whatsapp = ["twilio", "flask"]
|
||||
weaviate = ["weaviate-client"]
|
||||
qdrant = ["qdrant-client"]
|
||||
huggingface_hub=["huggingface_hub"]
|
||||
cohere = ["cohere"]
|
||||
together = ["together"]
|
||||
huggingface_hub=["huggingface_hub"]
|
||||
milvus = ["pymilvus"]
|
||||
dataloaders=[
|
||||
"youtube-transcript-api",
|
||||
@@ -226,4 +227,4 @@ mistralai = ["langchain-mistralai"]
|
||||
[tool.poetry.group.docs.dependencies]
|
||||
|
||||
[tool.poetry.scripts]
|
||||
ec = "embedchain.cli:cli"
|
||||
ec = "embedchain.cli:cli"
|
||||
@@ -0,0 +1,18 @@
|
||||
|
||||
from unittest.mock import patch
|
||||
from embedchain.config import BaseEmbedderConfig
|
||||
from embedchain.embedder.huggingface import HuggingFaceEmbedder
|
||||
|
||||
|
||||
def test_huggingface_embedder_with_model(monkeypatch):
|
||||
config = BaseEmbedderConfig(model="test-model", model_kwargs={"param": "value"})
|
||||
with patch('embedchain.embedder.huggingface.HuggingFaceEmbeddings') as mock_embeddings:
|
||||
embedder = HuggingFaceEmbedder(config=config)
|
||||
assert embedder.config.model == "test-model"
|
||||
assert embedder.config.model_kwargs == {"param": "value"}
|
||||
mock_embeddings.assert_called_once_with(
|
||||
model_name="test-model",
|
||||
model_kwargs={"param": "value"}
|
||||
)
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ from embedchain.llm.anthropic import AnthropicLlm
|
||||
@pytest.fixture
|
||||
def anthropic_llm():
|
||||
os.environ["ANTHROPIC_API_KEY"] = "test_api_key"
|
||||
config = BaseLlmConfig(temperature=0.5, model="gpt2")
|
||||
config = BaseLlmConfig(temperature=0.5, model="claude-instant-1", token_usage=False)
|
||||
return AnthropicLlm(config)
|
||||
|
||||
|
||||
@@ -20,7 +20,7 @@ def test_get_llm_model_answer(anthropic_llm):
|
||||
prompt = "Test Prompt"
|
||||
response = anthropic_llm.get_llm_model_answer(prompt)
|
||||
assert response == "Test Response"
|
||||
mock_method.assert_called_once_with(prompt=prompt, config=anthropic_llm.config)
|
||||
mock_method.assert_called_once_with(prompt, anthropic_llm.config)
|
||||
|
||||
|
||||
def test_get_messages(anthropic_llm):
|
||||
@@ -31,3 +31,24 @@ def test_get_messages(anthropic_llm):
|
||||
SystemMessage(content="Test System Prompt", additional_kwargs={}),
|
||||
HumanMessage(content="Test Prompt", additional_kwargs={}, example=False),
|
||||
]
|
||||
|
||||
|
||||
def test_get_llm_model_answer_with_token_usage(anthropic_llm):
|
||||
test_config = BaseLlmConfig(
|
||||
temperature=anthropic_llm.config.temperature, model=anthropic_llm.config.model, token_usage=True
|
||||
)
|
||||
anthropic_llm.config = test_config
|
||||
with patch.object(
|
||||
AnthropicLlm, "_get_answer", return_value=("Test Response", {"input_tokens": 1, "output_tokens": 2})
|
||||
) as mock_method:
|
||||
prompt = "Test Prompt"
|
||||
response, token_info = anthropic_llm.get_llm_model_answer(prompt)
|
||||
assert response == "Test Response"
|
||||
assert token_info == {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": 2,
|
||||
"total_tokens": 3,
|
||||
"total_cost": 1.265e-05,
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
mock_method.assert_called_once_with(prompt, anthropic_llm.config)
|
||||
|
||||
@@ -9,7 +9,7 @@ from embedchain.llm.cohere import CohereLlm
|
||||
@pytest.fixture
|
||||
def cohere_llm_config():
|
||||
os.environ["COHERE_API_KEY"] = "test_api_key"
|
||||
config = BaseLlmConfig(model="gptd-instruct-tft", max_tokens=50, temperature=0.7, top_p=0.8)
|
||||
config = BaseLlmConfig(model="command-r", max_tokens=100, temperature=0.7, top_p=0.8, token_usage=False)
|
||||
yield config
|
||||
os.environ.pop("COHERE_API_KEY")
|
||||
|
||||
@@ -36,10 +36,35 @@ def test_get_llm_model_answer(cohere_llm_config, mocker):
|
||||
assert answer == "Test answer"
|
||||
|
||||
|
||||
def test_get_llm_model_answer_with_token_usage(cohere_llm_config, mocker):
|
||||
test_config = BaseLlmConfig(
|
||||
temperature=cohere_llm_config.temperature,
|
||||
max_tokens=cohere_llm_config.max_tokens,
|
||||
top_p=cohere_llm_config.top_p,
|
||||
model=cohere_llm_config.model,
|
||||
token_usage=True,
|
||||
)
|
||||
mocker.patch(
|
||||
"embedchain.llm.cohere.CohereLlm._get_answer",
|
||||
return_value=("Test answer", {"input_tokens": 1, "output_tokens": 2}),
|
||||
)
|
||||
|
||||
llm = CohereLlm(test_config)
|
||||
answer, token_info = llm.get_llm_model_answer("Test query")
|
||||
|
||||
assert answer == "Test answer"
|
||||
assert token_info == {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": 2,
|
||||
"total_tokens": 3,
|
||||
"total_cost": 3.5e-06,
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
|
||||
|
||||
def test_get_answer_mocked_cohere(cohere_llm_config, mocker):
|
||||
mocked_cohere = mocker.patch("embedchain.llm.cohere.Cohere")
|
||||
mock_instance = mocked_cohere.return_value
|
||||
mock_instance.invoke.return_value = "Mocked answer"
|
||||
mocked_cohere = mocker.patch("embedchain.llm.cohere.ChatCohere")
|
||||
mocked_cohere.return_value.invoke.return_value.content = "Mocked answer"
|
||||
|
||||
llm = CohereLlm(cohere_llm_config)
|
||||
prompt = "Test query"
|
||||
|
||||
@@ -24,7 +24,7 @@ def test_mistralai_llm_init(monkeypatch):
|
||||
|
||||
|
||||
def test_get_llm_model_answer(monkeypatch, mistralai_llm_config):
|
||||
def mock_get_answer(prompt, config):
|
||||
def mock_get_answer(self, prompt, config):
|
||||
return "Generated Text"
|
||||
|
||||
monkeypatch.setattr(MistralAILlm, "_get_answer", mock_get_answer)
|
||||
@@ -36,7 +36,7 @@ def test_get_llm_model_answer(monkeypatch, mistralai_llm_config):
|
||||
|
||||
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")
|
||||
monkeypatch.setattr(MistralAILlm, "_get_answer", lambda self, prompt, config: "Generated Text")
|
||||
llm = MistralAILlm(config=mistralai_llm_config)
|
||||
result = llm.get_llm_model_answer("test prompt")
|
||||
|
||||
@@ -44,7 +44,7 @@ def test_get_llm_model_answer_with_system_prompt(monkeypatch, mistralai_llm_conf
|
||||
|
||||
|
||||
def test_get_llm_model_answer_empty_prompt(monkeypatch, mistralai_llm_config):
|
||||
monkeypatch.setattr(MistralAILlm, "_get_answer", lambda prompt, config: "Generated Text")
|
||||
monkeypatch.setattr(MistralAILlm, "_get_answer", lambda self, prompt, config: "Generated Text")
|
||||
llm = MistralAILlm(config=mistralai_llm_config)
|
||||
result = llm.get_llm_model_answer("")
|
||||
|
||||
@@ -53,8 +53,35 @@ def test_get_llm_model_answer_empty_prompt(monkeypatch, mistralai_llm_config):
|
||||
|
||||
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")
|
||||
monkeypatch.setattr(MistralAILlm, "_get_answer", lambda self, 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_with_token_usage(monkeypatch, mistralai_llm_config):
|
||||
test_config = BaseLlmConfig(
|
||||
temperature=mistralai_llm_config.temperature,
|
||||
max_tokens=mistralai_llm_config.max_tokens,
|
||||
top_p=mistralai_llm_config.top_p,
|
||||
model=mistralai_llm_config.model,
|
||||
token_usage=True,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
MistralAILlm,
|
||||
"_get_answer",
|
||||
lambda self, prompt, config: ("Generated Text", {"prompt_tokens": 1, "completion_tokens": 2}),
|
||||
)
|
||||
|
||||
llm = MistralAILlm(test_config)
|
||||
answer, token_info = llm.get_llm_model_answer("Test query")
|
||||
|
||||
assert answer == "Generated Text"
|
||||
assert token_info == {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": 2,
|
||||
"total_tokens": 3,
|
||||
"total_cost": 7.5e-07,
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
|
||||
+122
-4
@@ -1,5 +1,6 @@
|
||||
import os
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from langchain.callbacks.streaming_stdout import StreamingStdOutCallbackHandler
|
||||
|
||||
@@ -7,15 +8,27 @@ from embedchain.config import BaseLlmConfig
|
||||
from embedchain.llm.openai import OpenAILlm
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def config():
|
||||
@pytest.fixture()
|
||||
def env_config():
|
||||
os.environ["OPENAI_API_KEY"] = "test_api_key"
|
||||
os.environ["OPENAI_API_BASE"] = "https://api.openai.com/v1/engines/"
|
||||
yield
|
||||
os.environ.pop("OPENAI_API_KEY")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def config(env_config):
|
||||
config = BaseLlmConfig(
|
||||
temperature=0.7, max_tokens=50, top_p=0.8, stream=False, system_prompt="System prompt", model="gpt-3.5-turbo"
|
||||
temperature=0.7,
|
||||
max_tokens=50,
|
||||
top_p=0.8,
|
||||
stream=False,
|
||||
system_prompt="System prompt",
|
||||
model="gpt-3.5-turbo",
|
||||
http_client_proxies=None,
|
||||
http_async_client_proxies=None,
|
||||
)
|
||||
yield config
|
||||
os.environ.pop("OPENAI_API_KEY")
|
||||
|
||||
|
||||
def test_get_llm_model_answer(config, mocker):
|
||||
@@ -49,6 +62,35 @@ def test_get_llm_model_answer_empty_prompt(config, mocker):
|
||||
mocked_get_answer.assert_called_once_with("", config)
|
||||
|
||||
|
||||
def test_get_llm_model_answer_with_token_usage(config, mocker):
|
||||
test_config = BaseLlmConfig(
|
||||
temperature=config.temperature,
|
||||
max_tokens=config.max_tokens,
|
||||
top_p=config.top_p,
|
||||
stream=config.stream,
|
||||
system_prompt=config.system_prompt,
|
||||
model=config.model,
|
||||
token_usage=True,
|
||||
)
|
||||
mocked_get_answer = mocker.patch(
|
||||
"embedchain.llm.openai.OpenAILlm._get_answer",
|
||||
return_value=("Test answer", {"prompt_tokens": 1, "completion_tokens": 2}),
|
||||
)
|
||||
|
||||
llm = OpenAILlm(test_config)
|
||||
answer, token_info = llm.get_llm_model_answer("Test query")
|
||||
|
||||
assert answer == "Test answer"
|
||||
assert token_info == {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": 2,
|
||||
"total_tokens": 3,
|
||||
"total_cost": 5.5e-06,
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
mocked_get_answer.assert_called_once_with("Test query", test_config)
|
||||
|
||||
|
||||
def test_get_llm_model_answer_with_streaming(config, mocker):
|
||||
config.stream = True
|
||||
mocked_openai_chat = mocker.patch("embedchain.llm.openai.ChatOpenAI")
|
||||
@@ -75,6 +117,8 @@ def test_get_llm_model_answer_without_system_prompt(config, mocker):
|
||||
model_kwargs={"top_p": config.top_p},
|
||||
api_key=os.environ["OPENAI_API_KEY"],
|
||||
base_url=os.environ["OPENAI_API_BASE"],
|
||||
http_client=None,
|
||||
http_async_client=None,
|
||||
)
|
||||
|
||||
|
||||
@@ -93,6 +137,8 @@ def test_get_llm_model_answer_with_special_headers(config, mocker):
|
||||
api_key=os.environ["OPENAI_API_KEY"],
|
||||
base_url=os.environ["OPENAI_API_BASE"],
|
||||
default_headers={"test": "test"},
|
||||
http_client=None,
|
||||
http_async_client=None,
|
||||
)
|
||||
|
||||
|
||||
@@ -110,6 +156,8 @@ def test_get_llm_model_answer_with_model_kwargs(config, mocker):
|
||||
model_kwargs={"top_p": config.top_p, "response_format": {"type": "json_object"}},
|
||||
api_key=os.environ["OPENAI_API_KEY"],
|
||||
base_url=os.environ["OPENAI_API_BASE"],
|
||||
http_client=None,
|
||||
http_async_client=None,
|
||||
)
|
||||
|
||||
|
||||
@@ -136,8 +184,78 @@ def test_get_llm_model_answer_with_tools(config, mocker, mock_return, expected):
|
||||
model_kwargs={"top_p": config.top_p},
|
||||
api_key=os.environ["OPENAI_API_KEY"],
|
||||
base_url=os.environ["OPENAI_API_BASE"],
|
||||
http_client=None,
|
||||
http_async_client=None,
|
||||
)
|
||||
mocked_convert_to_openai_tool.assert_called_once_with({"test": "test"})
|
||||
mocked_json_output_tools_parser.assert_called_once()
|
||||
|
||||
assert answer == expected
|
||||
|
||||
|
||||
def test_get_llm_model_answer_with_http_client_proxies(env_config, mocker):
|
||||
mocked_openai_chat = mocker.patch("embedchain.llm.openai.ChatOpenAI")
|
||||
mock_http_client = mocker.Mock(spec=httpx.Client)
|
||||
mock_http_client_instance = mocker.Mock(spec=httpx.Client)
|
||||
mock_http_client.return_value = mock_http_client_instance
|
||||
|
||||
mocker.patch("httpx.Client", new=mock_http_client)
|
||||
|
||||
config = BaseLlmConfig(
|
||||
temperature=0.7,
|
||||
max_tokens=50,
|
||||
top_p=0.8,
|
||||
stream=False,
|
||||
system_prompt="System prompt",
|
||||
model="gpt-3.5-turbo",
|
||||
http_client_proxies="http://testproxy.mem0.net:8000",
|
||||
)
|
||||
|
||||
llm = OpenAILlm(config)
|
||||
llm.get_llm_model_answer("Test query")
|
||||
|
||||
mocked_openai_chat.assert_called_once_with(
|
||||
model=config.model,
|
||||
temperature=config.temperature,
|
||||
max_tokens=config.max_tokens,
|
||||
model_kwargs={"top_p": config.top_p},
|
||||
api_key=os.environ["OPENAI_API_KEY"],
|
||||
base_url=os.environ["OPENAI_API_BASE"],
|
||||
http_client=mock_http_client_instance,
|
||||
http_async_client=None,
|
||||
)
|
||||
mock_http_client.assert_called_once_with(proxies="http://testproxy.mem0.net:8000")
|
||||
|
||||
|
||||
def test_get_llm_model_answer_with_http_async_client_proxies(env_config, mocker):
|
||||
mocked_openai_chat = mocker.patch("embedchain.llm.openai.ChatOpenAI")
|
||||
mock_http_async_client = mocker.Mock(spec=httpx.AsyncClient)
|
||||
mock_http_async_client_instance = mocker.Mock(spec=httpx.AsyncClient)
|
||||
mock_http_async_client.return_value = mock_http_async_client_instance
|
||||
|
||||
mocker.patch("httpx.AsyncClient", new=mock_http_async_client)
|
||||
|
||||
config = BaseLlmConfig(
|
||||
temperature=0.7,
|
||||
max_tokens=50,
|
||||
top_p=0.8,
|
||||
stream=False,
|
||||
system_prompt="System prompt",
|
||||
model="gpt-3.5-turbo",
|
||||
http_async_client_proxies={"http://": "http://testproxy.mem0.net:8000"},
|
||||
)
|
||||
|
||||
llm = OpenAILlm(config)
|
||||
llm.get_llm_model_answer("Test query")
|
||||
|
||||
mocked_openai_chat.assert_called_once_with(
|
||||
model=config.model,
|
||||
temperature=config.temperature,
|
||||
max_tokens=config.max_tokens,
|
||||
model_kwargs={"top_p": config.top_p},
|
||||
api_key=os.environ["OPENAI_API_KEY"],
|
||||
base_url=os.environ["OPENAI_API_BASE"],
|
||||
http_client=None,
|
||||
http_async_client=mock_http_async_client_instance,
|
||||
)
|
||||
mock_http_async_client.assert_called_once_with(proxies={"http://": "http://testproxy.mem0.net:8000"})
|
||||
|
||||
@@ -9,7 +9,7 @@ from embedchain.llm.together import TogetherLlm
|
||||
@pytest.fixture
|
||||
def together_llm_config():
|
||||
os.environ["TOGETHER_API_KEY"] = "test_api_key"
|
||||
config = BaseLlmConfig(model="togethercomputer/RedPajama-INCITE-7B-Base", max_tokens=50, temperature=0.7, top_p=0.8)
|
||||
config = BaseLlmConfig(model="together-ai-up-to-3b", max_tokens=50, temperature=0.7, top_p=0.8)
|
||||
yield config
|
||||
os.environ.pop("TOGETHER_API_KEY")
|
||||
|
||||
@@ -36,10 +36,36 @@ def test_get_llm_model_answer(together_llm_config, mocker):
|
||||
assert answer == "Test answer"
|
||||
|
||||
|
||||
def test_get_llm_model_answer_with_token_usage(together_llm_config, mocker):
|
||||
test_config = BaseLlmConfig(
|
||||
temperature=together_llm_config.temperature,
|
||||
max_tokens=together_llm_config.max_tokens,
|
||||
top_p=together_llm_config.top_p,
|
||||
model=together_llm_config.model,
|
||||
token_usage=True,
|
||||
)
|
||||
mocker.patch(
|
||||
"embedchain.llm.together.TogetherLlm._get_answer",
|
||||
return_value=("Test answer", {"prompt_tokens": 1, "completion_tokens": 2}),
|
||||
)
|
||||
|
||||
llm = TogetherLlm(test_config)
|
||||
answer, token_info = llm.get_llm_model_answer("Test query")
|
||||
|
||||
assert answer == "Test answer"
|
||||
assert token_info == {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": 2,
|
||||
"total_tokens": 3,
|
||||
"total_cost": 3e-07,
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
|
||||
|
||||
def test_get_answer_mocked_together(together_llm_config, mocker):
|
||||
mocked_together = mocker.patch("embedchain.llm.together.Together")
|
||||
mocked_together = mocker.patch("embedchain.llm.together.ChatTogether")
|
||||
mock_instance = mocked_together.return_value
|
||||
mock_instance.invoke.return_value = "Mocked answer"
|
||||
mock_instance.invoke.return_value.content = "Mocked answer"
|
||||
|
||||
llm = TogetherLlm(together_llm_config)
|
||||
prompt = "Test query"
|
||||
|
||||
@@ -24,7 +24,32 @@ def test_get_llm_model_answer(vertexai_llm):
|
||||
prompt = "Test Prompt"
|
||||
response = vertexai_llm.get_llm_model_answer(prompt)
|
||||
assert response == "Test Response"
|
||||
mock_method.assert_called_once_with(prompt=prompt, config=vertexai_llm.config)
|
||||
mock_method.assert_called_once_with(prompt, vertexai_llm.config)
|
||||
|
||||
|
||||
def test_get_llm_model_answer_with_token_usage(vertexai_llm):
|
||||
test_config = BaseLlmConfig(
|
||||
temperature=vertexai_llm.config.temperature,
|
||||
max_tokens=vertexai_llm.config.max_tokens,
|
||||
top_p=vertexai_llm.config.top_p,
|
||||
model=vertexai_llm.config.model,
|
||||
token_usage=True,
|
||||
)
|
||||
vertexai_llm.config = test_config
|
||||
with patch.object(
|
||||
VertexAILlm,
|
||||
"_get_answer",
|
||||
return_value=("Test Response", {"prompt_token_count": 1, "candidates_token_count": 2}),
|
||||
):
|
||||
response, token_info = vertexai_llm.get_llm_model_answer("Test Query")
|
||||
assert response == "Test Response"
|
||||
assert token_info == {
|
||||
"prompt_tokens": 1,
|
||||
"completion_tokens": 2,
|
||||
"total_tokens": 3,
|
||||
"total_cost": 3.75e-07,
|
||||
"cost_currency": "USD",
|
||||
}
|
||||
|
||||
|
||||
@patch("embedchain.llm.vertex_ai.ChatVertexAI")
|
||||
|
||||
@@ -124,7 +124,6 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
db = WeaviateDB()
|
||||
app_config = AppConfig(collect_metrics=False)
|
||||
App(config=app_config, db=db, embedding_model=embedder)
|
||||
db.BATCH_SIZE = 1
|
||||
|
||||
documents = ["This is test document"]
|
||||
metadatas = [None]
|
||||
@@ -132,7 +131,7 @@ class TestWeaviateDb(unittest.TestCase):
|
||||
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_mock.configure.assert_called_once_with(batch_size=100, timeout_retries=3)
|
||||
weaviate_client_batch_enter_mock.add_data_object.assert_any_call(
|
||||
data_object={"text": documents[0]}, class_name="Embedchain_store_1536_metadata", vector=[1, 2, 3]
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user