|
4 | 4 | from collections import ( |
5 | 5 | Counter, |
6 | 6 | ) |
| 7 | +from types import ( |
| 8 | + SimpleNamespace, |
| 9 | +) |
7 | 10 |
|
8 | 11 | import mock |
9 | 12 | import numpy as np |
|
27 | 30 |
|
28 | 31 |
|
29 | 32 | class TestTrajsExplorationReport(unittest.TestCase): |
| 33 | + def test_adaptive_cutoff_and_convergence(self): |
| 34 | + model_devi = DeviManagerStd() |
| 35 | + model_devi.add( |
| 36 | + DeviManager.MAX_DEVI_F, |
| 37 | + np.array([0.10, 0.20, 0.30, 0.40, 0.90]), |
| 38 | + ) |
| 39 | + report = ExplorationReportAdaptiveLower( |
| 40 | + level_f_hi=0.80, |
| 41 | + numb_candi_f=2, |
| 42 | + rate_candi_f=0.0, |
| 43 | + n_checked_steps=3, |
| 44 | + conv_tolerance=0.05, |
| 45 | + ) |
| 46 | + |
| 47 | + report.record(model_devi) |
| 48 | + |
| 49 | + self.assertEqual(report.candi, {(0, 2), (0, 3)}) |
| 50 | + self.assertEqual(report.accur, {(0, 0), (0, 1)}) |
| 51 | + self.assertEqual(report.failed, [(0, 4)]) |
| 52 | + self.assertAlmostEqual(report.level_f_lo, 0.30) |
| 53 | + history = [ |
| 54 | + SimpleNamespace(level_f_lo=0.36), |
| 55 | + SimpleNamespace(level_f_lo=0.34), |
| 56 | + ] |
| 57 | + self.assertTrue(report.converged(history)) |
| 58 | + |
30 | 59 | def test_fv(self): |
31 | 60 | model_devi = DeviManagerStd() |
32 | 61 | model_devi.add( |
@@ -88,7 +117,7 @@ class MockedReport: |
88 | 117 | self.assertFalse(ter.converged([mr, mr1, mr])) |
89 | 118 | self.assertTrue(ter.converged([mr1, mr, mr])) |
90 | 119 |
|
91 | | - picked = ter.get_candidate_ids(2) |
| 120 | + picked = ter.get_candidate_ids(2, clear=False) |
92 | 121 | npicked = 0 |
93 | 122 | self.assertEqual(len(picked), 2) |
94 | 123 | for ii in range(2): |
@@ -198,32 +227,34 @@ def test_f_inv_pop(self): |
198 | 227 | ) |
199 | 228 |
|
200 | 229 | def faked_choices( |
201 | | - candi, |
202 | | - weights=None, |
203 | | - k=0, |
| 230 | + a, # numb_candi |
| 231 | + size=None, # numb_select |
| 232 | + replace=False, # non-repeative sampling |
| 233 | + p=None, # normalized prob |
204 | 234 | ): |
205 | 235 | # hist: 2bins, 0.1-0.4 5candi, 0.4-0.7 7candi |
206 | 236 | # only return those with mdf 0.1-0.4 |
207 | | - self.assertEqual(len(weights), 12) |
208 | | - self.assertEqual(len(candi), 12) |
209 | | - ret = [] |
210 | | - for ii in range(len(candi)): |
| 237 | + candi = ter.candi_picked |
| 238 | + self.assertEqual(a, 12) |
| 239 | + self.assertEqual(len(p), 12) |
| 240 | + ret_indices = [] |
| 241 | + for ii in range(a): |
211 | 242 | tidx, fidx = candi[ii] |
212 | 243 | this_mdf = md_f[tidx][fidx] |
213 | 244 | if this_mdf < 0.4: |
214 | | - self.assertAlmostEqual(weights[ii], 1.0 / 5.0) |
215 | | - ret.append(candi[ii]) |
| 245 | + self.assertAlmostEqual(p[ii], 0.1) # 1/5 / 2.0 |
| 246 | + ret_indices.append(ii) |
216 | 247 | else: |
217 | | - self.assertAlmostEqual(weights[ii], 1.0 / 7.0) |
218 | | - return ret |
| 248 | + self.assertAlmostEqual(p[ii], 1.0 / 14.0) # 1/7 / 2.0 |
| 249 | + return ret_indices |
219 | 250 |
|
220 | 251 | ter.record(model_devi) |
221 | | - with mock.patch("random.choices", faked_choices): |
222 | | - picked = ter.get_candidate_ids(11) |
223 | | - self.assertFalse(ter.converged([])) |
224 | 252 | self.assertEqual(ter.candi, expected_cand) |
225 | 253 | self.assertEqual(ter.accur, expected_accu) |
226 | 254 | self.assertEqual(set(ter.failed), expected_fail) |
| 255 | + with mock.patch("numpy.random.choice", faked_choices): |
| 256 | + picked = ter.get_candidate_ids(11) |
| 257 | + self.assertFalse(ter.converged([])) |
227 | 258 | self.assertEqual(len(picked), 2) |
228 | 259 | self.assertEqual(sorted(picked[0]), [1, 3]) |
229 | 260 | self.assertEqual(sorted(picked[1]), [1, 5, 7]) |
|
0 commit comments