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