Skip to content

Commit 86acbb5

Browse files
committed
fix: params_test
the params test used pytest style instead of unittest style
1 parent 6b2f0fb commit 86acbb5

1 file changed

Lines changed: 159 additions & 170 deletions

File tree

tests/params_test.py

Lines changed: 159 additions & 170 deletions
Original file line numberDiff line numberDiff line change
@@ -1,175 +1,164 @@
11
"""Test if setting params via the named subset works correctly."""
22

3-
import pytest
3+
import unittest
44

5-
from lymph import models
65
from lymph.types import ExtraParamsError
76

8-
from .fixtures import (
9-
RNG,
10-
)
11-
12-
13-
def test_set_named_params_default_behavior(
14-
binary_unilateral_model: models.Unilateral,
15-
) -> None:
16-
"""Ensure `set_named_params` works as `set_params` when no `named_params` set."""
17-
params = binary_unilateral_model.get_params(as_dict=True)
18-
new_params = {param: RNG.uniform() for param in params.keys()}
19-
binary_unilateral_model.set_named_params(**new_params)
20-
assert new_params == binary_unilateral_model.get_params(as_dict=True)
21-
22-
23-
def test_named_params_setter(
24-
binary_unilateral_model: models.Unilateral,
25-
binary_bilateral_model: models.Midline,
26-
) -> None:
27-
"""Check that setting `named_params` works correctly."""
28-
with pytest.raises(ValueError):
29-
binary_unilateral_model.named_params = ["invalid identifier!"]
30-
31-
with pytest.raises(ValueError):
32-
binary_unilateral_model.named_params = 123
33-
34-
params = binary_unilateral_model.get_params(as_dict=True).keys()
35-
params_subset = [param for param in params if RNG.uniform() > 0.5]
36-
binary_unilateral_model.named_params = params_subset
37-
38-
for stored, subset in zip(
39-
binary_unilateral_model.named_params,
40-
params_subset,
41-
strict=True,
42-
):
43-
assert stored == subset
44-
45-
binary_bilateral_model.named_params = ["ipsi_spread"]
46-
assert binary_bilateral_model.named_params == ["ipsi_spread"]
47-
48-
49-
def test_set_named_params_named_easy_subset(
50-
binary_unilateral_model: models.Unilateral,
51-
) -> None:
52-
"""Ensure `set_named_params` works correctly with an easy subset.
53-
54-
An "easy subset" is a literal subset of the params.
55-
"""
56-
params = binary_unilateral_model.get_params(as_dict=True)
57-
new_params = {param: RNG.uniform() for param in params.keys()}
58-
params_subset = {k: RNG.uniform() for k in params if RNG.uniform() > 0.5}
59-
60-
binary_unilateral_model.set_params(**new_params)
61-
binary_unilateral_model.named_params = params_subset.keys()
62-
binary_unilateral_model.set_named_params(**params_subset)
63-
64-
for param, new_val in new_params.items():
65-
stored_params = binary_unilateral_model.get_params(as_dict=True)
66-
if param in params_subset:
67-
assert params_subset[param] == stored_params[param]
68-
else:
69-
assert new_val == stored_params[param]
70-
71-
assert set(params_subset.keys()) == set(binary_unilateral_model.named_params)
72-
73-
74-
def test_set_named_params_raises(
75-
binary_unilateral_model: models.Unilateral,
76-
) -> None:
77-
"""Ensure `set_named_params` raises when provided with invalid keys."""
78-
binary_unilateral_model.named_params = ["spread"]
79-
with pytest.raises(ExtraParamsError):
80-
binary_unilateral_model.set_named_params(invalid=RNG.uniform())
81-
82-
83-
def test_set_named_params_allows_global_alias_not_named(
84-
binary_unilateral_model: models.Unilateral,
85-
) -> None:
86-
"""Allow global keys like `spread` even if not in `named_params`."""
87-
params = binary_unilateral_model.get_params(as_dict=True)
88-
new_params = {param: RNG.uniform() for param in params.keys()}
89-
first_lnl = list(binary_unilateral_model.graph.lnls.keys())[0]
90-
first_lnl_param = f"Tto{first_lnl}_spread"
91-
92-
binary_unilateral_model.set_params(**new_params)
93-
binary_unilateral_model.named_params = [first_lnl_param]
94-
spread_val = RNG.uniform()
95-
binary_unilateral_model.set_named_params(spread=spread_val)
96-
97-
stored_params = binary_unilateral_model.get_params(as_dict=True)
98-
for param, stored_param in stored_params.items():
99-
if "spread" in param:
100-
assert stored_param == spread_val
101-
102-
103-
def test_set_named_params_hard_subset(
104-
binary_unilateral_model: models.Unilateral,
105-
) -> None:
106-
"""Ensure `set_named_params` works correctly with a hard subset.
107-
108-
A "hard subset" is a subset that includes "global params". I.e., `spread` would
109-
not be a literal subset, because those are named something like `TtoII_spread`. But
110-
the `set_params()` method does accept it and will set all spread params with the
111-
provided value. It should be possible to set the `named_params` to such names and
112-
then set them with the `set_named_params()` method.
113-
"""
114-
params = binary_unilateral_model.get_params(as_dict=True)
115-
new_params = {param: RNG.uniform() for param in params.keys()}
116-
first_lnl = list(binary_unilateral_model.graph.lnls.keys())[0]
117-
first_lnl_param = f"Tto{first_lnl}_spread"
118-
params_subset = {k: RNG.uniform() for k in ["spread", first_lnl_param]}
119-
120-
binary_unilateral_model.set_params(**new_params)
121-
binary_unilateral_model.named_params = params_subset.keys()
122-
binary_unilateral_model.set_named_params(**params_subset)
123-
124-
stored_params = binary_unilateral_model.get_params(as_dict=True)
125-
for param, new_val, stored_param in zip(
126-
params.keys(),
127-
new_params.values(),
128-
stored_params.values(),
129-
strict=True,
130-
):
131-
if param == first_lnl_param:
132-
assert params_subset[first_lnl_param] == stored_param
133-
elif "spread" in param:
134-
assert params_subset["spread"] == stored_param
135-
else:
136-
assert new_val == stored_param
137-
138-
139-
def test_get_named_params_hard_subset(
140-
binary_unilateral_model: models.Unilateral,
141-
) -> None:
142-
"""Check that getting globals like `spread` works correctly."""
143-
params = binary_unilateral_model.get_params(as_dict=True)
144-
new_params = {param: RNG.uniform() for param in params.keys()}
145-
first_lnl = list(binary_unilateral_model.graph.lnls.keys())[0]
146-
first_lnl_param = f"Tto{first_lnl}_spread"
147-
params_subset = {k: RNG.uniform() for k in ["spread", first_lnl_param]}
148-
149-
binary_unilateral_model.set_params(**new_params)
150-
binary_unilateral_model.named_params = params_subset.keys()
151-
binary_unilateral_model.set_named_params(**params_subset)
152-
153-
stored_params = binary_unilateral_model.get_named_params()
154-
assert params_subset == stored_params
155-
156-
157-
def test_set_global_params_for_side(
158-
binary_bilateral_model: models.Bilateral,
159-
) -> None:
160-
"""Check that setting e.g. `"ipsi_spread"` works as global param to ipsi side."""
161-
params = binary_bilateral_model.get_params(as_dict=True)
162-
new_params = {param: RNG.uniform() for param in params.keys()}
163-
164-
binary_bilateral_model.named_params = ["ipsi_spread"]
165-
binary_bilateral_model.set_params(**new_params)
166-
ipsi_spread_val = RNG.uniform()
167-
binary_bilateral_model.set_named_params(ipsi_spread=ipsi_spread_val)
168-
169-
ipsi_stored_params = binary_bilateral_model.ipsi.get_params(as_dict=True)
170-
171-
for param, stored_param in ipsi_stored_params.items():
172-
if "spread" in param:
173-
assert stored_param == ipsi_spread_val
174-
else:
175-
assert stored_param in new_params.values()
7+
from . import fixtures
8+
from .fixtures import RNG
9+
10+
11+
class UnilateralNamedParamsTestCase(
12+
fixtures.BinaryUnilateralModelMixin,
13+
unittest.TestCase,
14+
):
15+
"""Tests for named params on unilateral models."""
16+
17+
def test_set_named_params_default_behavior(self) -> None:
18+
"""Ensure `set_named_params` works as `set_params` when no `named_params` set."""
19+
params = self.model.get_params(as_dict=True)
20+
new_params = {param: RNG.uniform() for param in params.keys()}
21+
self.model.set_named_params(**new_params)
22+
self.assertEqual(new_params, self.model.get_params(as_dict=True))
23+
24+
def test_named_params_setter(self) -> None:
25+
"""Check that setting `named_params` works correctly."""
26+
with self.assertRaises(ValueError):
27+
self.model.named_params = ["invalid identifier!"]
28+
29+
with self.assertRaises(ValueError):
30+
self.model.named_params = 123
31+
32+
params = self.model.get_params(as_dict=True).keys()
33+
params_subset = [param for param in params if RNG.uniform() > 0.5]
34+
self.model.named_params = params_subset
35+
36+
for stored, subset in zip(
37+
self.model.named_params,
38+
params_subset,
39+
strict=True,
40+
):
41+
self.assertEqual(stored, subset)
42+
43+
def test_set_named_params_named_easy_subset(self) -> None:
44+
"""Ensure `set_named_params` works correctly with an easy subset.
45+
46+
An "easy subset" is a literal subset of the params.
47+
"""
48+
params = self.model.get_params(as_dict=True)
49+
new_params = {param: RNG.uniform() for param in params.keys()}
50+
params_subset = {k: RNG.uniform() for k in params if RNG.uniform() > 0.5}
51+
52+
self.model.set_params(**new_params)
53+
self.model.named_params = params_subset.keys()
54+
self.model.set_named_params(**params_subset)
55+
56+
for param, new_val in new_params.items():
57+
stored_params = self.model.get_params(as_dict=True)
58+
if param in params_subset:
59+
self.assertEqual(params_subset[param], stored_params[param])
60+
else:
61+
self.assertEqual(new_val, stored_params[param])
62+
63+
self.assertEqual(set(params_subset.keys()), set(self.model.named_params))
64+
65+
def test_set_named_params_raises(self) -> None:
66+
"""Ensure `set_named_params` raises when provided with invalid keys."""
67+
self.model.named_params = ["spread"]
68+
with self.assertRaises(ExtraParamsError):
69+
self.model.set_named_params(invalid=RNG.uniform())
70+
71+
def test_set_named_params_allows_global_alias_not_named(self) -> None:
72+
"""Allow global keys like `spread` even if not in `named_params`."""
73+
params = self.model.get_params(as_dict=True)
74+
new_params = {param: RNG.uniform() for param in params.keys()}
75+
first_lnl = list(self.model.graph.lnls.keys())[0]
76+
first_lnl_param = f"Tto{first_lnl}_spread"
77+
78+
self.model.set_params(**new_params)
79+
self.model.named_params = [first_lnl_param]
80+
spread_val = RNG.uniform()
81+
self.model.set_named_params(spread=spread_val)
82+
83+
stored_params = self.model.get_params(as_dict=True)
84+
for param, stored_param in stored_params.items():
85+
if "spread" in param:
86+
self.assertEqual(stored_param, spread_val)
87+
88+
def test_set_named_params_hard_subset(self) -> None:
89+
"""Ensure `set_named_params` works correctly with a hard subset.
90+
91+
A "hard subset" is a subset that includes "global params". I.e., `spread` would
92+
not be a literal subset, because those are named something like `TtoII_spread`. But
93+
the `set_params()` method does accept it and will set all spread params with the
94+
provided value. It should be possible to set the `named_params` to such names and
95+
then set them with the `set_named_params()` method.
96+
"""
97+
params = self.model.get_params(as_dict=True)
98+
new_params = {param: RNG.uniform() for param in params.keys()}
99+
first_lnl = list(self.model.graph.lnls.keys())[0]
100+
first_lnl_param = f"Tto{first_lnl}_spread"
101+
params_subset = {k: RNG.uniform() for k in ["spread", first_lnl_param]}
102+
103+
self.model.set_params(**new_params)
104+
self.model.named_params = params_subset.keys()
105+
self.model.set_named_params(**params_subset)
106+
107+
stored_params = self.model.get_params(as_dict=True)
108+
for param, new_val, stored_param in zip(
109+
params.keys(),
110+
new_params.values(),
111+
stored_params.values(),
112+
strict=True,
113+
):
114+
if param == first_lnl_param:
115+
self.assertEqual(params_subset[first_lnl_param], stored_param)
116+
elif "spread" in param:
117+
self.assertEqual(params_subset["spread"], stored_param)
118+
else:
119+
self.assertEqual(new_val, stored_param)
120+
121+
def test_get_named_params_hard_subset(self) -> None:
122+
"""Check that getting globals like `spread` works correctly."""
123+
params = self.model.get_params(as_dict=True)
124+
new_params = {param: RNG.uniform() for param in params.keys()}
125+
first_lnl = list(self.model.graph.lnls.keys())[0]
126+
first_lnl_param = f"Tto{first_lnl}_spread"
127+
params_subset = {k: RNG.uniform() for k in ["spread", first_lnl_param]}
128+
129+
self.model.set_params(**new_params)
130+
self.model.named_params = params_subset.keys()
131+
self.model.set_named_params(**params_subset)
132+
133+
stored_params = self.model.get_named_params()
134+
self.assertEqual(params_subset, stored_params)
135+
136+
137+
class BilateralNamedParamsTestCase(
138+
fixtures.BilateralModelMixin,
139+
unittest.TestCase,
140+
):
141+
"""Tests for named params on bilateral models."""
142+
143+
def test_named_params_setter(self) -> None:
144+
"""Check that setting `named_params` works correctly."""
145+
self.model.named_params = ["ipsi_spread"]
146+
self.assertEqual(self.model.named_params, ["ipsi_spread"])
147+
148+
def test_set_global_params_for_side(self) -> None:
149+
"""Check that setting e.g. `"ipsi_spread"` works as global param to ipsi side."""
150+
params = self.model.get_params(as_dict=True)
151+
new_params = {param: RNG.uniform() for param in params.keys()}
152+
153+
self.model.named_params = ["ipsi_spread"]
154+
self.model.set_params(**new_params)
155+
ipsi_spread_val = RNG.uniform()
156+
self.model.set_named_params(ipsi_spread=ipsi_spread_val)
157+
158+
ipsi_stored_params = self.model.ipsi.get_params(as_dict=True)
159+
160+
for param, stored_param in ipsi_stored_params.items():
161+
if "spread" in param:
162+
self.assertEqual(stored_param, ipsi_spread_val)
163+
else:
164+
self.assertIn(stored_param, new_params.values())

0 commit comments

Comments
 (0)