Back to home page

EIC code displayed by LXR

 
 

    


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:]))