choupi-faq-api / embeddings.py
Loren's picture
Add metadata even when similarity < threshold
bd3d656
Raw
History Blame Contribute Delete
7.04 kB
"""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()