Download ablation_study_squad.py from Madras1/AetherMap: direct link, hf CLI and curl.
- Browser
- Download file 21.9 kB
-
https://huggingface.co/spaces/Madras1/AetherMap/resolve/main/ablation_study_squad.py
- Command line
-
hf download hf://spaces/Madras1/AetherMap/ablation_study_squad.py
-
curl -L -o ablation_study_squad.py https://huggingface.co/spaces/Madras1/AetherMap/resolve/main/ablation_study_squad.py
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 | |
| ) | |