"""Model factory: shared semantic + drywall-material heads over several backbones. The original prototype hard-coded MobileNetV3-Large/DeepLabV3. That is small and fast but well behind modern open segmentation models. This module keeps the two-head design (semantic coverage + independent drywall substrate) but lets you choose the encoder: * ``deeplabv3_mobilenet_v3_large`` - the original; edge/latency oriented. * ``deeplabv3_resnet50`` / ``deeplabv3_resnet101`` - stronger torchvision baselines (ASPP + ResNet), no extra dependencies. * ``segformer_b0`` .. ``segformer_b5`` - modern transformer segmentation from Hugging Face (requires ``transformers``); the recommended accuracy option. All variants return ``{"semantic": (B, C, H, W), "drywall": (B, 2, H, W)}`` at input resolution, so ``train.py`` and ``predict.py`` are architecture-agnostic. Normalisation is shared (ImageNet mean/std), which SegFormer also expects. """ import torch from torch import nn from torch.nn import functional as F from torchvision.models import MobileNet_V3_Large_Weights, ResNet50_Weights, ResNet101_Weights from torchvision.models.segmentation import ( deeplabv3_mobilenet_v3_large, deeplabv3_resnet50, deeplabv3_resnet101, ) # torchvision DeepLabV3 variants: name -> (builder, pretrained backbone weights). DEEPLAB_ARCHITECTURES = { "deeplabv3_mobilenet_v3_large": (deeplabv3_mobilenet_v3_large, MobileNet_V3_Large_Weights.DEFAULT), "deeplabv3_resnet50": (deeplabv3_resnet50, ResNet50_Weights.DEFAULT), "deeplabv3_resnet101": (deeplabv3_resnet101, ResNet101_Weights.DEFAULT), } # SegFormer variants: name -> Hugging Face id of a pretrained encoder/decoder. SEGFORMER_ARCHITECTURES = { "segformer_b0": "nvidia/segformer-b0-finetuned-ade-512-512", "segformer_b1": "nvidia/segformer-b1-finetuned-ade-512-512", "segformer_b2": "nvidia/segformer-b2-finetuned-ade-512-512", "segformer_b3": "nvidia/segformer-b3-finetuned-ade-512-512", "segformer_b4": "nvidia/segformer-b4-finetuned-ade-512-512", # b5 has no 512-input ADE20K checkpoint on the Hub; the only b5 release is 640. "segformer_b5": "nvidia/segformer-b5-finetuned-ade-640-640", } DEFAULT_ARCH = "deeplabv3_mobilenet_v3_large" def available_architectures(): return sorted(list(DEEPLAB_ARCHITECTURES) + list(SEGFORMER_ARCHITECTURES)) def _drywall_head(channels, hidden=128, dropout=0.1): return nn.Sequential( nn.Conv2d(channels, hidden, kernel_size=3, padding=1, bias=False), nn.BatchNorm2d(hidden), nn.ReLU(inplace=True), nn.Dropout2d(dropout), nn.Conv2d(hidden, 2, kernel_size=1), ) class WallPaintNet(nn.Module): """torchvision DeepLabV3 semantic head plus an independent drywall head.""" def __init__(self, num_semantic_classes, pretrained_backbone=False, arch=DEFAULT_ARCH): super().__init__() if arch not in DEEPLAB_ARCHITECTURES: raise ValueError(f"unknown DeepLab architecture {arch!r}; choose from {sorted(DEEPLAB_ARCHITECTURES)}") builder, default_weights = DEEPLAB_ARCHITECTURES[arch] weights = default_weights if pretrained_backbone else None self.arch = arch self.segmenter = builder(weights=None, weights_backbone=weights, num_classes=num_semantic_classes) was_training = self.segmenter.training self.segmenter.eval() with torch.inference_mode(): features = self.segmenter.backbone(torch.zeros(1, 3, 128, 128))["out"] self.segmenter.train(was_training) self.drywall_head = _drywall_head(features.shape[1]) def forward(self, x): size = x.shape[-2:] features = self.segmenter.backbone(x) semantic = self.segmenter.classifier(features["out"]) drywall = self.drywall_head(features["out"]) return { "semantic": F.interpolate(semantic, size=size, mode="bilinear", align_corners=False), "drywall": F.interpolate(drywall, size=size, mode="bilinear", align_corners=False), } class SegformerPaintNet(nn.Module): """Hugging Face SegFormer semantic head plus an independent drywall head. SegFormer predicts logits at 1/4 resolution; the last encoder hidden state feeds the drywall head. Heavy enough to matter, but a genuine modern transformer backbone rather than a 2021 MobileNet. """ def __init__(self, num_semantic_classes, pretrained_backbone=False, arch="segformer_b2"): super().__init__() if arch not in SEGFORMER_ARCHITECTURES: raise ValueError(f"unknown SegFormer architecture {arch!r}; choose from {sorted(SEGFORMER_ARCHITECTURES)}") try: from transformers import SegformerConfig, SegformerForSemanticSegmentation except ImportError as exc: # pragma: no cover - optional dependency raise ImportError("SegFormer architectures require `pip install transformers`") from exc hf_id = SEGFORMER_ARCHITECTURES[arch] self.arch = arch if pretrained_backbone: self.segmenter = SegformerForSemanticSegmentation.from_pretrained( hf_id, num_labels=num_semantic_classes, ignore_mismatched_sizes=True) else: config = SegformerConfig.from_pretrained(hf_id, num_labels=num_semantic_classes) self.segmenter = SegformerForSemanticSegmentation(config) self.drywall_head = _drywall_head(self.segmenter.config.hidden_sizes[-1]) def forward(self, x): size = x.shape[-2:] outputs = self.segmenter(pixel_values=x, output_hidden_states=True) semantic = F.interpolate(outputs.logits, size=size, mode="bilinear", align_corners=False) drywall = self.drywall_head(outputs.hidden_states[-1]) drywall = F.interpolate(drywall, size=size, mode="bilinear", align_corners=False) return {"semantic": semantic, "drywall": drywall} def build_model(arch, num_semantic_classes, pretrained_backbone=False): """Construct the dual-head model for ``arch`` (see :func:`available_architectures`).""" if arch in DEEPLAB_ARCHITECTURES: return WallPaintNet(num_semantic_classes, pretrained_backbone, arch) if arch in SEGFORMER_ARCHITECTURES: return SegformerPaintNet(num_semantic_classes, pretrained_backbone, arch) raise ValueError(f"unknown architecture {arch!r}; choose from {available_architectures()}")