File size: 6,188 Bytes
5dab1e8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
"""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)