Update models.py
Alessandro committed
Sep 20, 2024 at 21:55 UTC
7eaf3c758218d7fdefd03425a053e4ffab6b9255
1 file changed
+6
-1
models.py
+6
-1
@@ -8,6 +8,7 @@ from langchain_anthropic import ChatAnthropic
8
from langchain_groq import ChatGroq
9
from langchain_huggingface import HuggingFaceEmbeddings
10
from langchain_google_genai import GoogleGenerativeAI, HarmBlockThreshold, HarmCategory
11
+from langchain_mistralai import ChatMistralAI
12
from pydantic.v1.types import SecretStr
13
14
@@ -69,6 +70,10 @@ def get_azure_openai_embedding(deployment_name:str, api_key=get_api_key("openai_
70
def get_google_chat(model_name:str, api_key=get_api_key("google"), temperature=DEFAULT_TEMPERATURE):
71
return GoogleGenerativeAI(model=model_name, temperature=temperature, google_api_key=api_key, safety_settings={HarmCategory.HARM_CATEGORY_DANGEROUS_CONTENT: HarmBlockThreshold.BLOCK_NONE }) # type: ignore
72
73
+# Mistral models
74
+def get_mistral_chat(model_name:str, api_key=get_api_key("mistral"), temperature=DEFAULT_TEMPERATURE):
75
+ return ChatMistralAI(model=model_name, temperature=temperature, api_key=api_key) # type: ignore
76
+
77
# Groq models
78
def get_groq_chat(model_name:str, api_key=get_api_key("groq"), temperature=DEFAULT_TEMPERATURE):
79
return ChatGroq(model_name=model_name, temperature=temperature, api_key=api_key) # type: ignore
@@ -78,4 +83,4 @@ def get_openrouter_chat(model_name: str, api_key=get_api_key("openrouter"), temp
83
return ChatOpenAI(api_key=api_key, model=model_name, temperature=temperature, base_url=base_url) # type: ignore
84
85
def get_openrouter_embedding(model_name: str, api_key=get_api_key("openrouter"), base_url=os.getenv("OPEN_ROUTER_BASE_URL") or "https://openrouter.ai/api/v1"):
81
- return OpenAIEmbeddings(model=model_name, api_key=api_key, base_url=base_url) # type: ignore
\ No newline at end of file
86
+ return OpenAIEmbeddings(model=model_name, api_key=api_key, base_url=base_url) # type: ignore