Spaces:
Sleeping
Sleeping
| """Gestion des embeddings avec ChromaDB et Hugging Face Transformers.""" | |
| import os | |
| # Règle d’or : toute variable d’environnement qui influence le cache Hugging Face doit être | |
| # définie avant d’importer datasets ou transformers, sinon elle sera ignorée. | |
| cache_dir = "/tmp" | |
| os.makedirs(cache_dir, exist_ok=True) | |
| # Rediriger le cache HF globalement | |
| os.environ["HF_HOME"] = cache_dir | |
| os.environ["HF_DATASETS_CACHE"] = os.path.join(cache_dir, "datasets") | |
| os.environ["TRANSFORMERS_CACHE"] = os.path.join(cache_dir, "transformers") | |
| import logging | |
| from typing import Optional | |
| from models import QuestionInput, AnswerOutput | |
| import chromadb | |
| from sentence_transformers import SentenceTransformer | |
| import torch | |
| from config import ( | |
| CHROMA_COLLECTION_NAME, | |
| MODEL_NAME, | |
| SIMILARITY_THRESHOLD, | |
| ) | |
| from faq_loader import FAQEntry | |
| logger = logging.getLogger(__name__) | |
| class EmbeddingManager: | |
| """Gestionnaire des embeddings et de ChromaDB.""" | |
| def __init__( | |
| self, | |
| model_name: str = MODEL_NAME, | |
| collection_name: str = CHROMA_COLLECTION_NAME, | |
| ) -> None: | |
| """ | |
| Initialiser le gestionnaire d'embeddings. | |
| Args: | |
| model_name: Nom du modèle Sentence-Transformers à utiliser | |
| collection_name: Nom de la collection ChromaDB | |
| """ | |
| # Creation du Sentence transformer model | |
| logger.info(f"Initialisation du modèle d'embeddings: {model_name}") | |
| device = "cuda" if torch.cuda.is_available() else "cpu" | |
| self.model = SentenceTransformer(MODEL_NAME, device=device) | |
| # Initialiser ChromaDB en mémoire (EphemeralClient) | |
| # Idéal pour HF Spaces: pas de persistance disque, plus rapide au startup | |
| logger.info("Initialisation ChromaDB en mémoire (EphemeralClient)") | |
| self.client: chromadb.Client = chromadb.EphemeralClient() | |
| # Créer ou récupérer la collection | |
| self.collection = self.client.get_or_create_collection( | |
| name=collection_name, | |
| embedding_function=None, | |
| metadata={"hnsw:space": "ip"} # Utiliser produit scalaire => normaliser les embeddings | |
| ) | |
| self.collection_name: str = collection_name | |
| logger.info(f"✓ EmbeddingManager initialisé (collection: {collection_name})") | |
| def populate_collection(self, faq_entries: list[FAQEntry]) -> None: | |
| """ | |
| Remplir la collection ChromaDB avec les embeddings des FAQ. | |
| Args: | |
| faq_entries: Liste d'objets FAQEntry à indexer | |
| """ | |
| logger.info(f"Population de ChromaDB avec {len(faq_entries)} entrées...") | |
| # Vider la collection existante (optionnel, mais plus propre) | |
| # Récupérer l'ID de collection pour supprimer les anciens documents | |
| existing_ids = self.collection.get()["ids"] | |
| if existing_ids: | |
| logger.info(f"Suppression de {len(existing_ids)} anciens documents...") | |
| self.collection.delete(ids=existing_ids) | |
| # Générer les embeddings pour toutes les formulations | |
| formulations: list[str] = ["passage: " + entry.formulation for entry in faq_entries] | |
| logger.info("Génération des embeddings...") | |
| embeddings = self.model.encode(formulations, show_progress_bar=True, | |
| convert_to_numpy=True, normalize_embeddings=True) | |
| # Ajouter à ChromaDB | |
| logger.info("Ajout des documents à ChromaDB...") | |
| self.collection.add( | |
| ids=[f"faq_{i}" for i in range(len(faq_entries))], | |
| embeddings=embeddings.tolist(), | |
| documents=formulations, | |
| metadatas=[ | |
| { | |
| "theme": entry.theme, | |
| "response": entry.response, | |
| } | |
| for entry in faq_entries | |
| ], | |
| ) | |
| logger.info(f"✓ ChromaDB peuplée avec {len(faq_entries)} entrées") | |
| def search_similar_faq(self, payload: QuestionInput) -> AnswerOutput: | |
| """ | |
| Rechercher dans la FAQ la formulation laplus similaire à une question donnée. | |
| Args: | |
| question: Question à chercher | |
| threshold: Seuil de confiance minimum pour retourner une réponse | |
| Returns: | |
| Tuple (formulation, theme, response, similarity_score) ou None si aucun match | |
| """ | |
| # Récupérer la question et le seuil depuis le payload | |
| question = payload.question | |
| threshold = payload.threshold if payload.threshold is not None else SIMILARITY_THRESHOLD | |
| # Générer l'embedding de la question | |
| question_embedding = self.model.encode(["query: " + question], convert_to_numpy=True, | |
| normalize_embeddings=True).tolist() | |
| # Rechercher dans ChromaDB | |
| results = self.collection.query( | |
| query_embeddings=question_embedding, | |
| n_results=1, | |
| include=["documents", "metadatas", "distances"] | |
| ) | |
| if not results["ids"] or not results["ids"][0]: | |
| logger.warning(f"Aucun résultat trouvé pour: {question}") | |
| return AnswerOutput( | |
| question=question, | |
| formulation="unknown", | |
| answer="Je ne dispose pas d'une réponse suffisamment pertinente pour cette question.", | |
| theme="unknown", | |
| similarity_score=0.0, | |
| confidence=False, | |
| ) | |
| # Extraire le meilleur résultat | |
| # ChromaDB retourne les produits scalaires, on convertit en similarité cosinus | |
| distance = results["distances"][0][0] | |
| similarity_score: float = 1 - distance | |
| # Extraire les métadonnées | |
| metadata = results["metadatas"][0][0] | |
| formulation = results["documents"][0][0] | |
| theme = metadata["theme"] | |
| response = metadata["response"] | |
| # Vérifier le seuil | |
| if similarity_score < threshold: | |
| logger.info( | |
| f"Score de similarité {similarity_score:.3f} < seuil {threshold}. " | |
| f"Question: {question}" | |
| ) | |
| return AnswerOutput( | |
| question=question, | |
| formulation=formulation, | |
| answer="Je ne dispose pas d'une réponse suffisamment pertinente pour cette question.", | |
| theme=theme, | |
| similarity_score=round(similarity_score, 4), | |
| confidence=False, | |
| ) | |
| logger.info( | |
| f"Match trouvé - Theme: {theme}, Score: {similarity_score:.3f}" | |
| ) | |
| return AnswerOutput( | |
| question=question, | |
| formulation=formulation, | |
| answer=response, | |
| theme=theme, | |
| similarity_score=round(similarity_score, 4), | |
| confidence=True, | |
| ) | |
| def get_collection_size(self) -> int: | |
| """ | |
| Obtenir le nombre de documents dans la collection. | |
| Returns: | |
| Nombre de documents | |
| """ | |
| return self.collection.count() | |