test_strategy_morris.py 7.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199
  1. """P5-M4: unit tests for MorrisStrategy.
  2. Covers: happy path (sensitivity ranking on linear function), boundary
  3. (empty params / single trajectory / odd n_levels auto-corrected), anomaly
  4. (unknown point_id report / missing objective metric), null inputs, registry
  5. integration, and state() field contract.
  6. All source is ASCII only. Run: python scripts/test_strategy_morris.py
  7. exit 0 = PASS.
  8. """
  9. import os
  10. import sys
  11. import unittest
  12. _ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
  13. sys.path.insert(0, os.path.join(_ROOT, "src"))
  14. from afmcore.strategies import ( # noqa: E402
  15. MorrisStrategy,
  16. get_strategy,
  17. is_registered,
  18. list_strategy_kinds,
  19. )
  20. def _linear_objective(params):
  21. """y = 2*a + 0.5*b -> parameter 'a' should rank higher than 'b'."""
  22. return 2.0 * params.get("a", 0.0) + 0.5 * params.get("b", 0.0)
  23. class TestMorrisHappyPath(unittest.TestCase):
  24. def test_trajectory_point_count(self):
  25. s = MorrisStrategy(
  26. parameters=[{"name": "a", "min_value": 0, "max_value": 1},
  27. {"name": "b", "min_value": 0, "max_value": 1}],
  28. n_trajectories=5, n_levels=4, objective_metric="y", rng_seed=1,
  29. )
  30. pts = s.select_next(1000)
  31. # n_trajectories * (n_params + 1) = 5 * 3 = 15
  32. self.assertEqual(len(pts), 15)
  33. def test_step_zero_has_no_changed_param(self):
  34. s = MorrisStrategy(
  35. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  36. n_trajectories=2, n_levels=4, objective_metric="y", rng_seed=2,
  37. )
  38. pts = s.select_next(1000)
  39. step0 = [p for p in pts if p["step"] == 0]
  40. self.assertEqual(len(step0), 2)
  41. for p in step0:
  42. self.assertIsNone(p["changed_param"])
  43. def test_sensitivity_ranking_linear(self):
  44. s = MorrisStrategy(
  45. parameters=[{"name": "a", "min_value": 0, "max_value": 1},
  46. {"name": "b", "min_value": 0, "max_value": 1}],
  47. n_trajectories=8, n_levels=4, objective_metric="y", rng_seed=3,
  48. )
  49. pts = s.select_next(1000)
  50. for p in pts:
  51. s.report(p["point_id"], {"y": _linear_objective(p["params"])}, "ok")
  52. st = s.state()
  53. ranking = st["sensitivity_ranking"]
  54. self.assertEqual(len(ranking), 2)
  55. # 'a' should rank first (higher mu_star)
  56. self.assertEqual(ranking[0]["parameter"], "a")
  57. self.assertGreater(ranking[0]["mu_star"], ranking[1]["mu_star"])
  58. # linear function -> sigma near zero
  59. self.assertAlmostEqual(ranking[0]["sigma"], 0.0, places=6)
  60. def test_converged_after_all_reported(self):
  61. s = MorrisStrategy(
  62. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  63. n_trajectories=3, n_levels=4, objective_metric="y", rng_seed=4,
  64. )
  65. # take only one point first -> pending still non-empty -> not converged
  66. first = s.select_next(1)
  67. self.assertEqual(len(first), 1)
  68. self.assertFalse(s.is_converged())
  69. # take the rest and report all
  70. rest = s.select_next(1000)
  71. all_pts = first + rest
  72. for p in all_pts:
  73. s.report(p["point_id"], {"y": 1.0}, "ok")
  74. self.assertTrue(s.is_converged())
  75. class TestMorrisBoundary(unittest.TestCase):
  76. def test_empty_parameters(self):
  77. s = MorrisStrategy(parameters=[], n_trajectories=5, objective_metric="y")
  78. pts = s.select_next(1000)
  79. self.assertEqual(pts, [])
  80. self.assertTrue(s.is_converged())
  81. st = s.state()
  82. self.assertEqual(st["n_parameters"], 0)
  83. self.assertEqual(st["sensitivity_ranking"], [])
  84. def test_none_parameters(self):
  85. s = MorrisStrategy(parameters=None, n_trajectories=3, objective_metric="y")
  86. self.assertEqual(s.select_next(1000), [])
  87. def test_single_trajectory(self):
  88. s = MorrisStrategy(
  89. parameters=[{"name": "a", "min_value": 0, "max_value": 1},
  90. {"name": "b", "min_value": 0, "max_value": 1}],
  91. n_trajectories=1, n_levels=4, objective_metric="y", rng_seed=5,
  92. )
  93. pts = s.select_next(1000)
  94. self.assertEqual(len(pts), 3) # 1 * (2+1)
  95. def test_odd_n_levels_auto_corrected(self):
  96. s = MorrisStrategy(
  97. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  98. n_trajectories=1, n_levels=5, objective_metric="y", rng_seed=6,
  99. )
  100. self.assertEqual(s.n_levels, 6) # odd -> next even
  101. self.assertAlmostEqual(s._delta, 6 / (2 * 5), places=6)
  102. def test_params_within_bounds(self):
  103. s = MorrisStrategy(
  104. parameters=[{"name": "a", "min_value": 0.5, "max_value": 2.0, "step": 0.1},
  105. {"name": "b", "min_value": 10, "max_value": 20}],
  106. n_trajectories=4, n_levels=4, objective_metric="y", rng_seed=7,
  107. )
  108. pts = s.select_next(1000)
  109. for p in pts:
  110. self.assertGreaterEqual(p["params"]["a"], 0.5)
  111. self.assertLessEqual(p["params"]["a"], 2.0)
  112. self.assertGreaterEqual(p["params"]["b"], 10)
  113. self.assertLessEqual(p["params"]["b"], 20)
  114. class TestMorrisAnomaly(unittest.TestCase):
  115. def test_report_unknown_point_id_no_crash(self):
  116. s = MorrisStrategy(
  117. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  118. n_trajectories=2, n_levels=4, objective_metric="y", rng_seed=8,
  119. )
  120. s.select_next(1000)
  121. # should not raise
  122. s.report(999999, {"y": 1.0}, "ok")
  123. # state() must remain callable and well-formed
  124. st = s.state()
  125. self.assertIn("sensitivity_ranking", st)
  126. def test_missing_objective_metric_yields_zero_sensitivity(self):
  127. s = MorrisStrategy(
  128. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  129. n_trajectories=3, n_levels=4, objective_metric="nonexistent", rng_seed=9,
  130. )
  131. pts = s.select_next(1000)
  132. for p in pts:
  133. s.report(p["point_id"], {"y": 1.0}, "ok") # wrong metric key
  134. st = s.state()
  135. for r in st["sensitivity_ranking"]:
  136. self.assertEqual(r["mu_star"], 0.0)
  137. self.assertEqual(r["n_effects"], 0)
  138. def test_failed_status_points_excluded_from_effects(self):
  139. s = MorrisStrategy(
  140. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  141. n_trajectories=2, n_levels=4, objective_metric="y", rng_seed=10,
  142. )
  143. pts = s.select_next(1000)
  144. # report first point as failed, rest as ok
  145. s.report(pts[0]["point_id"], {"y": 999.0}, "failed")
  146. for p in pts[1:]:
  147. s.report(p["point_id"], {"y": _linear_objective(p["params"])}, "ok")
  148. st = s.state()
  149. # failed point's trajectory may be partially excluded; no crash
  150. self.assertIn("sensitivity_ranking", st)
  151. class TestMorrisRegistry(unittest.TestCase):
  152. def test_registered(self):
  153. self.assertTrue(is_registered("morris"))
  154. self.assertIn("morris", list_strategy_kinds())
  155. def test_get_strategy_creates_instance(self):
  156. s = get_strategy("morris", parameters=[{"name": "x", "min_value": 0, "max_value": 1}],
  157. n_trajectories=2, objective_metric="y", rng_seed=11)
  158. self.assertIsInstance(s, MorrisStrategy)
  159. self.assertEqual(s.kind, "morris")
  160. def test_state_field_contract(self):
  161. s = MorrisStrategy(
  162. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  163. n_trajectories=2, n_levels=4, objective_metric="y", rng_seed=12,
  164. )
  165. st = s.state()
  166. for key in ("kind", "batch_size", "n_parameters", "n_trajectories",
  167. "n_levels", "delta", "objective_metric", "total_points",
  168. "pending", "reported", "sensitivity_ranking", "key_parameters"):
  169. self.assertIn(key, st, "missing state field: %s" % key)
  170. self.assertEqual(st["kind"], "morris")
  171. if __name__ == "__main__":
  172. unittest.main(verbosity=2)