Skip to content

Commit 45d4504

Browse files
authored
Merge pull request #12 from NetherlandsForensicInstitute/serialize-models
BSZ-213: Add functionality to serialize/deserialize LR models for given mark/score types
2 parents 3b97aee + d5d0fab commit 45d4504

8 files changed

Lines changed: 174 additions & 79 deletions

File tree

.gitignore

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
# Project specific
33
lrsystem_output/
44
.testmondata
5+
tests/test_model_storage/
56

67
# Byte-compiled / optimized / DLL files
78
__pycache__/

lrmodule/__init__.py

Lines changed: 7 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -6,34 +6,32 @@
66
from lir.lrsystems.lrsystems import LRSystem
77

88
from lrmodule import persistence
9-
from lrmodule.data import get_dataset_id
109
from lrmodule.data_types import ModelSettings
1110
from lrmodule.lrsystem import get_trained_model
1211

1312

14-
def get_model(settings: ModelSettings, training_data: FeatureData, cache_dir: Path | None) -> LRSystem:
13+
def get_model(settings: ModelSettings, training_data: FeatureData, model_storage_path: Path | None) -> LRSystem:
1514
"""
1615
Obtain a model by loading it from disk, or by fitting it from training data.
1716
1817
:param settings: model settings
1918
:param training_data: training data
20-
:param cache_dir: cache dir
19+
:param model_storage_path: path where trained LR models are stored
2120
:return: a fitted LR system
2221
"""
23-
dataset_id = get_dataset_id(training_data)
24-
model = None if not cache_dir else persistence.load_model(settings, dataset_id, cache_dir)
22+
model = None if not model_storage_path else persistence.load_model(settings, model_storage_path)
2523
if not model:
2624
model = get_trained_model(settings, training_data)
27-
if cache_dir:
28-
persistence.save_model(model, settings, dataset_id, cache_dir)
25+
if model_storage_path:
26+
persistence.save_model(model, settings, model_storage_path)
2927
return model
3028

3129

3230
def calculate_llrs(
33-
features: np.ndarray, settings: ModelSettings, training_data: FeatureData, cache_dir: Path | None
31+
features: np.ndarray, settings: ModelSettings, training_data: FeatureData, model_storage_path: Path | None
3432
) -> LLRData:
3533
"""Calculate LLRs after fitting a model with a training set."""
36-
model = get_model(settings, training_data, cache_dir)
34+
model = get_model(settings, training_data, model_storage_path)
3735
return model.apply(FeatureData(features=features))
3836

3937

lrmodule/data.py

Lines changed: 0 additions & 12 deletions
This file was deleted.

lrmodule/persistence.py

Lines changed: 33 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1,25 +1,46 @@
1-
from hashlib import sha256
1+
import os
2+
import pickle
23
from pathlib import Path
34

45
from lir.lrsystems.lrsystems import LRSystem
56

67
from lrmodule.data_types import ModelSettings
78

89

9-
def _get_model_dirname(settings: ModelSettings, dataset_id: str) -> str:
10-
h = sha256()
11-
h.update(str(settings).encode("utf8"))
12-
h.update(dataset_id.encode("utf8"))
13-
return h.hexdigest()
10+
def _get_model_filename(settings: ModelSettings) -> str:
11+
"""Construct model filename based on mark and score type."""
12+
mark_type = settings.mark_type.value
13+
score_type = settings.score_type.value
1414

15+
return f"{mark_type}_{score_type}_model.pkl"
1516

16-
def load_model(settings: ModelSettings, dataset_id: str, cache_dir: Path) -> LRSystem | None:
17+
18+
def load_model(settings: ModelSettings, model_storage_path: Path) -> LRSystem:
1719
"""Load previously cached model."""
18-
_ = cache_dir / _get_model_dirname(settings, dataset_id) / "model.pkl"
19-
raise NotImplementedError
20+
model_filename = _get_model_filename(settings)
21+
model_file_path = model_storage_path / model_filename
22+
23+
mark_type = settings.mark_type.value
24+
score_type = settings.score_type.value
25+
26+
if not model_file_path.exists():
27+
raise FileNotFoundError(f"No model found for mark type '{mark_type}', score type: '{score_type}'.")
2028

29+
try:
30+
with open(model_file_path, "rb") as f:
31+
# It is assumed exclusively `LRSystem` models will be loaded, which are considered safe
32+
return pickle.load(f) # noqa: S301
33+
except Exception:
34+
raise RuntimeError(
35+
f"Could not load model from .pkl file for mark type '{mark_type}', score type: '{score_type}'"
36+
)
2137

22-
def save_model(model: LRSystem, settings: ModelSettings, dataset_id: str, cache_dir: Path) -> None:
38+
39+
def save_model(model: LRSystem, settings: ModelSettings, model_storage_path: Path) -> None:
2340
"""Save a model to disk."""
24-
_ = cache_dir / _get_model_dirname(settings, dataset_id)
25-
raise NotImplementedError
41+
model_filename = _get_model_filename(settings)
42+
model_file_path = model_storage_path / model_filename
43+
44+
os.makedirs(os.path.dirname(model_file_path), exist_ok=True)
45+
with open(model_file_path, "wb") as f:
46+
f.write(pickle.dumps(model))

tests/conftest.py

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,32 @@
1+
import pytest
2+
from lir.data.datasets.synthesized_normal_binary import SynthesizedNormalBinaryData, SynthesizedNormalDataClass
3+
from lir.data.models import FeatureData
4+
from lir.lrsystems.lrsystems import LRSystem
5+
6+
from lrmodule.data_types import MarkType, ModelSettings, ScoreType
7+
from lrmodule.lrsystem import load_lrsystem
8+
9+
10+
@pytest.fixture
11+
def sample_feature_data() -> FeatureData:
12+
"""Provide FeatureData collection of synthesized normal binary data."""
13+
data = SynthesizedNormalBinaryData(
14+
data_classes={
15+
0: SynthesizedNormalDataClass(mean=-1, std=1, size=100),
16+
1: SynthesizedNormalDataClass(mean=1, std=1, size=100),
17+
},
18+
seed=0,
19+
)
20+
data = data.get_instances()
21+
data = data.replace(features=data.features.flatten())
22+
23+
return data
24+
25+
26+
@pytest.fixture
27+
def trained_lr_system(sample_feature_data: FeatureData) -> LRSystem:
28+
"""Provide a basic trained LR system model based on specific settings and data."""
29+
lrsystem = load_lrsystem(ModelSettings(MarkType.FIRING_PIN_IMPRESSION, ScoreType.ACCF))
30+
lrsystem.fit(sample_feature_data)
31+
32+
return lrsystem

tests/test_data.py

Lines changed: 0 additions & 12 deletions
This file was deleted.

tests/test_lrsystem.py

Lines changed: 4 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -1,27 +1,7 @@
1-
from lir.data.datasets.synthesized_normal_binary import SynthesizedNormalBinaryData, SynthesizedNormalDataClass
2-
from lrmodule.data_types import MarkType, ModelSettings, ScoreType
3-
from lrmodule.lrsystem import load_lrsystem
1+
from lir.data.models import FeatureData
2+
from lir.lrsystems.lrsystems import LRSystem
43

54

6-
def test_load_lrsystem():
7-
load_lrsystem(ModelSettings(MarkType.FIRING_PIN_IMPRESSION, ScoreType.ACCF))
8-
9-
10-
def test_run_lrsystem():
11-
lrsystem = load_lrsystem(ModelSettings(MarkType.FIRING_PIN_IMPRESSION, ScoreType.ACCF))
12-
data = SynthesizedNormalBinaryData(
13-
data_classes={
14-
0: SynthesizedNormalDataClass(mean=-1, std=1, size=100),
15-
1: SynthesizedNormalDataClass(mean=1, std=1, size=100),
16-
},
17-
seed=0,
18-
)
19-
data = data.get_instances()
20-
data = data.replace(features=data.features.flatten())
21-
llrs = lrsystem.fit(data).apply(data)
5+
def test_run_lrsystem(trained_lr_system: LRSystem, sample_feature_data: FeatureData):
6+
llrs = trained_lr_system.apply(sample_feature_data)
227
assert llrs.features.shape == (200, 3)
23-
24-
25-
if __name__ == "__main__":
26-
test_load_lrsystem()
27-
test_run_lrsystem()

tests/test_persistence.py

Lines changed: 97 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,20 +1,107 @@
1+
import os
2+
import shutil
13
from pathlib import Path
4+
from pickle import UnpicklingError
5+
from unittest import mock
26

37
import pytest
8+
from lir.data.models import FeatureData
9+
from lir.lrsystems.lrsystems import LRSystem
10+
411
from lrmodule import ModelSettings
512
from lrmodule.data_types import MarkType, ScoreType
6-
from lrmodule.lrsystem import load_lrsystem
7-
from lrmodule.persistence import load_model, save_model
13+
from lrmodule.persistence import load_model, save_model, _get_model_filename
14+
15+
16+
MODEL_STORAGE_PATH = Path(__file__).parent / "test_model_storage"
17+
18+
19+
@pytest.fixture(autouse=True)
20+
def clear_test_model_storage_directory():
21+
"""Clean up 'test_model_storage' directory before running each test.
22+
23+
This ensures a fresh environment for each test. The generated artifacts are not
24+
cleaned up after each test to allow easy debugging of the generated pickle files.
25+
"""
26+
if MODEL_STORAGE_PATH.exists():
27+
shutil.rmtree(MODEL_STORAGE_PATH)
28+
MODEL_STORAGE_PATH.mkdir(parents=True)
29+
30+
31+
def test_serialize_trained_lr_system(trained_lr_system: LRSystem):
32+
"""Check that a trained LR system can be serialized."""
33+
# Given that we have a trained LR system
34+
settings = ModelSettings(MarkType.FIRING_PIN_IMPRESSION, ScoreType.ACCF)
35+
mark_type = settings.mark_type.value
36+
score_type = settings.score_type.value
37+
38+
# When we serialize the LR system
39+
save_model(trained_lr_system, settings, MODEL_STORAGE_PATH)
40+
41+
# There should be a file we can load
42+
model_filename = _get_model_filename(settings)
43+
model_file_path = MODEL_STORAGE_PATH / model_filename
44+
assert model_file_path.exists()
45+
46+
47+
def test_deserialize_trained_lr_system(trained_lr_system: LRSystem, sample_feature_data: FeatureData):
48+
"""Check that a deserialized, trained LR system yields exactly the same results."""
49+
# Given that we have a certain LR system serialized
50+
settings = ModelSettings(MarkType.FIRING_PIN_IMPRESSION, ScoreType.ACCF)
51+
save_model(trained_lr_system, settings, MODEL_STORAGE_PATH)
52+
53+
# When the model is deserialized
54+
deserialized_model = load_model(settings, MODEL_STORAGE_PATH)
55+
56+
# The deserialized model and the model it originated from should be of the same type of LR system
57+
assert type(trained_lr_system) == type(trained_lr_system)
58+
59+
# The calculated LLR output should be identical to the LR system output of the serialized model
60+
expected_llr_data = trained_lr_system.apply(sample_feature_data)
61+
deserialized_model_data = deserialized_model.apply(sample_feature_data)
62+
63+
assert deserialized_model_data == expected_llr_data
64+
65+
66+
@pytest.mark.parametrize('mark_type,score_type', [
67+
(MarkType.FIRING_PIN_IMPRESSION, ScoreType.CMC), # other score type
68+
(MarkType.BREECH_PIN_IMPRESSION, ScoreType.ACCF), # other mark type
69+
(MarkType.BREECH_PIN_IMPRESSION, ScoreType.CMC), # other mark and other score type
70+
])
71+
def test_deserialize_inexistent_lr_system(trained_lr_system: LRSystem, mark_type: MarkType, score_type: ScoreType):
72+
"""Check that an appropriate error is raised when there is no serialized model."""
73+
# Given that the LR model storage directory is empty
74+
assert os.listdir(MODEL_STORAGE_PATH) == []
75+
76+
# Given that we have a serialized model for a given type of `ModelSettings`
77+
settings = ModelSettings(MarkType.FIRING_PIN_IMPRESSION, ScoreType.ACCF)
78+
save_model(trained_lr_system, settings, MODEL_STORAGE_PATH)
79+
assert len(os.listdir(MODEL_STORAGE_PATH)) == 1
80+
81+
# When we try to deserialize a model for a different type of `ModelSettings`
82+
other_settings = ModelSettings(mark_type, score_type)
83+
84+
# An exception should be raised mentioning that we can't find that particular deserialized LR model
85+
with pytest.raises(FileNotFoundError) as exception_info:
86+
load_model(other_settings, MODEL_STORAGE_PATH)
87+
88+
# The exception should mention no models found for the requested mark/score types
89+
assert "No model found for mark type" in str(exception_info.value)
90+
assert other_settings.mark_type.value in str(exception_info.value)
91+
assert other_settings.score_type.value in str(exception_info.value)
892

993

10-
def test_persistence():
94+
def test_deserialize_from_invalid_pickle_file(trained_lr_system: LRSystem):
95+
"""Check that an appropriate error is raised when unable to unpickle serialized model."""
96+
# Given that we have a serialized model for a given type of `ModelSettings`
1197
settings = ModelSettings(MarkType.FIRING_PIN_IMPRESSION, ScoreType.ACCF)
98+
save_model(trained_lr_system, settings, MODEL_STORAGE_PATH)
1299

13-
# not implemented
14-
with pytest.raises(Exception):
15-
load_model(settings, "dataset_id", Path("/"))
100+
with mock.patch('pickle.load', side_effect=UnpicklingError("Some pickle error")):
101+
# When pickle can't load the given file, we expect an appropriate error to be raised
102+
with pytest.raises(RuntimeError) as exception_info:
103+
load_model(settings, MODEL_STORAGE_PATH)
16104

17-
# not implemented
18-
lrsystem = load_lrsystem(settings)
19-
with pytest.raises(Exception):
20-
save_model(lrsystem, settings, "dataset_id", Path("/"))
105+
assert "Could not load model from .pkl file for mark type" in str(exception_info.value)
106+
assert settings.mark_type.value in str(exception_info.value)
107+
assert settings.score_type.value in str(exception_info.value)

0 commit comments

Comments
 (0)