Spaces:
Sleeping
Sleeping
| """ | |
| Sensitivity Analysis for RAG Parameters | |
| ======================================== | |
| Teste l'accuracy du RAG sur différentes valeurs de paramètres, une à la fois, | |
| en gardant les autres paramètres à la valeur baseline. | |
| Usage: | |
| python sensitivity_analysis.py --input questions.xlsx | |
| python sensitivity_analysis.py --input questions.xlsx --output results.xlsx | |
| python sensitivity_analysis.py --input questions.xlsx --max-questions 20 # test rapide | |
| Format du fichier Excel d'entrΓ©e (colonnes requises): | |
| question : texte de la question | |
| correct_answer : lettre attendue (A/B/C/D) pour les QCM, mots-clΓ©s sΓ©parΓ©s par des virgules pour les questions ouvertes | |
| type : "QCM" ou "OPEN" (optionnel, auto-dΓ©tectΓ© si absent) | |
| """ | |
| import argparse | |
| import re | |
| import time | |
| import logging | |
| from pathlib import Path | |
| import pandas as pd | |
| logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s") | |
| logger = logging.getLogger(__name__) | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Grille de sensibilitΓ© β modifier ces valeurs pour changer les tests | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Configuration de référence (première valeur de chaque liste) | |
| BASELINE = { | |
| "temperature": 0.1, | |
| "max_tokens": 1024, | |
| "enable_reranking": True, | |
| "top_k_retrieval": 20, | |
| "top_k_reranked": 5, | |
| "enable_bm25": True, | |
| } | |
| # Pour chaque paramètre, liste des valeurs à tester | |
| # La valeur baseline ne sera pas re-testΓ©e sΓ©parΓ©ment (elle est dΓ©jΓ dans le run "baseline") | |
| SENSITIVITY_GRID = { | |
| "temperature": [0.05, 0.1, 0.3, 0.7, 1.5], | |
| "max_tokens": [512, 1024, 2048], | |
| "enable_reranking": [False, True], | |
| "top_k_retrieval": [5, 10, 15, 20], | |
| "top_k_reranked": [1,2,3,5,7,10], | |
| "enable_bm25": [False, True], | |
| } | |
| DELAY_BETWEEN_CALLS = 0.3 # secondes entre chaque appel API (Γ©vite le rate limiting) | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Helpers pour Γ©valuer les rΓ©ponses | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def detect_qcm(question: str) -> bool: | |
| """DΓ©tecte si une question est un QCM en cherchant des options A), B), C)...""" | |
| return bool(re.search(r'\b[A-D][).]\s', question)) | |
| def extract_letter(text: str) -> str: | |
| """Extrait la lettre de rΓ©ponse (A/B/C/D) depuis le texte de rΓ©ponse du LLM.""" | |
| if not text: | |
| return "" | |
| u = text.strip().upper() | |
| if u in ("A", "B", "C", "D"): | |
| return u | |
| # Lettre en dΓ©but : "A.", "A)", "A:" | |
| m = re.match(r'^([A-D])[.)\s:,]', u) | |
| if m: | |
| return m.group(1) | |
| # "la rΓ©ponse est A", "rΓ©ponse: A", "answer: A" | |
| m = re.search(r'(?:RΓPONSE|ANSWER|OPTION|CORRECT|LETTRE|CHOIX)\s*[:\s]\s*([A-D])\b', u) | |
| if m: | |
| return m.group(1) | |
| # Premier token standalone trouvΓ© | |
| m = re.search(r'\b([A-D])\b', u) | |
| return m.group(1) if m else "" | |
| def check_answer(got: str, expected: str, q_type: str) -> bool: | |
| """ | |
| Retourne True si la rΓ©ponse est correcte. | |
| QCM : compare la lettre extraite Γ la lettre attendue. | |
| OPEN : vΓ©rifie que tous les mots-clΓ©s attendus (sΓ©parΓ©s par virgule) apparaissent dans la rΓ©ponse. | |
| """ | |
| if q_type == "QCM": | |
| return extract_letter(got) == expected.strip().upper() | |
| # OPEN : matching par mots-clΓ©s | |
| keywords = [k.strip().lower() for k in expected.split(",") if k.strip()] | |
| if not keywords: | |
| return False # pas de mots-clΓ©s β rΓ©vision manuelle nΓ©cessaire | |
| return all(k in got.lower() for k in keywords) | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # ExΓ©cution d'un experiment (une config Γ toutes les questions) | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def run_experiment(df: pd.DataFrame, config: dict) -> pd.DataFrame: | |
| from app import rag_query # noqa: PLC0415 β import tardif voulu (init ChromaDB + reranker) | |
| rows = [] | |
| n = len(df) | |
| for i, row in enumerate(df.itertuples(index=False), 1): | |
| question = str(row.question) | |
| expected = str(row.correct_answer) | |
| q_type = str(row.type).upper() if hasattr(row, "type") else ( | |
| "QCM" if detect_qcm(question) else "OPEN" | |
| ) | |
| logger.info(f" [{i}/{n}] {question[:70]}...") | |
| t0 = time.perf_counter() | |
| try: | |
| result = rag_query( | |
| question, | |
| top_k=config["top_k_retrieval"], | |
| temperature=config["temperature"], | |
| max_tokens=config["max_tokens"], | |
| enable_reranking=config["enable_reranking"], | |
| top_k_reranked=config["top_k_reranked"], | |
| enable_bm25=config["enable_bm25"], | |
| ) | |
| answer = result.get("answer", "") | |
| correct = check_answer(answer, expected, q_type) | |
| total_token = result.get("total_token", 0) or 0 | |
| co2_grams = result.get("co2_grams") | |
| error = "" | |
| except Exception as e: | |
| answer = "" | |
| correct = False | |
| total_token = 0 | |
| co2_grams = None | |
| error = str(e) | |
| logger.error(f" Erreur sur la question {i}: {e}") | |
| elapsed_ms = round((time.perf_counter() - t0) * 1000) | |
| rows.append({ | |
| "question": question[:150], | |
| "type": q_type, | |
| "expected": expected, | |
| "got": answer[:300] if answer else "", | |
| "got_letter": extract_letter(answer) if q_type == "QCM" else "-", | |
| "correct": correct, | |
| "time_ms": elapsed_ms, | |
| "total_token": total_token, | |
| "co2_grams": co2_grams, | |
| "error": error, | |
| }) | |
| time.sleep(DELAY_BETWEEN_CALLS) | |
| return pd.DataFrame(rows) | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Calcul des mΓ©triques d'un experiment | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def compute_metrics(detail_df: pd.DataFrame) -> dict: | |
| qcm_mask = detail_df["type"] == "QCM" | |
| overall = detail_df["correct"].mean() * 100 | |
| qcm_acc = detail_df.loc[qcm_mask, "correct"].mean() * 100 if qcm_mask.any() else float("nan") | |
| open_acc = detail_df.loc[~qcm_mask, "correct"].mean() * 100 if (~qcm_mask).any() else float("nan") | |
| total_tokens = int(detail_df["total_token"].sum()) | |
| avg_tokens = round(detail_df["total_token"].mean(), 1) | |
| co2_values = detail_df["co2_grams"].dropna() | |
| total_co2 = round(co2_values.sum(), 6) if not co2_values.empty else "n/a" | |
| avg_co2 = round(co2_values.mean(), 6) if not co2_values.empty else "n/a" | |
| return { | |
| "accuracy_overall_%": round(overall, 1), | |
| "accuracy_qcm_%": round(qcm_acc, 1) if not pd.isna(qcm_acc) else "n/a", | |
| "accuracy_open_%": round(open_acc, 1) if not pd.isna(open_acc) else "n/a", | |
| "correct": int(detail_df["correct"].sum()), | |
| "total": len(detail_df), | |
| "avg_time_ms": int(detail_df["time_ms"].mean()), | |
| "total_tokens": total_tokens, | |
| "avg_tokens_q": avg_tokens, | |
| "total_co2_g": total_co2, | |
| "avg_co2_g_q": avg_co2, | |
| } | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Main | |
| # ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def main(): | |
| parser = argparse.ArgumentParser(description="RAG Sensitivity Analysis") | |
| parser.add_argument("--input", required=True, help="Fichier Excel d'entrΓ©e") | |
| parser.add_argument("--output", default="sensitivity_results.xlsx", help="Fichier Excel de sortie") | |
| parser.add_argument("--sheet", default=0, help="Nom ou index de la feuille d'entrΓ©e") | |
| parser.add_argument("--max-questions", type=int, default=None, help="Limite le nombre de questions (test rapide)") | |
| parser.add_argument("--baseline-only", action="store_true", help="ExΓ©cute uniquement la configuration baseline (pas de variations)") | |
| args = parser.parse_args() | |
| # ββ Lecture de l'Excel d'entrΓ©e ββββββββββββββββββββββββββββββββββββββββββ | |
| df = pd.read_excel(args.input, sheet_name=args.sheet) | |
| df.columns = df.columns.str.lower().str.strip() | |
| required_cols = {"question", "correct_answer"} | |
| missing = required_cols - set(df.columns) | |
| if missing: | |
| raise ValueError(f"Colonnes manquantes dans l'Excel : {missing}") | |
| if "type" not in df.columns: | |
| df["type"] = df["question"].apply(lambda q: "QCM" if detect_qcm(str(q)) else "OPEN") | |
| if args.max_questions: | |
| df = df.head(args.max_questions) | |
| logger.info(f"Mode test rapide : {args.max_questions} questions seulement.") | |
| counts = df["type"].value_counts().to_dict() | |
| logger.info(f"{len(df)} questions chargΓ©es β {counts}") | |
| # ββ Construction de la liste des experiments βββββββββββββββββββββββββββββ | |
| # Format : (nom_param, valeur_affichée, config_complète) | |
| experiments = [("baseline", "baseline", BASELINE)] | |
| if not args.baseline_only: | |
| for param, values in SENSITIVITY_GRID.items(): | |
| for v in values: | |
| if v == BASELINE.get(param): | |
| continue # dΓ©jΓ couvert par le baseline | |
| config = {**BASELINE, param: v} | |
| # Comparaison Γ©quitable pour enable_reranking=False : | |
| # on envoie le mΓͺme nombre de chunks au LLM qu'avec reranking | |
| if param == "enable_reranking" and v is False: | |
| config["top_k_retrieval"] = BASELINE["top_k_reranked"] | |
| experiments.append((param, str(v), config)) | |
| total = len(experiments) | |
| estimated_min = round(total * len(df) * (1 + DELAY_BETWEEN_CALLS) / 60, 1) | |
| logger.info(f"{total} experiments Γ {len(df)} questions β {estimated_min} min estimΓ©es") | |
| # ββ ExΓ©cution : collecte tous les rΓ©sultats en mΓ©moire βββββββββββββββββββ | |
| # On collecte d'abord TOUT, puis on Γ©crit l'Excel une seule fois. | |
| # Ainsi un crash en cours ne laisse pas un fichier Excel vide/corrompu. | |
| summary_rows = [] | |
| detail_sheets: dict[str, pd.DataFrame] = {} | |
| for idx, (param, value, config) in enumerate(experiments, 1): | |
| label = "baseline" if param == "baseline" else f"{param}={value}" | |
| logger.info(f"\n[{idx}/{total}] {label} | config={config}") | |
| try: | |
| detail_df = run_experiment(df, config) | |
| metrics = compute_metrics(detail_df) | |
| except Exception as e: | |
| logger.error(f" Experiment Γ©chouΓ© ({label}): {e}") | |
| continue | |
| summary_rows.append({ | |
| "parameter": param, | |
| "value": value, | |
| **metrics, | |
| **{f"cfg_{k}": v for k, v in config.items()}, | |
| }) | |
| detail_sheets[label[:31]] = detail_df | |
| logger.info(f" β accuracy={metrics['accuracy_overall_%']}% ({metrics['correct']}/{metrics['total']})") | |
| if not summary_rows: | |
| logger.error("Aucun experiment n'a abouti. VΓ©rifiez les dΓ©pendances (pip install -r requirements.txt).") | |
| return | |
| # ββ Γcriture de l'Excel βββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| summary_df = pd.DataFrame(summary_rows) | |
| with pd.ExcelWriter(args.output, engine="openpyxl") as writer: | |
| # Summary en premier | |
| summary_df.to_excel(writer, sheet_name="Summary", index=False) | |
| # DΓ©tails par experiment | |
| for sheet_name, detail_df in detail_sheets.items(): | |
| detail_df.to_excel(writer, sheet_name=sheet_name, index=False) | |
| logger.info(f"\nRΓ©sultats sauvegardΓ©s dans : {args.output}") | |
| logger.info("\n" + summary_df[ | |
| ["parameter", "value", "accuracy_overall_%", "accuracy_qcm_%", "correct", "total"] | |
| ].to_string(index=False)) | |
| if __name__ == "__main__": | |
| main() | |