Spaces:
Runtime error
Runtime error
| import os | |
| from functools import lru_cache | |
| from langchain.chat_models import AzureChatOpenAI, ChatOpenAI | |
| from langchain.chat_models.base import BaseChatModel | |
| from langchain.embeddings import OpenAIEmbeddings | |
| from langchain.schema import Document | |
| from langchain.vectorstores import Qdrant, VectorStore | |
| from edu_assistant.utils.qdrant_utils import load_qdrant_client | |
| def load_llm() -> BaseChatModel: | |
| if os.environ.get("AZURE_OPENAI"): | |
| llm = AzureChatOpenAI( | |
| openai_api_type="azure", | |
| openai_api_key=os.environ.get("AZURE_OPENAI_API_KEY"), | |
| openai_api_base=os.environ.get("AZURE_OPENAI_API_BASE"), | |
| openai_api_version="2023-05-15", | |
| deployment_name=os.environ.get("AZURE_OPENAI_DEPLOYMENT_ID", "gpt-35-turbo"), | |
| model="gpt-3.5-turbo", | |
| temperature=0, | |
| ) | |
| else: | |
| llm = ChatOpenAI( | |
| openai_api_key=os.environ.get("OPENAI_API_KEY"), | |
| openai_proxy=os.environ.get("OPENAI_PROXY", ""), | |
| model="gpt-3.5-turbo", | |
| ) | |
| return llm | |
| def load_gpt4_llm() -> BaseChatModel: | |
| llm = ChatOpenAI( | |
| openai_api_key=os.environ.get("OPENAI_API_KEY"), | |
| openai_proxy=os.environ.get("OPENAI_PROXY", ""), | |
| model="gpt-4", | |
| ) | |
| return llm | |
| def load_gpt4_flag() -> bool: | |
| return os.environ.get("CODEDOG_ENABLE_GPT4") is not None | |
| def load_embeddings(): | |
| if os.environ.get("AZURE_OPENAI"): | |
| embeddings = OpenAIEmbeddings( | |
| openai_api_type="azure", | |
| openai_api_key=os.environ.get("AZURE_OPENAI_API_KEY"), | |
| openai_api_base=os.environ.get("AZURE_OPENAI_API_BASE"), | |
| openai_api_version="2023-05-15", | |
| deployment=os.environ.get("AZURE_OPENAI_EMBEDDING_DEP_ID", ""), | |
| ) | |
| else: | |
| embeddings = OpenAIEmbeddings( | |
| openai_api_key=os.environ.get("OPENAI_API_KEY"), | |
| openai_proxy=os.environ.get("OPENAI_PROXY", ""), | |
| ) | |
| return embeddings | |
| def load_vectorstore(collection_name: str = "default") -> VectorStore: | |
| if os.environ.get("QDRANT_API"): | |
| client = load_qdrant_client() | |
| embeddings = load_embeddings() | |
| doc_store = Qdrant(client=client, collection_name=collection_name, embeddings=embeddings) | |
| return doc_store | |
| return None | |
| def escape_for_prompt(text: str) -> str: | |
| """escape cruly brackets in text for generate prompt. | |
| Args: | |
| text (str): input string. | |
| Returns: | |
| str: escaped string. | |
| """ | |
| return text.replace("{", "{{").replace("}", "}}") | |
| def shrink_docs(docs: list[Document], max_size=50): | |
| """shrink source docs content size for display. | |
| Args: | |
| docs (dict): Retrieval Chain returned docs. | |
| """ | |
| for doc in docs: | |
| doc.page_content = doc.page_content[:max_size] + ".." | |
| return docs | |