AetherMap / ablation_study_squad.py
Madras1's picture
Upload 71 files
971cb75 verified
Raw History Blame Contribute Delete
21.9 kB
# ==============================================================================
# AetherMap — ABLATION STUDY v2 (SQuAD EN)
# Usa SQuAD v1.1 original (inglês) com Q&A humanas
# Para rodar no Google Colab contra a API do HF Space
# ==============================================================================
# %%
# !pip install requests pandas numpy matplotlib datasets -q
# %%
import os
import json
import time
import random
import requests
import pandas as pd
import numpy as np
import tempfile
from typing import List, Dict, Any
from datasets import load_dataset
# ==============================================================================
# CONFIGURAÇÃO
# ==============================================================================
OPENROUTER_API_KEY = "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" # 👈 SUA KEY OPENROUTER
AETHERMAP_URL = "https://madras1-aethermap.hf.space"
# Modelo pro LLM-as-Judge (avaliador)
EVALUATOR_MODEL = "nvidia/nemotron-3-nano-30b-a3b:free"
DELAY_BETWEEN_LLM_CALLS = 2.0 # Segundos entre chamadas LLM
OPENROUTER_HEADERS = {
"Authorization": f"Bearer {OPENROUTER_API_KEY}",
"Content-Type": "application/json",
}
# Modos de ablation e seus componentes
ABLATION_MODES = {
"faiss_only": {"label": "FAISS Only", "faiss": True, "bm25": False, "rrf": False, "reranker": False, "expansion": False},
"bm25_only": {"label": "BM25 Only", "faiss": False, "bm25": True, "rrf": False, "reranker": False, "expansion": False},
"hybrid": {"label": "Hybrid (RRF)", "faiss": True, "bm25": True, "rrf": True, "reranker": False, "expansion": False},
"hybrid_rerank": {"label": "+ Reranker", "faiss": True, "bm25": True, "rrf": True, "reranker": True, "expansion": False},
"full": {"label": "Full Pipeline", "faiss": True, "bm25": True, "rrf": True, "reranker": True, "expansion": True},
}
print("✅ Configuração carregada")
print(f"🌐 API: {AETHERMAP_URL}")
print(f"🧑‍⚖️ Avaliador: {EVALUATOR_MODEL}")
print("📚 Dataset: SQuAD v1.1 (English)")
# %%
# ==============================================================================
# CARREGAR SQUAD PT-BR
# ==============================================================================
def carregar_squad_pt(n_contexts: int = 300, n_queries: int = 30, seed: int = 42) -> tuple:
"""
Carrega SQuAD v1.1 PT-BR e extrai contextos + Q&A pairs.
Returns:
contexts: Lista de parágrafos únicos para indexar
qa_pairs: Lista de {pergunta, resposta, contexto_id} para testar
"""
print("📥 Baixando SQuAD v1.1 (inglês) do HuggingFace...")
ds = load_dataset("rajpurkar/squad", split="validation") # validation = 10k exemplos, mais rápido
print(f" 📊 Total: {len(ds)} exemplos no dataset")
print(f" 📋 Colunas: {ds.column_names}")
# SQuAD original: id, title, context, question, answers
ctx_col = "context"
q_col = "question"
ans_col = "answers"
title_col = "title"
print(f" 🔍 Usando: context='{ctx_col}', question='{q_col}', answers='{ans_col}', title='{title_col}'")
# Extrair contextos únicos
context_set = {}
for item in ds:
ctx = str(item[ctx_col]).strip()
if ctx and len(ctx) > 50 and ctx not in context_set:
context_set[ctx] = {
"text": ctx,
"title": str(item.get(title_col, "")) if title_col else "",
}
all_contexts = list(context_set.values())
print(f" 📄 {len(all_contexts)} contextos únicos encontrados")
# Amostrar contextos
random.seed(seed)
sampled_contexts = random.sample(all_contexts, min(n_contexts, len(all_contexts)))
sampled_texts = set(c["text"] for c in sampled_contexts)
print(f" ✂️ {len(sampled_contexts)} contextos amostrados")
# Coletar Q&A pairs que pertencem aos contextos amostrados
qa_candidates = []
for item in ds:
ctx = str(item[ctx_col]).strip()
if ctx in sampled_texts:
# Extrair resposta (formato pode ser dict com "text" ou string direta)
if ans_col:
ans_raw = item[ans_col]
if isinstance(ans_raw, dict) and "text" in ans_raw:
answer_texts = ans_raw["text"]
if isinstance(answer_texts, list) and answer_texts:
answer = answer_texts[0].strip()
else:
answer = str(answer_texts).strip()
elif isinstance(ans_raw, list) and ans_raw:
answer = str(ans_raw[0]).strip()
else:
answer = str(ans_raw).strip()
else:
answer = ""
question = str(item[q_col]).strip()
if question and answer and len(answer) > 2:
qa_candidates.append({
"pergunta": question,
"resposta": answer,
"contexto": ctx[:200],
})
print(f" ❓ {len(qa_candidates)} Q&A pairs nos contextos amostrados")
# Amostrar queries diversificadas por contexto
random.shuffle(qa_candidates)
seen_contexts = set()
diverse_queries = []
for qa in qa_candidates:
ctx_key = qa["contexto"][:100]
if ctx_key not in seen_contexts:
diverse_queries.append(qa)
seen_contexts.add(ctx_key)
if len(diverse_queries) >= n_queries:
break
# Completar se necessário
if len(diverse_queries) < n_queries:
remaining = [q for q in qa_candidates if q not in diverse_queries]
diverse_queries.extend(remaining[:n_queries - len(diverse_queries)])
print(f" ✅ {len(diverse_queries)} queries selecionadas (diversificadas por contexto)")
return sampled_contexts, diverse_queries
# %%
# ==============================================================================
# FUNÇÕES AUXILIARES
# ==============================================================================
def call_llm(prompt: str, model: str, max_tokens: int = 1000) -> str:
"""Chama LLM via OpenRouter usando requests direto."""
time.sleep(DELAY_BETWEEN_LLM_CALLS)
try:
payload = {
"model": model,
"messages": [{"role": "user", "content": prompt}],
"max_tokens": max_tokens,
"temperature": 0.3,
"reasoning": {
"exclude": True # Reasoning vai pra campo separado, content fica limpo
},
}
resp = requests.post(
"https://openrouter.ai/api/v1/chat/completions",
headers=OPENROUTER_HEADERS,
json=payload,
timeout=90
)
if resp.status_code != 200:
print(f" ⚠️ LLM ({model}): HTTP {resp.status_code}")
print(f" 📝 {resp.text[:300]}")
return ""
data = resp.json()
choices = data.get("choices", [])
if not choices:
print(f" ⚠️ LLM ({model}): Sem choices")
print(f" 📝 Raw: {json.dumps(data)[:400]}")
return ""
message = choices[0].get("message", {})
content = message.get("content") or ""
if not content:
print(f" ⚠️ LLM ({model}): content vazio")
print(f" 📝 Message keys: {list(message.keys())}")
print(f" 📝 Raw: {json.dumps(message)[:400]}")
return ""
return content.strip()
except Exception as e:
print(f" ⚠️ Erro LLM ({model}): {type(e).__name__}: {e}")
time.sleep(5)
return ""
def avaliar_resposta(pergunta: str, resposta_rag: str, resposta_esperada: str) -> Dict[str, float]:
"""
Avalia qualidade da resposta do RAG comparando com ground truth do SQuAD.
Retorna scores de correção, completude e relevância (0-1).
"""
prompt = f"""Avalie a resposta do sistema comparando com a resposta de referência (ground truth humana).
PERGUNTA: {pergunta}
RESPOSTA DE REFERÊNCIA (humana): {resposta_esperada[:300]}
RESPOSTA DO SISTEMA (RAG): {resposta_rag[:400]}
Dê notas de 0.0 a 1.0 para:
- correcao: A informação está factualmente correta comparada à referência?
- completude: A resposta cobre todos os pontos da referência?
- relevancia: A resposta é diretamente relevante à pergunta?
Retorne APENAS JSON: {{"correcao": X.X, "completude": X.X, "relevancia": X.X}}"""
response = call_llm(prompt, EVALUATOR_MODEL, max_tokens=500)
try:
if "{" in response:
json_start = response.find("{")
json_end = response.rfind("}")
if json_end > json_start:
result = json.loads(response[json_start:json_end + 1])
return {
"correcao": float(result.get("correcao", 0.5)),
"completude": float(result.get("completude", 0.5)),
"relevancia": float(result.get("relevancia", 0.5)),
}
except Exception as e:
print(f" ⚠️ Erro parsing avaliação: {e}")
print(f" 📝 Response: {response[:200]}")
return {"correcao": 0.5, "completude": 0.5, "relevancia": 0.5}
# %%
# ==============================================================================
# UPLOAD DO DATASET
# ==============================================================================
def upload_contexts(contexts: List[Dict], text_column: str = "texto") -> str:
"""Faz upload dos contextos SQuAD para o AetherMap."""
# Criar CSV com os contextos
df = pd.DataFrame({
"texto": [c["text"] for c in contexts],
"titulo": [c["title"] for c in contexts],
})
with tempfile.NamedTemporaryFile(mode='w', suffix='.csv', delete=False, encoding='utf-8') as f:
df.to_csv(f, index=False)
temp_path = f.name
print(f"📤 Upload de {len(contexts)} contextos SQuAD...")
with open(temp_path, "rb") as f:
response = requests.post(
f"{AETHERMAP_URL}/process/",
files={"file": ("squad_contexts.csv", f)},
data={
"n_samples": len(contexts),
"text_column": text_column,
"fast_mode": "true"
},
timeout=300
)
os.unlink(temp_path)
if response.status_code != 200:
raise Exception(f"Erro no upload: {response.text[:200]}")
result = response.json()
job_id = result["job_id"]
meta = result["metadata"]
print(f"✅ Job criado: {job_id[:8]}...")
print(f" 📊 {meta['num_documents_processed']} docs | {meta['num_clusters_found']} clusters")
return job_id
# %%
# ==============================================================================
# BUSCA COM MODO DE ABLATION
# ==============================================================================
def search_with_mode(job_id: str, query: str, mode: str) -> Dict:
"""Executa busca com um modo de ablation específico."""
start = time.time()
try:
response = requests.post(
f"{AETHERMAP_URL}/search/",
data={
"query": query,
"job_id": job_id,
"ablation_mode": mode
},
timeout=300 # 5 min — Full Pipeline com 300 docs pode demorar
)
latency = time.time() - start
if response.status_code != 200:
return {"error": response.text[:100], "latency": latency}
result = response.json()
result["latency"] = latency
return result
except requests.exceptions.Timeout:
latency = time.time() - start
print(f" ⏰ Timeout após {latency:.0f}s")
return {"error": f"Timeout ({latency:.0f}s)", "latency": latency}
except Exception as e:
latency = time.time() - start
print(f" ⚠️ Erro: {type(e).__name__}: {e}")
return {"error": str(e)[:100], "latency": latency}
# %%
# ==============================================================================
# ABLATION STUDY PRINCIPAL
# ==============================================================================
def run_ablation_study(
n_contexts: int = 300,
n_queries: int = 30
):
"""
Executa ablation study completo com SQuAD PT-BR.
Args:
n_contexts: Número de contextos para indexar
n_queries: Número de queries para testar
"""
print("=" * 60)
print("🧪 AETHERMAP ABLATION STUDY v2 — SQuAD EN")
print("=" * 60)
# 1. Carregar SQuAD
print("\n📚 FASE 1: Carregando SQuAD PT-BR")
contexts, queries = carregar_squad_pt(n_contexts, n_queries)
# 2. Upload
print(f"\n📦 FASE 2: Upload dos contextos")
job_id = upload_contexts(contexts)
# 3. Mostrar queries
print(f"\n❓ FASE 3: {len(queries)} queries do SQuAD EN (ground truth humana)")
for i, q in enumerate(queries):
print(f" [{i+1}] Q: {q['pergunta'][:70]}...")
print(f" A: {q['resposta'][:70]}...")
# 4. Rodar cada modo
print(f"\n🧪 FASE 4: Testando {len(ABLATION_MODES)} modos de ablation...")
all_results = {}
for mode_key, mode_info in ABLATION_MODES.items():
mode_label = mode_info["label"]
print(f"\n{'─' * 40}")
print(f"🔬 Modo: {mode_label} ({mode_key})")
print(f" Componentes: ", end="")
components = []
if mode_info["faiss"]: components.append("FAISS")
if mode_info["bm25"]: components.append("BM25")
if mode_info["rrf"]: components.append("RRF")
if mode_info["reranker"]: components.append("Reranker")
if mode_info["expansion"]: components.append("QueryExpansion")
print(" + ".join(components))
scores = []
latencies = []
for i, q in enumerate(queries):
pergunta = q["pergunta"]
esperada = q["resposta"]
# Buscar
result = search_with_mode(job_id, pergunta, mode_key)
if "error" in result:
print(f" [{i+1}] ❌ Erro: {result['error'][:50]}")
scores.append({"correcao": 0, "completude": 0, "relevancia": 0})
latencies.append(result.get("latency", 0))
continue
resposta = result.get("summary", "")
latency = result.get("latency", 0)
latencies.append(latency)
# Avaliar
aval = avaliar_resposta(pergunta, resposta, esperada)
scores.append(aval)
avg = (aval["correcao"] + aval["completude"] + aval["relevancia"]) / 3
emoji = "✓" if avg >= 0.6 else "○"
print(f" [{i+1}] {emoji} C:{aval['correcao']:.2f} Com:{aval['completude']:.2f} R:{aval['relevancia']:.2f} | {latency:.1f}s")
print(f" Q: {pergunta[:80]}")
print(f" RAG: {resposta[:150]}...")
if avg < 0.3:
print(f" ⚠️ Esperado: {esperada[:100]}")
time.sleep(0.5) # Rate limiting suave
# Calcular médias
avg_correcao = np.mean([s["correcao"] for s in scores])
avg_completude = np.mean([s["completude"] for s in scores])
avg_relevancia = np.mean([s["relevancia"] for s in scores])
avg_score = (avg_correcao + avg_completude + avg_relevancia) / 3
avg_latency = np.mean(latencies)
all_results[mode_key] = {
"label": mode_label,
"avg_score": avg_score,
"correcao": avg_correcao,
"completude": avg_completude,
"relevancia": avg_relevancia,
"avg_latency": avg_latency,
"n_queries": len(queries),
"components": components,
}
print(f" 📊 Score médio: {avg_score:.3f} | Latência: {avg_latency:.1f}s")
# 5. Resultado Final
print_results(all_results)
plot_results(all_results)
return all_results
# %%
# ==============================================================================
# VISUALIZAÇÃO DOS RESULTADOS
# ==============================================================================
def print_results(results: Dict):
"""Imprime tabela comparativa dos resultados."""
print("\n" + "=" * 70)
print("📊 RESULTADO DO ABLATION STUDY — SQuAD EN")
print("=" * 70)
print(f"\n{'Modo':<20} {'Score':>7} {'Correção':>9} {'Complet.':>9} {'Relev.':>7} {'Latência':>9} {'Δ Score':>8}")
print("─" * 70)
baseline_score = results.get("faiss_only", {}).get("avg_score", 0)
best_score = max(r["avg_score"] for r in results.values())
for mode_key in ABLATION_MODES.keys():
if mode_key not in results:
continue
r = results[mode_key]
delta = r["avg_score"] - baseline_score
delta_str = f"+{delta:.3f}" if delta > 0 else f"{delta:.3f}"
trophy = " 🏆" if r["avg_score"] == best_score else ""
print(f"{r['label']:<20} {r['avg_score']:>7.3f} {r['correcao']:>9.3f} {r['completude']:>9.3f} {r['relevancia']:>7.3f} {r['avg_latency']:>8.1f}s {delta_str:>8}{trophy}")
# Insights
print(f"\n{'─' * 70}")
print("💡 INSIGHTS:")
if "faiss_only" in results and "hybrid" in results:
delta = results["hybrid"]["avg_score"] - results["faiss_only"]["avg_score"]
base = max(results['faiss_only']['avg_score'], 0.001)
print(f" Hybrid Search (RRF): {'+' if delta >= 0 else ''}{delta:.3f} vs FAISS only ({delta/base*100:+.1f}%)")
if "hybrid" in results and "hybrid_rerank" in results:
delta = results["hybrid_rerank"]["avg_score"] - results["hybrid"]["avg_score"]
base = max(results['hybrid']['avg_score'], 0.001)
print(f" + Reranker: {'+' if delta >= 0 else ''}{delta:.3f} vs Hybrid ({delta/base*100:+.1f}%)")
if "hybrid_rerank" in results and "full" in results:
delta = results["full"]["avg_score"] - results["hybrid_rerank"]["avg_score"]
base = max(results['hybrid_rerank']['avg_score'], 0.001)
print(f" + Query Expansion: {'+' if delta >= 0 else ''}{delta:.3f} vs Hybrid+Rerank ({delta/base*100:+.1f}%)")
# Comparar n_queries
n = list(results.values())[0]["n_queries"]
print(f"\n 📏 N = {n} queries | Dataset: SQuAD v1.1 EN (ground truth humana)")
def plot_results(results: Dict):
"""Gera gráfico de barras com os resultados."""
try:
import matplotlib.pyplot as plt
import matplotlib
matplotlib.rcParams['figure.facecolor'] = '#0d1117'
matplotlib.rcParams['axes.facecolor'] = '#161b22'
matplotlib.rcParams['text.color'] = '#c9d1d9'
matplotlib.rcParams['axes.labelcolor'] = '#c9d1d9'
matplotlib.rcParams['xtick.color'] = '#8b949e'
matplotlib.rcParams['ytick.color'] = '#8b949e'
except ImportError:
print("⚠️ matplotlib não disponível. Pulando gráfico.")
return
labels = [results[k]["label"] for k in ABLATION_MODES if k in results]
scores = [results[k]["avg_score"] for k in ABLATION_MODES if k in results]
latencies = [results[k]["avg_latency"] for k in ABLATION_MODES if k in results]
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 6))
# Cores neon progressivas
colors = ['#6366f1', '#818cf8', '#34d399', '#fbbf24', '#f472b6'][:len(labels)]
# Gráfico 1: Scores
bars1 = ax1.bar(labels, scores, color=colors, edgecolor='#30363d', linewidth=1.5)
ax1.set_title('Quality Score por Modo', fontsize=14, fontweight='bold', pad=15)
ax1.set_ylabel('Score Médio (0-1)')
ax1.set_ylim(0, 1.05)
ax1.grid(axis='y', alpha=0.15, color='#8b949e')
for bar, score in zip(bars1, scores):
ax1.text(bar.get_x() + bar.get_width()/2., bar.get_height() + 0.02,
f'{score:.3f}', ha='center', va='bottom', fontweight='bold', fontsize=11)
# Gráfico 2: Latência
bars2 = ax2.bar(labels, latencies, color=colors, edgecolor='#30363d', linewidth=1.5)
ax2.set_title('Latência por Modo', fontsize=14, fontweight='bold', pad=15)
ax2.set_ylabel('Latência Média (s)')
ax2.grid(axis='y', alpha=0.15, color='#8b949e')
for bar, lat in zip(bars2, latencies):
ax2.text(bar.get_x() + bar.get_width()/2., bar.get_height() + 0.1,
f'{lat:.1f}s', ha='center', va='bottom', fontweight='bold', fontsize=11)
plt.xticks(rotation=25, ha='right')
plt.suptitle('AetherMap RAG — Ablation Study (SQuAD EN)', fontsize=16, fontweight='bold', y=1.02, color='#f0f6fc')
plt.tight_layout()
plt.savefig('ablation_results_squad.png', dpi=150, bbox_inches='tight',
facecolor='#0d1117', edgecolor='none')
plt.show()
print("📈 Gráfico salvo em ablation_results_squad_en.png")
# %%
# ==============================================================================
# EXECUTAR
# ==============================================================================
resultados = run_ablation_study(
n_contexts=300,
n_queries=30
)