Skip to content

Commit 1d48ebd

Browse files
committed
Add unit tests for optimizer, model, and loss group configuration logging
1 parent 7663844 commit 1d48ebd

1 file changed

Lines changed: 306 additions & 0 deletions

File tree

Lines changed: 306 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,306 @@
1+
from types import SimpleNamespace
2+
3+
import pytest
4+
import torch
5+
6+
from virtual_stain_flow.models.model import BaseModel
7+
from virtual_stain_flow.vsf_logging.auto_loggers.model_config_logger import (
8+
AutoModelConfigLogger,
9+
)
10+
from virtual_stain_flow.vsf_logging.auto_loggers.optimizer_config_logger import (
11+
AutoOptimizerConfigLogger,
12+
)
13+
from virtual_stain_flow.vsf_logging.auto_loggers.loss_group_config_logger import (
14+
AutoLossGroupConfigLogger,
15+
)
16+
17+
18+
class _DummyLogger:
19+
def __init__(self):
20+
self.logged = []
21+
22+
def log_config(self, tag, config, stage=None):
23+
self.logged.append(
24+
{
25+
"tag": tag,
26+
"config": config,
27+
"stage": stage,
28+
}
29+
)
30+
31+
32+
class _FailingLogger(_DummyLogger):
33+
def log_config(self, tag, config, stage=None):
34+
raise RuntimeError("forced log failure")
35+
36+
37+
def _make_optimizer():
38+
model = torch.nn.Linear(4, 2)
39+
return torch.optim.Adam(model.parameters(), lr=1e-3)
40+
41+
42+
class _FakeLossGroup:
43+
def __init__(self, config):
44+
self._config = config
45+
46+
def get_config(self):
47+
return self._config
48+
49+
50+
class _FakeModel(BaseModel):
51+
def __init__(self, config):
52+
super().__init__()
53+
self._config = config
54+
55+
def forward(self, x):
56+
return x
57+
58+
def to_config(self):
59+
return self._config
60+
61+
@classmethod
62+
def from_config(cls, config):
63+
return cls(config)
64+
65+
66+
def test_discover_optimizers_supports_list_and_single():
67+
logger = _DummyLogger()
68+
auto_logger = AutoOptimizerConfigLogger(logger)
69+
70+
opt_a = _make_optimizer()
71+
opt_b = _make_optimizer()
72+
trainer = SimpleNamespace(optimizers=[opt_a], optimizer=opt_b)
73+
74+
optimizers = auto_logger.discover_optimizers(trainer)
75+
76+
assert optimizers == [opt_a, opt_b]
77+
78+
79+
def test_discover_optimizers_returns_empty_for_none_trainer():
80+
logger = _DummyLogger()
81+
auto_logger = AutoOptimizerConfigLogger(logger)
82+
83+
assert auto_logger.discover_optimizers(None) == []
84+
85+
86+
def test_log_optimizer_configs_sets_class_path_tags_and_artifacts(monkeypatch):
87+
logger = _DummyLogger()
88+
auto_logger = AutoOptimizerConfigLogger(logger)
89+
90+
captured_tags = {}
91+
92+
def fake_set_tag(key, value):
93+
captured_tags[key] = value
94+
95+
monkeypatch.setattr(
96+
"virtual_stain_flow.vsf_logging.auto_loggers.optimizer_config_logger.mlflow.set_tag",
97+
fake_set_tag,
98+
)
99+
100+
optimizer = _make_optimizer()
101+
trainer = SimpleNamespace(optimizer=optimizer)
102+
103+
auto_logger.log_optimizer_configs(trainer)
104+
105+
assert "optimizer.0.class_path" in captured_tags
106+
assert captured_tags["optimizer.0.class_path"].endswith("Adam")
107+
108+
assert len(logger.logged) == 1
109+
assert logger.logged[0]["tag"] == "optimizer_0"
110+
assert logger.logged[0]["config"]["class_path"].endswith("Adam")
111+
assert logger.logged[0]["config"]["defaults"]["lr"] == pytest.approx(1e-3)
112+
113+
114+
def test_log_optimizer_configs_skips_non_optimizer_entries(monkeypatch):
115+
logger = _DummyLogger()
116+
auto_logger = AutoOptimizerConfigLogger(logger)
117+
118+
captured_tags = {}
119+
120+
def fake_set_tag(key, value):
121+
captured_tags[key] = value
122+
123+
monkeypatch.setattr(
124+
"virtual_stain_flow.vsf_logging.auto_loggers.optimizer_config_logger.mlflow.set_tag",
125+
fake_set_tag,
126+
)
127+
128+
trainer = SimpleNamespace(optimizers=["not-an-optimizer"])
129+
130+
auto_logger.log_optimizer_configs(trainer)
131+
132+
assert captured_tags == {}
133+
assert logger.logged == []
134+
135+
136+
def test_log_optimizer_configs_swallows_log_config_failures(monkeypatch):
137+
logger = _FailingLogger()
138+
auto_logger = AutoOptimizerConfigLogger(logger)
139+
140+
def fake_set_tag(_key, _value):
141+
return None
142+
143+
monkeypatch.setattr(
144+
"virtual_stain_flow.vsf_logging.auto_loggers.optimizer_config_logger.mlflow.set_tag",
145+
fake_set_tag,
146+
)
147+
148+
trainer = SimpleNamespace(optimizer=_make_optimizer())
149+
150+
# Should not raise despite logger.log_config raising.
151+
auto_logger.log_optimizer_configs(trainer)
152+
153+
154+
def test_discover_models_prefers_models_list_over_single_model():
155+
logger = _DummyLogger()
156+
auto_logger = AutoModelConfigLogger(logger)
157+
158+
model_a = _FakeModel({"class_path": "pkg.ModelA", "init": {}})
159+
model_b = _FakeModel({"class_path": "pkg.ModelB", "init": {}})
160+
trainer = SimpleNamespace(_models=[model_a], model=model_b)
161+
162+
models = auto_logger._discover_models(trainer)
163+
164+
assert models == [model_a]
165+
166+
167+
def test_log_model_configs_sets_class_path_tag_and_artifact(monkeypatch):
168+
logger = _DummyLogger()
169+
auto_logger = AutoModelConfigLogger(logger)
170+
171+
captured_tags = {}
172+
173+
def fake_set_tag(key, value):
174+
captured_tags[key] = value
175+
176+
monkeypatch.setattr(
177+
"virtual_stain_flow.vsf_logging.auto_loggers.model_config_logger.mlflow.set_tag",
178+
fake_set_tag,
179+
)
180+
181+
model = _FakeModel({"class_path": "virtual_stain_flow.models.unet.UNet", "init": {"depth": 4}})
182+
trainer = SimpleNamespace(model=model)
183+
184+
auto_logger.log_model_configs(trainer)
185+
186+
assert captured_tags["model.0.class_path"].endswith("UNet")
187+
assert len(logger.logged) == 1
188+
assert logger.logged[0]["tag"] == "_FakeModel"
189+
assert logger.logged[0]["config"]["init"]["depth"] == 4
190+
191+
192+
def test_log_model_configs_skips_non_dict_configs(monkeypatch):
193+
logger = _DummyLogger()
194+
auto_logger = AutoModelConfigLogger(logger)
195+
196+
captured_tags = {}
197+
198+
def fake_set_tag(key, value):
199+
captured_tags[key] = value
200+
201+
monkeypatch.setattr(
202+
"virtual_stain_flow.vsf_logging.auto_loggers.model_config_logger.mlflow.set_tag",
203+
fake_set_tag,
204+
)
205+
206+
model = _FakeModel(["not", "a", "dict"])
207+
trainer = SimpleNamespace(model=model)
208+
209+
auto_logger.log_model_configs(trainer)
210+
211+
assert captured_tags == {}
212+
assert logger.logged == []
213+
214+
215+
def test_log_model_configs_swallows_log_config_failures(monkeypatch):
216+
logger = _FailingLogger()
217+
auto_logger = AutoModelConfigLogger(logger)
218+
219+
def fake_set_tag(_key, _value):
220+
return None
221+
222+
monkeypatch.setattr(
223+
"virtual_stain_flow.vsf_logging.auto_loggers.model_config_logger.mlflow.set_tag",
224+
fake_set_tag,
225+
)
226+
227+
model = _FakeModel({"class_path": "pkg.Model", "init": {}})
228+
trainer = SimpleNamespace(model=model)
229+
230+
# Should not raise despite logger.log_config raising.
231+
auto_logger.log_model_configs(trainer)
232+
233+
234+
def test_discover_loss_groups_supports_explicit_and_fallback_attrs():
235+
logger = _DummyLogger()
236+
auto_logger = AutoLossGroupConfigLogger(logger)
237+
238+
main_group = _FakeLossGroup([{"key": "MSELoss", "weight": 1.0}])
239+
gen_group = _FakeLossGroup([{"key": "L1Loss", "weight": 0.5}])
240+
trainer = SimpleNamespace(
241+
loss_groups={"main": main_group},
242+
_generator_loss_group=gen_group,
243+
)
244+
245+
loss_groups = auto_logger.discover_loss_groups(trainer)
246+
247+
assert set(loss_groups.keys()) == {"main", "generator"}
248+
assert loss_groups["main"] is main_group
249+
assert loss_groups["generator"] is gen_group
250+
251+
252+
def test_log_loss_group_configs_sets_tags_and_logs_config_artifact(monkeypatch):
253+
logger = _DummyLogger()
254+
auto_logger = AutoLossGroupConfigLogger(logger)
255+
256+
captured_tags = {}
257+
258+
def fake_set_tag(key, value):
259+
captured_tags[key] = value
260+
261+
monkeypatch.setattr(
262+
"virtual_stain_flow.vsf_logging.auto_loggers.loss_group_config_logger.mlflow.set_tag",
263+
fake_set_tag,
264+
)
265+
266+
group_items = [
267+
{"key": "MSELoss", "weight": 1.0},
268+
{"key": "L1Loss", "weight": 0.25},
269+
{"key": None, "weight": None},
270+
"ignored-non-dict-item",
271+
]
272+
trainer = SimpleNamespace(loss_groups={"main": _FakeLossGroup(group_items)})
273+
274+
auto_logger.log_loss_group_configs(trainer)
275+
276+
assert captured_tags["loss.main.0.name"] == "MSELoss"
277+
assert captured_tags["loss.main.0.weight"] == "1.0"
278+
assert captured_tags["loss.main.1.name"] == "L1Loss"
279+
assert captured_tags["loss.main.1.weight"] == "0.25"
280+
281+
assert len(logger.logged) == 1
282+
assert logger.logged[0]["tag"] == "loss_group_main"
283+
assert logger.logged[0]["config"]["group_name"] == "main"
284+
assert logger.logged[0]["config"]["items"] == group_items
285+
286+
287+
def test_log_loss_group_configs_skips_non_list_config(monkeypatch):
288+
logger = _DummyLogger()
289+
auto_logger = AutoLossGroupConfigLogger(logger)
290+
291+
captured_tags = {}
292+
293+
def fake_set_tag(key, value):
294+
captured_tags[key] = value
295+
296+
monkeypatch.setattr(
297+
"virtual_stain_flow.vsf_logging.auto_loggers.loss_group_config_logger.mlflow.set_tag",
298+
fake_set_tag,
299+
)
300+
301+
trainer = SimpleNamespace(loss_groups={"main": _FakeLossGroup({"not": "a-list"})})
302+
303+
auto_logger.log_loss_group_configs(trainer)
304+
305+
assert captured_tags == {}
306+
assert logger.logged == []

0 commit comments

Comments
 (0)