"""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)