File size: 1,882 Bytes
a78ad5d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
import torch.nn as nn

from .attention_unet import AttentionUNet
from .resunet import ResUNet
from .transunet import TransUNet
from .unet import UNet
from .unet3plus import UNet3Plus
from .unetplusplus import UNetPlusPlus

# ---------------------------------------------------------------------------
# Architecture registry
# ---------------------------------------------------------------------------

ARCHITECTURES: dict[str, type] = {
    "unet": UNet,
    "unetplusplus": UNetPlusPlus,
    "unet3plus": UNet3Plus,
    "attention_unet": AttentionUNet,
    "resunet": ResUNet,
    "transunet": TransUNet,
}

# Backbones available for all custom hierarchical architectures
BACKBONES: list[str] = ["default", "efficientnet", "convnext", "swin", "siglip"]

# Valid (architecture, backbone) combinations
VALID_COMBINATIONS: dict[str, list[str]] = {
    arch: BACKBONES for arch in ARCHITECTURES
}

# All architecture names
ALL_ARCHITECTURES: list[str] = list(ARCHITECTURES)


# ---------------------------------------------------------------------------
# create_model — unified factory
# ---------------------------------------------------------------------------

def create_model(
    architecture: str,
    backbone: str,
    in_channels: int = 3,
    out_channels: int = 1,
    **kwargs,
) -> nn.Module:
    valid_bbs = VALID_COMBINATIONS.get(architecture)
    if valid_bbs is None:
        raise ValueError(
            f"Unknown architecture '{architecture}'. "
            f"Choose from {ALL_ARCHITECTURES}"
        )
    if backbone not in valid_bbs:
        raise ValueError(
            f"Backbone '{backbone}' is not compatible with '{architecture}'. "
            f"Valid choices: {valid_bbs}"
        )
    return ARCHITECTURES[architecture](
        backbone_name=backbone,
        in_channels=in_channels,
        out_channels=out_channels,
        **kwargs,
    )