Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 14 additions & 13 deletions xopt/generators/bayesian/bax/algorithms.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,10 @@
from abc import ABC, abstractmethod
from typing import ClassVar, List
from typing import Any

import torch
from botorch.models.model import Model, ModelList
from pydantic import BaseModel, ConfigDict, Field, PositiveInt, computed_field
from torch import Tensor

from xopt.pydantic import XoptBaseModel


Expand Down Expand Up @@ -56,7 +55,7 @@ class Algorithm(XoptBaseModel, ABC):

Attributes
----------
name : ClassVar[str]
name : str
The name of the algorithm.
n_samples : PositiveInt
Number of execution paths to generate.
Expand All @@ -69,12 +68,12 @@ class Algorithm(XoptBaseModel, ABC):
Perform the virtual measurement and calculate objective values at the given inputs.
"""

name: ClassVar[str] = "base_algorithm"
name: str = Field(default="base_algorithm", frozen=True)
n_samples: PositiveInt = Field(
default=20, description="number of execution paths to generate"
)

@computed_field
@computed_field # type: ignore[prop-decorator]
@property
def class_path(self) -> str:
return f"{self.__class__.__module__}.{self.__class__.__name__}"
Expand Down Expand Up @@ -105,7 +104,7 @@ def perform_virtual_measurement(
x: Tensor,
bounds: Tensor,
n_samples: int,
tkwargs: dict = None,
tkwargs: dict[str, Any] | None = None,
) -> VirtualMeasurementResult:
"""
Evaluate the virtual objective at the given inputs.
Expand All @@ -120,7 +119,7 @@ def perform_virtual_measurement(
The bounds for the optimization.
n_samples : int
The number of samples to generate.
tkwargs : dict, optional
tkwargs : dict[str, Any] | None, optional
Additional keyword arguments for the evaluation.

Returns
Expand Down Expand Up @@ -148,7 +147,7 @@ class GridScanAlgorithm(Algorithm, ABC):
Create a mesh for evaluating posteriors on.
"""

name = "grid_scan_algorithm"
name: str = Field(default="grid_scan", frozen=True)
n_mesh_points: PositiveInt = Field(
default=10, description="number of mesh points along each axis"
)
Expand Down Expand Up @@ -200,13 +199,14 @@ class GridOptimize(GridScanAlgorithm):

Methods
-------
get_execution_paths(self, model: Model, bounds: Tensor) -> Tuple[Tensor, Tensor, Dict]
get_execution_paths(self, model: Model, bounds: Tensor) -> tuple[Tensor, Tensor, ExecutionPathsResult]
Get execution paths that minimize the objective function.
perform_virtual_measurement(self, model: Model, x: Tensor, bounds: Tensor, n_samples: int, tkwargs: dict = None) -> VirtualMeasurementResult
Evaluate the virtual measurement and calculate objective values (samples).
"""

observable_names_ordered: List[str] = Field(
name: str = Field(default="grid_optimize", frozen=True)
observable_names_ordered: list[str] = Field(
description="names of observable/objective models used in this algorithm",
)
minimize: bool = True
Expand All @@ -228,7 +228,7 @@ def execute(self, model: Model, bounds: Tensor) -> GridOptimizeResult:
Contains best_inputs, best_objective, input_execution_paths, output_execution_paths, and additional results.
"""
# build evaluation mesh
test_points = self.create_mesh(bounds)
test_points: Tensor = self.create_mesh(bounds)
if isinstance(model, ModelList):
test_points = test_points.to(model.models[0].train_targets)
else:
Expand Down Expand Up @@ -273,7 +273,7 @@ def perform_virtual_measurement(
x: Tensor,
bounds: Tensor,
n_samples: int,
tkwargs: dict = None,
tkwargs: dict[str, Any] | None = None,
) -> VirtualMeasurementResult:
"""
Perform the virtual measurement (samples).
Expand Down Expand Up @@ -318,6 +318,7 @@ class CurvatureGridOptimize(GridOptimize):
Perform the virtual measurement (samples) with curvature.
"""

name: str = Field(default="curvature_grid_optimize", frozen=True)
use_mean: bool = False

def perform_virtual_measurement(
Expand All @@ -326,7 +327,7 @@ def perform_virtual_measurement(
x: Tensor,
bounds: Tensor,
n_samples: int,
tkwargs: dict = None,
tkwargs: dict[str, Any] | None = None,
) -> VirtualMeasurementResult:
"""
Evaluate the virtual objective (samples) with curvature.
Expand Down
94 changes: 71 additions & 23 deletions xopt/generators/bayesian/bax_generator.py
Original file line number Diff line number Diff line change
@@ -1,23 +1,30 @@
from copy import deepcopy
import importlib
import logging
import pickle
from typing import Dict, List, Optional
from copy import deepcopy
from typing import Any, Hashable, Optional, cast

from botorch.models import ModelListGP, SingleTaskGP
from gpytorch import Module
from pydantic import (
Field,
SerializeAsAny,
ValidationInfo,
field_validator,
model_validator,
)

from pydantic.fields import ModelPrivateAttr, PrivateAttr
from xopt.errors import VOCSError
from xopt.generators.bayesian.bax.acquisition import ModelListExpectedInformationGain
from xopt.generators.bayesian.bax.algorithms import Algorithm, GridOptimize
from xopt.generators.bayesian.bayesian_generator import BayesianGenerator
from xopt.generators.bayesian.turbo import EntropyTurboController, SafetyTurboController
from xopt.generators.bayesian.turbo import (
EntropyTurboController,
SafetyTurboController,
TurboController,
)
from xopt.generators.bayesian.utils import validate_turbo_controller_center
from xopt.vocs import VOCS

logger = logging.getLogger()

Expand Down Expand Up @@ -56,48 +63,89 @@ class BaxGenerator(BayesianGenerator):
supports_no_objective: bool = True
supports_discrete_variables: bool = False
algorithm: SerializeAsAny[Algorithm] = Field(
description="algorithm evaluated in the BAX process"
default=GridOptimize(observable_names_ordered=[]),
description="algorithm evaluated in the BAX process",
)
algorithm_results: Optional[Dict] = Field(
algorithm_results: Optional[dict] = Field(
None, description="dictionary results from algorithm", exclude=True
)
algorithm_results_file: Optional[str] = Field(
None, description="file name to save algorithm results at every step"
)
_n_calls: int = 0
_compatible_turbo_controllers = [EntropyTurboController, SafetyTurboController]
_compatible_turbo_controllers: list[type[TurboController]] = PrivateAttr(
default=[EntropyTurboController, SafetyTurboController]
)

# NOTE: this is meant for use in Badger, TODO: add it to Xopt
_compatible_algorithms = [GridOptimize]
_compatible_algorithms: list[type[Algorithm]] = PrivateAttr(default=[GridOptimize])

@model_validator(mode="after")
def validate_model_after(self):
def validate_model_after(self) -> "BaxGenerator":
# validate turbo controller center if it exists
validate_turbo_controller_center(self)

return self

@field_validator("vocs", mode="after")
@classmethod
def validate_vocs(cls, v: VOCS, info: ValidationInfo) -> VOCS:
# Preserve inherited Bayesian VOCS validation behavior.
# v = super().validate_vocs(v, info)

# assert that the generator had no objectives
if not v.n_objectives == 0:
raise VOCSError("BAX generator only supports problems with no objectives")

return v

@field_validator("algorithm", mode="before")
def validate_algorithm(cls, v, info: ValidationInfo):
@classmethod
def validate_algorithm(cls, v: Any, info: ValidationInfo) -> Any:
if isinstance(v, dict):
try:
if "class_path" in v:
class_path = v.pop("class_path")
module_name, class_name = class_path.rsplit(".", 1)
except KeyError:
raise ValueError("Algorithm dictionary must contain 'class_path' key")

try:
algorithm_class = getattr(
importlib.import_module(module_name), class_name
try:
algorithm_class = getattr(
importlib.import_module(module_name), class_name
)
except ModuleNotFoundError:
raise ValueError(f"Cannot import '{module_name}.{class_name}'")
elif "name" in v:
name = v["name"]
algorithm_class = next(
(
c
for c in cls._compatible_algorithms.default
if c.model_fields["name"].default == name
),
None,
)
if algorithm_class is None:
raise ValueError(
f"Unknown algorithm name '{name}'. "
f"Provide one of {[c.model_fields['name'].default for c in cls._compatible_algorithms.default]} "
f"or supply 'class_path'."
)
else:
raise ValueError(
"Algorithm dictionary must contain 'class_path' or 'name' key"
)
except ModuleNotFoundError:
raise ValueError(f"Cannot import '{module_name}.{class_name}'")

v = algorithm_class.model_validate(v)

return v

def generate(self, n_candidates: int) -> List[Dict]:
@classmethod
def get_compatible_algorithms(cls) -> list[type[Algorithm]]:
compatible = cls._compatible_algorithms
compatible_list: list[type[Algorithm]] = []
if isinstance(compatible, ModelPrivateAttr):
compatible_list = cast(list[type[Algorithm]], compatible.get_default())
return compatible_list

def generate(self, n_candidates: int) -> list[dict[Hashable, Any]]:
"""
Generate a specified number of candidate samples.

Expand All @@ -108,19 +156,19 @@ def generate(self, n_candidates: int) -> List[Dict]:

Returns
-------
List[Dict]
list[dict[Hashable, Any]]
A list of dictionaries containing the generated samples.
"""
self._n_calls += 1
return super().generate(n_candidates)

def _get_acquisition(self, model) -> ModelListExpectedInformationGain:
def _get_acquisition(self, model: Module) -> ModelListExpectedInformationGain:
"""
Get the acquisition function.

Parameters
----------
model : Model
model : Module
The model to use for the acquisition function.

Returns
Expand Down
Loading
Loading