Spaces:
Running on Zero
Running on Zero
File size: 1,293 Bytes
5878871 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 | """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
|