File indexing completed on 2026-08-12 08:24:55
0001 """Run the advanced inline Ax configuration from starter_kit.ipynb in debug mode."""
0002
0003 from __future__ import annotations
0004
0005 import argparse
0006 import json
0007 import logging
0008 import math
0009 import os
0010 import sys
0011 from datetime import datetime, timezone
0012 from pathlib import Path
0013 from pprint import pformat
0014 from typing import Any
0015
0016 REPO_ROOT = Path(__file__).resolve().parents[2]
0017 if str(REPO_ROOT) not in sys.path:
0018 sys.path.insert(0, str(REPO_ROOT))
0019
0020 from aid2e.utilities import build_optimizer_from_config
0021 from aid2e.utilities.configurations import load_config
0022 from aid2e.utilities.configurations.optimizer_config import OptimizerConfiguration
0023
0024
0025 DEFAULT_CONFIG_PATH = Path(__file__).resolve().with_name("dtlz2_ax_optimizer_only.yml")
0026 LOGGER = logging.getLogger("examples.optimizers.ax_advanced_inline_debug")
0027
0028
0029 def parse_args(argv: list[str]) -> argparse.Namespace:
0030 parser = argparse.ArgumentParser(
0031 description=(
0032 "Exercise the inline advanced Ax setup from starter_kit.ipynb "
0033 "with debug logging and JSON output."
0034 )
0035 )
0036 parser.add_argument(
0037 "--config",
0038 type=Path,
0039 default=DEFAULT_CONFIG_PATH,
0040 help="Path to the DTLZ2 optimizer config YAML.",
0041 )
0042 parser.add_argument(
0043 "--output-dir",
0044 type=Path,
0045 default=Path.cwd(),
0046 help="Directory for JSON artifacts. Defaults to the current working directory.",
0047 )
0048 parser.add_argument(
0049 "--debug",
0050 action=argparse.BooleanOptionalAction,
0051 default=True,
0052 help="Enable verbose debug logging.",
0053 )
0054 return parser.parse_args(argv)
0055
0056
0057 def configure_logging(debug: bool) -> None:
0058 level = logging.DEBUG if debug else logging.INFO
0059 logging.basicConfig(
0060 level=level,
0061 format="%(asctime)s %(levelname)-8s [%(name)s] %(message)s",
0062 force=True,
0063 )
0064 for logger_name in ("aid2e", "ax", "botorch", "gpytorch"):
0065 logging.getLogger(logger_name).setLevel(level)
0066
0067
0068 def dtlz2_objectives(parameters: dict[str, Any]) -> dict[str, float]:
0069 """Compute the two-objective DTLZ2 values for the example design space."""
0070 x1 = float(parameters["DTLZ2_variables.x1"])
0071 tail = [
0072 float(parameters["DTLZ2_variables.x2"]),
0073 float(parameters["DTLZ2_variables.x3"]),
0074 float(parameters["DTLZ2_variables.x4"]),
0075 float(parameters["DTLZ2_variables.x5"]),
0076 ]
0077 g = sum((value - 0.5) ** 2 for value in tail)
0078 factor = 1.0 + g
0079 f1 = factor * math.cos(x1 * math.pi / 2.0)
0080 f2 = factor * math.sin(x1 * math.pi / 2.0)
0081 return {"f1": float(f1), "f2": float(f2)}
0082
0083
0084 def trial_to_dict(trial: Any) -> dict[str, Any]:
0085 return {
0086 "index": trial.index,
0087 "status": trial.status,
0088 "parameters": dict(trial.parameters),
0089 "metrics": dict(trial.metrics or {}),
0090 "metadata": dict(trial.metadata or {}),
0091 }
0092
0093
0094 def evaluate_candidates(
0095 optimizer: Any,
0096 candidates: list[dict[str, Any]],
0097 *,
0098 phase: str,
0099 ) -> list[dict[str, Any]]:
0100 """Evaluate locally and feed results back into the optimizer."""
0101 start_index = len(optimizer.get_trials()) - len(candidates)
0102 records: list[dict[str, Any]] = []
0103
0104 for offset, parameters in enumerate(candidates):
0105 trial_index = start_index + offset
0106 metrics = dtlz2_objectives(parameters)
0107 optimizer.update_with_results(
0108 trial_index=trial_index,
0109 parameters=parameters,
0110 metrics=metrics,
0111 )
0112 record = {
0113 "trial_index": trial_index,
0114 "phase": phase,
0115 "parameters": dict(parameters),
0116 "metrics": metrics,
0117 }
0118 LOGGER.debug("Recorded %s", pformat(record))
0119 records.append(record)
0120
0121 return records
0122
0123
0124 def build_inline_payload() -> dict[str, Any]:
0125 return {
0126 "name": "ax",
0127 "type": "bayesian",
0128 "parameters": {
0129 "initialization_strategy": "sobol",
0130 "generator": "BOTORCH_MODULAR",
0131 "generator_kwargs": {
0132 "surrogate_spec": {
0133 "model_configs": [
0134 {
0135 "botorch_model_class": "SaasFullyBayesianSingleTaskGP"
0136 }
0137 ]
0138 },
0139 "botorch_acqf_class": "qLogNoisyExpectedHypervolumeImprovement",
0140 },
0141 "objective_thresholds": {"f1": 1.0, "f2": 1.0},
0142 "n_initial_samples": 10,
0143 "n_iterations": 20,
0144 "batch_size": 4,
0145 "seed": 7,
0146 },
0147 }
0148
0149
0150 def main(argv: list[str]) -> int:
0151 args = parse_args(argv)
0152 configure_logging(args.debug)
0153
0154 config_path = args.config.resolve()
0155 output_dir = args.output_dir.resolve()
0156 output_dir.mkdir(parents=True, exist_ok=True)
0157
0158 LOGGER.info("Repository root: %s", REPO_ROOT)
0159 LOGGER.info("Config path: %s", config_path)
0160 LOGGER.info("Output directory: %s", output_dir)
0161 LOGGER.info("Debug logging enabled: %s", args.debug)
0162
0163 ax_full_config = load_config(str(config_path))
0164 ax_inline_optimizer_payload = build_inline_payload()
0165 ax_inline_optimizer_cfg = OptimizerConfiguration(**ax_inline_optimizer_payload)
0166 ax_inline_optimizer = build_optimizer_from_config(
0167 ax_full_config.problem,
0168 ax_inline_optimizer_cfg,
0169 )
0170
0171 optimizer_config = ax_inline_optimizer_cfg.parse_algorithm_params()
0172 if optimizer_config is None:
0173 raise RuntimeError("Failed to parse the inline Ax optimizer configuration.")
0174
0175 nodes = getattr(ax_inline_optimizer.generation_strategy, "nodes", None)
0176 generation_summary = {
0177 "name": ax_inline_optimizer.generation_strategy.name,
0178 "nodes": [node.name for node in nodes] if nodes else [],
0179 }
0180
0181 print("Inline Ax optimizer payload (SAAS surrogate + qLogNEHVI):")
0182 print(pformat(ax_inline_optimizer_payload))
0183 print("Generation strategy summary:")
0184 print(pformat(generation_summary))
0185
0186 records: list[dict[str, Any]] = []
0187 first_candidate: dict[str, Any] | None = None
0188
0189 remaining_init = optimizer_config.n_initial_samples
0190 init_round = 0
0191 while remaining_init > 0:
0192 init_round += 1
0193 current_batch = min(optimizer_config.batch_size, remaining_init)
0194 candidates = ax_inline_optimizer.suggest_candidates(n_candidates=current_batch)
0195 if first_candidate is None and candidates:
0196 first_candidate = dict(candidates[0])
0197 print("First Ax candidate from the inline config:")
0198 print(pformat(first_candidate))
0199 records.extend(
0200 evaluate_candidates(
0201 ax_inline_optimizer,
0202 candidates,
0203 phase=f"init-{init_round}",
0204 )
0205 )
0206 remaining_init -= current_batch
0207
0208 for iteration in range(optimizer_config.n_iterations):
0209 candidates = ax_inline_optimizer.suggest_candidates(
0210 n_candidates=optimizer_config.batch_size
0211 )
0212 records.extend(
0213 evaluate_candidates(
0214 ax_inline_optimizer,
0215 candidates,
0216 phase=f"iter-{iteration + 1}",
0217 )
0218 )
0219
0220 summary = {
0221 "n_trials": ax_inline_optimizer.get_optimization_results()["n_trials"],
0222 "pareto_points": len(ax_inline_optimizer.get_pareto_front()),
0223 "objective_names": [obj.name for obj in ax_full_config.problem.objectives],
0224 }
0225
0226 artifact = {
0227 "timestamp_utc": datetime.now(timezone.utc).isoformat(),
0228 "slurm_job_id": os.environ.get("SLURM_JOB_ID"),
0229 "cwd": str(Path.cwd()),
0230 "config_path": str(config_path),
0231 "output_dir": str(output_dir),
0232 "debug": args.debug,
0233 "inline_payload": ax_inline_optimizer_payload,
0234 "generation_strategy": generation_summary,
0235 "first_candidate": first_candidate,
0236 "records": records,
0237 "summary": summary,
0238 "trials": [trial_to_dict(trial) for trial in ax_inline_optimizer.get_trials()],
0239 }
0240
0241 job_suffix = os.environ.get("SLURM_JOB_ID", "local")
0242 artifact_path = output_dir / f"ax_advanced_inline_debug_{job_suffix}.json"
0243 artifact_path.write_text(json.dumps(artifact, indent=2), encoding="utf-8")
0244
0245 print("Optimization summary:")
0246 print(pformat(summary))
0247 print(f"Artifact written to: {artifact_path}")
0248 LOGGER.info("Artifact written to %s", artifact_path)
0249 return 0
0250
0251
0252 if __name__ == "__main__":
0253 raise SystemExit(main(sys.argv[1:]))