#!/usr/bin/env python3 """Comparação local: Attention Rollout vs GMAR-L1 vs GMAR-L2. Roda 2 imagens × 2 ataques (PGD, TGR) para cada método de agregação, medindo tempo total do bloco e valores de W1, JSD e entropy delta. Usage: python experiments/run_aggregation_comparison.py python experiments/run_aggregation_comparison.py --eps 0.03137 --steps 5 """ import argparse import sys import time from pathlib import Path import numpy as np import torch def _find_project_root() -> Path: current = Path(__file__).resolve().parent for parent in [current, *current.parents]: if (parent / "requirements.txt").exists(): return parent raise RuntimeError(f"Project root not found from {__file__}") PROJECT_ROOT = _find_project_root() sys.path.insert(0, str(PROJECT_ROOT)) from utils.attacks import PGDIterations, TGR from utils.metrics import compute_attention_entropy_delta, compute_attention_jsd, compute_attention_w1 from utils.model_loader import load_model_and_labels from utils.preprocessing import get_default_transform, preprocess_image from utils.seed import set_seed from utils.visualization import compute_attention_map def _list_images(images_dir: Path, n: int) -> list: exts = {".jpg", ".jpeg", ".png", ".bmp", ".webp"} files = [p for p in sorted(images_dir.iterdir()) if p.is_file() and p.suffix.lower() in exts] if not files: raise FileNotFoundError(f"Nenhuma imagem encontrada em {images_dir}") return files[:n] def _run_attack(attack, img_tensor, label): with torch.enable_grad(): adv_tensor, _ = attack(img_tensor, label) clean_attn = getattr(attack, "attentions_per_iter", [None])[0] adv_attn = getattr(attack, "attentions_per_iter", [None, None])[-1] return adv_tensor, clean_attn, adv_attn def _compute_metrics(method, clean_attn, adv_attn, model, img_tensor, adv_tensor, attn_kwargs): orig_map = compute_attention_map( method=method, attentions=clean_attn, model=model, image=img_tensor, **attn_kwargs, ) adv_map = compute_attention_map( method=method, attentions=adv_attn, model=model, image=adv_tensor, **attn_kwargs, ) return { "w1": compute_attention_w1(orig_map, adv_map), "jsd": compute_attention_jsd(orig_map, adv_map), "entropy": compute_attention_entropy_delta(orig_map, adv_map), } def main(): parser = argparse.ArgumentParser(description="Comparação Rollout vs GMAR (local, 2 imagens)") parser.add_argument("--model", type=str, default="hf-model://timm/vit_small_patch16_224.augreg_in1k") parser.add_argument("--images-dir", type=str, default="data/sample_images") parser.add_argument("--n-images", type=int, default=2) parser.add_argument("--eps", type=float, default=0.03137) # 8/255 parser.add_argument("--steps", type=int, default=10) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--discard-ratio", type=float, default=0.9) parser.add_argument("--head-fusion", type=str, default="max") parser.add_argument("--alpha", type=float, default=0.5) args = parser.parse_args() set_seed(args.seed) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Device: {device}") print(f"Carregando modelo: {args.model}") model, _, _, vit_cfg = load_model_and_labels(args.model, device=device) model.eval() images_dir = PROJECT_ROOT / args.images_dir image_paths = _list_images(images_dir, args.n_images) print(f"Imagens: {[p.name for p in image_paths]}") transform = get_default_transform(img_size=vit_cfg.img_size) eps = args.eps alpha = eps * 0.25 attacks_cfg = [ ("PGD", lambda m: PGDIterations(m, eps=eps, alpha=alpha, steps=args.steps, random_start=False, collect_images=False)), ("TGR", lambda m: TGR(m, eps=eps, steps=args.steps, decay=1.0, gamma_attn=0.25, gamma_qkv=0.75, gamma_mlp=0.5, collect_images=False)), ] methods = ["rollout", "gmar_l1", "gmar_l2"] attn_kwargs_rollout = {"discard_ratio": args.discard_ratio, "head_fusion": args.head_fusion} attn_kwargs_gmar = {"alpha": args.alpha} def _attn_kwargs(method): return attn_kwargs_rollout if method == "rollout" else attn_kwargs_gmar # ── Preparar tensores uma vez ──────────────────────────────────────────── samples = [] for img_path in image_paths: img_tensor = preprocess_image(str(img_path), transform=transform).to(device) with torch.no_grad(): out = model(img_tensor) logits = out.logits if hasattr(out, "logits") else out label = torch.tensor([logits.argmax().item()], device=device) samples.append((img_path.name, img_tensor, label)) # ── Rodar ataques uma vez e guardar tensores + atenções ───────────────── print("\nRodando ataques (uma vez, reutilizado para todos os métodos)...") attack_results = {} # (img_name, atk_name) -> (adv_tensor, clean_attn, adv_attn) for img_name, img_tensor, label in samples: for atk_name, atk_factory in attacks_cfg: set_seed(args.seed) attack = atk_factory(model) adv_tensor, clean_attn, adv_attn = _run_attack(attack, img_tensor, label) attack_results[(img_name, atk_name)] = (adv_tensor, img_tensor, clean_attn, adv_attn) print(f" {img_name} × {atk_name}: OK") # ── Loop por método ────────────────────────────────────────────────────── results = {} # method -> list of (img, atk, metrics_dict) for method in methods: print(f"\n{'─'*60}") print(f"Método: {method.upper()}") t_method_start = time.perf_counter() rows = [] for img_name, img_tensor, label in samples: for atk_name, _ in attacks_cfg: adv_tensor, orig_tensor, clean_attn, adv_attn = attack_results[(img_name, atk_name)] t0 = time.perf_counter() m = _compute_metrics( method=method, clean_attn=clean_attn, adv_attn=adv_attn, model=model, img_tensor=orig_tensor, adv_tensor=adv_tensor, attn_kwargs=_attn_kwargs(method), ) elapsed_ms = (time.perf_counter() - t0) * 1000 rows.append((img_name, atk_name, m, elapsed_ms)) print(f" {img_name} × {atk_name}: " f"W1={m['w1']:.4f} JSD={m['jsd']:.4f} ΔH={m['entropy']:+.4f} " f"({elapsed_ms:.0f} ms)") total_ms = (time.perf_counter() - t_method_start) * 1000 results[method] = {"rows": rows, "total_ms": total_ms} print(f" → Tempo total do bloco: {total_ms:.0f} ms") # ── Tabela comparativa ─────────────────────────────────────────────────── print(f"\n{'═'*72}") print("TABELA COMPARATIVA") print(f"{'═'*72}") # Header col_w = 12 print(f"{'Imagem':<16} {'Ataque':<6} {'Métrica':<10}", end="") for m in methods: print(f" {m.upper():>{col_w}}", end="") print() print("─" * 72) metric_labels = [("W1", "w1"), ("JSD", "jsd"), ("ΔEntropy", "entropy")] for img_name, _, __ in samples: for atk_name, _ in attacks_cfg: for metric_label, metric_key in metric_labels: print(f"{img_name:<16} {atk_name:<6} {metric_label:<10}", end="") for method in methods: row_data = next( (r for r in results[method]["rows"] if r[0] == img_name and r[1] == atk_name), None, ) val = row_data[2][metric_key] if row_data else float("nan") sign = "+" if val > 0 else "" print(f" {sign}{val:>{col_w}.4f}", end="") print() print() print("─" * 72) print(f"{'TEMPO TOTAL (ms)':<34}", end="") for method in methods: t = results[method]["total_ms"] print(f" {t:>{col_w}.0f}", end="") print() rollout_t = results["rollout"]["total_ms"] print(f"{'Overhead vs rollout':<34}", end="") for method in methods: t = results[method]["total_ms"] overhead = t / rollout_t if rollout_t > 0 else float("nan") print(f" {overhead:>{col_w}.2f}x", end="") print() print(f"{'═'*72}") print(f"\nConfig: eps={eps:.5f} steps={args.steps} " f"discard_ratio={args.discard_ratio} head_fusion={args.head_fusion} alpha={args.alpha}") if __name__ == "__main__": main()