constructelligence's picture
Upload test_models.py with huggingface_hub
185cf2c verified
Raw History Blame Contribute Delete
3.88 kB
#!/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()