import torch from torch import nn import numpy as np import cv2 import timm import torch.nn.functional as F def compute_CD_heatmap(img_bgr: np.ndarray, quality_factors=(30, 50, 70), smooth_ksize: int = 15) -> np.ndarray: residuals = [] for q in quality_factors: _, enc = cv2.imencode(".jpg", img_bgr, [int(cv2.IMWRITE_JPEG_QUALITY), q]) dec = cv2.imdecode(enc, cv2.IMREAD_COLOR) diff = (img_bgr.astype(np.float32) - dec.astype(np.float32)) ** 2 residuals.append(diff.mean(axis=2)) cd_map = np.mean(residuals, axis=0) cd_map = cv2.GaussianBlur(cd_map, (smooth_ksize, smooth_ksize), 0) return cd_map.astype(np.float32) def cd_to_tensor(cd_np: np.ndarray, device: torch.device) -> torch.Tensor: return torch.from_numpy(cd_np).unsqueeze(0).unsqueeze(0).to(device) class _Reassemble(nn.Module): def __init__(self, in_dim: int, out_dim: int, hw: int, scale_mode: str = "identity"): super().__init__() self.hw = hw self.proj = nn.Linear(in_dim, out_dim) if scale_mode == "up4": self.resample = nn.ConvTranspose2d(out_dim, out_dim, kernel_size=4, stride=4) elif scale_mode == "up2": self.resample = nn.ConvTranspose2d(out_dim, out_dim, kernel_size=2, stride=2) elif scale_mode == "down2": self.resample = nn.Conv2d(out_dim, out_dim, kernel_size=3, stride=2, padding=1) else: self.resample = nn.Identity() def forward(self, tokens: torch.Tensor) -> torch.Tensor: B = tokens.shape[0] x = self.proj(tokens) x = x.permute(0, 2, 1).reshape(B, -1, self.hw, self.hw) return self.resample(x) class _ResConvUnit(nn.Module): def __init__(self, features: int, use_bn: bool = True): super().__init__() bias = not use_bn self.block = nn.Sequential( nn.ReLU(inplace=False), nn.Conv2d(features, features, 3, padding=1, bias=bias), nn.BatchNorm2d(features) if use_bn else nn.Identity(), nn.ReLU(inplace=False), nn.Conv2d(features, features, 3, padding=1, bias=bias), nn.BatchNorm2d(features) if use_bn else nn.Identity(), ) def forward(self, x): return x + self.block(x) class _FusionBlock(nn.Module): def __init__(self, features: int, use_bn: bool = True): super().__init__() self.resunit_skip = _ResConvUnit(features, use_bn) self.resunit_out = _ResConvUnit(features, use_bn) self.out_conv = nn.Conv2d(features, features, kernel_size=1, bias=True) def forward(self, x: torch.Tensor, skip: torch.Tensor = None) -> torch.Tensor: if skip is not None: if x.shape[-2:] != skip.shape[-2:]: x = F.interpolate(x, size=skip.shape[-2:], mode="bilinear", align_corners=True) x = x + self.resunit_skip(skip) x = self.resunit_out(x) x = F.interpolate(x, scale_factor=2.0, mode="bilinear", align_corners=True) return self.out_conv(x) class DINOv2DPT(nn.Module): """ DINOv2 ViT-B/14 + DPT multi-scale fusion decoder + CD injection. Flow: img → DINOv2 (hooks@[2,5,8,11]) → Reassemble×4 → DPT Fusion → + CD feature (gated) → SegHead → bilinear upsample → logits """ _DMODEL = 768 def __init__(self, img_size=518, features=256, use_bn=True, hook_indices=(2, 5, 8, 11), unfreeze_blocks=3): super().__init__() self.img_size = img_size self.patch_size = 14 assert img_size % self.patch_size == 0 self.hw = img_size // self.patch_size self.hook_indices = list(hook_indices) D = self._DMODEL # ── Backbone ────────────────────────────────────────────────── self.backbone = timm.create_model( "vit_base_patch14_dinov2.lvd142m", pretrained=True, num_classes=0, global_pool="", img_size=img_size, ) for p in self.backbone.parameters(): p.requires_grad = False for blk in self.backbone.blocks[-unfreeze_blocks:]: for p in blk.parameters(): p.requires_grad = True if hasattr(self.backbone, "norm"): for p in self.backbone.norm.parameters(): p.requires_grad = True # ── Feature hooks ───────────────────────────────────────────── self._feat_store: dict[int, torch.Tensor] = {} self._hook_handles = [] for slot, blk_idx in enumerate(self.hook_indices): def _make_hook(s): def _hook(module, inp, out): self._feat_store[s] = out[:, 1:, :] # strip CLS return _hook self._hook_handles.append( self.backbone.blocks[blk_idx].register_forward_hook(_make_hook(slot)) ) # ── Reassemble (4 scales) ────────────────────────────────────── scale_modes = ["up4", "up2", "identity", "down2"] self.reassemble = nn.ModuleList([ _Reassemble(D, features, self.hw, scale_modes[i]) for i in range(4) ]) self.layer_rn = nn.ModuleList([ nn.Conv2d(features, features, kernel_size=1, bias=False) for _ in range(4) ]) # ── DPT Fusion (deep → shallow) ──────────────────────────────── self.fusion4 = _FusionBlock(features, use_bn) self.fusion3 = _FusionBlock(features, use_bn) self.fusion2 = _FusionBlock(features, use_bn) self.fusion1 = _FusionBlock(features, use_bn) # ── CD injection ─────────────────────────────────────────────── self.cd_embed = nn.Sequential( nn.Conv2d(1, 32, kernel_size=3, padding=1), nn.GELU(), nn.Conv2d(32, features, kernel_size=1), ) self.cd_gate = nn.Parameter(torch.zeros(1)) # ── Segmentation head ───────────────────────────────────────── self.head = nn.Sequential( nn.Conv2d(features, features // 2, 3, padding=1, bias=not use_bn), nn.BatchNorm2d(features // 2) if use_bn else nn.Identity(), nn.GELU(), nn.Dropout2d(0.1), nn.Conv2d(features // 2, 32, 3, padding=1), nn.GELU(), nn.Conv2d(32, 1, 1), ) def forward(self, x: torch.Tensor, cd_map: torch.Tensor = None) -> torch.Tensor: B = x.shape[0] self._feat_store.clear() _ = self.backbone.forward_features(x) feats = [] for i in range(4): spatial = self.reassemble[i](self._feat_store[i]) feats.append(self.layer_rn[i](spatial)) l1, l2, l3, l4 = feats path = self.fusion4(l4) path = self.fusion3(path, l3) path = self.fusion2(path, l2) path = self.fusion1(path, l1) if cd_map is not None: cd_feat = self.cd_embed( F.interpolate(cd_map.float(), size=path.shape[-2:], mode="bilinear", align_corners=False) ) path = path + torch.tanh(self.cd_gate) * cd_feat out = self.head(path) return F.interpolate(out, size=(self.img_size, self.img_size), mode="bilinear", align_corners=False) def remove_hooks(self): for h in self._hook_handles: h.remove() self._hook_handles.clear()