BELASRI commited on
Commit
7532f48
·
verified ·
1 Parent(s): 72d9c23

Déploiement initial : src/rag.py

Browse files
Files changed (1) hide show
  1. src/rag.py +387 -0
src/rag.py ADDED
@@ -0,0 +1,387 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import os
3
+ import re
4
+
5
+ # Cache HF explicite, à la racine du projet (portable local/HF Spaces — le
6
+ # disque du Space est accessible en écriture par défaut sur le tier gratuit,
7
+ # pas besoin de stockage persistant payant pour ça). Ne PAS forcer
8
+ # HF_HUB_OFFLINE=1 : ça bloquerait le tout premier déploiement (cache vide,
9
+ # mode offline = téléchargement impossible = le Space ne démarre jamais).
10
+ os.environ.setdefault("HF_HOME", os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), ".hf_cache"))
11
+
12
+ import faiss
13
+ import numpy as np
14
+ import torch
15
+ from dotenv import load_dotenv
16
+ from openai import OpenAI
17
+ from rank_bm25 import BM25Okapi
18
+ from sentence_transformers import SentenceTransformer
19
+
20
+ load_dotenv()
21
+
22
+ INDEX_PATH = "faiss_solva2/index.faiss"
23
+ MAPPING_PATH = "faiss_solva2/mapping.json"
24
+ EMBEDDING_MODEL_NAME = "Shitao/bge-m3" # miroir safetensors du même modèle que BAAI/bge-m3
25
+
26
+ # Fournisseurs LLM disponibles, tous compatibles OpenAI SDK (base_url + clé +
27
+ # nom de modèle changent, le reste du code est identique). Ne pas supprimer
28
+ # une entrée quand on change de défaut : ça permet de rebasculer en une ligne
29
+ # (provider="mistral"|"gemini"|...).
30
+ PROVIDERS = {
31
+ "gemini": {
32
+ "base_url": "https://generativelanguage.googleapis.com/v1beta/openai/",
33
+ "api_key_env": "GEMINI_API_KEY",
34
+ "model": "gemini-flash-latest", # "gemini-2.5-flash" renvoie 404 (fermé aux nouveaux comptes)
35
+ },
36
+ "mistral": {
37
+ "base_url": "https://api.mistral.ai/v1",
38
+ "api_key_env": "MISTRAL_API_KEY",
39
+ "model": "mistral-large-latest",
40
+ },
41
+ "groq": {
42
+ "base_url": "https://api.groq.com/openai/v1",
43
+ "api_key_env": "GROQ_API_KEY",
44
+ "model": "llama-3.3-70b-versatile",
45
+ },
46
+ "cerebras": {
47
+ "base_url": "https://api.cerebras.ai/v1",
48
+ "api_key_env": "CEREBRAS_API_KEY",
49
+ "model": "gpt-oss-120b", # llama-3.3-70b indisponible sur ce compte (404)
50
+ },
51
+ }
52
+
53
+ # Gemini retenu (Décision 008) : 1500 requêtes/jour, largement suffisant pour
54
+ # le pipeline conversationnel (3 appels LLM/question : reformulation +
55
+ # clarification + génération) qui saturait le tier gratuit Mistral.
56
+ DEFAULT_PROVIDER = "gemini"
57
+
58
+ # Décision 009 (Gemini->Mistral) puis Décision 012 (ajout de Groq en 3e,
59
+ # nouvelle clé validée) : chaîne de fallback automatique. Cerebras disponible
60
+ # dans PROVIDERS mais pas dans cette chaîne (pas testé en usage courant, pas
61
+ # de raison de l'y ajouter pour l'instant).
62
+ FALLBACK_ORDER = ["gemini", "mistral", "groq"]
63
+
64
+ _llm_clients = {} # un client OpenAI par provider, créé à la demande
65
+
66
+
67
+ def _get_llm_client(provider):
68
+ if provider not in _llm_clients:
69
+ cfg = PROVIDERS[provider]
70
+ _llm_clients[provider] = OpenAI(
71
+ api_key=os.environ[cfg["api_key_env"]],
72
+ base_url=cfg["base_url"],
73
+ timeout=30, # évite un blocage indéfini si l'API ne répond pas (défaut SDK ~10 min)
74
+ )
75
+ return _llm_clients[provider]
76
+
77
+
78
+ # Alias de compatibilité (utilisés par app.py et d'anciens scripts) : pointent
79
+ # vers le provider par défaut.
80
+ LLM_MODEL = PROVIDERS[DEFAULT_PROVIDER]["model"]
81
+ _llm_client = _get_llm_client(DEFAULT_PROVIDER)
82
+
83
+ # Gemini consomme une partie du budget max_tokens en raisonnement interne caché,
84
+ # même sur des prompts de classification triviaux (ex. "OUI"/"NON" a consommé
85
+ # ~150 tokens de réflexion avant le token de réponse visible). Pas de moyen
86
+ # documenté de désactiver ça via l'endpoint compatible OpenAI (reasoning_effort
87
+ # et extra_body google.thinking_config testés, tous deux rejetés en 400) — on
88
+ # compense en donnant assez de marge. Mistral et Groq n'ont pas ce
89
+ # comportement, donc pas besoin d'autant de marge de ce côté.
90
+ MAX_TOKENS_COURT = {"gemini": 2000, "mistral": 200, "groq": 200} # classification (juge OUI/NON, CLAIRE/AMBIGUE)
91
+ MAX_TOKENS_LONG = {"gemini": 4096, "mistral": 2048, "groq": 2048} # réponses substantielles (génération, reformulation)
92
+
93
+
94
+ # Codes déclenchant un fallback : quota/débit (429), facturation (402),
95
+ # authentification/clé invalide (401, 403 — et 400 : constaté empiriquement
96
+ # que l'endpoint Gemini répond 400 "Please pass a valid API key" pour une clé
97
+ # invalide, pas 401 comme la plupart des API OpenAI-compatibles).
98
+ CODES_FALLBACK = {400, 401, 402, 403, 429}
99
+
100
+
101
+ def appel_llm(prompt, taille="long", verbose=True, providers=None):
102
+ """Appel LLM unifié avec fallback automatique (Décision 009).
103
+
104
+ Essaie les providers de `providers` (par défaut FALLBACK_ORDER) dans
105
+ l'ordre. Sur une erreur de quota/débit/authentification (CODES_FALLBACK),
106
+ bascule silencieusement (visible seulement si verbose) sur le suivant. Si
107
+ tous échouent, ou sur toute autre erreur, l'exception remonte telle
108
+ quelle — pas de 3e fallback, pas de boucle.
109
+ """
110
+ providers = providers or FALLBACK_ORDER
111
+ max_tokens_par_taille = MAX_TOKENS_COURT if taille == "court" else MAX_TOKENS_LONG
112
+
113
+ derniere_erreur = None
114
+ for i, provider in enumerate(providers):
115
+ try:
116
+ client = _get_llm_client(provider)
117
+ model_name = PROVIDERS[provider]["model"]
118
+ completion = client.chat.completions.create(
119
+ model=model_name,
120
+ messages=[{"role": "user", "content": prompt}],
121
+ temperature=0,
122
+ max_tokens=max_tokens_par_taille[provider],
123
+ )
124
+ if i > 0 and verbose:
125
+ print(f"[fallback] {providers[i - 1]} indisponible -> bascule sur {provider}")
126
+ return completion.choices[0].message.content
127
+ except Exception as e:
128
+ derniere_erreur = e
129
+ status_code = getattr(e, "status_code", None)
130
+ declenche_fallback = status_code in CODES_FALLBACK
131
+ if verbose:
132
+ print(f"[appel_llm] échec sur {provider} (status={status_code}) : {str(e)[:120]}")
133
+ if not declenche_fallback:
134
+ raise
135
+ continue
136
+
137
+ raise RuntimeError(
138
+ f"Tous les fournisseurs LLM ont échoué ({', '.join(providers)}). "
139
+ f"Dernière erreur : {derniere_erreur}"
140
+ )
141
+
142
+
143
+ # Dictionnaire fixe (pas de LLM) : sigles courants de Solvabilité II.
144
+ # "SCR" apparaît tel quel dans le corpus (38 occurrences, ex. "capital de
145
+ # solvabilité requis (SCR)"), mais les autres sigles n'y figurent JAMAIS —
146
+ # le texte officiel les épelle toujours en toutes lettres. D'où l'intérêt
147
+ # de l'expansion : rapprocher la question posée (en jargon métier) du
148
+ # vocabulaire réellement utilisé dans le texte.
149
+ ACRONYMES = {
150
+ "SCR": "capital de solvabilité requis",
151
+ "MCR": "minimum de capital requis",
152
+ "ORSA": "évaluation interne des risques et de la solvabilité",
153
+ "BE": "meilleure estimation",
154
+ "PT": "provisions techniques",
155
+ "FP": "fonds propres",
156
+ "VaR": "valeur à risque",
157
+ "SFCR": "rapport sur la solvabilité et la situation financière",
158
+ "EIOPA": "Autorité européenne des assurances et des pensions professionnelles",
159
+ "AEAPP": "Autorité européenne des assurances et des pensions professionnelles",
160
+ "ACPR": "Autorité de contrôle prudentiel et de résolution",
161
+ }
162
+
163
+
164
+ def expand_query(question):
165
+ expanded = question
166
+ for sigle, forme_longue in ACRONYMES.items():
167
+ pattern = re.compile(rf"\b{re.escape(sigle)}\b", re.IGNORECASE)
168
+ if pattern.search(expanded):
169
+ expanded = pattern.sub(f"{sigle} ({forme_longue})", expanded, count=1)
170
+ return expanded
171
+
172
+
173
+ PROMPT_TEMPLATE = """Tu es un assistant spécialisé sur la directive Solvabilité II
174
+ (2009/138/CE). Tu réponds UNIQUEMENT à partir des articles fournis
175
+ ci-dessous.
176
+ RÈGLES ABSOLUES :
177
+ 1. Réponds uniquement d'après les articles fournis.
178
+ 2. Cite TOUJOURS l'article source au format [Article N].
179
+ 3. Si la question contient un chiffre, un pourcentage, un seuil ou une
180
+ borne inexact(e), et que les articles fournis donnent la valeur
181
+ correcte, NE REFUSE PAS : signale explicitement l'écart avec le
182
+ chiffre mentionné dans la question, indique la valeur correcte, et
183
+ cite l'article source. Ceci prime sur la règle 4. Si tu corriges une
184
+ valeur, ne commence JAMAIS ta réponse par la phrase de refus.
185
+ Commence directement par la correction, ex : "Le SCR est calibré à
186
+ 99,5 %, et non 99,9 % [Article X]."
187
+ 4. Si l'information n'est pas dans les articles fournis, réponds
188
+ exactement : "Je ne trouve pas cette information dans les
189
+ articles fournis."
190
+ 5. N'invente jamais. Ne complète jamais de mémoire.
191
+
192
+ ARTICLES FOURNIS :
193
+ {contexte}
194
+
195
+ QUESTION : {question}
196
+ RÉPONSE :"""
197
+
198
+ # --- Chargement une fois au chargement du module ---
199
+ _index = faiss.read_index(INDEX_PATH)
200
+ with open(MAPPING_PATH, encoding="utf-8") as _f:
201
+ _mapping = json.load(_f)
202
+
203
+ _embedding_model = None # chargement paresseux
204
+
205
+
206
+ def _get_embedding_model():
207
+ global _embedding_model
208
+ if _embedding_model is None:
209
+ device = "cuda" if torch.cuda.is_available() else "cpu"
210
+ _embedding_model = SentenceTransformer(EMBEDDING_MODEL_NAME, device=device)
211
+ _embedding_model.max_seq_length = 1024
212
+ return _embedding_model
213
+
214
+
215
+ def _rechercher(question, k):
216
+ model = _get_embedding_model()
217
+ q_emb = model.encode([question], normalize_embeddings=True, convert_to_numpy=True).astype(np.float32)
218
+ _, idxs = _index.search(q_emb, k)
219
+ return [_mapping[i] for i in idxs[0]]
220
+
221
+
222
+ def _rechercher_dense_idx(question, k):
223
+ model = _get_embedding_model()
224
+ q_emb = model.encode([question], normalize_embeddings=True, convert_to_numpy=True).astype(np.float32)
225
+ _, idxs = _index.search(q_emb, k)
226
+ return list(idxs[0])
227
+
228
+
229
+ # --- Retrieval hybride EN TEST (mode="hybrid" dans poser_question) : BM25 en
230
+ # complément du dense BGE-M3, ne remplace pas le retrieval actuel tant que non
231
+ # validé sur le golden set. ---
232
+ def _tokeniser_bm25(texte):
233
+ return re.findall(r"\w+", texte.lower())
234
+
235
+
236
+ _bm25_corpus_tokens = [_tokeniser_bm25(f"{a['titre']} {a['texte']}") for a in _mapping]
237
+ _bm25 = BM25Okapi(_bm25_corpus_tokens)
238
+
239
+
240
+ def _rechercher_bm25_idx(question, k):
241
+ scores = _bm25.get_scores(_tokeniser_bm25(question))
242
+ return list(np.argsort(scores)[::-1][:k])
243
+
244
+
245
+ def retrieval_hybride(question, k=5, k_fusion=20, poids_dense=0.7, poids_bm25=0.3, rrf_k=60):
246
+ """Fusion par Reciprocal Rank Fusion (RRF) du dense (BGE-M3/FAISS) et de
247
+ BM25, pondérée 70% dense / 30% BM25 par défaut."""
248
+ dense_idx = _rechercher_dense_idx(question, k_fusion)
249
+ bm25_idx = _rechercher_bm25_idx(question, k_fusion)
250
+
251
+ scores_fusion = {}
252
+ for rang, idx in enumerate(dense_idx, start=1):
253
+ idx = int(idx)
254
+ scores_fusion[idx] = scores_fusion.get(idx, 0.0) + poids_dense * (1 / (rrf_k + rang))
255
+ for rang, idx in enumerate(bm25_idx, start=1):
256
+ idx = int(idx)
257
+ scores_fusion[idx] = scores_fusion.get(idx, 0.0) + poids_bm25 * (1 / (rrf_k + rang))
258
+
259
+ tries = sorted(scores_fusion.items(), key=lambda x: x[1], reverse=True)[:k]
260
+ return [_mapping[idx] for idx, _ in tries]
261
+
262
+
263
+ def _construire_contexte(articles):
264
+ blocs = []
265
+ for a in articles:
266
+ blocs.append(f"[Article {a['numero_article']} — {a['titre']}]\n{a['texte']}\n")
267
+ return "\n".join(blocs)
268
+
269
+
270
+ def _est_un_refus(reponse):
271
+ return reponse.strip().startswith("Je ne trouve pas") and "[Article" not in reponse
272
+
273
+
274
+ JUGE_PROMPT_TEMPLATE = """Voici une question et des articles de loi récupérés pour y répondre.
275
+
276
+ QUESTION : {question}
277
+
278
+ ARTICLES RÉCUPÉRÉS :
279
+ {contexte}
280
+
281
+ Ces articles contiennent-ils de quoi répondre à la question ?
282
+ Réponds STRICTEMENT sous l'une de ces deux formes, rien d'autre :
283
+ OUI
284
+ NON : <reformulation courte de la requête de recherche>"""
285
+
286
+
287
+ def _juger_retrieval(question, articles, verbose=True):
288
+ contexte = _construire_contexte(articles)
289
+ prompt = JUGE_PROMPT_TEMPLATE.format(question=question, contexte=contexte)
290
+
291
+ reponse = appel_llm(prompt, taille="court", verbose=verbose).strip()
292
+
293
+ if reponse.upper().startswith("OUI"):
294
+ return True, None
295
+ if reponse.upper().startswith("NON"):
296
+ partie = reponse.split(":", 1)
297
+ reformulation = partie[1].strip() if len(partie) > 1 else None
298
+ return False, reformulation
299
+ # réponse hors format attendu : par prudence, on considère le retrieval suffisant
300
+ # (évite de déclencher une reformulation sur une base non fiable)
301
+ return True, None
302
+
303
+
304
+ CLARIFICATION_JUGE_PROMPT = """Cette question sur Solvabilité II est-elle assez précise pour y répondre, ou trop vague/ambiguë ?
305
+ Réponds uniquement par un seul mot : CLAIRE ou AMBIGUE.
306
+
307
+ QUESTION : {question}"""
308
+
309
+ CLARIFICATION_REFORMULATION_PROMPT = """Cette question sur Solvabilité II est trop vague ou ambiguë pour y répondre directement :
310
+
311
+ QUESTION : {question}
312
+
313
+ Propose UNE seule question de clarification précise, en français, pour aider
314
+ l'utilisateur à préciser sa demande. Réponds uniquement avec cette question
315
+ de clarification, rien d'autre."""
316
+
317
+
318
+ def clarifier(question, verbose=True):
319
+ decision = appel_llm(
320
+ CLARIFICATION_JUGE_PROMPT.format(question=question), taille="court", verbose=verbose
321
+ ).strip().upper()
322
+
323
+ if "AMBIG" not in decision:
324
+ return {"ambigue": False, "message": None}
325
+
326
+ message = appel_llm(
327
+ CLARIFICATION_REFORMULATION_PROMPT.format(question=question), taille="long", verbose=verbose
328
+ ).strip()
329
+ return {"ambigue": True, "message": message}
330
+
331
+
332
+ def poser_question(question, k=5, verbose=True, self_eval=False, clarify=False, mode="dense"):
333
+ if clarify:
334
+ resultat_clarif = clarifier(question, verbose=verbose)
335
+ if resultat_clarif["ambigue"]:
336
+ if verbose:
337
+ print(f"QUESTION : {question}")
338
+ print(f"[clarification] question jugée AMBIGUË — pas de retrieval, pas de génération")
339
+ print(f"CLARIFICATION DEMANDÉE : {resultat_clarif['message']}")
340
+ print()
341
+ return {"type": "clarification", "message": resultat_clarif["message"]}
342
+
343
+ requete_recherche = expand_query(question)
344
+ articles = retrieval_hybride(requete_recherche, k) if mode == "hybrid" else _rechercher(requete_recherche, k)
345
+
346
+ if self_eval:
347
+ suffisant, reformulation = _juger_retrieval(question, articles, verbose=verbose)
348
+ if not suffisant and reformulation:
349
+ if verbose:
350
+ print(f"[self-eval] retrieval jugé insuffisant, reformulation : \"{reformulation}\"")
351
+ reformulation_expandue = expand_query(reformulation)
352
+ articles = (
353
+ retrieval_hybride(reformulation_expandue, k)
354
+ if mode == "hybrid"
355
+ else _rechercher(reformulation_expandue, k)
356
+ )
357
+ elif not suffisant and verbose:
358
+ print("[self-eval] retrieval jugé insuffisant, mais pas de reformulation exploitable — on garde le résultat initial")
359
+
360
+ contexte = _construire_contexte(articles)
361
+ prompt = PROMPT_TEMPLATE.format(contexte=contexte, question=question)
362
+
363
+ reponse = appel_llm(prompt, taille="long", verbose=verbose)
364
+
365
+ if verbose:
366
+ print(f"QUESTION : {question}")
367
+ print(f"RÉPONSE : {reponse}")
368
+ refus = _est_un_refus(reponse)
369
+ print(f"REFUS DÉTECTÉ : {refus}")
370
+ print("SOURCES CONSULTÉES :")
371
+ for a in articles:
372
+ print(f" - Article {a['numero_article']} — {a['titre']}")
373
+ print()
374
+
375
+ return reponse
376
+
377
+
378
+ if __name__ == "__main__":
379
+ questions_test = [
380
+ "Comment calcule-t-on le minimum de capital requis ?",
381
+ "Comment sont calculées les provisions techniques ?",
382
+ "Quel est le taux de TVA sur les croissants ?",
383
+ ]
384
+
385
+ for q in questions_test:
386
+ poser_question(q)
387
+ print("-" * 70)