constructelligence's picture
Upload models.py with huggingface_hub
870d5df verified
Raw History Blame Contribute Delete
10.8 kB
"""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, M, H, W)}`` at
input resolution, where M is the number of material classes (3 with wallpaper;
2 for checkpoints trained before wallpaper existed), so ``train.py`` and ``predict.py`` are architecture-agnostic.
Models built with condition classes also return ``"condition": (B, K, H, W)`` -
surface condition (sound / paint failure / rust) from a third head on the same
features; checkpoints from before that head load without it.
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,
)
from label_schema import CONDITION_CLASSES, LEGACY_MATERIAL_CLASSES, MATERIAL_CLASSES
# 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, num_material_classes, 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, num_material_classes, 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,
num_material_classes=len(MATERIAL_CLASSES), num_condition_classes=len(CONDITION_CLASSES)):
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], num_material_classes)
self.condition_head = (_drywall_head(features.shape[1], num_condition_classes)
if num_condition_classes else None)
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"])
out = {
"semantic": F.interpolate(semantic, size=size, mode="bilinear", align_corners=False),
"drywall": F.interpolate(drywall, size=size, mode="bilinear", align_corners=False),
}
if self.condition_head is not None:
out["condition"] = F.interpolate(self.condition_head(features["out"]), size=size, mode="bilinear",
align_corners=False)
return out
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",
num_material_classes=len(MATERIAL_CLASSES), num_condition_classes=len(CONDITION_CLASSES)):
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], num_material_classes)
# Peeling edges and rust streaks are fine detail: the condition head reads the
# highest-resolution encoder stage (1/4 scale) rather than the coarsest one.
self.condition_head = (_drywall_head(self.segmenter.config.hidden_sizes[0], num_condition_classes)
if num_condition_classes else None)
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)
out = {"semantic": semantic, "drywall": drywall}
if self.condition_head is not None:
out["condition"] = F.interpolate(self.condition_head(outputs.hidden_states[0]), size=size,
mode="bilinear", align_corners=False)
return out
def build_model(arch, num_semantic_classes, pretrained_backbone=False, num_material_classes=len(MATERIAL_CLASSES),
num_condition_classes=len(CONDITION_CLASSES)):
"""Construct the model for ``arch`` (see :func:`available_architectures`); 0 condition classes = no condition head."""
if arch in DEEPLAB_ARCHITECTURES:
return WallPaintNet(num_semantic_classes, pretrained_backbone, arch, num_material_classes, num_condition_classes)
if arch in SEGFORMER_ARCHITECTURES:
return SegformerPaintNet(num_semantic_classes, pretrained_backbone, arch, num_material_classes,
num_condition_classes)
raise ValueError(f"unknown architecture {arch!r}; choose from {available_architectures()}")
def checkpoint_material_classes(checkpoint):
"""Material classes a checkpoint was trained with (pre-wallpaper checkpoints have two)."""
names = tuple(checkpoint.get("material_classes") or ())
if names not in (MATERIAL_CLASSES, LEGACY_MATERIAL_CLASSES):
raise ValueError(f"checkpoint material classes {names!r} are not a known schema; retrain with drywall_masks")
return names
def checkpoint_condition_classes(checkpoint):
"""Condition classes a checkpoint was trained with; () for checkpoints from before the condition head."""
names = tuple(checkpoint.get("condition_classes") or ())
if names not in ((), CONDITION_CLASSES):
raise ValueError(f"checkpoint condition classes {names!r} are not a known schema")
return names
def build_from_checkpoint(checkpoint, num_semantic_classes):
"""Rebuild and load the model a checkpoint describes: 2 or 3 materials, with or without a condition head."""
model = build_model(str(checkpoint.get("arch", DEFAULT_ARCH)), num_semantic_classes, pretrained_backbone=False,
num_material_classes=len(checkpoint_material_classes(checkpoint)),
num_condition_classes=len(checkpoint_condition_classes(checkpoint)))
model.load_state_dict(checkpoint["model"])
return model
def warm_start(model, checkpoint):
"""Initialise ``model`` from a checkpoint of the same arch, growing the material head if needed.
A pre-wallpaper checkpoint has a 2-way material classifier; its rows are
copied into the new 3-way head and the wallpaper row keeps its fresh init,
so facade/wall knowledge survives while wallpaper is learned. A checkpoint
without a condition head leaves that head at its fresh init. Returns the
names of tensors that could not be copied (empty when everything matched).
"""
target = model.state_dict()
skipped = []
for name, value in checkpoint["model"].items():
if name not in target:
skipped.append(name)
elif target[name].shape == value.shape:
target[name] = value
elif name.startswith("drywall_head.") and value.shape[1:] == target[name].shape[1:] \
and value.shape[0] < target[name].shape[0]:
target[name][:value.shape[0]] = value
else:
skipped.append(name)
model.load_state_dict(target)
return skipped