aeye-backend / aeye_next /models /dense_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
6.19 kB
"""Dense localization decoder for A-EYE PatchGuard.
This model keeps the image-level body and the coarse PatchGuard head compatible
with models 54-61, then adds a small convolutional decoder that refines the
16x16 patch logits into a denser mask (default 32x32). The dense head is trained
directly against mask-derived targets, so localization can improve without
throwing away the stable image decision learned by model 49/60/61.
"""
from __future__ import annotations
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from aeye_next.models.patchguard import PatchGuardDetector
class _ConvBlock(nn.Module):
def __init__(self, channels: int, dropout: float = 0.0) -> None:
super().__init__()
groups = max(1, min(8, channels // 16))
self.net = nn.Sequential(
nn.Conv2d(channels, channels, kernel_size=3, padding=1, bias=False),
nn.GroupNorm(groups, channels),
nn.GELU(),
nn.Dropout2d(dropout),
nn.Conv2d(channels, channels, kernel_size=3, padding=1, bias=False),
nn.GroupNorm(groups, channels),
)
self.act = nn.GELU()
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.act(x + self.net(x))
class DensePatchGuardDetector(PatchGuardDetector):
def __init__(
self,
*args,
decoder_grid: int = 32,
decoder_channels: int = 192,
decoder_blocks: int = 3,
decoder_dropout: float = 0.04,
**kwargs,
) -> None:
super().__init__(*args, **kwargs)
if decoder_grid < 16 or decoder_grid % 16 != 0:
raise ValueError("decoder_grid must be a multiple of 16 and >= 16")
self.decoder_grid = int(decoder_grid)
self.decoder_channels = int(decoder_channels)
clip_hidden = int(self.clip.config.hidden_size)
forensic_map_dim = 192
self.clip_map_proj = nn.Sequential(
nn.LayerNorm(clip_hidden),
nn.Linear(clip_hidden, decoder_channels),
)
self.forensic_map_proj = nn.Conv2d(forensic_map_dim, decoder_channels, kernel_size=1)
self.dense_decoder = nn.Sequential(
*[_ConvBlock(decoder_channels, dropout=decoder_dropout) for _ in range(int(decoder_blocks))]
)
self.mask_delta_head = nn.Sequential(
nn.Conv2d(decoder_channels, decoder_channels // 2, kernel_size=3, padding=1),
nn.GELU(),
nn.Conv2d(decoder_channels // 2, 1, kernel_size=1),
)
# Start as "model 61 upsampled": the new decoder initially contributes
# almost nothing, then learns boundary/detail corrections.
nn.init.zeros_(self.mask_delta_head[-1].weight)
nn.init.zeros_(self.mask_delta_head[-1].bias)
def _dense_logits(
self,
tokens: torch.Tensor,
forensic_map: torch.Tensor,
coarse_logits: torch.Tensor,
grid: int,
) -> torch.Tensor:
bsz = tokens.shape[0]
clip_map = self.clip_map_proj(tokens).transpose(1, 2).reshape(
bsz, self.decoder_channels, grid, grid
)
fmap = F.interpolate(
forensic_map.float(), size=(grid, grid), mode="bilinear", align_corners=False
)
fused = clip_map + self.forensic_map_proj(fmap)
if self.decoder_grid != grid:
fused = F.interpolate(
fused, size=(self.decoder_grid, self.decoder_grid),
mode="bilinear", align_corners=False,
)
decoded = self.dense_decoder(fused)
delta = self.mask_delta_head(decoded).squeeze(1)
coarse = F.interpolate(
coarse_logits.unsqueeze(1),
size=(self.decoder_grid, self.decoder_grid),
mode="bilinear",
align_corners=False,
).squeeze(1)
return coarse + delta
def forward(self, x: torch.Tensor) -> dict[str, torch.Tensor]:
raw = self._to_raw_rgb(x)
tokens = self._clip_tokens(x).float()
pooled = self.semantic_pool(tokens)
semantic = self.semantic_proj(pooled)
forensic_map = self._forensic_spatial(raw)
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)
grid = int(math.sqrt(tokens.shape[1]))
fmap = F.interpolate(
forensic_map.float(), size=(grid, grid), mode="bilinear", align_corners=False
)
fmap_tokens = fmap.flatten(2).transpose(1, 2)
coarse_logits = self.patch_head(torch.cat([tokens, fmap_tokens], dim=-1)).squeeze(-1)
coarse_logits = coarse_logits.view(-1, grid, grid)
patch_logits = self._dense_logits(tokens, forensic_map, coarse_logits, grid)
flat = patch_logits.flatten(1)
k = max(1, int(round(flat.shape[1] * self.patch_topk_frac)))
patch_summary = flat.topk(k, dim=1).values.mean(dim=1)
patch_peak = flat.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,
"coarse_patch_logits": coarse_logits,
"patch_summary": patch_summary,
"uncertainty_logit": self.uncertainty_head(features).squeeze(1),
"features": features,
}
def param_summary(self) -> str:
base = super().param_summary()
dense = sum(
p.numel()
for name, p in self.named_parameters()
if name.startswith(("clip_map_proj.", "forensic_map_proj.", "dense_decoder.", "mask_delta_head."))
)
return f"{base} dense_grid={self.decoder_grid} dense_params={dense:,}"
def build_dense_patchguard(cfg: dict) -> DensePatchGuardDetector:
mcfg = dict(cfg.get("model", {}))
mcfg.pop("type", None)
return DensePatchGuardDetector(**mcfg)