| |
| |
| |
| |
|
|
| import os |
| from dotenv import load_dotenv |
|
|
| load_dotenv() |
|
|
| IS_HF_SPACES: bool = bool(os.getenv("SPACE_ID", "")) |
|
|
|
|
| class EmbeddingStore: |
| def __init__(self, model_name=None): |
| if IS_HF_SPACES: |
| from langchain_huggingface import HuggingFaceEmbeddings |
|
|
| hf_model = model_name or os.getenv( |
| "EMBED_MODEL", "sentence-transformers/all-MiniLM-L6-v2" |
| ) |
|
|
| if "/" not in hf_model: |
| hf_model = f"sentence-transformers/{hf_model}" |
|
|
| self.embeddings = HuggingFaceEmbeddings( |
| model_name=hf_model, |
| model_kwargs={"device": "cpu"}, |
| encode_kwargs={"normalize_embeddings": True}, |
| ) |
| else: |
| from langchain_ollama import OllamaEmbeddings |
|
|
| ollama_model = model_name or os.getenv("EMBED_MODEL", "nomic-embed-text") |
| self.embeddings = OllamaEmbeddings(model=ollama_model) |
|
|
| |
| def embed_documents(self, texts: list) -> list: |
| return self.embeddings.embed_documents(texts) |
|
|
| def embed_query(self, text: str) -> list: |
| return self.embeddings.embed_query(text) |
|
|
| |
| def embed_chunks(self, texts: list) -> list: |
| return self.embed_documents(texts) |
|
|