"""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