Spaces:
Sleeping
Sleeping
| """Standalone A11_CA prebackbone (defect + golden reference -> enriched image).""" | |
| from __future__ import annotations | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| __all__ = [ | |
| "ConvBNAct", | |
| "LocalContrastNorm", | |
| "FixedHaarBands", | |
| "EncoderAttentionCoordStem2d", | |
| "A11CoordinateEncoderAttentionPreBackbone", | |
| "build_prebackbone", | |
| ] | |
| class ConvBNAct(nn.Module): | |
| def __init__( | |
| self, | |
| c1: int, | |
| c2: int, | |
| k: int = 3, | |
| s: int = 1, | |
| p: int | None = None, | |
| groups: int = 1, | |
| act: bool = True, | |
| ): | |
| super().__init__() | |
| if p is None: | |
| p = k // 2 | |
| self.conv = nn.Conv2d(c1, c2, k, s, p, groups=groups, bias=False) | |
| self.bn = nn.BatchNorm2d(c2) | |
| self.act = nn.SiLU(inplace=True) if act else nn.Identity() | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.act(self.bn(self.conv(x))) | |
| class LocalContrastNorm(nn.Module): | |
| """Lightweight no-parameter local contrast normalization.""" | |
| def __init__(self, kernel_size: int = 7, eps: float = 1e-4): | |
| super().__init__() | |
| self.kernel_size = kernel_size | |
| self.eps = eps | |
| self.pad = kernel_size // 2 | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| mean = F.avg_pool2d(x, self.kernel_size, stride=1, padding=self.pad) | |
| var = F.avg_pool2d((x - mean) ** 2, self.kernel_size, stride=1, padding=self.pad) | |
| return (x - mean) / torch.sqrt(var + self.eps) | |
| class FixedHaarBands(nn.Module): | |
| """Fixed Haar wavelet decomposition at 1/2 resolution (LL, LH, HL, HH per channel).""" | |
| def __init__(self, channels: int = 3): | |
| super().__init__() | |
| self.channels = channels | |
| ll = torch.tensor([[1, 1], [1, 1]], dtype=torch.float32) / 2.0 | |
| lh = torch.tensor([[-1, -1], [1, 1]], dtype=torch.float32) / 2.0 | |
| hl = torch.tensor([[-1, 1], [-1, 1]], dtype=torch.float32) / 2.0 | |
| hh = torch.tensor([[1, -1], [-1, 1]], dtype=torch.float32) / 2.0 | |
| weight = torch.stack([ll, lh, hl, hh], dim=0).view(4, 1, 2, 2) | |
| weight = weight.repeat(channels, 1, 1, 1) | |
| self.register_buffer("weight", weight) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return F.conv2d(x, self.weight, stride=2, padding=0, groups=self.channels) | |
| class EncoderAttentionCoordStem2d(nn.Module): | |
| """H/W pooled coordinate modulation; returns feat * attn_h * attn_w.""" | |
| def __init__(self, hidden: int) -> None: | |
| super().__init__() | |
| ch = max(hidden // 8, 8) | |
| self.pool_h = nn.AdaptiveAvgPool2d((None, 1)) | |
| self.pool_w = nn.AdaptiveAvgPool2d((1, None)) | |
| self.conv1 = nn.Conv2d(hidden, ch, kernel_size=1, bias=False) | |
| self.bn1 = nn.BatchNorm2d(ch) | |
| self.act = nn.SiLU(inplace=True) | |
| self.conv_h = nn.Conv2d(ch, hidden, kernel_size=1, bias=True) | |
| self.conv_w = nn.Conv2d(ch, hidden, kernel_size=1, bias=True) | |
| def forward(self, feat: torch.Tensor) -> torch.Tensor: | |
| _, _, h, w = feat.shape | |
| xh = self.pool_h(feat) | |
| xw = self.pool_w(feat).permute(0, 1, 3, 2) | |
| coord = torch.cat([xh, xw], dim=2) | |
| coord = self.act(self.bn1(self.conv1(coord))) | |
| ah, aw = torch.split(coord, [h, w], dim=2) | |
| aw = aw.permute(0, 1, 3, 2) | |
| mh = torch.sigmoid(self.conv_h(ah)) | |
| mw = torch.sigmoid(self.conv_w(aw)) | |
| return feat * mh * mw | |
| class A11CoordinateEncoderAttentionPreBackbone(nn.Module): | |
| """ | |
| A11_CA: defect + golden -> enriched = defect + alpha * gate * delta. | |
| Cues: Haar bands, signed low-res residual, morphology; encoder + channel/spatial gates. | |
| """ | |
| def __init__( | |
| self, | |
| channels: int = 3, | |
| hidden: int = 24, | |
| use_lcn: bool = True, | |
| alpha_init: float = 0.08, | |
| ): | |
| super().__init__() | |
| self.channels = channels | |
| self.hidden = hidden | |
| self.lcn = LocalContrastNorm(kernel_size=7) if use_lcn else nn.Identity() | |
| self.haar = FixedHaarBands(channels=channels) | |
| in_ch = channels * 12 | |
| self.encoder = nn.Sequential( | |
| ConvBNAct(in_ch, hidden, k=1, s=1), | |
| ConvBNAct(hidden, hidden, k=3, s=1, groups=hidden), | |
| ConvBNAct(hidden, hidden, k=1, s=1), | |
| ConvBNAct(hidden, hidden, k=3, s=1, groups=hidden), | |
| ConvBNAct(hidden, hidden, k=1, s=1), | |
| ) | |
| gate_hidden = max(hidden // 8, 4) | |
| self.channel_gate = nn.Sequential( | |
| nn.AdaptiveAvgPool2d(1), | |
| nn.Conv2d(hidden, gate_hidden, kernel_size=1, bias=True), | |
| nn.SiLU(inplace=True), | |
| nn.Conv2d(gate_hidden, hidden, kernel_size=1, bias=True), | |
| nn.Sigmoid(), | |
| ) | |
| self.rgb_delta = nn.Sequential( | |
| ConvBNAct(hidden, hidden, k=1, s=1), | |
| nn.Conv2d(hidden, channels, kernel_size=1, bias=True), | |
| nn.Tanh(), | |
| ) | |
| self.spatial_gate_replacement = nn.Sequential( | |
| EncoderAttentionCoordStem2d(hidden), | |
| nn.Conv2d(hidden, 1, kernel_size=1, bias=True), | |
| nn.Sigmoid(), | |
| ) | |
| self.alpha = nn.Parameter(torch.tensor(float(alpha_init))) | |
| def forward(self, defect: torch.Tensor, golden: torch.Tensor) -> torch.Tensor: | |
| if defect.shape != golden.shape: | |
| raise ValueError( | |
| f"A11_CA expects same shape for defect and golden tensors, " | |
| f"got {tuple(defect.shape)} vs {tuple(golden.shape)}" | |
| ) | |
| defect_n = self.lcn(defect) | |
| golden_n = self.lcn(golden) | |
| bd = self.haar(defect_n) | |
| bg = self.haar(golden_n) | |
| defect_lr = F.avg_pool2d(defect_n, kernel_size=2, stride=2) | |
| golden_lr = F.avg_pool2d(golden_n, kernel_size=2, stride=2) | |
| signed_lr = defect_lr - golden_lr | |
| pos_lr = F.relu(signed_lr) | |
| neg_lr = F.relu(-signed_lr) | |
| morph_pos = F.max_pool2d(pos_lr, kernel_size=3, stride=1, padding=1) | |
| morph_neg = F.max_pool2d(neg_lr, kernel_size=3, stride=1, padding=1) | |
| x = torch.cat([bd, bg, pos_lr, neg_lr, morph_pos, morph_neg], dim=1) | |
| feat = self.encoder(x) | |
| feat = feat * self.channel_gate(feat) | |
| delta_lr = self.rgb_delta(feat) | |
| gate_lr = self.spatial_gate_replacement(feat) | |
| gate = F.interpolate(gate_lr, size=defect.shape[2:], mode="bilinear", align_corners=False) | |
| delta = F.interpolate(delta_lr, size=defect.shape[2:], mode="bilinear", align_corners=False) | |
| enriched = defect + self.alpha * gate * delta | |
| self._debug = { | |
| "defect": defect.detach(), | |
| "golden": golden.detach(), | |
| "gate": gate.detach(), | |
| "delta": delta.detach(), | |
| "alpha": float(self.alpha.detach().item()), | |
| "enriched": enriched.detach(), | |
| } | |
| return enriched | |
| _REGISTRY: dict[str, type[nn.Module]] = { | |
| "A11_CA": A11CoordinateEncoderAttentionPreBackbone, | |
| } | |
| def build_prebackbone(name: str | None, channels: int = 3, **kwargs) -> nn.Module | None: | |
| if not name: | |
| return None | |
| key = str(name).upper() | |
| if key not in _REGISTRY: | |
| supported = ", ".join(sorted(_REGISTRY)) | |
| raise ValueError(f"Unsupported prebackbone '{name}'. Supported: {supported}") | |
| return _REGISTRY[key](channels=channels, **kwargs) | |