Back to home page

EIC code displayed by LXR

 
 

    


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

0001 """Tests for Ax optimizer configuration parsing and validation."""
0002 
0003 import pytest
0004 
0005 from aid2e.optimizers.ax import AxOptimizerConfig
0006 from aid2e.utilities.configurations.optimizer_config import OptimizerConfiguration
0007 
0008 
0009 class TestAxConfigurationLoading:
0010     """Tests for loading AxOptimizerConfig from dict and YAML-shaped payloads."""
0011 
0012     def test_load_ax_config_from_dict(self):
0013         """Test loading AxOptimizerConfig from a dictionary."""
0014         config_dict = {
0015             "initialization_strategy": "sobol",
0016             "generator": "BOTORCH_MODULAR",
0017             "generator_kwargs": {
0018                 "botorch_acqf_class": "qLogNoisyExpectedHypervolumeImprovement",
0019             },
0020             "generator_gen_kwargs": {
0021                 "model_gen_options": {"optimizer_kwargs": {"sequential": False}}
0022             },
0023             "objective_thresholds": {"f1": 1.0, "f2": 1.0},
0024             "n_initial_samples": 10,
0025             "n_iterations": 50,
0026             "batch_size": 5,
0027             "seed": 42,
0028         }
0029 
0030         config = AxOptimizerConfig(**config_dict)
0031 
0032         assert config.initialization_strategy == "sobol"
0033         assert config.generator == "BOTORCH_MODULAR"
0034         assert config.generator_kwargs["botorch_acqf_class"] == (
0035             "qLogNoisyExpectedHypervolumeImprovement"
0036         )
0037         assert config.generator_gen_kwargs["model_gen_options"]["optimizer_kwargs"][
0038             "sequential"
0039         ] is False
0040         assert config.objective_thresholds == {"f1": 1.0, "f2": 1.0}
0041         assert config.n_initial_samples == 10
0042         assert config.n_iterations == 50
0043         assert config.batch_size == 5
0044         assert config.seed == 42
0045 
0046     def test_load_optimizer_config_with_ax(self):
0047         """Test loading OptimizerConfiguration with Ax optimizer payload."""
0048         config = OptimizerConfiguration(
0049             name="ax",
0050             type="Bayesian",
0051             parameters={
0052                 "initialization_strategy": "sobol",
0053                 "generator": "BOTORCH_MODULAR",
0054                 "n_initial_samples": 10,
0055                 "n_iterations": 50,
0056                 "batch_size": 5,
0057                 "seed": 42,
0058             },
0059         )
0060 
0061         assert config.name == "ax"
0062         assert config.parameters["generator"] == "BOTORCH_MODULAR"
0063 
0064 
0065 class TestAxConfigurationDefaults:
0066     """Tests for Ax configuration defaults."""
0067 
0068     def test_sobol_is_default_initialization(self):
0069         """Test that Sobol is the default initialization strategy."""
0070         config = AxOptimizerConfig()
0071         assert config.initialization_strategy == "sobol"
0072 
0073     def test_mbm_is_default_generator(self):
0074         """Test that Modular BoTorch is the default generator."""
0075         config = AxOptimizerConfig()
0076         assert config.generator == "BOTORCH_MODULAR"
0077 
0078     def test_default_generator_kwargs_are_empty(self):
0079         """Test that generator kwargs default to an empty mapping."""
0080         config = AxOptimizerConfig()
0081         assert config.generator_kwargs == {}
0082         assert config.generator_gen_kwargs == {}
0083 
0084     def test_reasonable_default_iterations(self):
0085         """Test that default iterations are reasonable."""
0086         config = AxOptimizerConfig()
0087         assert config.n_iterations >= 10
0088         assert config.n_iterations <= 1000
0089 
0090     def test_reasonable_default_initial_samples(self):
0091         """Test that default initial samples are reasonable."""
0092         config = AxOptimizerConfig()
0093         assert config.n_initial_samples >= 1
0094         assert config.n_initial_samples <= 100
0095 
0096 
0097 class TestAxConfigurationValidation:
0098     """Tests for Ax configuration validation."""
0099 
0100     def test_positive_n_initial_samples_required(self):
0101         """Test that n_initial_samples must be positive."""
0102         with pytest.raises(ValueError):
0103             AxOptimizerConfig(n_initial_samples=0)
0104 
0105         with pytest.raises(ValueError):
0106             AxOptimizerConfig(n_initial_samples=-1)
0107 
0108     def test_positive_n_iterations_required(self):
0109         """Test that n_iterations must be positive."""
0110         with pytest.raises(ValueError):
0111             AxOptimizerConfig(n_iterations=0)
0112 
0113         with pytest.raises(ValueError):
0114             AxOptimizerConfig(n_iterations=-5)
0115 
0116     def test_positive_batch_size_required(self):
0117         """Test that batch_size must be positive."""
0118         with pytest.raises(ValueError):
0119             AxOptimizerConfig(batch_size=0)
0120 
0121         with pytest.raises(ValueError):
0122             AxOptimizerConfig(batch_size=-1)
0123 
0124     def test_none_seed_allowed(self):
0125         """Test that seed can be None for non-deterministic results."""
0126         config = AxOptimizerConfig(seed=None)
0127         assert config.seed is None
0128 
0129     def test_integer_seed_allowed(self):
0130         """Test that seed can be an integer."""
0131         config = AxOptimizerConfig(seed=42)
0132         assert config.seed == 42
0133         assert isinstance(config.seed, int)
0134 
0135     def test_legacy_fields_are_rejected(self):
0136         """Test that retired Ax config fields fail fast."""
0137         with pytest.raises(ValueError, match="legacy fields"):
0138             AxOptimizerConfig(
0139                 surrogate_model="saasbo",
0140                 acquisition_function="qnehvi",
0141             )
0142 
0143     def test_invalid_generator_is_rejected(self):
0144         """Test that unsupported Ax generators fail validation."""
0145         with pytest.raises(ValueError, match="Unsupported Ax generator"):
0146             AxOptimizerConfig(generator="SAASBO")
0147 
0148 
0149 class TestAxConfigurationDocumentation:
0150     """Tests verifying proper field descriptions for Ax configuration."""
0151 
0152     def test_ax_config_has_docstring(self):
0153         """Test that AxOptimizerConfig has a docstring."""
0154         assert AxOptimizerConfig.__doc__ is not None
0155         assert len(AxOptimizerConfig.__doc__) > 0
0156 
0157     def test_ax_config_field_descriptions(self):
0158         """Test that AxOptimizerConfig fields have descriptions."""
0159         model_fields = AxOptimizerConfig.model_fields
0160 
0161         assert "initialization_strategy" in model_fields
0162         assert "generator" in model_fields
0163         assert "generator_kwargs" in model_fields
0164 
0165         assert model_fields["initialization_strategy"].description
0166         assert model_fields["generator"].description
0167         assert model_fields["generator_kwargs"].description
0168 
0169 
0170 class TestAxConfigurationExamples:
0171     """Example-driven tests for the supported config surface."""
0172 
0173     def test_batch_configuration_round_trips(self):
0174         """Test various batch configurations for the MBM-first config."""
0175         batch_sizes = [1, 5, 10, 20]
0176 
0177         for batch_size in batch_sizes:
0178             config = AxOptimizerConfig(
0179                 initialization_strategy="sobol",
0180                 generator="BOTORCH_MODULAR",
0181                 batch_size=batch_size,
0182             )
0183             assert config.batch_size == batch_size
0184 
0185     def test_generator_kwargs_preserve_yaml_friendly_strings(self):
0186         """Test that config preserves raw string symbols until runtime resolution."""
0187         config = AxOptimizerConfig(
0188             generator_kwargs={
0189                 "botorch_acqf_class": "qLogNoisyExpectedImprovement",
0190                 "surrogate_spec": {
0191                     "model_configs": [{"botorch_model_class": "SingleTaskGP"}]
0192                 },
0193             }
0194         )
0195         assert config.generator_kwargs["botorch_acqf_class"] == (
0196             "qLogNoisyExpectedImprovement"
0197         )
0198         assert config.generator_kwargs["surrogate_spec"]["model_configs"][0][
0199             "botorch_model_class"
0200         ] == "SingleTaskGP"
0201 
0202 
0203 if __name__ == "__main__":
0204     pytest.main([__file__, "-v"])