test_strategy_surrogate.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255
  1. """P5-M4: unit tests for SurrogateGuidedStrategy.
  2. Covers: happy path (initial LHS + surrogate-guided convergence on bowl
  3. function), boundary (empty params / budget exhaustion / n_initial > budget),
  4. anomaly (unknown point_id / missing objective), null inputs, budget-adaptive
  5. batch sizing, both maximize and minimize directions, registry integration,
  6. and state() field contract.
  7. All source is ASCII only. Run: python scripts/test_strategy_surrogate.py
  8. exit 0 = PASS.
  9. """
  10. import os
  11. import sys
  12. import unittest
  13. _ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
  14. sys.path.insert(0, os.path.join(_ROOT, "src"))
  15. from afmcore.strategies import ( # noqa: E402
  16. SurrogateGuidedStrategy,
  17. get_strategy,
  18. is_registered,
  19. list_strategy_kinds,
  20. )
  21. def _bowl(params):
  22. """Convex bowl: y = (a-0.5)^2 + (b-0.5)^2, minimum at (0.5, 0.5)."""
  23. return (params.get("a", 0.0) - 0.5) ** 2 + (params.get("b", 0.0) - 0.5) ** 2
  24. def _hill(params):
  25. """Inverse bowl: y = 1 - bowl, maximum at (0.5, 0.5)."""
  26. return 1.0 - _bowl(params)
  27. class TestSurrogateHappyPath(unittest.TestCase):
  28. def test_initial_lhs_batch(self):
  29. s = SurrogateGuidedStrategy(
  30. parameters=[{"name": "a", "min_value": 0, "max_value": 1},
  31. {"name": "b", "min_value": 0, "max_value": 1}],
  32. objective_metric="y", objective_direction="minimize",
  33. n_initial=8, batch_size=4, budget=20, rng_seed=1,
  34. )
  35. self.assertEqual(s.state()["phase"], "initial")
  36. batch = s.select_next(1000)
  37. self.assertEqual(len(batch), 8)
  38. for p in batch:
  39. self.assertIn("point_id", p)
  40. self.assertIn("params", p)
  41. self.assertIn("a", p["params"])
  42. self.assertIn("b", p["params"])
  43. def test_convergence_on_bowl_minimize(self):
  44. s = SurrogateGuidedStrategy(
  45. parameters=[{"name": "a", "min_value": 0, "max_value": 1},
  46. {"name": "b", "min_value": 0, "max_value": 1}],
  47. objective_metric="y", objective_direction="minimize",
  48. n_initial=10, batch_size=3, max_batch_size=5, budget=40,
  49. n_candidates=60, rng_seed=2,
  50. )
  51. # run full budget
  52. used = 0
  53. while used < 40:
  54. batch = s.select_next()
  55. if not batch:
  56. break
  57. for p in batch:
  58. s.report(p["point_id"], {"y": _bowl(p["params"])}, "ok")
  59. used += 1
  60. best = min(r["metrics"]["y"] for r in s._done.values() if r["status"] == "ok")
  61. # should find a point reasonably close to the minimum (0.0)
  62. self.assertLess(best, 0.05)
  63. self.assertEqual(s.state()["phase"], "surrogate_guided")
  64. def test_maximize_direction(self):
  65. s = SurrogateGuidedStrategy(
  66. parameters=[{"name": "a", "min_value": 0, "max_value": 1},
  67. {"name": "b", "min_value": 0, "max_value": 1}],
  68. objective_metric="y", objective_direction="maximize",
  69. n_initial=8, batch_size=3, budget=25, n_candidates=50, rng_seed=3,
  70. )
  71. used = 0
  72. while used < 25:
  73. batch = s.select_next()
  74. if not batch:
  75. break
  76. for p in batch:
  77. s.report(p["point_id"], {"y": _hill(p["params"])}, "ok")
  78. used += 1
  79. best = max(r["metrics"]["y"] for r in s._done.values() if r["status"] == "ok")
  80. # hill maximum is 1.0
  81. self.assertGreater(best, 0.95)
  82. class TestSurrogateBoundary(unittest.TestCase):
  83. def test_empty_parameters(self):
  84. s = SurrogateGuidedStrategy(parameters=[], n_initial=5, budget=10, objective_metric="y")
  85. batch = s.select_next(1000)
  86. self.assertEqual(batch, [])
  87. self.assertTrue(s.is_converged())
  88. def test_none_parameters(self):
  89. s = SurrogateGuidedStrategy(parameters=None, n_initial=5, budget=10, objective_metric="y")
  90. self.assertEqual(s.select_next(1000), [])
  91. def test_budget_exhaustion(self):
  92. s = SurrogateGuidedStrategy(
  93. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  94. objective_metric="y", objective_direction="minimize",
  95. n_initial=3, batch_size=2, budget=5, rng_seed=4,
  96. )
  97. used = 0
  98. while True:
  99. batch = s.select_next()
  100. if not batch:
  101. break
  102. for p in batch:
  103. s.report(p["point_id"], {"y": p["params"]["a"] ** 2}, "ok")
  104. used += 1
  105. self.assertLessEqual(used, 5)
  106. self.assertTrue(s.is_converged())
  107. self.assertEqual(s.state()["phase"], "exhausted")
  108. def test_n_initial_greater_than_budget(self):
  109. s = SurrogateGuidedStrategy(
  110. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  111. objective_metric="y", n_initial=20, batch_size=4, budget=5, rng_seed=5,
  112. )
  113. batch = s.select_next(1000)
  114. # initial pending is 20 but budget is 5; select_next serves pending
  115. # regardless (pending points were already generated)
  116. self.assertEqual(len(batch), 20)
  117. def test_adaptive_batch_size_within_bounds(self):
  118. s = SurrogateGuidedStrategy(
  119. parameters=[{"name": "a", "min_value": 0, "max_value": 1},
  120. {"name": "b", "min_value": 0, "max_value": 1}],
  121. objective_metric="y", objective_direction="minimize",
  122. n_initial=6, batch_size=2, max_batch_size=6, budget=30,
  123. n_candidates=40, rng_seed=6,
  124. )
  125. # initial
  126. batch = s.select_next(1000)
  127. for p in batch:
  128. s.report(p["point_id"], {"y": _bowl(p["params"])}, "ok")
  129. # surrogate batches
  130. for _ in range(4):
  131. b = s.select_next()
  132. if not b:
  133. break
  134. self.assertGreaterEqual(len(b), 1)
  135. self.assertLessEqual(len(b), 6)
  136. for p in b:
  137. s.report(p["point_id"], {"y": _bowl(p["params"])}, "ok")
  138. class TestSurrogateAnomaly(unittest.TestCase):
  139. def test_report_unknown_point_id_no_crash(self):
  140. s = SurrogateGuidedStrategy(
  141. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  142. objective_metric="y", n_initial=3, budget=10, rng_seed=7,
  143. )
  144. s.select_next(1000)
  145. s.report(999999, {"y": 1.0}, "ok") # should not raise
  146. # state() must remain callable and well-formed
  147. st = s.state()
  148. self.assertIn("surrogate", st)
  149. def test_missing_objective_metric_excluded_from_training(self):
  150. s = SurrogateGuidedStrategy(
  151. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  152. objective_metric="nonexistent", n_initial=3, budget=10, rng_seed=8,
  153. )
  154. batch = s.select_next(1000)
  155. for p in batch:
  156. s.report(p["point_id"], {"y": 0.5}, "ok") # wrong metric
  157. # surrogate should have 0 training points (no objective values)
  158. train = s._training_data()
  159. self.assertEqual(len(train), 0)
  160. def test_failed_status_excluded_from_training(self):
  161. s = SurrogateGuidedStrategy(
  162. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  163. objective_metric="y", n_initial=4, budget=10, rng_seed=9,
  164. )
  165. batch = s.select_next(1000)
  166. s.report(batch[0]["point_id"], {"y": 999.0}, "failed")
  167. for p in batch[1:]:
  168. s.report(p["point_id"], {"y": 0.5}, "ok")
  169. train = s._training_data()
  170. self.assertEqual(len(train), 3) # failed excluded
  171. class TestSurrogateRegistry(unittest.TestCase):
  172. def test_registered(self):
  173. self.assertTrue(is_registered("surrogate_guided"))
  174. self.assertIn("surrogate_guided", list_strategy_kinds())
  175. def test_get_strategy_creates_instance(self):
  176. s = get_strategy("surrogate_guided",
  177. parameters=[{"name": "x", "min_value": 0, "max_value": 1}],
  178. objective_metric="y", budget=10, rng_seed=10)
  179. self.assertIsInstance(s, SurrogateGuidedStrategy)
  180. self.assertEqual(s.kind, "surrogate_guided")
  181. def test_state_field_contract(self):
  182. s = SurrogateGuidedStrategy(
  183. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  184. objective_metric="y", n_initial=3, budget=10, rng_seed=11,
  185. )
  186. st = s.state()
  187. for key in ("kind", "batch_size", "max_batch_size", "budget", "used_budget",
  188. "remaining_budget", "n_initial", "n_parameters", "objective_metric",
  189. "objective_direction", "phase", "pending", "reported",
  190. "last_batch_size", "surrogate", "idw_power", "kappa"):
  191. self.assertIn(key, st, "missing state field: %s" % key)
  192. self.assertEqual(st["kind"], "surrogate_guided")
  193. self.assertIn("n_train", st["surrogate"])
  194. class TestSurrogateIDWInternals(unittest.TestCase):
  195. def test_idw_prediction_interpolation(self):
  196. s = SurrogateGuidedStrategy(
  197. parameters=[{"name": "a", "min_value": 0, "max_value": 1}],
  198. objective_metric="y", n_initial=0, budget=10, rng_seed=12,
  199. )
  200. train = [((0.0,), 10.0), ((1.0,), 20.0)]
  201. pred, unc = s._idw_predict((0.5,), train)
  202. # midpoint should be between 10 and 20
  203. self.assertGreater(pred, 10.0)
  204. self.assertLess(pred, 20.0)
  205. self.assertGreater(unc, 0.0)
  206. def test_distance_calculation(self):
  207. d = SurrogateGuidedStrategy._distance((0.0, 0.0), (3.0, 4.0))
  208. self.assertAlmostEqual(d, 5.0, places=6)
  209. def test_normalize_denormalize_roundtrip(self):
  210. s = SurrogateGuidedStrategy(
  211. parameters=[{"name": "a", "min_value": 0, "max_value": 10},
  212. {"name": "b", "min_value": -5, "max_value": 5}],
  213. objective_metric="y", n_initial=0, budget=10, rng_seed=13,
  214. )
  215. phys = {"a": 5.0, "b": 0.0}
  216. norm = s._normalize(phys)
  217. self.assertAlmostEqual(norm[0], 0.5, places=6)
  218. self.assertAlmostEqual(norm[1], 0.5, places=6)
  219. back = s._denormalize(norm)
  220. self.assertAlmostEqual(back["a"], 5.0, places=6)
  221. self.assertAlmostEqual(back["b"], 0.0, places=6)
  222. if __name__ == "__main__":
  223. unittest.main(verbosity=2)