"""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()