Back to home page

EIC code displayed by LXR

 
 

    


File indexing completed on 2026-08-12 08:24:57

0001 """Constraint integration tests across DesignConfig, SearchSpace, and Ax."""
0002 
0003 import pytest
0004 from pydantic import ValidationError
0005 
0006 from aid2e.optimizers.ax.config import AxOptimizerConfig
0007 from aid2e.optimizers.ax.optimizer import AxOptimizer
0008 from aid2e.optimizers.ax import optimizer as ax_optimizer_module
0009 from aid2e.optimizers.base import SearchSpace
0010 from aid2e.utilities.configurations.design_config import DesignConfig, ParameterConstraint
0011 
0012 
0013 AX_NODE_RUNTIME_AVAILABLE = ax_optimizer_module.AX_NODE_STRATEGY_AVAILABLE
0014 
0015 
0016 def _design_config_with_constraints(constraints, include_z: bool = False) -> DesignConfig:
0017     """Build a validated DesignConfig for constraint tests."""
0018     parameters = {
0019         "x": {"value": 0.5, "bounds": [0.0, 1.0]},
0020         "y": {"value": 0.5, "bounds": [0.0, 1.0]},
0021     }
0022     if include_z:
0023         parameters["z"] = {"value": 0.5, "bounds": [0.0, 1.0]}
0024 
0025     return DesignConfig(
0026         design_parameters={"group1": {"parameters": parameters}},
0027         parameter_constraints=constraints,
0028     )
0029 
0030 
0031 class TestConstraintSyntaxValidation:
0032     """Test constraint syntax validation in DesignConfig."""
0033 
0034     def test_valid_constraint_accepted(self):
0035         """Valid constraints should be accepted during DesignConfig creation."""
0036         config = _design_config_with_constraints(
0037             [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0038         )
0039         assert len(config.parameter_constraints) == 1
0040         assert config.parameter_constraints[0].name == "sum_limit"
0041 
0042     def test_unknown_parameter_rejected(self):
0043         """Constraints referencing unknown parameters should be rejected."""
0044         with pytest.raises(ValidationError, match="Invalid constraint.*[Uu]nknown parameter"):
0045             DesignConfig(
0046                 design_parameters={
0047                     "group1": {"parameters": {"x": {"value": 0.5, "bounds": [0.0, 1.0]}}}
0048                 },
0049                 parameter_constraints=[
0050                     {
0051                         "name": "bad_constraint",
0052                         "rule": "group1.x + group1.unknown <= 1.0",
0053                     }
0054                 ],
0055             )
0056 
0057     def test_syntax_error_rejected(self):
0058         """Constraints with invalid syntax should be rejected."""
0059         with pytest.raises(ValidationError, match="Invalid constraint.*syntax"):
0060             DesignConfig(
0061                 design_parameters={
0062                     "group1": {"parameters": {"x": {"value": 0.5, "bounds": [0.0, 1.0]}}}
0063                 },
0064                 parameter_constraints=[{"name": "bad_syntax", "rule": "group1.x +* 1.0"}],
0065             )
0066 
0067 
0068 class TestParameterConstraintMethods:
0069     """Test ParameterConstraint helper methods."""
0070 
0071     def test_extract_parameter_names(self):
0072         """Test extraction of parameter names from constraint rules."""
0073         constraint = ParameterConstraint(
0074             name="test",
0075             rule="tracker.x + magnet.y + detector.z <= 10.0",
0076         )
0077 
0078         param_names = constraint.extract_parameter_names()
0079         assert param_names == {"tracker.x", "magnet.y", "detector.z"}
0080 
0081     def test_validate_syntax_valid(self):
0082         """Test syntax validation for valid constraints."""
0083         constraint = ParameterConstraint(
0084             name="test",
0085             rule="group1.x + group1.y <= 1.5",
0086         )
0087 
0088         valid_params = {"group1.x", "group1.y"}
0089         is_valid, error = constraint.validate_syntax(valid_params)
0090         assert is_valid
0091         assert error is None
0092 
0093     def test_validate_syntax_unknown_param(self):
0094         """Test syntax validation rejects unknown parameters."""
0095         constraint = ParameterConstraint(
0096             name="test",
0097             rule="group1.x + group1.unknown <= 1.5",
0098         )
0099 
0100         valid_params = {"group1.x"}
0101         is_valid, error = constraint.validate_syntax(valid_params)
0102         assert not is_valid
0103         assert "unknown parameter" in error.lower()
0104 
0105     def test_evaluate_constraint(self):
0106         """Test runtime constraint evaluation."""
0107         constraint = ParameterConstraint(
0108             name="sum_limit",
0109             rule="group1.x + group1.y <= 1.5",
0110         )
0111 
0112         assert constraint.evaluate({"group1.x": 0.5, "group1.y": 0.8}) is True
0113         assert constraint.evaluate({"group1.x": 1.0, "group1.y": 0.6}) is False
0114 
0115 
0116 class TestSearchSpaceConstraints:
0117     """Test SearchSpace constraint handling."""
0118 
0119     def test_constraints_passed_from_design_config(self):
0120         """SearchSpace should receive constraints from DesignConfig."""
0121         design_config = _design_config_with_constraints(
0122             [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0123         )
0124         search_space = SearchSpace.from_design_config(design_config)
0125 
0126         assert len(search_space.constraints) == 1
0127         assert search_space.constraints[0].name == "sum_limit"
0128         assert search_space.constraints[0].rule == "group1.x + group1.y <= 1.5"
0129 
0130     def test_check_constraints_method(self):
0131         """Test SearchSpace.validate() for runtime validation."""
0132         design_config = _design_config_with_constraints(
0133             [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0134         )
0135         search_space = SearchSpace.from_design_config(design_config)
0136 
0137         is_valid, errors = search_space.validate({"group1.x": 0.5, "group1.y": 0.8})
0138         assert is_valid
0139         assert len(errors) == 0
0140 
0141         is_valid, errors = search_space.validate({"group1.x": 1.0, "group1.y": 0.8})
0142         assert not is_valid
0143         assert len(errors) > 0
0144         assert "sum_limit" in errors[0]
0145 
0146 
0147 class TestAxOptimizerConstraints:
0148     """Test Ax optimizer constraint integration."""
0149 
0150     def test_ax_requires_node_runtime_for_constraint_integration(self):
0151         """Ax constraint integration should fail fast without node runtime support."""
0152         design_config = _design_config_with_constraints(
0153             [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0154         )
0155         search_space = SearchSpace.from_design_config(design_config)
0156         ax_config = AxOptimizerConfig(n_initial_samples=5, seed=42)
0157 
0158         if AX_NODE_RUNTIME_AVAILABLE:
0159             optimizer = AxOptimizer(
0160                 search_space=search_space,
0161                 config=ax_config,
0162                 objective_names=["f1"],
0163             )
0164             assert optimizer.search_space.constraints
0165             return
0166 
0167         with pytest.raises(RuntimeError, match="node-based generation API required"):
0168             AxOptimizer(
0169                 search_space=search_space,
0170                 config=ax_config,
0171                 objective_names=["f1"],
0172             )
0173 
0174     @pytest.mark.skipif(
0175         not AX_NODE_RUNTIME_AVAILABLE,
0176         reason="Installed Ax runtime lacks required node-based generation APIs.",
0177     )
0178     def test_constraint_parsing_to_ax_format(self):
0179         """Test parsing of constraint rules to Ax ParameterConstraint format."""
0180         design_config = _design_config_with_constraints(
0181             [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0182         )
0183         search_space = SearchSpace.from_design_config(design_config)
0184 
0185         optimizer = AxOptimizer(
0186             search_space=search_space,
0187             config=AxOptimizerConfig(n_initial_samples=5, seed=42),
0188             objective_names=["f1"],
0189         )
0190 
0191         constraint = design_config.parameter_constraints[0]
0192         ax_constraint = optimizer._parse_constraint_to_ax(constraint)
0193 
0194         assert ax_constraint is not None
0195         assert ax_constraint.constraint_dict == {"group1.x": 1.0, "group1.y": 1.0}
0196         assert ax_constraint.bound == 1.5
0197 
0198     @pytest.mark.skipif(
0199         not AX_NODE_RUNTIME_AVAILABLE,
0200         reason="Installed Ax runtime lacks required node-based generation APIs.",
0201     )
0202     def test_ax_enforces_constraints(self):
0203         """Test that Ax enforces constraints during candidate generation."""
0204         design_config = _design_config_with_constraints(
0205             [{"name": "sum_limit", "rule": "group1.x + group1.y <= 1.5"}]
0206         )
0207         search_space = SearchSpace.from_design_config(design_config)
0208 
0209         optimizer = AxOptimizer(
0210             search_space=search_space,
0211             config=AxOptimizerConfig(n_initial_samples=10, seed=42),
0212             objective_names=["f1"],
0213         )
0214 
0215         candidates = optimizer.suggest_candidates(n_candidates=20)
0216         violations = []
0217         for i, candidate in enumerate(candidates):
0218             sum_val = candidate["group1.x"] + candidate["group1.y"]
0219             if sum_val > 1.5:
0220                 violations.append((i, sum_val))
0221 
0222         assert len(violations) == 0
0223 
0224     @pytest.mark.skipif(
0225         not AX_NODE_RUNTIME_AVAILABLE,
0226         reason="Installed Ax runtime lacks required node-based generation APIs.",
0227     )
0228     def test_multiple_constraints(self):
0229         """Test Ax with multiple constraints."""
0230         design_config = _design_config_with_constraints(
0231             [
0232                 {"name": "sum_xy", "rule": "group1.x + group1.y <= 1.2"},
0233                 {"name": "sum_xz", "rule": "group1.x + group1.z <= 1.3"},
0234             ],
0235             include_z=True,
0236         )
0237         search_space = SearchSpace.from_design_config(design_config)
0238 
0239         optimizer = AxOptimizer(
0240             search_space=search_space,
0241             config=AxOptimizerConfig(n_initial_samples=10, seed=42),
0242             objective_names=["f1"],
0243         )
0244 
0245         candidates = optimizer.suggest_candidates(n_candidates=20)
0246         violations = []
0247         for i, candidate in enumerate(candidates):
0248             x = candidate["group1.x"]
0249             y = candidate["group1.y"]
0250             z = candidate["group1.z"]
0251             if x + y > 1.2:
0252                 violations.append((i, "sum_xy"))
0253             if x + z > 1.3:
0254                 violations.append((i, "sum_xz"))
0255 
0256         assert len(violations) == 0
0257 
0258     @pytest.mark.skipif(
0259         not AX_NODE_RUNTIME_AVAILABLE,
0260         reason="Installed Ax runtime lacks required node-based generation APIs.",
0261     )
0262     def test_greater_than_constraint(self):
0263         """Test Ax with >= constraint converted to negated <=."""
0264         design_config = _design_config_with_constraints(
0265             [{"name": "min_sum", "rule": "group1.x + group1.y >= 0.3"}]
0266         )
0267         search_space = SearchSpace.from_design_config(design_config)
0268 
0269         optimizer = AxOptimizer(
0270             search_space=search_space,
0271             config=AxOptimizerConfig(n_initial_samples=10, seed=42),
0272             objective_names=["f1"],
0273         )
0274 
0275         constraint = design_config.parameter_constraints[0]
0276         ax_constraint = optimizer._parse_constraint_to_ax(constraint)
0277 
0278         assert ax_constraint is not None
0279         assert ax_constraint.constraint_dict == {"group1.x": -1.0, "group1.y": -1.0}
0280         assert ax_constraint.bound == -0.3
0281 
0282         candidates = optimizer.suggest_candidates(n_candidates=20)
0283         violations = []
0284         for i, candidate in enumerate(candidates):
0285             if candidate["group1.x"] + candidate["group1.y"] < 0.3:
0286                 violations.append(i)
0287 
0288         assert len(violations) == 0