Spaces:
Runtime error
Runtime error
| 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() | |