#!/usr/bin/env python3 """Architecture-factory tests. Skipped automatically when torch is unavailable. Run with ``python3 test_models.py``. Validates that every torchvision backbone builds, runs a forward pass, and returns both heads at input resolution, and that an unknown architecture is rejected. """ import sys def main(): try: import torch except ImportError: print("torch not installed; skipping model architecture tests") return from models import DEFAULT_ARCH, available_architectures, build_model architectures = available_architectures() assert DEFAULT_ARCH in architectures assert {"deeplabv3_mobilenet_v3_large", "deeplabv3_resnet50", "segformer_b2"} <= set(architectures) for arch in ("deeplabv3_mobilenet_v3_large", "deeplabv3_resnet50"): model = build_model(arch, 10, pretrained_backbone=False).eval() with torch.inference_mode(): output = model(torch.zeros(1, 3, 128, 128)) assert output["semantic"].shape == (1, 10, 128, 128), (arch, output["semantic"].shape) assert output["drywall"].shape == (1, 3, 128, 128), (arch, output["drywall"].shape) assert output["condition"].shape == (1, 3, 128, 128), (arch, output["condition"].shape) print(f"ok {arch}") # Checkpoints from before the wallpaper class (two material outputs) must still load. from label_schema import LEGACY_MATERIAL_CLASSES, MATERIAL_CLASSES from models import build_from_checkpoint for names in (LEGACY_MATERIAL_CLASSES, MATERIAL_CLASSES): # Checkpoints from before the condition head carry no condition_classes and no head weights. source = build_model(DEFAULT_ARCH, 10, num_material_classes=len(names), num_condition_classes=0) checkpoint = {"arch": DEFAULT_ARCH, "material_classes": names, "model": source.state_dict()} model = build_from_checkpoint(checkpoint, 10).eval() with torch.inference_mode(): output = model(torch.zeros(1, 3, 64, 64)) assert output["drywall"].shape[1] == len(names) and "condition" not in output from label_schema import CONDITION_CLASSES source = build_model(DEFAULT_ARCH, 10) checkpoint = {"arch": DEFAULT_ARCH, "material_classes": MATERIAL_CLASSES, "condition_classes": CONDITION_CLASSES, "model": source.state_dict()} with torch.inference_mode(): assert build_from_checkpoint(checkpoint, 10).eval()(torch.zeros(1, 3, 64, 64))["condition"].shape[1] == 3 print("ok checkpoints with and without the condition head load") # Warm start grows a 2-way material head to 3 ways, keeping the old rows. from models import warm_start old = build_model(DEFAULT_ARCH, 10, num_material_classes=2, num_condition_classes=0) new = build_model(DEFAULT_ARCH, 10, num_material_classes=3) fresh_condition = new.condition_head[-1].weight.detach().clone() assert warm_start(new, {"model": old.state_dict()}) == [] assert torch.equal(new.condition_head[-1].weight, fresh_condition) # new head keeps its own init w_old, w_new = old.drywall_head[-1].weight, new.drywall_head[-1].weight assert torch.equal(w_new[:2], w_old) and w_new.shape[0] == 3 print("ok warm start grows the material head") try: build_from_checkpoint({"arch": DEFAULT_ARCH, "material_classes": ("x",), "model": {}}, 10) except ValueError: print("ok legacy and wallpaper material heads load; unknown schema rejected") else: raise AssertionError("unknown material schema should raise ValueError") try: build_model("not_a_real_arch", 10) except ValueError: print("ok unknown architecture rejected") else: raise AssertionError("unknown architecture should raise ValueError") print("model architecture tests passed") if __name__ == "__main__": main()