Aller au contenu principal

Corrective RAG

Quand le retrieval ne suffit pas

Dans un RAG classique, vous récupérez les chunks les plus similaires et les injectez dans le prompt. Mais que se passe-t-il si les documents retrouvés ne sont pas pertinents pour la question ? Le modèle génère une réponse basée sur du contexte inadapté, ce qui produit des hallucinations.

Le Corrective RAG (CRAG) ajoute une étape d’évaluation après le retrieval : un LLM vérifie si chaque document retrouvé est réellement pertinent. Si les documents sont jugés non pertinents, le système effectue un fallback vers une recherche web pour trouver de meilleures informations.

Architecture du Corrective RAG

Le flux se décompose en quatre étapes :

  1. Retrieval : rechercher les documents les plus similaires dans le vector store
  2. Grading : évaluer la pertinence de chaque document par rapport à la question
  3. Décision : si au moins un document est non pertinent, lancer une recherche web
  4. Génération : produire la réponse avec les meilleurs documents disponibles

Le grader de documents

Le coeur du Corrective RAG est le grader : un LLM qui évalue chaque document retrouvé.

from mistralai import Mistral
import json

client = Mistral(api_key=api_key)

def grade_document(question, document):
    """Évaluer si un document est pertinent pour une question."""
    prompt = f"""Vous etes un evaluateur de pertinence. Etant donne une question
et un document, determinez si le document contient des informations pertinentes
pour repondre a la question.

Repondez UNIQUEMENT avec un objet JSON : {{"score": "yes"}} ou {{"score": "no"}}

Question : {question}
Document : {document}

Evaluation JSON :"""

    response = client.chat.complete(
        model="mistral-large-latest",
        messages=[{"role": "user", "content": prompt}],
        temperature=0,
        response_format={"type": "json_object"},
    )

    result = json.loads(response.choices[0].message.content)
    return result.get("score", "no") == "yes"

# Test
question = "Quelles sont les étapes du RAG ?"
doc_pertinent = "Le RAG se décompose en retrieval, augmentation et génération."
doc_non_pertinent = "La recette du tiramisu nécessite du mascarpone."

print(grade_document(question, doc_pertinent))      # True
print(grade_document(question, doc_non_pertinent))   # False

Évaluer tous les documents retrouvés

def grade_documents(question, retrieved_chunks):
    """Évaluer et filtrer les documents retrouvés."""
    relevant_docs = []
    irrelevant_count = 0

    for chunk in retrieved_chunks:
        is_relevant = grade_document(question, chunk["text"])
        if is_relevant:
            relevant_docs.append(chunk)
        else:
            irrelevant_count += 1

    need_web_search = irrelevant_count > 0

    print(f"  Pertinents: {len(relevant_docs)}/{len(retrieved_chunks)}")
    print(f"  Recherche web nécessaire: {need_web_search}")

    return relevant_docs, need_web_search

Fallback vers la recherche web

Quand les documents locaux ne suffisent pas, le système effectue une recherche web :

def web_search_fallback(question, n_results=3):
    """Recherche web de secours quand les documents locaux sont insuffisants."""
    # En production, utilisez une API de recherche (Tavily, Serper, etc.)
    # Ici, un exemple simplifié avec l'API Mistral agents

    prompt = f"""Recherchez sur le web des informations pour répondre à cette question :
{question}

Fournissez une réponse factuelle basée sur les résultats de recherche."""

    response = client.chat.complete(
        model="mistral-large-latest",
        messages=[{"role": "user", "content": prompt}],
        temperature=0.1,
    )

    # Retourner comme un "document" web
    return [{
        "text": response.choices[0].message.content,
        "source": "web_search",
    }]

Pipeline Corrective RAG complet

class CorrectiveRAG:
    """Pipeline Corrective RAG avec évaluation et fallback web."""

    def __init__(self, chunks, embeddings, index):
        self.chunks = chunks
        self.embeddings = embeddings
        self.index = index

    def retrieve(self, question, k=4):
        """Étape 1 : retrieval standard."""
        response = client.embeddings.create(
            model="mistral-embed",
            inputs=[question],
        )
        query_vec = np.array([response.data[0].embedding], dtype=np.float32)
        distances, indices = self.index.search(query_vec, k)

        results = []
        for dist, idx in zip(distances[0], indices[0]):
            if idx >= 0:
                results.append({
                    "text": self.chunks[idx]["text"],
                    "source": self.chunks[idx].get("source", "local"),
                })
        return results

    def grade(self, question, documents):
        """Étape 2 : évaluer la pertinence."""
        relevant = []
        for doc in documents:
            if grade_document(question, doc["text"]):
                relevant.append(doc)
        need_web = len(relevant) < len(documents)
        return relevant, need_web

    def generate(self, question, documents):
        """Étape 3 : générer la réponse."""
        context = "\n\n---\n\n".join([d["text"] for d in documents])
        prompt = f"""Contexte :\n{context}\n\nQuestion : {question}"""

        response = client.chat.complete(
            model="mistral-large-latest",
            messages=[
                {"role": "system", "content": "Repondez en vous basant sur le contexte."},
                {"role": "user", "content": prompt},
            ],
            temperature=0.1,
        )
        return response.choices[0].message.content

    def ask(self, question, k=4):
        """Pipeline complet."""
        print(f"\n=== Question : {question} ===")

        # 1. Retrieval
        print("1. Retrieval...")
        docs = self.retrieve(question, k=k)

        # 2. Grading
        print("2. Grading...")
        relevant_docs, need_web = self.grade(question, docs)

        # 3. Web search fallback si nécessaire
        if need_web:
            print("3. Fallback web search...")
            web_docs = web_search_fallback(question)
            relevant_docs.extend(web_docs)

        # 4. Génération
        print("4. Génération...")
        answer = self.generate(question, relevant_docs)

        return {
            "answer": answer,
            "sources_used": len(relevant_docs),
            "web_search_used": need_web,
        }

# Utilisation
crag = CorrectiveRAG(chunks, embeddings, index)
result = crag.ask("Comment Mistral implémente-t-il les embeddings ?")
print(f"\nRéponse : {result['answer']}")
print(f"Web search utilisé : {result['web_search_used']}")

Points clés à retenir

  • Le Corrective RAG ajoute une évaluation de pertinence après le retrieval
  • Un LLM sert de “grader” pour juger chaque document : pertinent ou non
  • Si des documents sont jugés non pertinents, un fallback vers la recherche web se déclenche
  • Cette approche réduit significativement les hallucinations dues à un contexte inadapté
  • Le grading ajoute de la latence mais améliore fortement la fiabilité des réponses