elliotalderson000's picture
clean hf deploy
133f2f7
Raw
History Blame Contribute Delete
8 kB
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()