WaveSeg-demo / waveseg_model.py
Sarvarbek13's picture
Deploy WaveSeg demo
e6b6a82
Raw
History Blame Contribute Delete
11.4 kB
"""WaveSeg (A1) model definition - a frozen, self-contained copy for this
Hugging Face Space, so app/ can be deployed on its own without the training
repo's src/ tree as a dependency.
This is a verbatim copy of the classes in the training repo's
src/models/{fba,segformer_base,waveseg}.py, unmodified, to guarantee zero
drift from the exact architecture the checkpoints were trained with. The
deployed configuration is always: HaarDWT + gating ON + deep placement
(A1 = WaveSeg, ours) - see MODEL_CARD.md for why the boundary-frequency loss
and the other ablation axes (FFT, shallow placement, no-gate) are not part
of the shipped model. Those extra axes (FFTHighPass, use_gate=False,
placement="shallow") are kept here anyway rather than stripped out, since a
verbatim copy is lower-risk than a hand-simplified one.
"""
from __future__ import annotations
from typing import Tuple
import torch
import torch.nn.functional as F
from torch import nn
from transformers import SegformerForSemanticSegmentation
class HaarDWT(nn.Module):
"""Single-level 2D Haar discrete wavelet transform via a fixed,
non-learnable depthwise stride-2 convolution.
Splits each input channel into 4 orthonormal sub-bands (LL, LH, HL, HH),
each at half spatial resolution. The filters are registered as a buffer
(not a parameter) since this is a classical transform, not a learned
filter bank - it contributes 0 trainable parameters.
"""
def __init__(self, in_channels: int) -> None:
super().__init__()
self.in_channels = in_channels
ll = torch.tensor([[1.0, 1.0], [1.0, 1.0]])
lh = torch.tensor([[1.0, -1.0], [1.0, -1.0]])
hl = torch.tensor([[1.0, 1.0], [-1.0, -1.0]])
hh = torch.tensor([[1.0, -1.0], [-1.0, 1.0]])
filters = torch.stack([ll, lh, hl, hh], dim=0) * 0.5 # (4, 2, 2), orthonormal
# Depthwise grouped conv: each input channel gets its own copy of the
# 4 filters -> weight shape (4*C, 1, 2, 2), consumed with groups=C so
# group g (= input channel g) produces output channels [4g : 4g+4].
weight = filters.unsqueeze(1).repeat(in_channels, 1, 1, 1)
self.register_buffer("weight", weight)
def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""Decompose ``x`` into (LL, LH, HL, HH), each shape ``(B, C, H/2, W/2)``.
Args:
x: Input feature map, shape ``(B, C, H, W)``, with ``C`` equal to
``in_channels`` and ``H, W`` even.
"""
out = F.conv2d(x, self.weight, stride=2, groups=self.in_channels) # (B, 4*C, H/2, W/2)
b, _, h, w = out.shape
out = out.view(b, self.in_channels, 4, h, w)
ll, lh, hl, hh = out[:, :, 0], out[:, :, 1], out[:, :, 2], out[:, :, 3]
return ll, lh, hl, hh
class FFTHighPass(nn.Module):
"""Single-map FFT high-pass filter - the A4 ablation's alternative to the
Haar DWT. Not used by A1 (WaveSeg, ours uses ``freq_transform="dwt"``);
kept here only because this file is a verbatim copy of the training repo.
"""
def __init__(self, cutoff_ratio: float = 0.25) -> None:
super().__init__()
if not 0.0 < cutoff_ratio < 1.0:
raise ValueError(f"cutoff_ratio must be in (0, 1), got {cutoff_ratio}")
self.cutoff_ratio = cutoff_ratio
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Return the high-pass-filtered map, same shape as ``x``.
Args:
x: Input feature map, shape ``(B, C, H, W)``.
"""
h, w = x.shape[-2:]
cy, cx = h // 2, w // 2
ry = max(1, int(round(h * self.cutoff_ratio / 2)))
rx = max(1, int(round(w * self.cutoff_ratio / 2)))
with torch.autocast(device_type=x.device.type, enabled=False):
spectrum = torch.fft.fftshift(torch.fft.fft2(x.float()), dim=(-2, -1))
low_freq_mask = torch.ones(h, w, device=x.device, dtype=torch.float32)
low_freq_mask[cy - ry : cy + ry, cx - rx : cx + rx] = 0.0
spectrum = spectrum * low_freq_mask
high_pass = torch.fft.ifft2(torch.fft.ifftshift(spectrum, dim=(-2, -1))).real
return high_pass.to(x.dtype)
class FrequencyBoundaryAdapter(nn.Module):
"""FBA: gates a decoder feature map using a boundary-attention map derived
from a high-frequency decomposition of that feature map.
``out = feat * (1 + alpha * attn)`` when ``use_gate=True`` (A1/WaveSeg -
the deployed configuration). ``use_gate=False`` and ``freq_transform="fft"``
are other ablations, not used here; kept only for a verbatim copy.
"""
def __init__(
self,
in_channels: int,
alpha_init: float = 0.1,
use_gate: bool = True,
freq_transform: str = "dwt",
fft_cutoff_ratio: float = 0.25,
) -> None:
super().__init__()
if freq_transform not in ("dwt", "fft"):
raise ValueError(f"freq_transform must be 'dwt' or 'fft', got {freq_transform!r}")
self.freq_transform = freq_transform
self.use_gate = use_gate
if freq_transform == "dwt":
self.dwt = HaarDWT(in_channels)
self.gate = nn.Conv2d(in_channels * 3, 1, kernel_size=1)
else:
self.fft_highpass = FFTHighPass(cutoff_ratio=fft_cutoff_ratio)
self.gate = nn.Conv2d(in_channels, 1, kernel_size=1)
self.alpha = nn.Parameter(torch.tensor(float(alpha_init)))
def forward(self, feat: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""Gate a decoder feature map with a frequency-derived boundary attention map.
Args:
feat: Decoder feature map, shape ``(B, C, H, W)``. ``H, W`` must
be even when ``freq_transform="dwt"``.
Returns:
``(out, attn)``: ``out`` has the same shape as ``feat``; ``attn``
is the boundary attention map, shape ``(B, 1, H, W)``, at
``feat``'s resolution.
"""
h, w = feat.shape[-2:]
if self.freq_transform == "dwt":
_, lh, hl, hh = self.dwt(feat)
high_freq = torch.cat([lh, hl, hh], dim=1) # (B, 3C, H/2, W/2)
attn = torch.sigmoid(self.gate(high_freq)) # (B, 1, H/2, W/2)
attn = F.interpolate(attn, size=(h, w), mode="bilinear", align_corners=False)
else:
high_freq = self.fft_highpass(feat) # (B, C, H, W), already full resolution
attn = torch.sigmoid(self.gate(high_freq)) # (B, 1, H, W)
out = feat * (1.0 + self.alpha * attn) if self.use_gate else feat
return out, attn
class SegFormerBaseline(nn.Module):
"""SegFormer-B0 wrapped for binary segmentation at full input resolution.
HuggingFace's SegFormer decode head predicts at 1/4 input resolution, so
this wrapper bilinearly upsamples logits back to the input size.
"""
def __init__(self, pretrained_name: str = "nvidia/mit-b0", num_classes: int = 1) -> None:
super().__init__()
# token=False: this is a public checkpoint that never needs auth - a
# stale cached HF token can otherwise turn a no-auth-needed request
# into a 401 (see the training repo's CLAUDE.md).
self.model = SegformerForSemanticSegmentation.from_pretrained(
pretrained_name,
num_labels=num_classes,
ignore_mismatched_sizes=True,
token=False,
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Run the encoder+decoder and upsample logits to the input resolution."""
h, w = x.shape[-2:]
logits = self.model(pixel_values=x).logits
return F.interpolate(logits, size=(h, w), mode="bilinear", align_corners=False)
class WaveSeg(nn.Module):
"""SegFormer-B0 baseline with an FBA gating a decoder/encoder feature map.
Deployed configuration (A1, WaveSeg ours): ``use_gate=True``,
``freq_transform="dwt"``, ``placement="deep"`` - see
``build_waveseg_a1()`` below. HuggingFace's ``SegformerDecodeHead``
always ends with ``self.classifier``, a plain 1x1 conv applied last; we
swap it for ``nn.Identity``, gate the pre-classifier features with the
FBA, then apply the real classifier ourselves.
"""
def __init__(
self,
pretrained_name: str = "nvidia/mit-b0",
num_classes: int = 1,
alpha_init: float = 0.1,
use_gate: bool = True,
freq_transform: str = "dwt",
fft_cutoff_ratio: float = 0.25,
placement: str = "deep",
) -> None:
super().__init__()
if placement not in ("deep", "shallow"):
raise ValueError(f"placement must be 'deep' or 'shallow', got {placement!r}")
self.placement = placement
self.baseline = SegFormerBaseline(pretrained_name=pretrained_name, num_classes=num_classes)
if placement == "deep":
decode_head = self.baseline.model.decode_head
fba_channels = decode_head.classifier.in_channels
self.classifier = decode_head.classifier
decode_head.classifier = nn.Identity()
else: # shallow: not used by A1, kept for a verbatim copy
fba_channels = self.baseline.model.config.hidden_sizes[0]
self.fba = FrequencyBoundaryAdapter(
in_channels=fba_channels,
alpha_init=alpha_init,
use_gate=use_gate,
freq_transform=freq_transform,
fft_cutoff_ratio=fft_cutoff_ratio,
)
def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""Run the encoder+decoder, gate with the FBA, then classify.
Args:
x: Input batch, shape ``(B, 3, H, W)``.
Returns:
``(logits, attn)``: segmentation logits ``(B, num_classes, H, W)``
and the FBA's boundary attention map ``(B, 1, H, W)``, both
upsampled to input resolution.
"""
h, w = x.shape[-2:]
if self.placement == "deep":
decoder_feat = self.baseline.model(pixel_values=x).logits # pre-classifier features
gated_feat, attn = self.fba(decoder_feat)
logits = self.classifier(gated_feat)
else: # shallow
outputs = self.baseline.model.segformer(pixel_values=x, output_hidden_states=True, return_dict=True)
hidden_states = list(outputs.hidden_states)
gated_stage0, attn = self.fba(hidden_states[0])
hidden_states[0] = gated_stage0
logits = self.baseline.model.decode_head(tuple(hidden_states))
logits = F.interpolate(logits, size=(h, w), mode="bilinear", align_corners=False)
attn = F.interpolate(attn, size=(h, w), mode="bilinear", align_corners=False)
return logits, attn
def build_waveseg_a1(pretrained_name: str = "nvidia/mit-b0", num_classes: int = 1) -> WaveSeg:
"""Build WaveSeg with the locked A1 configuration: FBA gating (DWT, deep
placement), no boundary loss involved (that's a training-time concern -
architecturally A1 and A3 are identical; only the loss used to train
them differs, which is why this one class serves the deployed model).
"""
return WaveSeg(
pretrained_name=pretrained_name,
num_classes=num_classes,
alpha_init=0.1,
use_gate=True,
freq_transform="dwt",
placement="deep",
)