SingleViewHeigthEstimation / tests /test_model_loading.py
akashch1512's picture
fix: match v5 checkpoint architecture and benchmark scaled inference
5878871 verified
Raw History Blame Contribute Delete
1.29 kB
"""Checkpoint configuration must be restored and mismatches must fail."""
import sys
import types
import pytest
import torch
from infer.predict import load_model
@pytest.mark.parametrize("mismatch", [False, True])
def test_checkpoint_loading_restores_config_and_rejects_mismatch(tmp_path, monkeypatch, mismatch):
class FakeNet(torch.nn.Module):
def __init__(self, cfg):
super().__init__()
assert cfg.detail_branch is True
assert cfg.encoder_feature_indices == (1, 2, 3, 4)
self.weight = torch.nn.Parameter(torch.zeros(1))
monkeypatch.setitem(sys.modules, "models.heads", types.SimpleNamespace(DepthWizardNet=FakeNet))
path = tmp_path / "best.pt"
state = {"wrong" if mismatch else "weight": torch.tensor([7.0])}
torch.save({"model": state, "epoch": 7, "encoder_included": True,
"config": {"detail_branch": True, "encoder_feature_indices": [1, 2, 3, 4]}}, path)
if mismatch:
with pytest.raises(RuntimeError, match="does not match"):
load_model(str(path), torch.device("cpu"))
else:
model, _, cfg = load_model(str(path), torch.device("cpu"))
assert model.weight.item() == 7.0
assert cfg.checkpoint_epoch == 7 and cfg.checkpoint_tensor_count == 1