Skip to content
Draft
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
69 changes: 40 additions & 29 deletions bax_algorithms/emittance.py
Original file line number Diff line number Diff line change
@@ -1,22 +1,22 @@
from pydantic import Field, PositiveInt
from typing import Optional
import torch
from torch import Tensor
import importlib.util
from typing import Any, Callable, Optional

from gpytorch.kernels.rbf_kernel import RBFKernel
from bax_algorithms.pathwise.base import PathwiseOptimization
from bax_algorithms.pathwise.sampling import draw_product_kernel_post_paths
from botorch.sampling.pathwise.posterior_samplers import draw_matheron_paths
import numpy as np
import torch
from botorch.models.model import Model
from botorch.sampling.pathwise.posterior_samplers import draw_matheron_paths
from gpytorch.kernels.rbf_kernel import RBFKernel
from pydantic import Field, PositiveInt
from scipy.optimize import minimize
from torch import Tensor
from xopt.generators.bayesian.bax.algorithms import (
Algorithm,
OptimizationAlgorithmResult,
VirtualMeasurementResult,
)

import numpy as np
from scipy.optimize import minimize
import importlib.util
from bax_algorithms.pathwise.base import PathwiseOptimization
from bax_algorithms.pathwise.sampling import draw_product_kernel_post_paths


def bmag_func(bb, ab, bl, al):
Expand Down Expand Up @@ -461,34 +461,35 @@ class VirtualEmittanceMeasurementResult(VirtualMeasurementResult):


class EmittanceAlgorithm(Algorithm):
name: str = Field(default="emittance", frozen=True)
x_key: str = Field(
None,
"",
description="key designating the beamsize squared output in x from evaluate function",
)
y_key: str = Field(
None,
"",
description="key designating the beamsize squared output in y from evaluate function",
)
energy: float = Field(1.0, description="Beam energy in [eV]")
q_len: float = Field(
description="the longitudinal thickness of the measurement quadrupole"
0.0, description="the longitudinal thickness of the measurement quadrupole"
)
rmat_x: Tensor = Field(
rmat_x: Tensor | None = Field(
None, description="tensor shape 2x2 containing downstream rmat for x dimension"
)
rmat_y: Tensor = Field(
rmat_y: Tensor | None = Field(
None, description="tensor shape 2x2 containing downstream rmat for y dimension"
)
twiss0_x: Tensor = Field(
twiss0_x: Tensor | None = Field(
None,
description="1d tensor length 2 containing design x-twiss: [beta0_x, alpha0_x] (for bmag)",
)
twiss0_y: Tensor = Field(
twiss0_y: Tensor | None = Field(
None,
description="1d tensor length 2 containing design y-twiss: [beta0_y, alpha0_y] (for bmag)",
)
meas_dim: int = Field(
None,
0,
description="index identifying the measurement quad dimension in the model",
)
n_steps_measurement_param: int = Field(
Expand All @@ -502,7 +503,7 @@ class EmittanceAlgorithm(Algorithm):
True,
description="Whether to multiply the emit by the bmag to get virtual objective.",
)
results: dict = Field(
results: dict[str, Any] = Field(
{}, description="Dictionary to store results from emittance calculcation"
)
maxiter_fit: int = Field(
Expand Down Expand Up @@ -530,8 +531,13 @@ def y_idx(self) -> int:
return self.observable_names_ordered.index(self.y_key)

def perform_virtual_measurement(
self, model, x, bounds, tkwargs: dict = None, n_samples: int = None
):
self,
model: Model,
x: Tensor,
bounds: Tensor,
n_samples: int | None = None,
tkwargs: dict[str, Any] | None = None,
) -> VirtualEmittanceMeasurementResult:
"""
inputs:
model: a botorch ModelListGP
Expand Down Expand Up @@ -600,7 +606,7 @@ def perform_virtual_measurement(

def get_meas_scan_inputs(
self, x_tuning: Tensor, bounds: Tensor, tkwargs: dict = None
):
) -> Tensor:
"""
A function that generates the inputs for virtual emittance measurement scans at the tuning
configurations specified by x_tuning.
Expand Down Expand Up @@ -649,8 +655,13 @@ def get_meas_scan_inputs(
return x

def evaluate_posterior_emittance(
self, model, x_tuning, bounds, tkwargs: dict = None, n_samples: int = None
):
self,
model: Model,
x_tuning: Tensor,
bounds: Tensor,
tkwargs: dict[str, Any] | None = None,
n_samples: int | None = None,
) -> tuple[Tensor, Tensor]:
"""
inputs:
x_tuning: tensor shape n_points x (n_dim-1) specifying points in the **tuning** space
Expand Down Expand Up @@ -795,7 +806,7 @@ class PathwiseMinimizeEmittance(EmittanceAlgorithm, PathwiseOptimization):
description="Number of sample batches to optimize, with each batch containing self.n_samples",
)

def execute(self, model: Model, bounds: Tensor) -> Tensor:
def execute(self, model: Model, bounds: Tensor) -> OptimizationAlgorithmResult:
best_tuning_inputs_list = []
best_objective_list = []
best_scan_inputs_list = []
Expand Down Expand Up @@ -842,8 +853,8 @@ def execute(self, model: Model, bounds: Tensor) -> Tensor:

return algorithm_result

def draw_sample_functions_list(self, model):
sample_funcs_list = []
def draw_sample_functions_list(self, model: Model) -> list[Callable[..., Any]]:
sample_funcs_list: list[Callable[..., Any]] = []
for m in model.models:
if isinstance(model.models[0].covar_module, RBFKernel):
sample_funcs = draw_matheron_paths(
Expand All @@ -856,7 +867,7 @@ def draw_sample_functions_list(self, model):
sample_funcs_list += [sample_funcs]
return sample_funcs_list

def _get_optimization_indeces(self, bounds) -> Tensor:
def _get_optimization_indeces(self, bounds: Tensor) -> Tensor:
"""
Get indeces specifying parameters for virtual objective optimization.
"""
Expand Down
14 changes: 8 additions & 6 deletions bax_algorithms/pathwise/base.py
Original file line number Diff line number Diff line change
@@ -1,15 +1,17 @@
# to be added to basic algorithms in Xopt

from abc import abstractmethod
from bax_algorithms.pathwise.optimize import VirtualOptimizer, DifferentialEvolution
from collections.abc import Callable
from typing import List

import torch
from botorch.models.model import Model, ModelList
from botorch.sampling.pathwise.posterior_samplers import draw_matheron_paths
from pydantic import Field
from xopt.generators.bayesian.bax.algorithms import Algorithm
from torch import Tensor
import torch
from typing import List
from collections.abc import Callable
from xopt.generators.bayesian.bax.algorithms import Algorithm

from bax_algorithms.pathwise.optimize import DifferentialEvolution, VirtualOptimizer


class PathwiseOptimization(Algorithm):
Expand Down Expand Up @@ -39,7 +41,7 @@ class PathwiseOptimization(Algorithm):
Get the bounds for virtual optimization.
"""

name = "pathwise_optimization"
name: str = Field("pathwise_optimization", frozen=True)
optimizer: VirtualOptimizer = Field(
DifferentialEvolution(), description="Optimizer for virtual objective."
)
Expand Down
11 changes: 6 additions & 5 deletions bax_algorithms/solenoid_alignment.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,14 @@
import torch
from torch import Tensor
from botorch.models.model import Model
from pydantic import Field, PositiveInt
from bax_algorithms.pathwise.base import PathwiseOptimization
from torch import Tensor
from xopt.generators.bayesian.bax.algorithms import (
OptimizationAlgorithmResult,
VirtualMeasurementResult,
)

from bax_algorithms.pathwise.base import PathwiseOptimization


class VirtualAlignmentMeasurementResult(VirtualMeasurementResult):
misalignment_x: Tensor = Field(
Expand All @@ -19,13 +20,13 @@ class VirtualAlignmentMeasurementResult(VirtualMeasurementResult):


class PathwiseSolenoidAlignment(PathwiseOptimization):
name: str = Field("PathwiseSolenoidAlignment", frozen=True)
name: str = Field("pathwise_solenoid_alignment", frozen=True)
x_key: str = Field(
None,
"",
description="key designating the centroid position in x from evaluate function",
)
y_key: str = Field(
None,
"",
description="key designating the centroid poisition in y from evaluate function",
)
meas_dim: int = Field(
Expand Down
28 changes: 20 additions & 8 deletions bax_algorithms/utils.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,17 @@
from xopt.generators.bayesian.bax_generator import BaxGenerator
from collections.abc import Callable
from typing import Any

import torch
from botorch.models.gp_regression import SingleTaskGP
from botorch.models.model_list_gp_regression import ModelListGP
from xopt.generators.bayesian.bax_generator import BaxGenerator

from bax_algorithms.pathwise.optimize import VirtualOptimizer
import torch


def get_bax_model_and_bounds(generator: BaxGenerator):
def get_bax_model_and_bounds(
generator: BaxGenerator,
) -> tuple[ModelListGP, torch.Tensor]:
bax_model_ids = [
generator.vocs.output_names.index(name)
for name in generator.algorithm.observable_names_ordered
Expand All @@ -19,11 +25,13 @@ def get_bax_model_and_bounds(generator: BaxGenerator):
return bax_model, generator._get_optimization_bounds()


def get_bax_mean_prediction(generator: BaxGenerator, mean_optimizer: VirtualOptimizer):
def get_bax_mean_prediction(
generator: BaxGenerator, mean_optimizer: VirtualOptimizer
) -> torch.Tensor:
model, bounds = get_bax_model_and_bounds(generator)

def get_mean_function(model):
def func(x):
def get_mean_function(model: ModelListGP) -> Callable[[torch.Tensor], torch.Tensor]:
def func(x: torch.Tensor) -> torch.Tensor:
gp_mean = model.posterior(x).mean
return gp_mean

Expand All @@ -45,7 +53,9 @@ def func(x):
return best_inputs.squeeze(0)


def tuning_input_tensor_to_dict(generator, x_tuning):
def tuning_input_tensor_to_dict(
generator: BaxGenerator, x_tuning: torch.Tensor
) -> dict[str, Any]:
"""
Converts a single set of tuning parameters to a dictionary for input to Xopt

Expand All @@ -62,7 +72,9 @@ def tuning_input_tensor_to_dict(generator, x_tuning):
return x_tuning_dict


def uniform_random_sample_in_bounds(n_samples, bounds):
def uniform_random_sample_in_bounds(
n_samples: int, bounds: torch.Tensor
) -> torch.Tensor:
ndim = len(bounds.T)

# uniform sample, rescaled, and shifted to cover the domain
Expand Down
10 changes: 6 additions & 4 deletions bax_algorithms/visualize.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,15 @@
import os
import pickle
from typing import List

import torch
from matplotlib import pyplot as plt
from typing import List
from xopt.generators.bayesian.bax_generator import BaxGenerator
from xopt.generators.bayesian.visualize import (
_generate_input_mesh,
_get_reference_point,
)
from xopt.generators.bayesian.bax_generator import BaxGenerator

from bax_algorithms.utils import get_bax_model_and_bounds


Expand Down Expand Up @@ -219,7 +221,7 @@ def plot_bax_objective_convergence(
if file_name.startswith(file_prefix)
]
file_names = sorted(file_names, key=lambda x: int(x[prefix_len:-ext_len]))
file_paths = [os.path.abspath(file_name) for file_name in file_names]
file_paths = [os.path.join(directory, file_name) for file_name in file_names]

results_dicts = []
aggregated_results_dict = {}
Expand Down Expand Up @@ -278,7 +280,7 @@ def plot_bax_input_convergence(
if file_name.startswith(file_prefix)
]
file_names = sorted(file_names, key=lambda x: int(x[prefix_len:-ext_len]))
file_paths = [os.path.abspath(file_name) for file_name in file_names]
file_paths = [os.path.join(directory, file_name) for file_name in file_names]

results_dicts = []
aggregated_results_dict = {}
Expand Down
Loading