diff --git a/brain/app/api/deps.py b/brain/app/api/deps.py index 8d58b6c..4506fdd 100644 --- a/brain/app/api/deps.py +++ b/brain/app/api/deps.py @@ -142,10 +142,13 @@ def get_notebook_chat_use_case( def get_notebook_deep_use_case( llm: Annotated[LLMProvider, Depends(get_llm_provider)], + embedder: Annotated[object, Depends(get_embedding_provider)], settings: Annotated[Settings, Depends(get_settings)], ) -> NotebookDeepUseCase: return NotebookDeepUseCase( llm=llm, batch_tokens=settings.import_chunk_tokens, map_concurrency=settings.llm_map_concurrency, + embedder=embedder, + summary_filter=settings.deep_summary_filter, ) diff --git a/brain/app/application/notebook_deep.py b/brain/app/application/notebook_deep.py index ca8df1c..5a04f08 100644 --- a/brain/app/application/notebook_deep.py +++ b/brain/app/application/notebook_deep.py @@ -30,6 +30,28 @@ logger = logging.getLogger(__name__) _NO_MATCH = "RAS" _MAP_TEMPERATURE = 0.2 +# --- Index de résumés (pré-filtrage des lots) -------------------------------- +# Sans index : CHAQUE question relit TOUT le document (1 appel LLM par lot). +# Avec : les résumés de lots (construits UNE fois, cache disque) sont comparés +# à la question par embedding, et seuls les lots plausiblement pertinents sont +# relus. Sélection volontairement CONSERVATRICE (on préfère relire un lot de +# trop que rater une mention) ; désactivable via deep_summary_filter=False. +_SUMMARY_PROMPT = """Résume l'EXTRAIT ci-dessous en 4 à 8 puces factuelles : lieux, PNJ et +créatures nommés, objets notables, évènements, règles particulières. Pas d'analyse, pas +d'introduction — uniquement les puces, pour servir d'index de recherche. + +--- EXTRAIT --- +{excerpt} +--- FIN EXTRAIT --- + +Résumé :""" + +# Un lot est gardé si son score est proche du meilleur (marge) OU bon dans +# l'absolu ; et on garde toujours au moins _MIN_KEPT lots. +_SELECT_MARGIN = 0.10 +_SELECT_FLOOR = 0.5 +_MIN_KEPT = 3 + _MAP_PROMPT = """Voici un EXTRAIT d'un document. Extrais UNIQUEMENT les informations pertinentes pour répondre à la question ci-dessous. Conserve les détails utiles et indique les numéros de page (format « p. X »). Si l'extrait ne contient RIEN de @@ -65,12 +87,21 @@ Réponds en français.""" class NotebookDeepUseCase: def __init__( - self, llm: LLMProvider, batch_tokens: int = 10000, map_concurrency: int = 1 + self, + llm: LLMProvider, + batch_tokens: int = 10000, + map_concurrency: int = 1, + embedder=None, + summary_filter: bool = True, ) -> None: self._llm = llm self._batch_tokens = max(2000, batch_tokens) # Lots MAP traités par vagues de cette taille (parallélisme LLM). self._map_concurrency = max(1, map_concurrency) + # EmbeddingProvider (duck typing) pour l'index de résumés ; None = pas + # de pré-filtrage (plein scan, comportement historique). + self._embedder = embedder + self._summary_filter = summary_filter async def stream( self, @@ -90,24 +121,45 @@ class NotebookDeepUseCase: # une relance conversationnelle, il faut y résoudre les références # implicites, sinon les lots sont filtrés sur un texte sans sujet. question = await standalone_question(self._llm, messages) - chunks: list[dict] = [] + # Lots PAR SOURCE (l'index de résumés est caché par source). + per_source: list[tuple[str, list[dict]]] = [] for sid in source_ids: - chunks.extend(vector_store.all_chunks(sid)) - if not chunks: + chunks = vector_store.all_chunks(sid) + for batch in self._group(chunks): + per_source.append((sid, batch)) + if not per_source: yield {"type": "token", "token": "Aucune source indexée à analyser."} yield {"type": "done"} return - batches = self._group(chunks) - total = len(batches) + # Pré-filtrage par index de résumés (best-effort : tout échec → plein scan). + selected: set[int] | None = None + if self._summary_filter and self._embedder is not None: + try: + async for ev_or_result in self._select_batches(per_source, question): + if isinstance(ev_or_result, dict): + yield ev_or_result # progress de construction de l'index + else: + selected = ev_or_result + except Exception as exc: # noqa: BLE001 — le filtre ne doit jamais bloquer + logger.warning("Index de résumés ignoré (échec) : %s", exc) + selected = None + if selected is not None: + logger.info( + "Analyse approfondie : %s/%s lot(s) retenus via l'index de résumés.", + len(selected), len(per_source)) + + indices = sorted(selected) if selected is not None else list(range(len(per_source))) + total = len(indices) notes: list[str] = [] # Lots traités par VAGUES parallèles ; les notes restent dans l'ordre du # document (gather préserve l'ordre des tâches de la vague). for start in range(0, total, self._map_concurrency): yield {"type": "progress", "current": start, "total": total} - wave = batches[start:start + self._map_concurrency] + wave = indices[start:start + self._map_concurrency] results = await asyncio.gather( - *(self._map_batch(question, b) for b in wave), return_exceptions=True) + *(self._map_batch(question, per_source[i][1]) for i in wave), + return_exceptions=True) for j, res in enumerate(results): if isinstance(res, LLMProviderError): logger.warning( @@ -145,6 +197,63 @@ class NotebookDeepUseCase: )} yield {"type": "done"} + # --- Index de résumés ------------------------------------------------------ + + async def _select_batches(self, per_source: list[tuple[str, list[dict]]], question: str): + """Générateur : yield des évènements `progress` pendant la construction de + l'index (1ère analyse d'une source), puis le set des indices retenus — + ou None si le filtre n'apporte rien (tous retenus).""" + # 1. Charge/construit les résumés par source (cache disque). + by_sid: dict[str, list[int]] = {} + for i, (sid, _) in enumerate(per_source): + by_sid.setdefault(sid, []).append(i) + vectors: list[list[float] | None] = [None] * len(per_source) + + to_build = [] + for sid, idxs in by_sid.items(): + cached = vector_store.load_summaries(sid, self._batch_tokens) + if cached is not None and len(cached) == len(idxs): + for i, entry in zip(idxs, cached): + vectors[i] = entry.get("vector") + else: + to_build.append((sid, idxs)) + + total_build = sum(len(idxs) for _, idxs in to_build) + done_build = 0 + for sid, idxs in to_build: + summaries: list[str] = [] + for start in range(0, len(idxs), self._map_concurrency): + yield {"type": "progress", "current": done_build, "total": total_build} + wave = idxs[start:start + self._map_concurrency] + results = await asyncio.gather( + *(self._summarize_batch(per_source[i][1]) for i in wave)) + summaries.extend(results) + done_build += len(wave) + vecs = await self._embedder.embed(summaries, kind="document") + entries = [{"summary": s, "vector": v} for s, v in zip(summaries, vecs)] + vector_store.save_summaries(sid, self._batch_tokens, entries) + for i, entry in zip(idxs, entries): + vectors[i] = entry["vector"] + + # 2. Score de chaque lot face à la question, sélection conservatrice. + qv = (await self._embedder.embed([question], kind="query"))[0] + scores = [ + vector_store.cosine_similarity(qv, v) if v else 0.0 + for v in vectors + ] + best = max(scores) + keep = {i for i, s in enumerate(scores) if s >= best - _SELECT_MARGIN or s >= _SELECT_FLOOR} + floor = min(_MIN_KEPT, len(scores)) + if len(keep) < floor: + keep = set(sorted(range(len(scores)), key=lambda i: -scores[i])[:floor]) + yield keep if len(keep) < len(scores) else None + + async def _summarize_batch(self, batch: list[dict]) -> str: + excerpt = "\n\n".join(c.get("text", "").strip() for c in batch) + raw = await generate_with_retry( + self._llm, _SUMMARY_PROMPT.format(excerpt=excerpt), temperature=_MAP_TEMPERATURE) + return (raw or "").strip() + async def _map_batch(self, question: str, batch: list[dict]) -> str: """Phase MAP d'un lot : extrait les infos pertinentes ('' si RAS).""" excerpt = "\n\n".join( diff --git a/brain/app/core/config.py b/brain/app/core/config.py index aa11395..58a9db0 100644 --- a/brain/app/core/config.py +++ b/brain/app/core/config.py @@ -83,6 +83,12 @@ class Settings(BaseSettings): # mais prompt plus long. 8 par défaut (montable jusqu'à ~20 sur grand contexte). rag_top_k: int = 8 + # Analyse approfondie : pré-filtrage des lots via un index de résumés + # (construit une fois par source, cache disque). Les questions ciblées ne + # relisent que les lots plausiblement pertinents (3-5x moins d'appels) ; + # False = relire TOUT le document à chaque question (exhaustivité maximale). + deep_summary_filter: bool = True + # 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 diff --git a/brain/app/infrastructure/vector_store.py b/brain/app/infrastructure/vector_store.py index 380b714..b9a13f7 100644 --- a/brain/app/infrastructure/vector_store.py +++ b/brain/app/infrastructure/vector_store.py @@ -80,6 +80,42 @@ def exists(source_id: str) -> bool: def delete(source_id: str) -> None: _CACHE.pop(source_id, None) _path(source_id).unlink(missing_ok=True) + _summaries_path(source_id).unlink(missing_ok=True) + + +# --- Index de résumés (analyse approfondie) ---------------------------------- +# Cache disque des résumés PAR LOT d'une source : construit paresseusement à la +# première analyse approfondie, réutilisé ensuite pour ne relire que les lots +# pertinents. Invalidé avec la source (delete) et si batch_tokens change. + + +def _summaries_path(source_id: str) -> Path: + safe = _SAFE_ID.sub("_", str(source_id)) + return _STORE_DIR / f"{safe}.summaries.json" + + +def save_summaries(source_id: str, batch_tokens: int, entries: list[dict]) -> None: + """Persiste les résumés de lots ({"summary": str, "vector": [...]}).""" + _STORE_DIR.mkdir(parents=True, exist_ok=True) + payload = {"batch_tokens": int(batch_tokens), "entries": entries} + _summaries_path(source_id).write_text( + json.dumps(payload, ensure_ascii=False), encoding="utf-8") + + +def load_summaries(source_id: str, batch_tokens: int) -> list[dict] | None: + """Résumés de lots d'une source, ou None si absents / construits avec une + autre taille de lot (le découpage ne correspondrait plus).""" + p = _summaries_path(source_id) + if not p.exists(): + return None + try: + data = json.loads(p.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + return None + if not isinstance(data, dict) or data.get("batch_tokens") != int(batch_tokens): + return None + entries = data.get("entries") + return entries if isinstance(entries, list) else None def _load(source_id: str) -> list[dict]: @@ -139,6 +175,11 @@ def _chunk_words(chunk: dict) -> frozenset[str]: return words +# Alias public du cosinus (réutilisé par l'index de résumés de l'analyse +# approfondie — même métrique que la recherche). +cosine_similarity = _cosine + + def search( source_ids: list[str], query_vector: list[float],