| """Perceptual budget: LPIPS penalty + DCT low-pass projection + stealth metrics. |
| |
| The plain attack is effective but can be VISIBLE (PSNR ~26 dB). The perceptual |
| budget trades a little raw strength for genuine imperceptibility via two |
| mechanisms used together: |
| |
| - LPIPS penalty: a differentiable perceptual-distance term subtracted from the |
| adversarial objective, so the optimizer is pushed toward perturbations the |
| human eye can't see (much better aligned with perception than raw L-inf). |
| - DCT low-pass projection: constrain the perturbation to low spatial frequencies. |
| JPEG quantizes high-frequency DCT coefficients hardest, so keeping energy in the |
| low band is what actually survives recompression (complements the EOT JPEG-STE). |
| |
| Plus PSNR / SSIM / CIEDE2000-p95 metric helpers for reporting the stealth axis of |
| the Pareto curve. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import functools |
|
|
| import torch |
| import torch.nn.functional as F |
|
|
|
|
| |
| class LPIPSPenalty: |
| """Thin wrapper around the `lpips` package (AlexNet backbone by default).""" |
|
|
| def __init__(self, net: str = "alex", device: str = "cuda", |
| dtype: torch.dtype = torch.float32): |
| import lpips |
| self.model = lpips.LPIPS(net=net).to(device).eval() |
| for p in self.model.parameters(): |
| p.requires_grad_(False) |
| self.device = device |
| self.dtype = dtype |
|
|
| def distance(self, adv: torch.Tensor, clean: torch.Tensor) -> torch.Tensor: |
| """LPIPS distance; inputs are (B,3,H,W) in [0,1]. Differentiable in `adv`.""" |
| a = (adv.to(self.dtype) * 2 - 1) |
| c = (clean.to(self.dtype) * 2 - 1) |
| return self.model(a, c).mean() |
|
|
|
|
| |
| @functools.lru_cache(maxsize=8) |
| def _dct_matrix(n: int, device: str, dtype_str: str) -> torch.Tensor: |
| """Orthonormal DCT-II basis matrix D (n x n).""" |
| dtype = getattr(torch, dtype_str) |
| k = torch.arange(n, device=device, dtype=torch.float64).view(-1, 1) |
| m = torch.arange(n, device=device, dtype=torch.float64).view(1, -1) |
| d = torch.cos(torch.pi * (2 * m + 1) * k / (2 * n)) |
| d *= torch.sqrt(torch.tensor(2.0 / n, dtype=torch.float64)) |
| d[0] *= torch.sqrt(torch.tensor(0.5, dtype=torch.float64)) |
| return d.to(dtype) |
|
|
|
|
| def _dct2(x: torch.Tensor, dh: torch.Tensor, dw: torch.Tensor) -> torch.Tensor: |
| return torch.einsum("ij,bcjk,lk->bcil", dh, x, dw) |
|
|
|
|
| def _idct2(c: torch.Tensor, dh: torch.Tensor, dw: torch.Tensor) -> torch.Tensor: |
| return torch.einsum("ji,bcjk,kl->bcil", dh, c, dw) |
|
|
|
|
| def dct_lowpass(delta: torch.Tensor, keep_frac: float) -> torch.Tensor: |
| """Project a perturbation onto its lowest-frequency DCT band. |
| |
| keep_frac in (0,1]: fraction of DCT rows/cols (from the DC corner) retained. |
| keep_frac>=1 is a no-op. Differentiable (pure linear ops). |
| """ |
| if keep_frac >= 1.0: |
| return delta |
| _, _, h, w = delta.shape |
| dh = _dct_matrix(h, str(delta.device), "float32").to(delta.dtype) |
| dw = _dct_matrix(w, str(delta.device), "float32").to(delta.dtype) |
| coeff = _dct2(delta.float(), dh.float(), dw.float()) |
| kh, kw = max(1, int(h * keep_frac)), max(1, int(w * keep_frac)) |
| mask = torch.zeros_like(coeff) |
| mask[:, :, :kh, :kw] = 1.0 |
| return _idct2(coeff * mask, dh.float(), dw.float()).to(delta.dtype) |
|
|
|
|
| |
| @torch.no_grad() |
| def psnr(adv: torch.Tensor, clean: torch.Tensor) -> float: |
| mse = F.mse_loss(adv.clamp(0, 1), clean.clamp(0, 1)).item() |
| return 99.0 if mse < 1e-12 else 10.0 * torch.log10(torch.tensor(1.0 / mse)).item() |
|
|
|
|
| @torch.no_grad() |
| def ssim(adv: torch.Tensor, clean: torch.Tensor, win: int = 11) -> float: |
| a = adv.clamp(0, 1).mean(1, keepdim=True) |
| b = clean.clamp(0, 1).mean(1, keepdim=True) |
| c1, c2 = 0.01 ** 2, 0.03 ** 2 |
| pad = win // 2 |
| k = torch.ones(1, 1, win, win, device=a.device) / (win * win) |
| mu_a = F.conv2d(a, k, padding=pad) |
| mu_b = F.conv2d(b, k, padding=pad) |
| va = F.conv2d(a * a, k, padding=pad) - mu_a ** 2 |
| vb = F.conv2d(b * b, k, padding=pad) - mu_b ** 2 |
| vab = F.conv2d(a * b, k, padding=pad) - mu_a * mu_b |
| s = ((2 * mu_a * mu_b + c1) * (2 * vab + c2)) / \ |
| ((mu_a ** 2 + mu_b ** 2 + c1) * (va + vb + c2)) |
| return float(s.mean().item()) |
|
|
|
|
| @torch.no_grad() |
| def _rgb_to_lab(x: torch.Tensor) -> torch.Tensor: |
| |
| m = (x > 0.04045).float() |
| lin = m * ((x + 0.055) / 1.055) ** 2.4 + (1 - m) * (x / 12.92) |
| r, g, b = lin[:, 0], lin[:, 1], lin[:, 2] |
| xw, yw, zw = 0.95047, 1.0, 1.08883 |
| X = (0.4124 * r + 0.3576 * g + 0.1805 * b) / xw |
| Y = (0.2126 * r + 0.7152 * g + 0.0722 * b) / yw |
| Z = (0.0193 * r + 0.1192 * g + 0.9505 * b) / zw |
| xyz = torch.stack([X, Y, Z], 1) |
| d = (xyz > 0.008856).float() |
| f = d * xyz.clamp(min=1e-6) ** (1 / 3) + (1 - d) * (7.787 * xyz + 16 / 116) |
| L = 116 * f[:, 1] - 16 |
| a = 500 * (f[:, 0] - f[:, 1]) |
| bb = 200 * (f[:, 1] - f[:, 2]) |
| return torch.stack([L, a, bb], 1) |
|
|
|
|
| @torch.no_grad() |
| def delta_e_p95(adv: torch.Tensor, clean: torch.Tensor) -> float: |
| """95th-percentile CIE76 color difference (fast proxy for CIEDE2000-p95).""" |
| la = _rgb_to_lab(adv.clamp(0, 1)) |
| lb = _rgb_to_lab(clean.clamp(0, 1)) |
| de = ((la - lb) ** 2).sum(1).sqrt().flatten() |
| return float(torch.quantile(de, 0.95).item()) |
|
|
|
|
| @torch.no_grad() |
| def stealth_metrics(adv: torch.Tensor, clean: torch.Tensor, |
| lpips_fn: "LPIPSPenalty | None" = None) -> dict: |
| out = {"psnr": round(psnr(adv, clean), 2), |
| "ssim": round(ssim(adv, clean), 4), |
| "deltaE_p95": round(delta_e_p95(adv, clean), 3)} |
| if lpips_fn is not None: |
| out["lpips"] = round(float(lpips_fn.distance(adv, clean).item()), 4) |
| return out |
|
|