diff --git a/brain/app/api/deps.py b/brain/app/api/deps.py index 4506fdd..a60ae73 100644 --- a/brain/app/api/deps.py +++ b/brain/app/api/deps.py @@ -136,8 +136,10 @@ def get_notebook_rag_use_case( def get_notebook_chat_use_case( llm: Annotated[LLMProvider, Depends(get_llm_provider)], rag: Annotated[NotebookRagUseCase, Depends(get_notebook_rag_use_case)], + settings: Annotated[Settings, Depends(get_settings)], ) -> NotebookChatUseCase: - return NotebookChatUseCase(rag=rag, llm=llm) # type: ignore[arg-type] + return NotebookChatUseCase( + rag=rag, llm=llm, rerank_enabled=settings.rag_rerank) # type: ignore[arg-type] def get_notebook_deep_use_case( diff --git a/brain/app/application/notebook_chat.py b/brain/app/application/notebook_chat.py index 32701ab..2db8b51 100644 --- a/brain/app/application/notebook_chat.py +++ b/brain/app/application/notebook_chat.py @@ -10,6 +10,7 @@ from typing import AsyncIterator from app.application.notebook_rag import NotebookRagUseCase from app.application.query_rewrite import standalone_question +from app.application.rerank import pool_size, rerank from app.domain.models import ChatMessage from app.domain.ports import LLMChatProvider @@ -58,9 +59,13 @@ Réponds en français, de façon utile et concise. Mets le texte explicatif AVAN class NotebookChatUseCase: - def __init__(self, rag: NotebookRagUseCase, llm: LLMChatProvider) -> None: + def __init__( + self, rag: NotebookRagUseCase, llm: LLMChatProvider, rerank_enabled: bool = False + ) -> None: self._rag = rag self._llm = llm + # Reranking LLM d'un pool élargi avant injection (voir app.application.rerank). + self._rerank_enabled = rerank_enabled async def stream( self, @@ -76,7 +81,14 @@ class NotebookChatUseCase: # le sujet → on le résout depuis l'historique (best-effort, 1 appel léger, # uniquement à partir du 2e tour). La réponse, elle, voit tout l'historique. search_query = await standalone_question(self._llm, messages) - passages = await self._rag.retrieve(source_ids, search_query, top_k=top_k) + if self._rerank_enabled: + # Pool élargi → notation LLM → top_k final (meilleure précision sur + # les questions ambiguës, au prix d'un appel avant le premier token). + pool = await self._rag.retrieve( + source_ids, search_query, top_k=pool_size(top_k)) + passages = await rerank(self._llm, search_query, pool, top_k) + else: + passages = await self._rag.retrieve(source_ids, search_query, top_k=top_k) # Évènement 'sources' AVANT le premier token : l'UI peut afficher les # pages utilisées (« 📖 p. 12, 47 ») dès le début de la réponse. yield {"type": "sources", "sources": [ diff --git a/brain/app/application/rerank.py b/brain/app/application/rerank.py new file mode 100644 index 0000000..c7e6147 --- /dev/null +++ b/brain/app/application/rerank.py @@ -0,0 +1,73 @@ +"""Reranking LLM des passages RAG (chat des ateliers). + +Le cosinus classe par similarité de SURFACE ; sur les questions ambiguës, des +passages proches lexicalement mais inutiles passent devant l'extrait qui répond +vraiment. Le reranking récupère un POOL élargi (ex. 3× top_k), fait noter la +pertinence de chaque extrait par le LLM en UN appel, et garde les top_k mieux +notés. Coût : ~1 appel LLM avant le premier token — opt-in via RAG_RERANK. +""" +from __future__ import annotations + +import logging + +from app.application.llm_json import load_json_object + +logger = logging.getLogger(__name__) + +# Taille du pool élargi : multiple du top_k demandé, plafonné (le prompt de +# notation doit rester raisonnable même avec rag_top_k élevé). +POOL_FACTOR = 3 +POOL_MAX = 24 + +# Un extrait long n'a pas besoin d'être noté en entier : tronquer borne le +# prompt sans changer le jugement de pertinence. +_EXCERPT_CHARS = 600 + +_RERANK_PROMPT = """Tu évalues la PERTINENCE d'extraits d'un document pour répondre à une question. +Note chaque extrait de 0 (sans rapport) à 10 (répond directement), indépendamment des autres. + +QUESTION : {question} + +{passages} + +Réponds UNIQUEMENT par un objet JSON : {{"scores": [note_extrait_1, note_extrait_2, ...]}} +Le tableau doit contenir EXACTEMENT {count} notes, dans l'ordre des extraits.""" + + +def pool_size(top_k: int) -> int: + """Taille du pool à récupérer avant reranking.""" + return min(max(top_k * POOL_FACTOR, top_k), POOL_MAX) + + +async def rerank(llm, question: str, passages: list[dict], top_k: int) -> list[dict]: + """Renvoie les `top_k` passages les mieux notés par le LLM (tri stable : + à note égale, l'ordre cosinus d'origine est préservé). + + BEST-EFFORT : échec LLM, JSON invalide ou nombre de notes incohérent → + on renvoie simplement les `top_k` premiers du classement cosinus. + """ + if len(passages) <= top_k: + return passages + numbered = "\n\n".join( + f"--- EXTRAIT {i + 1} ---\n{(p.get('text') or '')[:_EXCERPT_CHARS]}" + for i, p in enumerate(passages) + ) + prompt = _RERANK_PROMPT.format( + question=question, passages=numbered, count=len(passages)) + try: + raw = await llm.generate(prompt, temperature=0.0) + except Exception as exc: # noqa: BLE001 — un chat dégradé vaut mieux que pas de chat + logger.warning("Reranking ignoré (échec LLM) : %s", exc) + return passages[:top_k] + parsed, _ = load_json_object(raw) + scores = parsed.get("scores") if isinstance(parsed, dict) else None + if not isinstance(scores, list) or len(scores) != len(passages): + logger.warning("Reranking ignoré (notes inexploitables).") + return passages[:top_k] + try: + scored = [(float(s), i) for i, s in enumerate(scores)] + except (TypeError, ValueError): + logger.warning("Reranking ignoré (notes non numériques).") + return passages[:top_k] + order = sorted(range(len(passages)), key=lambda i: (-scored[i][0], i)) + return [passages[i] for i in order[:top_k]] diff --git a/brain/app/core/config.py b/brain/app/core/config.py index 58a9db0..603d7ad 100644 --- a/brain/app/core/config.py +++ b/brain/app/core/config.py @@ -89,6 +89,13 @@ class Settings(BaseSettings): # False = relire TOUT le document à chaque question (exhaustivité maximale). deep_summary_filter: bool = True + # Reranking LLM du chat atelier : recupere un pool elargi (3x top_k, max 24) + # puis fait NOTER la pertinence de chaque extrait par le LLM avant d'injecter + # les top_k meilleurs. Meilleure precision sur les questions ambigues, MAIS + # +1 appel LLM avant le premier token (quelques secondes sur un petit modele + # local). Desactive par defaut ; recommande avec un provider cloud rapide. + rag_rerank: bool = False + # Cosinus minimal pour qu'un extrait soit injecté dans le prompt du chat # atelier : en dessous, l'extrait n'a aucun rapport avec la question → mieux # vaut moins d'extraits que du bruit. Défaut conservateur (0.30) : les paires