RefDiffNet / a11_ca.py
ShashwatGupta23's picture
Cursor
RefDiffNet: standalone A11_CA prebackbone Gradio Space
5673379
Raw
History Blame Contribute Delete
7.38 kB
"""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)