aeye-backend / zero_shot_v4.py
wonjun12's picture
Deploy A-EYE app (Expo web) + hybrid backend (model 49 verdict + 63 heatmap)
5dab1e8 verified
Raw
History Blame Contribute Delete
14.7 kB
"""Model 49 (ZeroShotV4Detector) — self-contained copy for the web backend.
Ported verbatim from the model repo's ``src/models/zero_shot_v4.py``. The only
change is that ``CLIP_MEAN`` / ``CLIP_STD`` are inlined here instead of importing
``src.data.transforms`` (which is not part of the web backend).
- frozen CLIP ViT-L/14 intermediate patch tokens (semantic/texture cues)
- trainable forensic residual CNN (sensor/compression/noise evidence)
- trainable radial FFT branch (frequency statistics)
The frozen CLIP backbone is NOT stored in the checkpoint; it is downloaded from
Hugging Face on first construction.
"""
from __future__ import annotations
from collections import OrderedDict
from typing import Iterable
import torch
import torch.nn as nn
import torch.nn.functional as F
# Inlined from src/data/transforms.py (CLIP normalization constants).
CLIP_MEAN = [0.48145466, 0.4578275, 0.40821073]
CLIP_STD = [0.26862954, 0.26130258, 0.27577711]
try:
from transformers import CLIPVisionModelWithProjection
except ImportError: # pragma: no cover
CLIPVisionModelWithProjection = None
try:
from transformers import SiglipVisionModel
except ImportError: # pragma: no cover
SiglipVisionModel = None
CLIP_BACKBONES = {
"clip-vit-b-32": "openai/clip-vit-base-patch32",
"clip-vit-b-16": "openai/clip-vit-base-patch16",
"clip-vit-l-14": "openai/clip-vit-large-patch14",
}
# Stronger frozen backbones (SigLIP: no CLS token, no pre_layrnorm, 0.5/0.5 norm)
SIGLIP_BACKBONES = {
"siglip2-large-256": "google/siglip2-large-patch16-256",
"siglip-large-256": "google/siglip-large-patch16-256",
"siglip-so400m-384": "google/siglip-so400m-patch14-384",
}
class GradientReverseFn(torch.autograd.Function):
@staticmethod
def forward(ctx, x: torch.Tensor, strength: float) -> torch.Tensor:
ctx.strength = float(strength)
return x.view_as(x)
@staticmethod
def backward(ctx, grad_output: torch.Tensor):
return -ctx.strength * grad_output, None
def gradient_reverse(x: torch.Tensor, strength: float = 1.0) -> torch.Tensor:
return GradientReverseFn.apply(x, strength)
class AttentionPool(nn.Module):
def __init__(self, dim: int):
super().__init__()
self.score = nn.Linear(dim, 1)
nn.init.trunc_normal_(self.score.weight, std=0.02)
nn.init.zeros_(self.score.bias)
def forward(self, tokens: torch.Tensor) -> torch.Tensor:
weights = torch.softmax(self.score(tokens).squeeze(-1), dim=-1)
return torch.sum(tokens * weights.unsqueeze(-1), dim=1)
class ConvNeXtMiniBlock(nn.Module):
def __init__(self, dim: int, drop: float = 0.0):
super().__init__()
self.dwconv = nn.Conv2d(dim, dim, kernel_size=7, padding=3, groups=dim)
self.norm = nn.GroupNorm(1, dim)
self.pw1 = nn.Conv2d(dim, dim * 4, kernel_size=1)
self.act = nn.GELU()
self.pw2 = nn.Conv2d(dim * 4, dim, kernel_size=1)
self.drop = nn.Dropout2d(drop) if drop > 0 else nn.Identity()
def forward(self, x: torch.Tensor) -> torch.Tensor:
residual = x
x = self.dwconv(x)
x = self.norm(x)
x = self.pw1(x)
x = self.act(x)
x = self.drop(x)
x = self.pw2(x)
return residual + x
class ForensicResidualBranch(nn.Module):
def __init__(self, out_dim: int = 256, dropout: float = 0.12):
super().__init__()
self.stem = nn.Sequential(
nn.Conv2d(6, 48, kernel_size=5, stride=2, padding=2, bias=False),
nn.BatchNorm2d(48),
nn.GELU(),
)
self.stage1 = nn.Sequential(
ConvNeXtMiniBlock(48, drop=dropout * 0.25),
nn.Conv2d(48, 96, kernel_size=3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(96),
nn.GELU(),
)
self.stage2 = nn.Sequential(
ConvNeXtMiniBlock(96, drop=dropout * 0.35),
nn.Conv2d(96, 160, kernel_size=3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(160),
nn.GELU(),
)
self.stage3 = nn.Sequential(
ConvNeXtMiniBlock(160, drop=dropout * 0.5),
nn.Conv2d(160, 192, kernel_size=3, stride=2, padding=1, bias=False),
nn.BatchNorm2d(192),
nn.GELU(),
ConvNeXtMiniBlock(192, drop=dropout * 0.5),
)
self.pool = nn.AdaptiveAvgPool2d(1)
self.proj = nn.Sequential(
nn.Dropout(p=dropout),
nn.Linear(192, out_dim),
nn.LayerNorm(out_dim),
)
self._init_weights()
def _init_weights(self) -> None:
for module in self.modules():
if isinstance(module, nn.Conv2d):
nn.init.kaiming_normal_(module.weight, mode="fan_out", nonlinearity="relu")
elif isinstance(module, (nn.BatchNorm2d, nn.GroupNorm)):
nn.init.ones_(module.weight)
nn.init.zeros_(module.bias)
elif isinstance(module, nn.Linear):
nn.init.trunc_normal_(module.weight, std=0.02)
nn.init.zeros_(module.bias)
def forward(self, raw: torch.Tensor) -> torch.Tensor:
low = F.avg_pool2d(raw, kernel_size=5, stride=1, padding=2)
residual = raw - low
x = torch.cat([residual, residual.abs()], dim=1)
x = self.stem(x)
x = self.stage1(x)
x = self.stage2(x)
x = self.stage3(x)
return self.proj(self.pool(x).flatten(1))
class RadialFFTBranch(nn.Module):
def __init__(
self,
image_size: int = 224,
bins: int = 48,
out_dim: int = 192,
dropout: float = 0.12,
):
super().__init__()
self.image_size = int(image_size)
self.bins = int(bins)
masks = self._make_radial_masks(self.image_size, self.bins)
self.register_buffer("masks", masks)
in_dim = 3 * self.bins + 3
self.mlp = nn.Sequential(
nn.LayerNorm(in_dim),
nn.Linear(in_dim, max(256, out_dim * 2)),
nn.GELU(),
nn.Dropout(p=dropout),
nn.Linear(max(256, out_dim * 2), out_dim),
nn.LayerNorm(out_dim),
)
for module in self.modules():
if isinstance(module, nn.Linear):
nn.init.trunc_normal_(module.weight, std=0.02)
nn.init.zeros_(module.bias)
@staticmethod
def _make_radial_masks(size: int, bins: int) -> torch.Tensor:
axis = torch.linspace(-1.0, 1.0, size)
yy, xx = torch.meshgrid(axis, axis, indexing="ij")
rr = torch.sqrt(xx.square() + yy.square()).clamp(max=1.0)
edges = torch.linspace(0.0, 1.0, bins + 1)
masks = []
for idx in range(bins):
mask = ((rr >= edges[idx]) & (rr < edges[idx + 1])).float()
denom = mask.sum().clamp_min(1.0)
masks.append(mask / denom)
return torch.stack(masks, dim=0)
def forward(self, raw: torch.Tensor) -> torch.Tensor:
freq = torch.fft.fftshift(torch.fft.fft2(raw, norm="ortho"), dim=(-2, -1))
mag = torch.log1p(torch.abs(freq))
radial = torch.einsum("bchw,nhw->bcn", mag, self.masks)
radial = radial.flatten(1)
h = max(1, self.bins // 4)
high = radial.view(raw.shape[0], 3, self.bins)[:, :, -h:].mean(dim=-1)
low = radial.view(raw.shape[0], 3, self.bins)[:, :, :h].mean(dim=-1).clamp_min(1e-6)
ratio = torch.log1p(high / low)
return self.mlp(torch.cat([radial, ratio], dim=1))
def _deep_head(in_dim: int, out_dim: int, dropout: float) -> nn.Sequential:
hidden = max(512, in_dim // 2)
return nn.Sequential(
nn.LayerNorm(in_dim),
nn.Linear(in_dim, hidden),
nn.GELU(),
nn.Dropout(p=dropout),
nn.Linear(hidden, max(256, hidden // 2)),
nn.GELU(),
nn.Dropout(p=dropout * 0.75),
nn.Linear(max(256, hidden // 2), out_dim),
)
class ZeroShotV4Detector(nn.Module):
def __init__(
self,
clip_backbone: str = "clip-vit-l-14",
clip_layer: int = 18,
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 = 1.0,
freeze_clip: bool = True,
):
super().__init__()
self.clip_backbone = clip_backbone
self.is_siglip = clip_backbone in SIGLIP_BACKBONES
if self.is_siglip:
if SiglipVisionModel is None:
raise ImportError("transformers SiglipVisionModel is required for siglip backbones")
self.hf_name = SIGLIP_BACKBONES[clip_backbone]
elif clip_backbone in CLIP_BACKBONES:
if CLIPVisionModelWithProjection is None:
raise ImportError("transformers is required for ZeroShotV4Detector")
self.hf_name = CLIP_BACKBONES[clip_backbone]
else:
raise ValueError(f"Unsupported backbone: {clip_backbone}")
self.clip_layer = int(clip_layer)
self.image_size = int(image_size)
self.num_classes = int(num_classes)
self.num_sources = int(num_sources)
self.source_grl_lambda = float(source_grl_lambda)
self.freeze_clip = bool(freeze_clip)
if self.is_siglip:
self.clip = SiglipVisionModel.from_pretrained(self.hf_name)
else:
self.clip = CLIPVisionModelWithProjection.from_pretrained(self.hf_name)
hidden = int(self.clip.config.hidden_size)
layers = int(self.clip.config.num_hidden_layers)
self.clip_layer = max(1, min(self.clip_layer, layers))
if self.freeze_clip:
for param in self.clip.parameters():
param.requires_grad_(False)
self.semantic_pool = AttentionPool(hidden)
self.semantic_proj = nn.Sequential(
nn.LayerNorm(hidden),
nn.Linear(hidden, semantic_dim),
nn.GELU(),
nn.Dropout(p=dropout * 0.5),
nn.LayerNorm(semantic_dim),
)
self.forensic_branch = ForensicResidualBranch(out_dim=forensic_dim, dropout=dropout * 0.5)
self.frequency_branch = RadialFFTBranch(
image_size=image_size,
bins=fft_bins,
out_dim=frequency_dim,
dropout=dropout * 0.5,
)
self.register_buffer("clip_mean", torch.tensor(CLIP_MEAN).view(1, 3, 1, 1))
self.register_buffer("clip_std", torch.tensor(CLIP_STD).view(1, 3, 1, 1))
fused_dim = semantic_dim + forensic_dim + frequency_dim
self.fused_dim = fused_dim
self.head = _deep_head(fused_dim, num_classes, dropout)
self.source_head = _deep_head(fused_dim, self.num_sources, dropout * 0.75)
self.uncertainty_head = nn.Sequential(
nn.LayerNorm(fused_dim),
nn.Linear(fused_dim, 1),
)
self.contrastive_proj = nn.Sequential(
nn.LayerNorm(fused_dim),
nn.Linear(fused_dim, 256),
nn.GELU(),
nn.Linear(256, 128),
)
self._init_output_layers()
def _init_output_layers(self) -> None:
for module in (self.head[-1], self.source_head[-1], self.uncertainty_head[-1]):
if isinstance(module, nn.Linear):
nn.init.trunc_normal_(module.weight, std=0.02)
nn.init.zeros_(module.bias)
def _to_raw_rgb(self, x: torch.Tensor) -> torch.Tensor:
return (x * self.clip_std + self.clip_mean).clamp(0.0, 1.0)
def semantic_features(self, x: torch.Tensor) -> torch.Tensor:
context = torch.no_grad() if self.freeze_clip else torch.enable_grad()
with context:
vision = self.clip.vision_model
if self.is_siglip:
# SigLIP wants 0.5/0.5 normalization, has no pre_layrnorm and no CLS token
inp = (self._to_raw_rgb(x) - 0.5) / 0.5
hidden = vision.embeddings(pixel_values=inp)
else:
hidden = vision.embeddings(pixel_values=x)
hidden = vision.pre_layrnorm(hidden)
for idx, layer in enumerate(vision.encoder.layers, start=1):
# transformers>=4.5x requires causal_attention_mask positionally;
# the CLIP vision path uses neither mask (both None) -> identical math.
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
tokens = hidden if self.is_siglip else hidden[:, 1:]
pooled = self.semantic_pool(tokens.float())
return self.semantic_proj(pooled)
def forward_features(self, x: torch.Tensor) -> torch.Tensor:
raw = self._to_raw_rgb(x)
semantic = self.semantic_features(x)
forensic = self.forensic_branch(raw)
frequency = self.frequency_branch(raw)
return torch.cat([semantic, forensic, frequency], dim=1)
def forward(self, x: torch.Tensor) -> dict[str, torch.Tensor]:
features = self.forward_features(x)
source_features = gradient_reverse(features, self.source_grl_lambda)
return {
"logits": self.head(features),
"source_logits": self.source_head(source_features),
"uncertainty_logit": self.uncertainty_head(features).squeeze(1),
"contrastive_features": F.normalize(self.contrastive_proj(features), dim=-1),
"features": features,
}
def trainable_state_dict(self) -> "OrderedDict[str, torch.Tensor]":
trainable_names = {
name for name, param in self.named_parameters()
if param.requires_grad and not name.startswith("clip.")
}
return OrderedDict(
(name, value.detach().cpu())
for name, value in self.state_dict().items()
if name in trainable_names or not name.startswith("clip.")
)
def load_trainable_state_dict(self, state: dict[str, torch.Tensor]) -> None:
current = self.state_dict()
matched = {
name: value for name, value in state.items()
if name in current and current[name].shape == value.shape
}
current.update(matched)
self.load_state_dict(current, strict=False)
def trainable_parameters(self) -> Iterable[nn.Parameter]:
return (param for param in self.parameters() if param.requires_grad)