|
1 | 1 | """Test if setting params via the named subset works correctly.""" |
2 | 2 |
|
3 | | -import pytest |
| 3 | +import unittest |
4 | 4 |
|
5 | | -from lymph import models |
6 | 5 | from lymph.types import ExtraParamsError |
7 | 6 |
|
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