Upload attention_unet+convnext model, code, and model card
Browse files- README.md +160 -0
- config.json +27 -0
- loss.py +94 -0
- model.safetensors +3 -0
- models/__init__.py +62 -0
- models/attention_unet.py +44 -0
- models/backbones.py +317 -0
- models/blocks.py +140 -0
- models/resunet.py +47 -0
- models/transunet.py +71 -0
- models/unet.py +44 -0
- models/unet3plus.py +120 -0
- models/unetplusplus.py +83 -0
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
|