aeye-backend / aeye_next /models /patchguard.py
wonjun12's picture
Deploy A-EYE app (Expo web) + hybrid backend (model 49 verdict + 63 heatmap)
5dab1e8 verified
Raw
History Blame Contribute Delete
7.23 kB
"""Model 54 "PatchGuard" (v13): decision-linked patch-level detector.
Subclasses the champion ZeroShotV4Detector (model 49) WITHOUT modifying it:
- identical frozen CLIP ViT-L/14 + forensic residual + radial FFT body,
so 49's trainable checkpoint warm-starts every shared module.
- NEW patch head: per-token MLP over [CLIP patch token (1024), upsampled
forensic stage-3 feature (192)] -> one logit per 14px patch (16x16 grid).
- Image decision fuses both streams:
z_img = (logit_fake - logit_real) + gamma * mean(top-k patch logits)
gamma starts at 0, so at warm-start the image decision equals model 49
exactly; training learns how much patch evidence to trust.
The 16x16 sigmoid patch map IS the heatmap: the decision and the explanation
come from the same forward pass and the same features.
"""
from __future__ import annotations
import math
import sys
from pathlib import Path
import torch
import torch.nn as nn
import torch.nn.functional as F
ROOT = Path(__file__).resolve().parent.parent.parent
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
# Space copy: the repo's src/models/zero_shot_v4.py ships as the root-level
# standalone zero_shot_v4.py (identical class; state-dict compatible).
from zero_shot_v4 import ZeroShotV4Detector # noqa: E402 (import only, never modified)
class PatchGuardDetector(ZeroShotV4Detector):
def __init__(
self,
clip_backbone: str = "clip-vit-l-14",
clip_layer: int = 13,
semantic_dim: int = 512,
forensic_dim: int = 256,
frequency_dim: int = 192,
fft_bins: int = 48,
image_size: int = 224,
num_classes: int = 2,
num_sources: int = 2,
dropout: float = 0.25,
source_grl_lambda: float = 0.0,
freeze_clip: bool = True,
patch_hidden: int = 256,
patch_topk_frac: float = 0.25,
):
super().__init__(
clip_backbone=clip_backbone,
clip_layer=clip_layer,
semantic_dim=semantic_dim,
forensic_dim=forensic_dim,
frequency_dim=frequency_dim,
fft_bins=fft_bins,
image_size=image_size,
num_classes=num_classes,
num_sources=num_sources,
dropout=dropout,
source_grl_lambda=source_grl_lambda,
freeze_clip=freeze_clip,
)
if self.is_siglip:
raise ValueError("PatchGuard targets the CLIP champion body (model 49)")
clip_hidden = int(self.clip.config.hidden_size)
forensic_map_dim = 192 # forensic_branch stage3 channel count
self.patch_topk_frac = float(patch_topk_frac)
self.patch_head = nn.Sequential(
nn.LayerNorm(clip_hidden + forensic_map_dim),
nn.Linear(clip_hidden + forensic_map_dim, patch_hidden),
nn.GELU(),
nn.Dropout(p=dropout * 0.5),
nn.Linear(patch_hidden, 1),
)
nn.init.trunc_normal_(self.patch_head[-1].weight, std=0.02)
nn.init.constant_(self.patch_head[-1].bias, -2.0) # start patches near "real"
self.gamma = nn.Parameter(torch.zeros(1))
# peak-patch fusion: lets a single very-confident AI patch raise the image
# score (helps small localized inpaints that the top-k mean dilutes). Init
# 0 -> warm-starting from 54-58 is numerically identical at the start, and
# old checkpoints (no gamma_max key) load with this term disabled.
self.gamma_max = nn.Parameter(torch.zeros(1))
# --- feature extraction (no modification of the parent class) -----------
def _clip_tokens(self, x: torch.Tensor) -> torch.Tensor:
"""Patch tokens of the frozen CLIP at self.clip_layer, shape (B, N, H)."""
context = torch.no_grad() if self.freeze_clip else torch.enable_grad()
with context:
vision = self.clip.vision_model
hidden = vision.embeddings(pixel_values=x)
hidden = vision.pre_layrnorm(hidden)
for idx, layer in enumerate(vision.encoder.layers, start=1):
# Newer transformers CLIP encoder layers require
# causal_attention_mask explicitly even for vision use.
layer_out = layer(hidden, attention_mask=None, causal_attention_mask=None)
hidden = layer_out[0] if isinstance(layer_out, (tuple, list)) else layer_out
if idx >= self.clip_layer:
break
return hidden[:, 1:] # drop CLS
def _forensic_spatial(self, raw: torch.Tensor) -> torch.Tensor:
"""Stage-3 spatial map of the forensic branch, shape (B, 192, h, w)."""
branch = self.forensic_branch
low = F.avg_pool2d(raw, kernel_size=5, stride=1, padding=2)
residual = raw - low
x = torch.cat([residual, residual.abs()], dim=1)
x = branch.stem(x)
x = branch.stage1(x)
x = branch.stage2(x)
return branch.stage3(x)
# --- forward -------------------------------------------------------------
def forward(self, x: torch.Tensor) -> dict[str, torch.Tensor]:
raw = self._to_raw_rgb(x)
tokens = self._clip_tokens(x).float() # (B, N, 1024)
pooled = self.semantic_pool(tokens)
semantic = self.semantic_proj(pooled)
forensic_map = self._forensic_spatial(raw) # (B, 192, 14, 14)
forensic_vec = self.forensic_branch.proj(
self.forensic_branch.pool(forensic_map).flatten(1)
)
frequency = self.frequency_branch(raw)
features = torch.cat([semantic, forensic_vec, frequency], dim=1)
logits = self.head(features) # (B, 2)
grid = int(math.sqrt(tokens.shape[1])) # 16 for ViT-L/14 @224
fmap = F.interpolate(
forensic_map.float(), size=(grid, grid), mode="bilinear", align_corners=False
)
fmap_tokens = fmap.flatten(2).transpose(1, 2) # (B, N, 192)
patch_logits = self.patch_head(torch.cat([tokens, fmap_tokens], dim=-1)).squeeze(-1)
k = max(1, int(round(patch_logits.shape[1] * self.patch_topk_frac)))
patch_summary = patch_logits.topk(k, dim=1).values.mean(dim=1)
patch_peak = patch_logits.max(dim=1).values
z_img = (
(logits[:, 1] - logits[:, 0])
+ self.gamma.squeeze() * patch_summary
+ self.gamma_max.squeeze() * patch_peak
)
return {
"logits": logits,
"z_img": z_img,
"patch_logits": patch_logits.view(-1, grid, grid),
"patch_summary": patch_summary,
"uncertainty_logit": self.uncertainty_head(features).squeeze(1),
"features": features,
}
@torch.no_grad()
def predict_proba(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Returns (P(AI) per image, patch probability map (B, g, g))."""
out = self.forward(x)
return torch.sigmoid(out["z_img"]), torch.sigmoid(out["patch_logits"])
def build_patchguard(cfg: dict) -> PatchGuardDetector:
mcfg = dict(cfg.get("model", {}))
mcfg.pop("type", None)
return PatchGuardDetector(**mcfg)