#!/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, 2, 128, 128), (arch, output["drywall"].shape) print(f"ok {arch}") 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()