ViTViz / scripts /plot_extra_analyses.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
14.1 kB
#!/usr/bin/env python3
"""4 figuras adicionais pro probe ε × métricas (Análises A.2–A.5 do plano).
Lê o CSV merged do probe (output de `merge_probe_csvs.py`) e gera:
A.2 — `{asr,ssim,lpips,psnr}_vs_eps_lines.png` (4 PNGs):
Curvas mean ± IC95 por ataque. Versão "limpa" dos boxplots.
A.3 — `heatmap_model_attack_eps8_{asr,ssim}.png` (2 PNGs):
Matriz 4×4 modelo × ataque, células coloridas por mean métrica a ε=8/255.
A.4 — `efficiency_per_attack_eps8.png` (1 PNG):
Bar chart: eficiência ASR / (1−SSIM) por ataque a ε=8/255.
A.5 — `asr_vs_eps_by_mask.png` (1 PNG):
Boxplot ASR ε × ataque, facetado por has_mask.
Total: 8 figuras adicionais.
Usage:
python scripts/plot_extra_analyses.py \\
--csv results/raw/probe_merged_no_tgr.csv \\
--out results/figures/probe_eps_curve_final/ \\
--eps-ref 8
"""
from __future__ import annotations
import argparse
import sys
from pathlib import Path
def _project_root() -> Path:
cur = Path(__file__).resolve().parent
for p in [cur, *cur.parents]:
if (p / "requirements.txt").exists():
return p
raise RuntimeError("project root not found")
PROJECT_ROOT = _project_root()
PALETTE = {
"FGSM": "#e41a1c", "PGD": "#377eb8", "MIM": "#4daf4a",
"TGR": "#984ea3", "SAGA": "#ff7f00",
}
ATTACK_ORDER = ["FGSM", "PGD", "MIM", "SAGA"] # TGR descartado pelo TCC
def _model_short_name(name: str) -> str:
"""'ViT-S/16 · ImageNet-1k' → 'ViT-S/16'."""
return name.split(" ·")[0].strip() if " ·" in name else name.strip()
def _setup_matplotlib():
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
return plt
# ─── A.2: curvas mean ± IC95 ──────────────────────────────────────────────────
def plot_lines_eps_metric(df, metric: str, ylabel: str, title: str,
out_path: Path, ylim=None) -> None:
plt = _setup_matplotlib()
import numpy as np
eps_values = sorted(df["eps_255"].unique())
fig, ax = plt.subplots(figsize=(8, 5))
for atk in ATTACK_ORDER:
sub = df[df["attack"] == atk]
if sub.empty:
continue
means, lo, hi = [], [], []
for e in eps_values:
vals = sub.loc[sub["eps_255"] == e, metric].dropna().values
if len(vals) == 0:
means.append(np.nan); lo.append(np.nan); hi.append(np.nan)
continue
m = float(np.mean(vals))
sem = float(np.std(vals, ddof=1)) / max(np.sqrt(len(vals)), 1)
means.append(m)
lo.append(m - 1.96 * sem)
hi.append(m + 1.96 * sem)
ax.plot(eps_values, means, marker="o", color=PALETTE[atk],
linewidth=2, label=atk)
ax.fill_between(eps_values, lo, hi, alpha=0.18, color=PALETTE[atk])
ax.set_xlabel("ε∞ (×1/255)")
ax.set_ylabel(ylabel)
ax.set_title(title)
ax.set_xticks(eps_values)
if ylim:
ax.set_ylim(*ylim)
ax.grid(alpha=0.3)
ax.legend(loc="best", fontsize=10)
fig.tight_layout()
fig.savefig(out_path, dpi=150)
plt.close(fig)
print(f" ✓ {out_path.name}")
# ─── A.3: heatmap modelo × ataque a ε=8 ───────────────────────────────────────
def plot_heatmap_model_attack(df, metric: str, eps_ref: int,
cmap: str, fmt: str,
title: str, out_path: Path,
vmin=None, vmax=None) -> None:
plt = _setup_matplotlib()
import numpy as np
sub = df[df["eps_255"] == eps_ref].copy()
sub["model_short"] = sub["model"].apply(_model_short_name)
pivot = sub.pivot_table(
index="model_short", columns="attack", values=metric, aggfunc="mean"
)
# Reorder columns
cols = [a for a in ATTACK_ORDER if a in pivot.columns]
pivot = pivot[cols]
# Reorder rows (S/16, S/32, B/32, B/16 — paper order)
desired_rows = ["ViT-S/16", "ViT-S/32", "ViT-B/32", "ViT-B/16"]
pivot = pivot.reindex([r for r in desired_rows if r in pivot.index])
fig, ax = plt.subplots(figsize=(7, 5))
im = ax.imshow(pivot.values, aspect="auto", cmap=cmap, vmin=vmin, vmax=vmax)
# Annotations
for i in range(pivot.shape[0]):
for j in range(pivot.shape[1]):
v = pivot.values[i, j]
if np.isnan(v):
txt = "—"
else:
txt = format(v, fmt)
# cor adaptativa
cell_color = im.cmap(im.norm(v)) if not np.isnan(v) else (1, 1, 1, 1)
lum = 0.299 * cell_color[0] + 0.587 * cell_color[1] + 0.114 * cell_color[2]
text_color = "white" if lum < 0.5 else "black"
ax.text(j, i, txt, ha="center", va="center", color=text_color, fontsize=11)
ax.set_xticks(range(len(pivot.columns)))
ax.set_xticklabels(pivot.columns)
ax.set_yticks(range(len(pivot.index)))
ax.set_yticklabels(pivot.index)
ax.set_title(title)
cbar = plt.colorbar(im, ax=ax, fraction=0.04, pad=0.04)
cbar.set_label(metric.upper())
fig.tight_layout()
fig.savefig(out_path, dpi=150)
plt.close(fig)
print(f" ✓ {out_path.name}")
# ─── A.4: bar chart de eficiência ASR / (1-SSIM) ──────────────────────────────
def plot_efficiency_bar(df, eps_ref: int, out_path: Path) -> None:
plt = _setup_matplotlib()
sub = df[df["eps_255"] == eps_ref]
rows = []
for atk in ATTACK_ORDER:
atk_sub = sub[sub["attack"] == atk]
if atk_sub.empty:
continue
mean_asr = float(atk_sub["asr"].mean())
mean_ssim = float(atk_sub["ssim"].mean())
denom = max(1.0 - mean_ssim, 1e-4) # evitar div por 0
eff = mean_asr / denom
rows.append({"attack": atk, "asr": mean_asr, "ssim": mean_ssim,
"efficiency": eff})
if not rows:
print(f" ⚠️ sem dados a ε={eps_ref}/255 — pulando efficiency bar")
return
rows.sort(key=lambda r: r["efficiency"], reverse=True)
attacks = [r["attack"] for r in rows]
effs = [r["efficiency"] for r in rows]
colors = [PALETTE[a] for a in attacks]
fig, ax = plt.subplots(figsize=(8, 5))
bars = ax.bar(attacks, effs, color=colors, edgecolor="black", alpha=0.85)
for bar, r in zip(bars, rows):
ax.text(bar.get_x() + bar.get_width() / 2, bar.get_height() * 1.02,
f"{r['efficiency']:.1f}\n(ASR={r['asr']:.2f}, SSIM={r['ssim']:.3f})",
ha="center", va="bottom", fontsize=9)
ax.set_ylabel("Eficiência = mean ASR / (1 − mean SSIM)")
ax.set_xlabel("Ataque")
ax.set_title(f"Eficiência por ataque a ε={eps_ref}/255 — quanto ASR por unidade de degradação visual")
ax.grid(axis="y", alpha=0.3)
ax.set_ylim(0, max(effs) * 1.25)
fig.tight_layout()
fig.savefig(out_path, dpi=150)
plt.close(fig)
print(f" ✓ {out_path.name}")
# ─── A.5: boxplot ASR ε × ataque, facetado por has_mask ───────────────────────
def plot_asr_by_mask(df, out_path: Path) -> None:
plt = _setup_matplotlib()
if "has_mask" not in df.columns:
print(f" ⚠️ has_mask ausente — pulando A.5")
return
eps_values = sorted(df["eps_255"].unique())
fig, axes = plt.subplots(1, 2, figsize=(15, 5.5), sharey=True)
titles = ["has_mask=1 (Guillaumin GT)", "has_mask=0 (IN-1k val)"]
n_attacks = len(ATTACK_ORDER)
box_width = 0.8 / n_attacks
for ax, mask_val, ttl in zip(axes, [1, 0], titles):
sub_mask = df[df["has_mask"] == mask_val]
legend_handles = []
for j, atk in enumerate(ATTACK_ORDER):
sub_atk = sub_mask[sub_mask["attack"] == atk]
if sub_atk.empty:
continue
data, positions = [], []
for i, e in enumerate(eps_values):
vals = sub_atk.loc[sub_atk["eps_255"] == e, "asr"].dropna().values
if len(vals) == 0:
continue
data.append(vals)
offset = (j - (n_attacks - 1) / 2) * box_width
positions.append(i + offset)
if not data:
continue
ax.boxplot(
data, positions=positions, widths=box_width * 0.85,
patch_artist=True, showfliers=False,
medianprops={"color": "black", "linewidth": 1.2},
boxprops={"facecolor": PALETTE[atk], "alpha": 0.7,
"edgecolor": PALETTE[atk]},
whiskerprops={"color": PALETTE[atk]},
capprops={"color": PALETTE[atk]},
)
legend_handles.append(plt.Rectangle(
(0, 0), 1, 1, fc=PALETTE[atk], alpha=0.7, label=atk
))
n_imgs = sub_mask["image"].nunique()
ax.set_xticks(range(len(eps_values)))
ax.set_xticklabels([f"{e}" for e in eps_values])
ax.set_xlabel("ε∞ (×1/255)")
ax.set_title(f"{ttl} — N={n_imgs} imgs")
ax.set_ylim(-0.05, 1.05)
ax.grid(axis="y", alpha=0.3)
if mask_val == 1:
ax.set_ylabel("ASR (per image)")
if legend_handles and mask_val == 0:
ax.legend(handles=legend_handles, loc="lower right",
fontsize=9, ncol=n_attacks)
fig.suptitle("ASR por ε × ataque, facetado por has_mask (Análise F)",
y=0.99, fontsize=12)
fig.tight_layout()
fig.savefig(out_path, dpi=150)
plt.close(fig)
# Diagnóstico de diferença entre os 2 grupos a ε=8/255
if 8 in eps_values:
sub8 = df[df["eps_255"] == 8]
for atk in ATTACK_ORDER:
sub_atk = sub8[sub8["attack"] == atk]
if sub_atk.empty:
continue
asr_with = sub_atk[sub_atk["has_mask"] == 1]["asr"].mean()
asr_without = sub_atk[sub_atk["has_mask"] == 0]["asr"].mean()
diff_pp = abs(asr_with - asr_without) * 100
warn = " ⚠️ confound!" if diff_pp > 5 else ""
print(f" {atk} a ε=8: with_mask ASR={asr_with:.3f} | "
f"without_mask ASR={asr_without:.3f} | "
f"diff={diff_pp:.1f}pp{warn}")
print(f" ✓ {out_path.name}")
# ─── main ─────────────────────────────────────────────────────────────────────
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
parser.add_argument("--csv", type=Path, required=True,
help="CSV merged do probe (output de merge_probe_csvs.py)")
parser.add_argument("--out", type=Path, required=True,
help="Diretório de saída (será criado).")
parser.add_argument("--eps-ref", type=int, default=8,
help="ε de referência pra heatmap + efficiency (default: 8)")
args = parser.parse_args()
if not args.csv.exists():
print(f"ERROR: CSV não encontrado: {args.csv}")
return 1
try:
import pandas as pd
except ImportError:
print("ERROR: pandas necessário")
return 1
args.out.mkdir(parents=True, exist_ok=True)
print(f"Lendo {args.csv} ...")
df = pd.read_csv(args.csv)
if "eps_255" not in df.columns:
df["eps_255"] = (df["epsilon"].astype(float) * 255).round().astype(int)
print(f" {len(df)} rows | "
f"modelos={df['model'].nunique()} | "
f"ataques={sorted(df['attack'].unique())} | "
f"ε={sorted(df['eps_255'].unique())} | "
f"imgs={df['image'].nunique()}")
print(f"\n=== A.2: Curvas mean ± IC95 ===")
plot_lines_eps_metric(df, "asr", "ASR (mean ± IC95)",
"Curva ε × ASR — linhas por ataque",
args.out / "asr_vs_eps_lines.png", ylim=(-0.05, 1.05))
plot_lines_eps_metric(df, "ssim", "SSIM (mean ± IC95)",
"Curva ε × SSIM — linhas por ataque",
args.out / "ssim_vs_eps_lines.png", ylim=(0.4, 1.02))
plot_lines_eps_metric(df, "lpips", "LPIPS (mean ± IC95)",
"Curva ε × LPIPS — linhas por ataque",
args.out / "lpips_vs_eps_lines.png")
plot_lines_eps_metric(df, "psnr", "PSNR dB (mean ± IC95)",
"Curva ε × PSNR — linhas por ataque",
args.out / "psnr_vs_eps_lines.png")
print(f"\n=== A.3: Heatmap modelo × ataque a ε={args.eps_ref}/255 ===")
plot_heatmap_model_attack(
df, metric="asr", eps_ref=args.eps_ref, cmap="Reds", fmt=".2f",
title=f"Mean ASR por modelo × ataque a ε={args.eps_ref}/255",
out_path=args.out / f"heatmap_model_attack_eps{args.eps_ref}_asr.png",
vmin=0, vmax=1,
)
plot_heatmap_model_attack(
df, metric="ssim", eps_ref=args.eps_ref, cmap="Blues", fmt=".3f",
title=f"Mean SSIM por modelo × ataque a ε={args.eps_ref}/255",
out_path=args.out / f"heatmap_model_attack_eps{args.eps_ref}_ssim.png",
vmin=0.5, vmax=1.0,
)
print(f"\n=== A.4: Bar chart de eficiência a ε={args.eps_ref}/255 ===")
plot_efficiency_bar(df, eps_ref=args.eps_ref,
out_path=args.out / f"efficiency_per_attack_eps{args.eps_ref}.png")
print(f"\n=== A.5: ASR ε × ataque facetado por has_mask ===")
plot_asr_by_mask(df, out_path=args.out / "asr_vs_eps_by_mask.png")
print(f"\n✓ Done. {len(list(args.out.glob('*.png')))} figuras em {args.out}")
return 0
if __name__ == "__main__":
sys.exit(main())