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