Spaces:
Running
Running
| """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, | |
| } | |
| 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) | |