Spaces:
Running on Zero
Running on Zero
Download tests/test_model_loading.py from akashch1512/SingleViewHeigthEstimation: direct link, hf CLI and curl.
- Browser
- Download file 1.29 kB
-
https://huggingface.co/spaces/akashch1512/SingleViewHeigthEstimation/resolve/main/tests/test_model_loading.py
- Command line
-
hf download hf://spaces/akashch1512/SingleViewHeigthEstimation/tests/test_model_loading.py
-
curl -L -o test_model_loading.py https://huggingface.co/spaces/akashch1512/SingleViewHeigthEstimation/resolve/main/tests/test_model_loading.py
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 | |
| 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 | |