Back to home page

EIC code displayed by LXR

 
 

    


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

0001 """Resolvers for YAML-friendly Ax / BoTorch symbol configuration.
0002 
0003 This module keeps the public Ax optimizer config serializable and YAML-friendly
0004 by resolving supported string symbols into the concrete Python classes that Ax's
0005 Modular BoTorch generator expects at runtime.
0006 """
0007 
0008 from __future__ import annotations
0009 
0010 from copy import deepcopy
0011 from typing import Any
0012 
0013 from ax.generators.torch.botorch_modular.acquisition import Acquisition
0014 from ax.generators.torch.botorch_modular.multi_acquisition import MultiAcquisition
0015 from ax.generators.torch.botorch_modular.surrogate import ModelConfig, SurrogateSpec
0016 from botorch.acquisition.analytic import LogExpectedImprovement, PosteriorMean
0017 from botorch.acquisition.logei import (
0018     qLogExpectedImprovement,
0019     qLogNoisyExpectedImprovement,
0020 )
0021 from botorch.acquisition.monte_carlo import (
0022     qExpectedImprovement,
0023     qNoisyExpectedImprovement,
0024 )
0025 from botorch.acquisition.multi_objective.logei import (
0026     qLogNoisyExpectedHypervolumeImprovement,
0027 )
0028 from botorch.acquisition.multi_objective.monte_carlo import (
0029     qNoisyExpectedHypervolumeImprovement,
0030 )
0031 from botorch.models import SingleTaskGP
0032 from botorch.models.fully_bayesian import SaasFullyBayesianSingleTaskGP
0033 from botorch.models.fully_bayesian_multitask import SaasFullyBayesianMultiTaskGP
0034 from botorch.models.map_saas import AdditiveMapSaasSingleTaskGP
0035 from botorch.models.multitask import MultiTaskGP
0036 from botorch.models.transforms.input import Normalize, Warp
0037 from botorch.models.transforms.outcome import Standardize
0038 from gpytorch.kernels import MaternKernel, RBFKernel, ScaleKernel
0039 from gpytorch.likelihoods import GaussianLikelihood
0040 from gpytorch.mlls import ExactMarginalLogLikelihood
0041 
0042 SUPPORTED_AX_GENERATORS = {"BOTORCH_MODULAR"}
0043 
0044 _AX_ACQUISITION_CLASSES = {
0045     "Acquisition": Acquisition,
0046     "MultiAcquisition": MultiAcquisition,
0047 }
0048 
0049 _BOTORCH_ACQF_CLASSES = {
0050     "LogExpectedImprovement": LogExpectedImprovement,
0051     "PosteriorMean": PosteriorMean,
0052     "qExpectedImprovement": qExpectedImprovement,
0053     "qNoisyExpectedImprovement": qNoisyExpectedImprovement,
0054     "qLogExpectedImprovement": qLogExpectedImprovement,
0055     "qLogNoisyExpectedImprovement": qLogNoisyExpectedImprovement,
0056     "qNoisyExpectedHypervolumeImprovement": qNoisyExpectedHypervolumeImprovement,
0057     "qLogNoisyExpectedHypervolumeImprovement": (
0058         qLogNoisyExpectedHypervolumeImprovement
0059     ),
0060 }
0061 
0062 _BOTORCH_MODEL_CLASSES = {
0063     "AdditiveMapSaasSingleTaskGP": AdditiveMapSaasSingleTaskGP,
0064     "MultiTaskGP": MultiTaskGP,
0065     "SaasFullyBayesianMultiTaskGP": SaasFullyBayesianMultiTaskGP,
0066     "SaasFullyBayesianSingleTaskGP": SaasFullyBayesianSingleTaskGP,
0067     "SingleTaskGP": SingleTaskGP,
0068 }
0069 
0070 _INPUT_TRANSFORM_CLASSES = {
0071     "Normalize": Normalize,
0072     "Warp": Warp,
0073 }
0074 
0075 _OUTCOME_TRANSFORM_CLASSES = {
0076     "Standardize": Standardize,
0077 }
0078 
0079 _COVAR_MODULE_CLASSES = {
0080     "MaternKernel": MaternKernel,
0081     "RBFKernel": RBFKernel,
0082     "ScaleKernel": ScaleKernel,
0083 }
0084 
0085 _LIKELIHOOD_CLASSES = {
0086     "GaussianLikelihood": GaussianLikelihood,
0087 }
0088 
0089 _MLL_CLASSES = {
0090     "ExactMarginalLogLikelihood": ExactMarginalLogLikelihood,
0091 }
0092 
0093 
0094 def validate_generator_name(generator_name: str) -> str:
0095     """Normalize and validate the configured Ax generator name."""
0096     normalized = str(generator_name).strip().upper()
0097     if normalized not in SUPPORTED_AX_GENERATORS:
0098         supported = ", ".join(sorted(SUPPORTED_AX_GENERATORS))
0099         raise ValueError(
0100             f"Unsupported Ax generator '{generator_name}'. "
0101             f"Supported values for this backend are: {supported}."
0102         )
0103     return normalized
0104 
0105 
0106 def resolve_generator_kwargs(
0107     *, generator_name: str, generator_kwargs: dict[str, Any] | None
0108 ) -> dict[str, Any]:
0109     """Resolve YAML-friendly generator kwargs into Ax runtime objects."""
0110     resolved = deepcopy(generator_kwargs or {})
0111     generator_name = validate_generator_name(generator_name)
0112     if generator_name != "BOTORCH_MODULAR":
0113         return resolved
0114 
0115     if "acquisition_class" in resolved:
0116         resolved["acquisition_class"] = _resolve_symbol(
0117             resolved["acquisition_class"],
0118             registry=_AX_ACQUISITION_CLASSES,
0119             kind="Ax acquisition class",
0120         )
0121 
0122     if "botorch_acqf_class" in resolved:
0123         resolved["botorch_acqf_class"] = _resolve_symbol(
0124             resolved["botorch_acqf_class"],
0125             registry=_BOTORCH_ACQF_CLASSES,
0126             kind="BoTorch acquisition function",
0127         )
0128 
0129     if "botorch_acqf_classes_with_options" in resolved:
0130         resolved["botorch_acqf_classes_with_options"] = [
0131             _resolve_botorch_acqf_entry(entry)
0132             for entry in resolved["botorch_acqf_classes_with_options"]
0133         ]
0134 
0135     if "surrogate_spec" in resolved:
0136         resolved["surrogate_spec"] = _resolve_surrogate_spec(resolved["surrogate_spec"])
0137 
0138     return resolved
0139 
0140 
0141 def _resolve_botorch_acqf_entry(entry: Any) -> tuple[type[Any], dict[str, Any]]:
0142     """Resolve one MultiAcquisition entry from YAML-friendly form."""
0143     if isinstance(entry, dict):
0144         class_value = entry.get("class")
0145         options = dict(entry.get("options") or {})
0146     elif isinstance(entry, (list, tuple)) and len(entry) == 2:
0147         class_value, options = entry
0148         options = dict(options or {})
0149     else:
0150         raise ValueError(
0151             "Entries in 'botorch_acqf_classes_with_options' must be provided as "
0152             "{'class': <name>, 'options': {...}} or [<name>, {...}]."
0153         )
0154 
0155     resolved_class = _resolve_symbol(
0156         class_value,
0157         registry=_BOTORCH_ACQF_CLASSES,
0158         kind="BoTorch acquisition function",
0159     )
0160     return resolved_class, options
0161 
0162 
0163 def _resolve_surrogate_spec(raw_value: Any) -> SurrogateSpec:
0164     """Resolve a SurrogateSpec payload from YAML-friendly form."""
0165     if isinstance(raw_value, SurrogateSpec):
0166         return raw_value
0167     if not isinstance(raw_value, dict):
0168         raise ValueError(
0169             "'surrogate_spec' must be a mapping or an Ax SurrogateSpec instance."
0170         )
0171 
0172     payload = deepcopy(raw_value)
0173     payload["model_configs"] = [
0174         _resolve_model_config(config) for config in payload.get("model_configs", [])
0175     ]
0176 
0177     metric_to_model_configs = payload.get("metric_to_model_configs") or {}
0178     if metric_to_model_configs:
0179         payload["metric_to_model_configs"] = {
0180             metric_name: [_resolve_model_config(config) for config in configs]
0181             for metric_name, configs in metric_to_model_configs.items()
0182         }
0183 
0184     return SurrogateSpec(**payload)
0185 
0186 
0187 def _resolve_model_config(raw_value: Any) -> ModelConfig:
0188     """Resolve a ModelConfig payload from YAML-friendly form."""
0189     if isinstance(raw_value, ModelConfig):
0190         return raw_value
0191     if not isinstance(raw_value, dict):
0192         raise ValueError(
0193             "Entries in 'model_configs' must be mappings or Ax ModelConfig instances."
0194         )
0195 
0196     payload = deepcopy(raw_value)
0197 
0198     if "botorch_model_class" in payload:
0199         payload["botorch_model_class"] = _resolve_symbol(
0200             payload["botorch_model_class"],
0201             registry=_BOTORCH_MODEL_CLASSES,
0202             kind="BoTorch model class",
0203         )
0204 
0205     if "mll_class" in payload:
0206         payload["mll_class"] = _resolve_symbol(
0207             payload["mll_class"],
0208             registry=_MLL_CLASSES,
0209             kind="GPyTorch marginal log likelihood class",
0210         )
0211 
0212     if "covar_module_class" in payload:
0213         payload["covar_module_class"] = _resolve_symbol(
0214             payload["covar_module_class"],
0215             registry=_COVAR_MODULE_CLASSES,
0216             kind="GPyTorch covariance module class",
0217         )
0218 
0219     if "likelihood_class" in payload:
0220         payload["likelihood_class"] = _resolve_symbol(
0221             payload["likelihood_class"],
0222             registry=_LIKELIHOOD_CLASSES,
0223             kind="GPyTorch likelihood class",
0224         )
0225 
0226     if "input_transform_classes" in payload and payload["input_transform_classes"] is not None:
0227         payload["input_transform_classes"] = [
0228             _resolve_symbol(
0229                 transform_class,
0230                 registry=_INPUT_TRANSFORM_CLASSES,
0231                 kind="BoTorch input transform class",
0232             )
0233             for transform_class in payload["input_transform_classes"]
0234         ]
0235 
0236     if (
0237         "outcome_transform_classes" in payload
0238         and payload["outcome_transform_classes"] is not None
0239     ):
0240         payload["outcome_transform_classes"] = [
0241             _resolve_symbol(
0242                 transform_class,
0243                 registry=_OUTCOME_TRANSFORM_CLASSES,
0244                 kind="BoTorch outcome transform class",
0245             )
0246             for transform_class in payload["outcome_transform_classes"]
0247         ]
0248 
0249     return ModelConfig(**payload)
0250 
0251 
0252 def _resolve_symbol(
0253     value: Any,
0254     *,
0255     registry: dict[str, Any],
0256     kind: str,
0257 ) -> Any:
0258     """Resolve one configured symbol against a fixed registry."""
0259     if not isinstance(value, str):
0260         return value
0261 
0262     if value in registry:
0263         return registry[value]
0264 
0265     supported = ", ".join(sorted(registry))
0266     raise ValueError(f"Unsupported {kind} '{value}'. Supported values: {supported}.")