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}.")