"""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)