My_App_Datack / sensitivity_analysis.py
ewannhugging's picture
last change
bf61a7b
Raw
History Blame Contribute Delete
14 kB
"""
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()