ViTViz / experiments /run_aggregation_comparison.py
lucasddmc's picture
feat(ui): descriΓ§Γ£o, crΓ©ditos e agradecimento de financiamento; corrige caixa de upload presa ao voltar para modelo default
b0e01a5
Raw
History Blame Contribute Delete
9.23 kB
#!/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()