andreribeiro87 commited on
Commit
9d81219
·
verified ·
1 Parent(s): f10b54f

Upload attention_unet+convnext model, code, and model card

Browse files
README.md ADDED
@@ -0,0 +1,160 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ library_name: pytorch
4
+ tags:
5
+ - image-segmentation
6
+ - medical
7
+ - polyp
8
+ - colonoscopy
9
+ - unet
10
+ - attention-unet
11
+ - convnext
12
+ datasets:
13
+ - andreribeiro87/kvasir-seg-augmented
14
+ metrics:
15
+ - dice
16
+ - iou
17
+ pipeline_tag: image-segmentation
18
+ ---
19
+
20
+ # Attention U-Net ConvNeXt — Polyp Segmentation (Best Test Dice)
21
+
22
+ Binary polyp segmentation model trained on [Kvasir-SEG](https://huggingface.co/datasets/Angelou0516/kvasir-seg).
23
+ **Highest Dice score on the test set** (0.9411) among all 24 architecture × backbone combinations evaluated in
24
+ the UNet-A benchmark sweep.
25
+
26
+ ## Model Description
27
+
28
+ | Property | Value |
29
+ |---|---|
30
+ | Architecture | **Attention U-Net** (gate-based skip connections) |
31
+ | Backbone | **ConvNeXt-Tiny** (ImageNet pre-trained via `timm`) |
32
+ | Input size | 256 × 256 × 3 |
33
+ | Output | 256 × 256 × 1 logit map (sigmoid → binary mask) |
34
+ | Parameters | ~133 MB |
35
+ | Loss | BCEDice (α = 0.5) |
36
+
37
+ ### Architecture Details
38
+
39
+ **Attention U-Net** (Oktay et al., 2018) augments the standard encoder-decoder with *attention gates* on every
40
+ skip connection. The gate computes a spatial attention coefficient from the decoder query and the encoder key,
41
+ suppressing activations in irrelevant background regions and focusing the model on lesion boundaries.
42
+
43
+ The **ConvNeXt-Tiny** backbone (pre-trained on ImageNet-1k) provides five resolution levels of feature maps.
44
+ ConvNeXt's depthwise convolution design gives excellent feature quality with relatively low memory usage,
45
+ making it well-suited for high-resolution segmentation.
46
+
47
+ ## Test Set Results
48
+
49
+ Evaluated on the fixed 53-image test partition of Kvasir-SEG (50 % of the original validation split, seed 42):
50
+
51
+ | Metric | Value |
52
+ |---|---|
53
+ | **Dice** | **0.9411** |
54
+ | **IoU** | **0.8888** |
55
+ | F1 | 0.9411 |
56
+ | Precision | 0.9618 |
57
+ | Recall | 0.9213 |
58
+ | Accuracy | 0.9803 |
59
+
60
+ ## Sweep Leaderboard (all 24 models)
61
+
62
+ | Rank | Model | Test Dice | Test IoU |
63
+ |------|-------|-----------|---------|
64
+ | **1** | **attention_unet_convnext (this model)** | **0.9411** | **0.8888** |
65
+ | 2 | unet3plus_convnext | 0.9395 | 0.8859 |
66
+ | 3 | unet_convnext | 0.9383 | 0.8838 |
67
+ | 4 | resunet_efficientnet | 0.9338 | 0.8759 |
68
+ | 5 | unet3plus_efficientnet | 0.9335 | 0.8753 |
69
+
70
+ ## Training Configuration
71
+
72
+ - **Optimiser:** AdamW
73
+ - **Learning rate:** 1e-3 (cosine decay, 5 % warmup)
74
+ - **Weight decay:** 1e-4
75
+ - **Batch size:** 64
76
+ - **Epochs:** 50
77
+ - **FP16:** enabled (A100 GPU)
78
+ - **Loss:** BCEDice
79
+ - **Dataset:** Kvasir-SEG augmented (5,368 train / 53 val / 53 test)
80
+ - **Augmentation:** random H/V flips, ±30° rotation, brightness/contrast/saturation ±20 %
81
+
82
+ > Note: This model was **not** subject to Optuna HPO. It is the raw sweep checkpoint.
83
+ > For the HPO-tuned variant see [andreribeiro87/unet3plus-efficientnet-kvasir-seg](https://huggingface.co/andreribeiro87/unet3plus-efficientnet-kvasir-seg).
84
+
85
+ ## How to Use
86
+
87
+ This model uses a custom PyTorch architecture. The model code is included in the repository.
88
+
89
+ ### Installation
90
+
91
+ ```bash
92
+ pip install torch torchvision timm safetensors
93
+ ```
94
+
95
+ ### Inference
96
+
97
+ ```python
98
+ import torch
99
+ from safetensors.torch import load_file
100
+ from torchvision.transforms import functional as TF
101
+ from PIL import Image
102
+
103
+ # 1. Clone or download the repo files
104
+ # git clone https://huggingface.co/andreribeiro87/attention-unet-convnext-kvasir-seg
105
+ # cd attention-unet-convnext-kvasir-seg
106
+
107
+ # 2. Import the model
108
+ from models import create_model
109
+
110
+ # 3. Instantiate and load weights
111
+ model = create_model("attention_unet", backbone="convnext")
112
+ state_dict = load_file("model.safetensors")
113
+ model.load_state_dict(state_dict)
114
+ model.eval()
115
+
116
+ # 4. Preprocess an image
117
+ image = Image.open("your_colonoscopy_image.jpg").convert("RGB")
118
+ x = TF.to_tensor(TF.resize(image, [256, 256])).unsqueeze(0) # (1, 3, 256, 256)
119
+
120
+ # 5. Predict
121
+ with torch.no_grad():
122
+ logit = model(x) # (1, 1, 256, 256)
123
+ mask = (logit.sigmoid() > 0.5).squeeze() # bool tensor (256, 256)
124
+
125
+ # 6. Convert to PIL
126
+ pred_mask = TF.to_pil_image(mask.float())
127
+ ```
128
+
129
+ ## Citation
130
+
131
+ If you use this model or dataset, please cite the original Kvasir-SEG paper:
132
+
133
+ ```bibtex
134
+ @inproceedings{jha2020kvasir,
135
+ title = {Kvasir-SEG: A Segmented Polyp Dataset},
136
+ author = {Jha, Debesh and Smedsrud, Pia H and Riegler, Michael A and Halvorsen, P{a}l
137
+ and de Lange, Thomas and Johansen, Dag and Johansen, H{a}vard D},
138
+ booktitle = {MultiMedia Modeling (MMM)},
139
+ year = {2020}
140
+ }
141
+ ```
142
+
143
+ ```bibtex
144
+ @article{oktay2018attention,
145
+ title = {Attention U-Net: Learning Where to Look for the Pancreas},
146
+ author = {Oktay, Ozan and Schlemper, Jo and Folgoc, Loic Le and Lee, Matthew and Heinrich,
147
+ Mattias and Misawa, Kazunari and Mori, Kensaku and McDonagh, Steven and
148
+ Hammerla, Nils Y and Kainz, Bernhard and others},
149
+ journal = {arXiv preprint arXiv:1804.03999},
150
+ year = {2018}
151
+ }
152
+ ```
153
+
154
+ ## Limitations
155
+
156
+ - Trained and evaluated exclusively on **Kvasir-SEG** (single-centre, single-modality).
157
+ Performance may degrade on other colonoscopy datasets or imaging conditions.
158
+ - Binary segmentation only; does not distinguish between polyp types or severity.
159
+ - Input resolution is fixed at **256 × 256**; very small polyps may not be fully captured.
160
+ - **Not validated for clinical use.** This is a research model.
config.json ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architecture": "attention_unet",
3
+ "backbone": "convnext",
4
+ "img_size": 256,
5
+ "in_channels": 3,
6
+ "out_channels": 1,
7
+ "training": {
8
+ "loss": "bce_dice",
9
+ "lr": 0.001,
10
+ "weight_decay": 0.0001,
11
+ "warmup_ratio": 0.05,
12
+ "scheduler": "cosine",
13
+ "batch_size": 64,
14
+ "epochs": 50,
15
+ "optimiser": "AdamW"
16
+ },
17
+ "test_metrics": {
18
+ "dice": 0.9411,
19
+ "iou": 0.8888,
20
+ "f1": 0.9411,
21
+ "precision": 0.9618,
22
+ "recall": 0.9213,
23
+ "accuracy": 0.9803
24
+ },
25
+ "dataset": "andreribeiro87/kvasir-seg-augmented",
26
+ "note": "Best test-Dice model from architecture sweep (no HPO applied)"
27
+ }
loss.py ADDED
@@ -0,0 +1,94 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+
5
+
6
+ class DiceFocalLoss(nn.Module):
7
+ """Combined Dice + Focal loss for class-imbalanced binary segmentation.
8
+
9
+ Focal loss down-weights easy negatives, forcing the network to focus on
10
+ hard/uncertain pixels — the main failure mode when plateauing near 0.1.
11
+
12
+ Args:
13
+ alpha: Focal weighting factor for positive class (0.25 typical).
14
+ gamma: Focal modulating exponent. Higher = more focus on hard
15
+ pixels. Tune in [0.5, 5.0] via Optuna.
16
+ dice_weight: Weight of the Dice component.
17
+ focal_weight: Weight of the Focal component.
18
+ smooth: Laplace smoothing for Dice denominator.
19
+ """
20
+
21
+ def __init__(
22
+ self,
23
+ alpha: float = 0.25,
24
+ gamma: float = 2.0,
25
+ dice_weight: float = 0.5,
26
+ focal_weight: float = 0.5,
27
+ smooth: float = 1e-6,
28
+ ) -> None:
29
+ super().__init__()
30
+ self.alpha = alpha
31
+ self.gamma = gamma
32
+ self.dice_weight = dice_weight
33
+ self.focal_weight = focal_weight
34
+ self.smooth = smooth
35
+
36
+ def forward(self, y_pred: torch.Tensor, y_true: torch.Tensor) -> torch.Tensor:
37
+ # ---- Focal part ------------------------------------------------
38
+ bce = F.binary_cross_entropy_with_logits(y_pred, y_true.float(), reduction="none")
39
+ pt = torch.exp(-bce)
40
+ focal = self.alpha * (1.0 - pt) ** self.gamma * bce
41
+ focal_loss = focal.mean()
42
+
43
+ # ---- Dice part -------------------------------------------------
44
+ pred_sig = torch.sigmoid(y_pred)
45
+ inter = (pred_sig * y_true).sum(dim=(2, 3))
46
+ dice_loss = 1.0 - (2.0 * inter + self.smooth) / (
47
+ pred_sig.sum(dim=(2, 3)) + y_true.sum(dim=(2, 3)) + self.smooth
48
+ )
49
+ dice_loss = dice_loss.mean()
50
+
51
+ return self.dice_weight * dice_loss + self.focal_weight * focal_loss
52
+
53
+
54
+ class DiceLoss(nn.Module):
55
+
56
+ def __init__(self, smooth: float = 1.0) -> None:
57
+ super().__init__()
58
+ self.smooth = smooth
59
+
60
+ def forward(self, y_pred: torch.Tensor, y_true: torch.Tensor) -> torch.Tensor:
61
+ assert y_pred.size() == y_true.size()
62
+ y_pred = y_pred[:, 0].contiguous().view(-1)
63
+ y_true = y_true[:, 0].contiguous().view(-1)
64
+ intersection = (y_pred * y_true).sum()
65
+ dsc = (2.0 * intersection + self.smooth) / (
66
+ y_pred.sum() + y_true.sum() + self.smooth
67
+ )
68
+ return 1.0 - dsc
69
+
70
+
71
+ class BCEDiceLoss(nn.Module):
72
+
73
+ def __init__(
74
+ self,
75
+ smooth: float = 1.0,
76
+ bce_weight: float = 0.5,
77
+ dice_weight: float = 0.5,
78
+ label_smoothing: float = 0.0,
79
+ ) -> None:
80
+ super().__init__()
81
+ self.bce_weight = bce_weight
82
+ self.dice_weight = dice_weight
83
+ self.label_smoothing = label_smoothing
84
+ self.dice = DiceLoss(smooth=smooth)
85
+
86
+ def forward(self, y_pred: torch.Tensor, y_true: torch.Tensor) -> torch.Tensor:
87
+ if self.label_smoothing > 0.0:
88
+ # Smooth labels towards 0.5: prevents overconfident BCE
89
+ y_bce = y_true * (1.0 - self.label_smoothing) + self.label_smoothing * 0.5
90
+ else:
91
+ y_bce = y_true
92
+ bce = F.binary_cross_entropy_with_logits(y_pred, y_bce)
93
+ dice = self.dice(torch.sigmoid(y_pred), y_true)
94
+ return self.bce_weight * bce + self.dice_weight * dice
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b0c29289995adfc04f79dc1c5a0f1c0e0f85b3b19f11f0fab4ee0dd9699a2a5d
3
+ size 139209476
models/__init__.py ADDED
@@ -0,0 +1,62 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch.nn as nn
2
+
3
+ from .attention_unet import AttentionUNet
4
+ from .resunet import ResUNet
5
+ from .transunet import TransUNet
6
+ from .unet import UNet
7
+ from .unet3plus import UNet3Plus
8
+ from .unetplusplus import UNetPlusPlus
9
+
10
+ # ---------------------------------------------------------------------------
11
+ # Architecture registry
12
+ # ---------------------------------------------------------------------------
13
+
14
+ ARCHITECTURES: dict[str, type] = {
15
+ "unet": UNet,
16
+ "unetplusplus": UNetPlusPlus,
17
+ "unet3plus": UNet3Plus,
18
+ "attention_unet": AttentionUNet,
19
+ "resunet": ResUNet,
20
+ "transunet": TransUNet,
21
+ }
22
+
23
+ # Backbones available for all custom hierarchical architectures
24
+ BACKBONES: list[str] = ["default", "efficientnet", "convnext", "swin", "siglip"]
25
+
26
+ # Valid (architecture, backbone) combinations
27
+ VALID_COMBINATIONS: dict[str, list[str]] = {
28
+ arch: BACKBONES for arch in ARCHITECTURES
29
+ }
30
+
31
+ # All architecture names
32
+ ALL_ARCHITECTURES: list[str] = list(ARCHITECTURES)
33
+
34
+
35
+ # ---------------------------------------------------------------------------
36
+ # create_model — unified factory
37
+ # ---------------------------------------------------------------------------
38
+
39
+ def create_model(
40
+ architecture: str,
41
+ backbone: str,
42
+ in_channels: int = 3,
43
+ out_channels: int = 1,
44
+ **kwargs,
45
+ ) -> nn.Module:
46
+ valid_bbs = VALID_COMBINATIONS.get(architecture)
47
+ if valid_bbs is None:
48
+ raise ValueError(
49
+ f"Unknown architecture '{architecture}'. "
50
+ f"Choose from {ALL_ARCHITECTURES}"
51
+ )
52
+ if backbone not in valid_bbs:
53
+ raise ValueError(
54
+ f"Backbone '{backbone}' is not compatible with '{architecture}'. "
55
+ f"Valid choices: {valid_bbs}"
56
+ )
57
+ return ARCHITECTURES[architecture](
58
+ backbone_name=backbone,
59
+ in_channels=in_channels,
60
+ out_channels=out_channels,
61
+ **kwargs,
62
+ )
models/attention_unet.py ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+
5
+ from .backbones import create_backbone
6
+ from .blocks import AttentionDecoderBlock
7
+
8
+
9
+ class AttentionUNet(nn.Module):
10
+ def __init__(
11
+ self,
12
+ backbone_name: str = "default",
13
+ in_channels: int = 3,
14
+ out_channels: int = 1,
15
+ **_kw,
16
+ ):
17
+ super().__init__()
18
+ self.backbone = create_backbone(backbone_name, in_channels)
19
+ channels = self.backbone.out_channels
20
+ skip_channels, bottleneck_ch = channels[:-1], channels[-1]
21
+
22
+ self.decoders = nn.ModuleList()
23
+ in_ch = bottleneck_ch
24
+ for s_ch in reversed(skip_channels):
25
+ self.decoders.append(AttentionDecoderBlock(in_ch, s_ch, s_ch))
26
+ in_ch = s_ch
27
+
28
+ self.head = nn.Conv2d(in_ch, out_channels, kernel_size=1)
29
+
30
+ def forward(
31
+ self, x: torch.Tensor | None = None, pixel_values: torch.Tensor | None = None, **_kw,
32
+ ) -> torch.Tensor:
33
+ if pixel_values is not None:
34
+ x = pixel_values
35
+ input_size = x.shape[2:]
36
+
37
+ skips, x = self.backbone(x)
38
+ for dec, skip in zip(self.decoders, reversed(skips)):
39
+ x = dec(x, skip)
40
+
41
+ x = self.head(x)
42
+ if x.shape[2:] != input_size:
43
+ x = F.interpolate(x, size=input_size, mode="bilinear", align_corners=False)
44
+ return x
models/backbones.py ADDED
@@ -0,0 +1,317 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import ssl
2
+ from abc import ABC, abstractmethod
3
+
4
+ import torch
5
+ import torch.nn as nn
6
+ import torch.nn.functional as F
7
+ from torchvision.models import (
8
+ ConvNeXt_Tiny_Weights,
9
+ EfficientNet_B0_Weights,
10
+ Swin_T_Weights,
11
+ convnext_tiny,
12
+ efficientnet_b0,
13
+ swin_t,
14
+ )
15
+ from torchvision.models.feature_extraction import create_feature_extractor
16
+
17
+ from .blocks import ConvBlock
18
+
19
+
20
+ def _disable_ssl_verification():
21
+ """Create an unverified SSL context for downloading pretrained weights."""
22
+ return ssl._create_unverified_context()
23
+
24
+
25
+ # Workaround for SSL certificate verification issues when downloading pretrained weights
26
+ _original_create_default_https_context = ssl._create_default_https_context
27
+
28
+
29
+ def _enable_unverified_ssl():
30
+ ssl._create_default_https_context = _disable_ssl_verification
31
+
32
+
33
+ def _restore_ssl():
34
+ ssl._create_default_https_context = _original_create_default_https_context
35
+
36
+
37
+ class Backbone(ABC):
38
+ """All backbones return (skip_features high→low, bottleneck)."""
39
+
40
+ @property
41
+ @abstractmethod
42
+ def out_channels(self) -> list[int]:
43
+ """Channel counts from highest-res skip to bottleneck (last element)."""
44
+ ...
45
+
46
+ @abstractmethod
47
+ def forward(self, x: torch.Tensor) -> tuple[list[torch.Tensor], torch.Tensor]:
48
+ ...
49
+
50
+
51
+ class DefaultBackbone(Backbone, nn.Module):
52
+ """Plain conv encoder identical to the classic U-Net."""
53
+
54
+ def __init__(self, in_channels: int = 3, block_cls: type | None = None, **_kw):
55
+ nn.Module.__init__(self)
56
+ block = block_cls or ConvBlock
57
+ self.enc1 = block(in_channels, 64)
58
+ self.enc2 = block(64, 128)
59
+ self.enc3 = block(128, 256)
60
+ self.enc4 = block(256, 512)
61
+ self.bottleneck = block(512, 1024)
62
+ self.pool = nn.MaxPool2d(2, 2)
63
+
64
+ @property
65
+ def out_channels(self) -> list[int]:
66
+ return [64, 128, 256, 512, 1024]
67
+
68
+ def forward(self, x: torch.Tensor):
69
+ skips = []
70
+ for enc in [self.enc1, self.enc2, self.enc3, self.enc4]:
71
+ x = enc(x)
72
+ skips.append(x)
73
+ x = self.pool(x)
74
+ return skips, self.bottleneck(x)
75
+
76
+
77
+ class EfficientNetBackbone(Backbone, nn.Module):
78
+ """EfficientNet-B0 pretrained encoder (ImageNet)."""
79
+
80
+ _RETURN_NODES = {
81
+ "features.1": "s1", # H/2, 16 ch
82
+ "features.2": "s2", # H/4, 24 ch
83
+ "features.3": "s3", # H/8, 40 ch
84
+ "features.5": "s4", # H/16, 112 ch
85
+ "features.8": "bottleneck", # H/32, 1280 ch
86
+ }
87
+
88
+ def __init__(self, in_channels: int = 3, **_kw):
89
+ nn.Module.__init__(self)
90
+ _enable_unverified_ssl()
91
+ try:
92
+ base = efficientnet_b0(weights=EfficientNet_B0_Weights.DEFAULT)
93
+ finally:
94
+ _restore_ssl()
95
+ if in_channels != 3:
96
+ old = base.features[0][0]
97
+ base.features[0][0] = nn.Conv2d(
98
+ in_channels, old.out_channels, old.kernel_size,
99
+ old.stride, old.padding, bias=False,
100
+ )
101
+ self.body = create_feature_extractor(base, return_nodes=self._RETURN_NODES)
102
+
103
+ @property
104
+ def out_channels(self) -> list[int]:
105
+ return [16, 24, 40, 112, 1280]
106
+
107
+ def forward(self, x: torch.Tensor):
108
+ f = self.body(x)
109
+ return [f["s1"], f["s2"], f["s3"], f["s4"]], f["bottleneck"]
110
+
111
+
112
+ class ConvNeXTBackbone(Backbone, nn.Module):
113
+ """ConvNeXT-Tiny pretrained encoder (ImageNet)."""
114
+
115
+ _RETURN_NODES = {
116
+ "features.1": "s1", # H/4, 96 ch
117
+ "features.3": "s2", # H/8, 192 ch
118
+ "features.5": "s3", # H/16, 384 ch
119
+ "features.7": "bottleneck", # H/32, 768 ch
120
+ }
121
+
122
+ def __init__(self, in_channels: int = 3, **_kw):
123
+ nn.Module.__init__(self)
124
+ _enable_unverified_ssl()
125
+ try:
126
+ base = convnext_tiny(weights=ConvNeXt_Tiny_Weights.DEFAULT)
127
+ finally:
128
+ _restore_ssl()
129
+ if in_channels != 3:
130
+ old = base.features[0][0]
131
+ base.features[0][0] = nn.Conv2d(
132
+ in_channels, old.out_channels, old.kernel_size,
133
+ old.stride, old.padding, bias=False,
134
+ )
135
+ self.body = create_feature_extractor(base, return_nodes=self._RETURN_NODES)
136
+
137
+ @property
138
+ def out_channels(self) -> list[int]:
139
+ return [96, 192, 384, 768]
140
+
141
+ def forward(self, x: torch.Tensor):
142
+ f = self.body(x)
143
+ return [f["s1"], f["s2"], f["s3"]], f["bottleneck"]
144
+
145
+
146
+ class SwinBackbone(Backbone, nn.Module):
147
+ """Swin Transformer Tiny pretrained encoder (ImageNet).
148
+
149
+ Produces hierarchical features at four scales — identical channel widths
150
+ to ConvNeXt-Tiny (96 → 192 → 384 → 768) — so it slots into every
151
+ existing decoder without any code changes.
152
+
153
+ Input (B, 3, H, W) → skips [(B,96,H/4,W/4), (B,192,H/8,W/8),
154
+ (B,384,H/16,W/16)], bottleneck (B,768,H/32,W/32)
155
+
156
+ Note: torchvision Swin outputs tensors in (B, H, W, C) layout;
157
+ the backbone permutes them to the standard (B, C, H, W) before returning.
158
+ """
159
+
160
+ _RETURN_NODES = {
161
+ "features.1": "s1", # H/4, 96 ch
162
+ "features.3": "s2", # H/8, 192 ch
163
+ "features.5": "s3", # H/16, 384 ch
164
+ "features.7": "bottleneck", # H/32, 768 ch
165
+ }
166
+
167
+ def __init__(self, in_channels: int = 3, **_kw):
168
+ nn.Module.__init__(self)
169
+ _enable_unverified_ssl()
170
+ try:
171
+ base = swin_t(weights=Swin_T_Weights.DEFAULT)
172
+ finally:
173
+ _restore_ssl()
174
+ if in_channels != 3:
175
+ # Replace the first patch-embedding conv
176
+ old = base.features[0][0]
177
+ base.features[0][0] = nn.Conv2d(
178
+ in_channels, old.out_channels,
179
+ kernel_size=old.kernel_size, stride=old.stride,
180
+ padding=old.padding, bias=False,
181
+ )
182
+ self.body = create_feature_extractor(base, return_nodes=self._RETURN_NODES)
183
+
184
+ @property
185
+ def out_channels(self) -> list[int]:
186
+ return [96, 192, 384, 768]
187
+
188
+ def forward(self, x: torch.Tensor):
189
+ f = self.body(x)
190
+ # Swin-T stores features as (B, H, W, C) → convert to (B, C, H, W)
191
+ s1 = f["s1"].permute(0, 3, 1, 2).contiguous()
192
+ s2 = f["s2"].permute(0, 3, 1, 2).contiguous()
193
+ s3 = f["s3"].permute(0, 3, 1, 2).contiguous()
194
+ bn = f["bottleneck"].permute(0, 3, 1, 2).contiguous()
195
+ return [s1, s2, s3], bn
196
+
197
+
198
+ class SigLIPBackbone(Backbone, nn.Module):
199
+ """SigLIP-Base/16 (93 M) pretrained vision encoder as a flat-ViT backbone.
200
+
201
+ Gemma 3 uses the larger ``google/siglip-so400m-patch14-384`` (400 M); this
202
+ class defaults to the practical ``google/siglip-base-patch16-224`` variant
203
+ (93 M) that fits comfortably alongside a UNet decoder. Swap MODEL_ID for
204
+ the Gemma 3 variant when VRAM permits.
205
+
206
+ Architecture note
207
+ -----------------
208
+ SigLIP is a pure Vision Transformer — every layer produces tokens at the
209
+ **same** spatial resolution (H/16 × W/16 = 14×14 for 224-px input). This
210
+ backbone therefore returns four feature maps that are **all at 14×14** but
211
+ at different semantic depths (layers 3 / 6 / 9 / 12 of the 12-layer ViT):
212
+
213
+ [z3, z6, z9] → skips (each B, 768, 14, 14)
214
+ z12 → bottleneck (B, 768, 14, 14)
215
+
216
+ This is intentionally different from the hierarchical CNN/Swin backbones.
217
+ Use the ``ViTUNet`` architecture (``models/vit_unet.py``) which handles the
218
+ flat-resolution skips by projecting and bilinearly upsampling them to match
219
+ each decoder stage. Standard UNet/AttentionUNet/ResUNet/TransUNet decoders
220
+ will NOT work correctly with this backbone.
221
+
222
+ Gemma 3 variant
223
+ ---------------
224
+ Replace MODEL_ID with ``"google/siglip-so400m-patch14-384"`` and set
225
+ ``PATCH_SIZE = 14``, ``INPUT_SIZE = 384`` to use the exact Gemma 3 encoder.
226
+ You will also need to adjust ``EXTRACT_LAYERS`` (the model has 27 layers).
227
+ """
228
+
229
+ MODEL_ID = "google/siglip-base-patch16-224"
230
+ PATCH_SIZE = 16
231
+ INPUT_SIZE = 224
232
+ # EXTRACT_LAYERS: quarter-intervals of the 12-layer ViT (1-indexed after embedding)
233
+ EXTRACT_LAYERS = (3, 6, 9, 12)
234
+
235
+ def __init__(self, in_channels: int = 3, **_kw):
236
+ nn.Module.__init__(self)
237
+ # google/siglip-base-patch16-224 is a full CLIP-style model (vision +
238
+ # text). Loading with SiglipVisionModel.from_pretrained fails because
239
+ # the repo config is a SiglipConfig, not a SiglipVisionConfig.
240
+ # Solution: load the full SiglipModel and keep only the vision encoder.
241
+ from transformers import SiglipModel
242
+
243
+ full = SiglipModel.from_pretrained(self.MODEL_ID)
244
+ self._vit = full.vision_model # SiglipVisionTransformer
245
+ self._hidden_dim: int = full.config.vision_config.hidden_size # 768
246
+ del full # free text encoder weights
247
+
248
+ self._patch_grid = self.INPUT_SIZE // self.PATCH_SIZE # 14
249
+
250
+ if in_channels != 3:
251
+ old = self._vit.embeddings.patch_embedding
252
+ self._vit.embeddings.patch_embedding = nn.Conv2d(
253
+ in_channels, self._hidden_dim,
254
+ kernel_size=self.PATCH_SIZE, stride=self.PATCH_SIZE, bias=False,
255
+ )
256
+
257
+ # Register forward hooks on the target transformer layers to capture
258
+ # intermediate hidden states. This is version-agnostic — it works
259
+ # regardless of whether the transformers library honours
260
+ # `output_hidden_states=True` via its **kwargs API.
261
+ self._hooked: dict[int, torch.Tensor] = {}
262
+ self._hook_handles: list = []
263
+ for layer_idx in self.EXTRACT_LAYERS:
264
+ layer = self._vit.encoder.layers[layer_idx - 1] # 0-indexed
265
+ handle = layer.register_forward_hook(self._make_hook(layer_idx))
266
+ self._hook_handles.append(handle)
267
+
268
+ def _make_hook(self, idx: int):
269
+ def _hook(module, inp, out):
270
+ # SiglipEncoderLayer returns a tuple; the first element is the
271
+ # hidden state tensor (B, N_patches, hidden_dim).
272
+ self._hooked[idx] = out[0] if isinstance(out, tuple) else out
273
+ return _hook
274
+
275
+ @property
276
+ def out_channels(self) -> list[int]:
277
+ return [self._hidden_dim] * len(self.EXTRACT_LAYERS) # [768, 768, 768, 768]
278
+
279
+ def forward(self, x: torch.Tensor):
280
+ B = x.shape[0]
281
+ if x.shape[2] != self.INPUT_SIZE or x.shape[3] != self.INPUT_SIZE:
282
+ x = F.interpolate(
283
+ x, size=(self.INPUT_SIZE, self.INPUT_SIZE),
284
+ mode="bilinear", align_corners=False,
285
+ )
286
+
287
+ self._hooked.clear()
288
+ self._vit(pixel_values=x) # hooks populate self._hooked
289
+
290
+ G = self._patch_grid # 14
291
+ features = []
292
+ for layer_idx in self.EXTRACT_LAYERS:
293
+ hs = self._hooked[layer_idx] # (B, 196, 768)
294
+ feat = (
295
+ hs.reshape(B, G, G, self._hidden_dim)
296
+ .permute(0, 3, 1, 2)
297
+ .contiguous()
298
+ ) # (B, 768, 14, 14)
299
+ features.append(feat)
300
+
301
+ # z3=features[0] (early/local), z12=features[3] (final/global)
302
+ return features[:-1], features[-1] # skips, bottleneck
303
+
304
+
305
+ BACKBONE_REGISTRY: dict[str, type] = {
306
+ "default": DefaultBackbone,
307
+ "efficientnet": EfficientNetBackbone,
308
+ "convnext": ConvNeXTBackbone,
309
+ "swin": SwinBackbone,
310
+ "siglip": SigLIPBackbone,
311
+ }
312
+
313
+
314
+ def create_backbone(name: str, in_channels: int = 3, **kwargs) -> nn.Module:
315
+ if name not in BACKBONE_REGISTRY:
316
+ raise ValueError(f"Unknown backbone '{name}'. Choose from {list(BACKBONE_REGISTRY)}")
317
+ return BACKBONE_REGISTRY[name](in_channels=in_channels, **kwargs)
models/blocks.py ADDED
@@ -0,0 +1,140 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+
5
+
6
+ class ConvBlock(nn.Module):
7
+ def __init__(self, in_channels: int, out_channels: int):
8
+ super().__init__()
9
+ self.block = nn.Sequential(
10
+ nn.Conv2d(in_channels, out_channels, 3, padding=1, bias=False),
11
+ nn.BatchNorm2d(out_channels),
12
+ nn.ReLU(inplace=True),
13
+ nn.Conv2d(out_channels, out_channels, 3, padding=1, bias=False),
14
+ nn.BatchNorm2d(out_channels),
15
+ nn.ReLU(inplace=True),
16
+ )
17
+
18
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
19
+ return self.block(x)
20
+
21
+
22
+ class ResBlock(nn.Module):
23
+ def __init__(self, in_channels: int, out_channels: int):
24
+ super().__init__()
25
+ self.conv = nn.Sequential(
26
+ nn.Conv2d(in_channels, out_channels, 3, padding=1, bias=False),
27
+ nn.BatchNorm2d(out_channels),
28
+ nn.ReLU(inplace=True),
29
+ nn.Conv2d(out_channels, out_channels, 3, padding=1, bias=False),
30
+ nn.BatchNorm2d(out_channels),
31
+ )
32
+ self.shortcut = (
33
+ nn.Sequential(
34
+ nn.Conv2d(in_channels, out_channels, 1, bias=False),
35
+ nn.BatchNorm2d(out_channels),
36
+ )
37
+ if in_channels != out_channels
38
+ else nn.Identity()
39
+ )
40
+ self.relu = nn.ReLU(inplace=True)
41
+
42
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
43
+ return self.relu(self.conv(x) + self.shortcut(x))
44
+
45
+
46
+ class AttentionGate(nn.Module):
47
+ """Additive attention gate: re-weights skip features using the decoder gate signal."""
48
+
49
+ def __init__(self, gate_channels: int, skip_channels: int):
50
+ super().__init__()
51
+ inter = max(1, skip_channels // 2)
52
+ self.W_gate = nn.Sequential(
53
+ nn.Conv2d(gate_channels, inter, 1, bias=False),
54
+ nn.BatchNorm2d(inter),
55
+ )
56
+ self.W_skip = nn.Sequential(
57
+ nn.Conv2d(skip_channels, inter, 1, bias=False),
58
+ nn.BatchNorm2d(inter),
59
+ )
60
+ self.psi = nn.Sequential(
61
+ nn.Conv2d(inter, 1, 1, bias=False),
62
+ nn.BatchNorm2d(1),
63
+ nn.Sigmoid(),
64
+ )
65
+ self.relu = nn.ReLU(inplace=True)
66
+
67
+ def forward(self, gate: torch.Tensor, skip: torch.Tensor) -> torch.Tensor:
68
+ g = self.W_gate(gate)
69
+ s = self.W_skip(skip)
70
+ if g.shape[2:] != s.shape[2:]:
71
+ g = F.interpolate(g, size=s.shape[2:], mode="bilinear", align_corners=False)
72
+ return skip * self.psi(self.relu(g + s))
73
+
74
+
75
+ class TransformerBlock(nn.Module):
76
+ def __init__(self, dim: int, num_heads: int = 8, mlp_ratio: float = 4.0, dropout: float = 0.1):
77
+ super().__init__()
78
+ self.norm1 = nn.LayerNorm(dim)
79
+ self.attn = nn.MultiheadAttention(dim, num_heads, dropout=dropout, batch_first=True)
80
+ self.norm2 = nn.LayerNorm(dim)
81
+ self.mlp = nn.Sequential(
82
+ nn.Linear(dim, int(dim * mlp_ratio)),
83
+ nn.GELU(),
84
+ nn.Dropout(dropout),
85
+ nn.Linear(int(dim * mlp_ratio), dim),
86
+ nn.Dropout(dropout),
87
+ )
88
+
89
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
90
+ h = self.norm1(x)
91
+ h, _ = self.attn(h, h, h)
92
+ x = x + h
93
+ x = x + self.mlp(self.norm2(x))
94
+ return x
95
+
96
+
97
+ # ---------------------------------------------------------------------------
98
+ # Decoder blocks (one per architecture variant)
99
+ # ---------------------------------------------------------------------------
100
+
101
+
102
+ class UNetDecoderBlock(nn.Module):
103
+ def __init__(self, in_ch: int, skip_ch: int, out_ch: int):
104
+ super().__init__()
105
+ self.up = nn.ConvTranspose2d(in_ch, out_ch, kernel_size=2, stride=2)
106
+ self.conv = ConvBlock(out_ch + skip_ch, out_ch)
107
+
108
+ def forward(self, x: torch.Tensor, skip: torch.Tensor) -> torch.Tensor:
109
+ x = self.up(x)
110
+ if x.shape[2:] != skip.shape[2:]:
111
+ x = F.interpolate(x, size=skip.shape[2:], mode="bilinear", align_corners=False)
112
+ return self.conv(torch.cat([x, skip], dim=1))
113
+
114
+
115
+ class AttentionDecoderBlock(nn.Module):
116
+ def __init__(self, in_ch: int, skip_ch: int, out_ch: int):
117
+ super().__init__()
118
+ self.up = nn.ConvTranspose2d(in_ch, out_ch, kernel_size=2, stride=2)
119
+ self.attn = AttentionGate(gate_channels=out_ch, skip_channels=skip_ch)
120
+ self.conv = ConvBlock(out_ch + skip_ch, out_ch)
121
+
122
+ def forward(self, x: torch.Tensor, skip: torch.Tensor) -> torch.Tensor:
123
+ x = self.up(x)
124
+ if x.shape[2:] != skip.shape[2:]:
125
+ x = F.interpolate(x, size=skip.shape[2:], mode="bilinear", align_corners=False)
126
+ skip = self.attn(gate=x, skip=skip)
127
+ return self.conv(torch.cat([x, skip], dim=1))
128
+
129
+
130
+ class ResDecoderBlock(nn.Module):
131
+ def __init__(self, in_ch: int, skip_ch: int, out_ch: int):
132
+ super().__init__()
133
+ self.up = nn.ConvTranspose2d(in_ch, out_ch, kernel_size=2, stride=2)
134
+ self.conv = ResBlock(out_ch + skip_ch, out_ch)
135
+
136
+ def forward(self, x: torch.Tensor, skip: torch.Tensor) -> torch.Tensor:
137
+ x = self.up(x)
138
+ if x.shape[2:] != skip.shape[2:]:
139
+ x = F.interpolate(x, size=skip.shape[2:], mode="bilinear", align_corners=False)
140
+ return self.conv(torch.cat([x, skip], dim=1))
models/resunet.py ADDED
@@ -0,0 +1,47 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+
5
+ from .backbones import create_backbone
6
+ from .blocks import ResBlock, ResDecoderBlock
7
+
8
+
9
+ class ResUNet(nn.Module):
10
+ def __init__(
11
+ self,
12
+ backbone_name: str = "default",
13
+ in_channels: int = 3,
14
+ out_channels: int = 1,
15
+ **_kw,
16
+ ):
17
+ super().__init__()
18
+ backbone_kwargs: dict = {}
19
+ if backbone_name == "default":
20
+ backbone_kwargs["block_cls"] = ResBlock
21
+ self.backbone = create_backbone(backbone_name, in_channels, **backbone_kwargs)
22
+ channels = self.backbone.out_channels
23
+ skip_channels, bottleneck_ch = channels[:-1], channels[-1]
24
+
25
+ self.decoders = nn.ModuleList()
26
+ in_ch = bottleneck_ch
27
+ for s_ch in reversed(skip_channels):
28
+ self.decoders.append(ResDecoderBlock(in_ch, s_ch, s_ch))
29
+ in_ch = s_ch
30
+
31
+ self.head = nn.Conv2d(in_ch, out_channels, kernel_size=1)
32
+
33
+ def forward(
34
+ self, x: torch.Tensor | None = None, pixel_values: torch.Tensor | None = None, **_kw,
35
+ ) -> torch.Tensor:
36
+ if pixel_values is not None:
37
+ x = pixel_values
38
+ input_size = x.shape[2:]
39
+
40
+ skips, x = self.backbone(x)
41
+ for dec, skip in zip(self.decoders, reversed(skips)):
42
+ x = dec(x, skip)
43
+
44
+ x = self.head(x)
45
+ if x.shape[2:] != input_size:
46
+ x = F.interpolate(x, size=input_size, mode="bilinear", align_corners=False)
47
+ return x
models/transunet.py ADDED
@@ -0,0 +1,71 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+
5
+ from .backbones import create_backbone
6
+ from .blocks import TransformerBlock, UNetDecoderBlock
7
+
8
+
9
+ class TransUNet(nn.Module):
10
+ def __init__(
11
+ self,
12
+ backbone_name: str = "default",
13
+ in_channels: int = 3,
14
+ out_channels: int = 1,
15
+ img_size: int = 256,
16
+ transformer_dim: int = 512,
17
+ num_heads: int = 8,
18
+ num_layers: int = 6,
19
+ mlp_ratio: float = 4.0,
20
+ **_kw,
21
+ ):
22
+ super().__init__()
23
+ self.backbone = create_backbone(backbone_name, in_channels)
24
+ channels = self.backbone.out_channels
25
+ skip_channels, bottleneck_ch = channels[:-1], channels[-1]
26
+
27
+ with torch.no_grad():
28
+ dummy = torch.zeros(1, in_channels, img_size, img_size)
29
+ _, dummy_bn = self.backbone(dummy)
30
+ num_patches = dummy_bn.shape[2] * dummy_bn.shape[3]
31
+
32
+ self.proj_in = nn.Conv2d(bottleneck_ch, transformer_dim, kernel_size=1)
33
+ self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, transformer_dim))
34
+ nn.init.trunc_normal_(self.pos_embed, std=0.02)
35
+ self.transformer = nn.Sequential(
36
+ *[TransformerBlock(transformer_dim, num_heads, mlp_ratio) for _ in range(num_layers)]
37
+ )
38
+ self.proj_out = nn.Conv2d(transformer_dim, bottleneck_ch, kernel_size=1)
39
+
40
+ self.decoders = nn.ModuleList()
41
+ in_ch = bottleneck_ch
42
+ for s_ch in reversed(skip_channels):
43
+ self.decoders.append(UNetDecoderBlock(in_ch, s_ch, s_ch))
44
+ in_ch = s_ch
45
+
46
+ self.head = nn.Conv2d(in_ch, out_channels, kernel_size=1)
47
+
48
+ def forward(
49
+ self, x: torch.Tensor | None = None, pixel_values: torch.Tensor | None = None, **_kw,
50
+ ) -> torch.Tensor:
51
+ if pixel_values is not None:
52
+ x = pixel_values
53
+ input_size = x.shape[2:]
54
+
55
+ skips, bottleneck = self.backbone(x)
56
+
57
+ B, _C, H, W = bottleneck.shape
58
+ t = self.proj_in(bottleneck) # B, D, H, W
59
+ t = t.flatten(2).transpose(1, 2) # B, N, D
60
+ t = t + self.pos_embed
61
+ t = self.transformer(t) # B, N, D
62
+ t = t.transpose(1, 2).view(B, -1, H, W) # B, D, H, W
63
+ x = self.proj_out(t) # B, bottleneck_ch, H, W
64
+
65
+ for dec, skip in zip(self.decoders, reversed(skips)):
66
+ x = dec(x, skip)
67
+
68
+ x = self.head(x)
69
+ if x.shape[2:] != input_size:
70
+ x = F.interpolate(x, size=input_size, mode="bilinear", align_corners=False)
71
+ return x
models/unet.py ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+
5
+ from .backbones import create_backbone
6
+ from .blocks import UNetDecoderBlock
7
+
8
+
9
+ class UNet(nn.Module):
10
+ def __init__(
11
+ self,
12
+ backbone_name: str = "default",
13
+ in_channels: int = 3,
14
+ out_channels: int = 1,
15
+ **_kw,
16
+ ):
17
+ super().__init__()
18
+ self.backbone = create_backbone(backbone_name, in_channels)
19
+ channels = self.backbone.out_channels
20
+ skip_channels, bottleneck_ch = channels[:-1], channels[-1]
21
+
22
+ self.decoders = nn.ModuleList()
23
+ in_ch = bottleneck_ch
24
+ for s_ch in reversed(skip_channels):
25
+ self.decoders.append(UNetDecoderBlock(in_ch, s_ch, s_ch))
26
+ in_ch = s_ch
27
+
28
+ self.head = nn.Conv2d(in_ch, out_channels, kernel_size=1)
29
+
30
+ def forward(
31
+ self, x: torch.Tensor | None = None, pixel_values: torch.Tensor | None = None, **_kw,
32
+ ) -> torch.Tensor:
33
+ if pixel_values is not None:
34
+ x = pixel_values
35
+ input_size = x.shape[2:]
36
+
37
+ skips, x = self.backbone(x)
38
+ for dec, skip in zip(self.decoders, reversed(skips)):
39
+ x = dec(x, skip)
40
+
41
+ x = self.head(x)
42
+ if x.shape[2:] != input_size:
43
+ x = F.interpolate(x, size=input_size, mode="bilinear", align_corners=False)
44
+ return x
models/unet3plus.py ADDED
@@ -0,0 +1,120 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+
5
+ from .backbones import create_backbone
6
+ from .blocks import ConvBlock
7
+
8
+
9
+ class UNet3Plus(nn.Module):
10
+ """UNet 3+ with full-scale skip connections.
11
+
12
+ Every decoder node aggregates features from ALL encoder levels (not just
13
+ the matching scale), giving each node simultaneous access to fine-grained
14
+ detail and deep semantic context. Each incoming stream is projected to
15
+ `inter_ch` channels before concatenation so total input channels are
16
+ predictable regardless of backbone channel widths.
17
+
18
+ Reference: Huang et al., "UNet 3+: A Full-Scale Connected UNet for
19
+ Medical Image Segmentation", ICASSP 2020.
20
+ """
21
+
22
+ def __init__(
23
+ self,
24
+ backbone_name: str = "default",
25
+ in_channels: int = 3,
26
+ out_channels: int = 1,
27
+ inter_ch: int = 64,
28
+ **_kw,
29
+ ):
30
+ super().__init__()
31
+ self.backbone = create_backbone(backbone_name, in_channels)
32
+ all_ch = list(self.backbone.out_channels) # [skip0, ..., skip_{D-1}, bottleneck]
33
+ D = len(all_ch) - 1 # number of skip levels = decoder levels
34
+ self._D = D
35
+ self._inter_ch = inter_ch
36
+ self._dec_ch = (D + 1) * inter_ch # fixed output channels for every decoder node
37
+
38
+ # --- projections: each incoming stream is independently projected to inter_ch ---
39
+ # enc_projs[k][j]: encoder/bottleneck level j → inter_ch, for use at decoder k
40
+ self.enc_projs = nn.ModuleList([
41
+ nn.ModuleList([
42
+ nn.Sequential(
43
+ nn.Conv2d(all_ch[j], inter_ch, 1, bias=False),
44
+ nn.BatchNorm2d(inter_ch),
45
+ nn.ReLU(inplace=True),
46
+ )
47
+ for j in range(D + 1)
48
+ ])
49
+ for k in range(D)
50
+ ])
51
+
52
+ # dec_projs[k][i]: prior decoder level i → inter_ch, for use at decoder k (i < k)
53
+ # All prior decoder outputs have self._dec_ch channels.
54
+ self.dec_projs = nn.ModuleList([
55
+ nn.ModuleList([
56
+ nn.Sequential(
57
+ nn.Conv2d(self._dec_ch, inter_ch, 1, bias=False),
58
+ nn.BatchNorm2d(inter_ch),
59
+ nn.ReLU(inplace=True),
60
+ )
61
+ for _ in range(k) # one projection per prior decoder level
62
+ ])
63
+ for k in range(D)
64
+ ])
65
+
66
+ # Fusion conv at each decoder level: (D+1+k)*inter_ch → dec_ch
67
+ self.dec_convs = nn.ModuleList([
68
+ ConvBlock((D + 1 + k) * inter_ch, self._dec_ch)
69
+ for k in range(D)
70
+ ])
71
+
72
+ self.head = nn.Conv2d(self._dec_ch, out_channels, kernel_size=1)
73
+
74
+ def forward(
75
+ self,
76
+ x: torch.Tensor | None = None,
77
+ pixel_values: torch.Tensor | None = None,
78
+ **_kw,
79
+ ) -> torch.Tensor:
80
+ if pixel_values is not None:
81
+ x = pixel_values
82
+ input_size = x.shape[2:]
83
+
84
+ skips, bottleneck = self.backbone(x)
85
+ D = self._D
86
+ # enc_feats: shallowest encoder first, bottleneck last (D+1 tensors)
87
+ enc_feats = list(skips) + [bottleneck]
88
+
89
+ dec_outputs: list[torch.Tensor] = []
90
+
91
+ # Build decoder from deepest (k=0) to shallowest (k=D-1).
92
+ # Decoder k targets the same spatial resolution as skips[D-1-k].
93
+ for k in range(D):
94
+ target = skips[D - 1 - k].shape[2:]
95
+ parts: list[torch.Tensor] = []
96
+
97
+ # All encoder / bottleneck streams
98
+ for j, feat in enumerate(enc_feats):
99
+ feat = _resize(feat, target)
100
+ parts.append(self.enc_projs[k][j](feat))
101
+
102
+ # All prior decoder streams (deeper → current scale = always upsample)
103
+ for i, df in enumerate(dec_outputs):
104
+ feat = F.interpolate(df, size=target, mode="bilinear", align_corners=False)
105
+ parts.append(self.dec_projs[k][i](feat))
106
+
107
+ dec_outputs.append(self.dec_convs[k](torch.cat(parts, dim=1)))
108
+
109
+ out = self.head(dec_outputs[-1])
110
+ if out.shape[2:] != input_size:
111
+ out = F.interpolate(out, size=input_size, mode="bilinear", align_corners=False)
112
+ return out
113
+
114
+
115
+ def _resize(feat: torch.Tensor, target: tuple[int, int] | torch.Size) -> torch.Tensor:
116
+ if feat.shape[2:] == torch.Size(target):
117
+ return feat
118
+ if feat.shape[2] > target[0]:
119
+ return F.adaptive_max_pool2d(feat, target)
120
+ return F.interpolate(feat, size=target, mode="bilinear", align_corners=False)
models/unetplusplus.py ADDED
@@ -0,0 +1,83 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+
5
+ from .backbones import create_backbone
6
+ from .blocks import ConvBlock
7
+
8
+
9
+ class UNetPlusPlus(nn.Module):
10
+ """UNet++ with dense nested skip connections.
11
+
12
+ Each node x[i][j] aggregates all previous same-scale nodes x[i][0..j-1]
13
+ plus an upsampled feature from the level below x[i+1][j-1], enabling
14
+ the decoder to learn progressively richer skip representations before
15
+ the final prediction at x[0][D].
16
+ See this: https://arxiv.org/abs/1807.10165
17
+ """
18
+
19
+ def __init__(
20
+ self,
21
+ backbone_name: str = "default",
22
+ in_channels: int = 3,
23
+ out_channels: int = 1,
24
+ **_kw,
25
+ ):
26
+ super().__init__()
27
+ self.backbone = create_backbone(backbone_name, in_channels)
28
+ channels = self.backbone.out_channels
29
+ skip_channels = list(channels[:-1]) # shallowest → deepest
30
+ bottleneck_ch = channels[-1]
31
+ self._D = len(skip_channels)
32
+
33
+ # all_ch[i] = channel count of the raw encoder node x[i][0]
34
+ all_ch = skip_channels + [bottleneck_ch]
35
+
36
+ # Dense nodes: x[i][j] for j > 0.
37
+ # x[i][j] concatenates: j same-scale predecessors + 1 upsampled from below.
38
+ # All intermediate node outputs keep skip_channels[i] channels.
39
+ self.nodes = nn.ModuleDict()
40
+ for j in range(1, self._D + 1):
41
+ for i in range(self._D - j + 1):
42
+ # channels coming from same scale: j tensors each with skip_channels[i] channels
43
+ from_same = j * skip_channels[i]
44
+ # channels coming from one level deeper (upsampled)
45
+ from_below = skip_channels[i + 1] if j > 1 else all_ch[i + 1]
46
+ self.nodes[f"{i}_{j}"] = ConvBlock(from_same + from_below, skip_channels[i])
47
+
48
+ self.head = nn.Conv2d(skip_channels[0], out_channels, kernel_size=1)
49
+
50
+ def forward(
51
+ self,
52
+ x: torch.Tensor | None = None,
53
+ pixel_values: torch.Tensor | None = None,
54
+ **_kw,
55
+ ) -> torch.Tensor:
56
+ if pixel_values is not None:
57
+ x = pixel_values
58
+ input_size = x.shape[2:]
59
+
60
+ skips, bottleneck = self.backbone(x)
61
+ D = self._D
62
+
63
+ # Initialise node cache with encoder outputs
64
+ cache: dict[tuple[int, int], torch.Tensor] = {}
65
+ for i, s in enumerate(skips):
66
+ cache[(i, 0)] = s
67
+ cache[(D, 0)] = bottleneck
68
+
69
+ # Fill the dense grid column by column (increasing j)
70
+ for j in range(1, D + 1):
71
+ for i in range(D - j + 1):
72
+ prev = [cache[(i, k)] for k in range(j)]
73
+ target_size = prev[0].shape[2:]
74
+ below = F.interpolate(
75
+ cache[(i + 1, j - 1)], size=target_size,
76
+ mode="bilinear", align_corners=False,
77
+ )
78
+ cache[(i, j)] = self.nodes[f"{i}_{j}"](torch.cat(prev + [below], dim=1))
79
+
80
+ out = self.head(cache[(0, D)])
81
+ if out.shape[2:] != input_size:
82
+ out = F.interpolate(out, size=input_size, mode="bilinear", align_corners=False)
83
+ return out