From a831ed4bc0c31ee51c49be2b782fd2e4fd79a998 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Wed, 10 Jun 2026 21:55:22 -0700 Subject: [PATCH 01/18] Start of Bax extension with updates fixing bo extension mypy errors --- .pre-commit-config.yaml | 9 +- pyproject.toml | 15 +-- .../gui/components/analysis_extensions.py | 46 +++++--- src/badger/gui/components/analysis_widget.py | 13 +-- .../gui/components/bax_visualizer/__init__.py | 0 .../components/bax_visualizer/bax_widget.py | 31 +++++ .../gui/components/bo_visualizer/bo_widget.py | 94 +++++++--------- .../components/bo_visualizer/plotting_area.py | 82 +++++++------- .../gui/components/bo_visualizer/types.py | 4 +- .../components/bo_visualizer/ui_components.py | 81 +++++++------ .../gui/components/extensions_palette.py | 68 ++++++++--- src/badger/gui/components/run_monitor.py | 106 +++++++++--------- src/badger/gui/pages/home_page.py | 83 +++++++------- src/badger/utils.py | 13 ++- 14 files changed, 366 insertions(+), 279 deletions(-) create mode 100644 src/badger/gui/components/bax_visualizer/__init__.py create mode 100644 src/badger/gui/components/bax_visualizer/bax_widget.py diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 8d919046..697dc9b0 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -2,7 +2,7 @@ # See https://pre-commit.com/hooks.html for more hooks repos: - repo: https://github.com/pre-commit/pre-commit-hooks.git - rev: v5.0.0 + rev: v6.0.0 hooks: - id: no-commit-to-branch - id: trailing-whitespace @@ -17,8 +17,13 @@ repos: exclude: "^(python/src/badger/_version.py)$" - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.12.2 + rev: v0.15.16 hooks: - id: ruff-check args: [--fix] - id: ruff-format + # - repo: https://github.com/pre-commit/mirrors-mypy + # rev: v2.1.0 + # hooks: + # - id: mypy + # args: [--strict, --ignore-missing-imports] diff --git a/pyproject.toml b/pyproject.toml index 0b1ad630..779d725d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,9 +6,7 @@ build-backend = "setuptools.build_meta" name = "badger-opt" description = "Interface for optimization of arbitrary problems in Python." readme = "README.md" -authors = [ - { name = "Zhe Zhang", email = "zhezhang@slac.stanford.edu" }, -] +authors = [{ name = "Zhe Zhang", email = "zhezhang@slac.stanford.edu" }] keywords = ["optimization", "machine learning", "GUI"] classifiers = [ "Development Status :: 4 - Beta", @@ -36,12 +34,7 @@ dynamic = ["version"] version_file = "src/badger/_version.py" [project.optional-dependencies] -dev = [ - "pytest", - "pytest-cov", - "pytest-qt", - "pytest-mock" -] +dev = ["pytest", "pytest-cov", "pytest-qt", "pytest-mock", "pyqt5-stubs"] [project.urls] Homepage = "https://github.com/xopt-org/Badger" @@ -60,12 +53,12 @@ include_package_data = true [tool.setuptools.packages.find] where = ["src"] -include = [ "badger", ] +include = ["badger"] namespaces = false [tool.ruff.lint] extend-select = ["TID252"] # Defaults + check imports -ignore = ["E722"] # Until bare except blocks get fixed +ignore = ["E722"] # Until bare except blocks get fixed [tool.pytest.ini_options] addopts = "--cov=badger/" diff --git a/src/badger/gui/components/analysis_extensions.py b/src/badger/gui/components/analysis_extensions.py index f23a886a..2b1a096d 100644 --- a/src/badger/gui/components/analysis_extensions.py +++ b/src/badger/gui/components/analysis_extensions.py @@ -1,27 +1,28 @@ -from typing import Optional, cast import logging +from typing import Optional, cast from PyQt5.QtCore import pyqtSignal -from PyQt5.QtWidgets import QDialog, QVBoxLayout from PyQt5.QtGui import QCloseEvent +from PyQt5.QtWidgets import QVBoxLayout, QWidget +from xopt import Generator +from xopt.generators.bayesian.bayesian_generator import BayesianGenerator +from xopt.generators.bayesian.mobo import MOBOGenerator + from badger.gui.components.analysis_widget import AnalysisWidget +from badger.gui.components.bax_visualizer.bax_widget import BaxWidget from badger.gui.components.bo_visualizer.bo_widget import BOPlotWidget from badger.gui.components.pf_viewer.pf_widget import ParetoFrontWidget from badger.routine import Routine -from xopt.generators.bayesian.bayesian_generator import BayesianGenerator -from xopt.generators.bayesian.mobo import MOBOGenerator -from xopt import Generator - logger = logging.getLogger(__name__) -class AnalysisExtension(QDialog): +class AnalysisExtension(QWidget): window_closed = pyqtSignal(object) - generator_type = Generator - widget = AnalysisWidget + generator_type: type[Generator] + widget: AnalysisWidget - def __init__(self, parent: Optional[QDialog] = None): + def __init__(self, parent: Optional[QWidget] = None): super().__init__(parent=parent) def update_window(self, routine: Routine) -> None: @@ -59,8 +60,6 @@ def update_extension( This method should be implemented to handle the update logic for the extension. """ - self.widget = cast(AnalysisWidget, self.widget) - self.widget.isValidRoutine(routine) self.widget.update_routine(routine, self.generator_type) @@ -70,7 +69,7 @@ def update_extension( self.widget.update_plots(requires_rebuild, interval=self.widget.update_interval) - def closeEvent(self, a0: Optional[QCloseEvent]) -> None: + def closeEvent(self, a0: QCloseEvent) -> None: self.window_closed.emit(self) super().closeEvent(a0) @@ -79,7 +78,7 @@ class ParetoFrontViewer(AnalysisExtension): def __init__( self, routine: Routine, - parent: Optional[QDialog] = None, + parent: Optional[QWidget] = None, ): super().__init__(parent=parent) @@ -94,12 +93,27 @@ class BOVisualizer(AnalysisExtension): def __init__( self, routine: Routine, - parent: Optional[QDialog] = None, + parent: Optional[QWidget] = None, ): super().__init__(parent=parent) self.initialize_extension( extension_widget=BOPlotWidget(routine=routine), extension_name="Bayesian Optimization Visualizer", - generator_type=BayesianGenerator, + generator_type=cast(type[Generator], BayesianGenerator), + ) + + +class BaxVisualizer(AnalysisExtension): + def __init__( + self, + routine: Routine, + parent: Optional[QWidget] = None, + ): + super().__init__(parent=parent) + + self.initialize_extension( + extension_widget=BaxWidget(routine=routine), + extension_name="Bax Visualizer", + generator_type=cast(type[Generator], BayesianGenerator), ) diff --git a/src/badger/gui/components/analysis_widget.py b/src/badger/gui/components/analysis_widget.py index 1b1179a6..6c50e0f6 100644 --- a/src/badger/gui/components/analysis_widget.py +++ b/src/badger/gui/components/analysis_widget.py @@ -1,19 +1,18 @@ +import logging from abc import abstractmethod from collections.abc import Callable from typing import Any, Optional -from PyQt5.QtWidgets import QDialog -from badger.gui.components.extension_utilities import HandledException -from badger.routine import Routine +from PyQt5.QtWidgets import QWidget from xopt import Generator -import logging - +from badger.gui.components.extension_utilities import HandledException +from badger.routine import Routine logger = logging.getLogger(__name__) -class AnalysisWidget(QDialog): +class AnalysisWidget(QWidget): routine: Routine generator: Generator parameters: dict[str, Any] = {} @@ -27,7 +26,7 @@ class AnalysisWidget(QDialog): def __init__( self, routine: Routine, - parent: Optional[QDialog] = None, + parent: Optional[QWidget] = None, ): super().__init__(parent=parent) self.routine = routine diff --git a/src/badger/gui/components/bax_visualizer/__init__.py b/src/badger/gui/components/bax_visualizer/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/src/badger/gui/components/bax_visualizer/bax_widget.py b/src/badger/gui/components/bax_visualizer/bax_widget.py new file mode 100644 index 00000000..aca10fca --- /dev/null +++ b/src/badger/gui/components/bax_visualizer/bax_widget.py @@ -0,0 +1,31 @@ +from typing import Optional + +from PyQt5.QtWidgets import QDialog +from xopt.generators.bayesian.bayesian_generator import BayesianGenerator + +from badger.gui.components.analysis_widget import AnalysisWidget +from badger.gui.components.extension_utilities import HandledException +from badger.routine import Routine + + +class BaxWidget(AnalysisWidget): # type: ignore[misc] + def __init__(self, routine: Routine, parent: Optional[QDialog] = None): + super().__init__(routine=routine, parent=parent) + + def initialize_widget(self) -> None: + pass + + def requires_reinitialization(self) -> bool: + return False + + def update_plots(self, requires_rebuild: bool, interval: int) -> None: + pass + + def setup_connections(self) -> None: + pass + + def isValidRoutine(self, routine: Routine) -> None: + if not isinstance(routine.generator, BayesianGenerator): + raise HandledException( + ValueError, "Bax Visualizer can only be used with a BayesianGenerator." + ) diff --git a/src/badger/gui/components/bo_visualizer/bo_widget.py b/src/badger/gui/components/bo_visualizer/bo_widget.py index 5991aaaa..34788394 100644 --- a/src/badger/gui/components/bo_visualizer/bo_widget.py +++ b/src/badger/gui/components/bo_visualizer/bo_widget.py @@ -1,14 +1,23 @@ +import logging from typing import Optional, cast + +from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( QHBoxLayout, - QWidget, - QVBoxLayout, QMessageBox, + QSizePolicy, QTableWidgetItem, + QVBoxLayout, + QWidget, ) -from PyQt5.QtWidgets import QSizePolicy +from xopt.generator import Generator +from xopt.generators.bayesian.bayesian_generator import BayesianGenerator +from xopt.vocs import select_best +from badger.gui.components.analysis_widget import AnalysisWidget +from badger.gui.components.bo_visualizer.plotting_area import PlottingArea from badger.gui.components.bo_visualizer.types import ConfigurableOptions +from badger.gui.components.bo_visualizer.ui_components import UIComponents from badger.gui.components.extension_utilities import ( HandledException, signal_logger, @@ -17,15 +26,6 @@ from badger.routine import Routine from badger.utils import BlockSignalsContext, create_archive_run_filename -from xopt.generator import Generator -from badger.gui.components.bo_visualizer.ui_components import UIComponents -from badger.gui.components.bo_visualizer.plotting_area import PlottingArea -from PyQt5.QtCore import Qt -from xopt.generators.bayesian.bayesian_generator import BayesianGenerator -from badger.gui.components.analysis_widget import AnalysisWidget - -import logging - logger = logging.getLogger(__name__) @@ -47,9 +47,11 @@ } -class BOPlotWidget(AnalysisWidget): - generator: BayesianGenerator # type: ignore - parameters: ConfigurableOptions = DEFAULT_PARAMETERS.copy() # type: ignore +class BOPlotWidget(AnalysisWidget): # type: ignore[misc] + generator: BayesianGenerator # pyright: ignore[reportIncompatibleVariableOverride] + parameters: ConfigurableOptions = DEFAULT_PARAMETERS.copy() + df_length: float = float("inf") + initialized: bool = False def __init__( self, @@ -115,19 +117,12 @@ def initialize_widget(self) -> None: self.parameters["include_variable_2"] = False self.parameters["variable_2"] = -1 - vocs_variables = cast( - dict[str, tuple[float, float]], - self.routine.vocs.variables, # type: ignore - ) + vocs_variables = self.routine.vocs.variables self.ui_components.initialize_variables(self.parameters, vocs_variables) self.ui_components.update_variables(self.parameters) - vocs_variables = cast( - dict[str, tuple[float, float]], - self.routine.vocs.variables, # type: ignore - ) # Initialize UI Components self.ui_components.initialize_ui_components( self.parameters, @@ -189,12 +184,11 @@ def setup_connections(self) -> None: # Reference inputs - if self.ui_components.reference_table is not None: - self.ui_components.reference_table.cellChanged.connect( - lambda: signal_logger("Updated 'reference_table'")( - lambda: self.on_reference_points_changed() - )() - ) + self.ui_components.reference_table.cellChanged.connect( + lambda: signal_logger("Updated 'reference_table'")( + lambda: self.on_reference_points_changed() + )() + ) self.ui_components.set_best_reference_point_button.clicked.connect( lambda: signal_logger("Set best reference points clicked")( @@ -238,7 +232,9 @@ def reset_widget(self) -> None: """ logger.debug("Resetting components of BOPlotWidget") self.ui_components.best_point_display.setText("") - self.parameters = DEFAULT_PARAMETERS.copy() # type: ignore + self.parameters = ( # pyright: ignore[reportIncompatibleVariableOverride] + DEFAULT_PARAMETERS.copy() + ) def requires_reinitialization(self) -> bool: # Check if the extension needs to be reinitialized @@ -279,7 +275,7 @@ def requires_reinitialization(self) -> bool: return False - def on_axis_selection_changed(self): + def on_axis_selection_changed(self) -> None: logger.debug("Axis selection changed") selected_variables: list[str] = [] @@ -354,15 +350,14 @@ def on_axis_selection_changed(self): if previous_selected_options != current_selected_options: logger.debug(f"Selected variables for plotting: {self.selected_variables}") # Update the reference point table based on the selected variables - if self.ui_components.reference_table is not None: - with BlockSignalsContext( - self.ui_components.reference_table, - ): - self.update_reference_point_table(self.selected_variables) + with BlockSignalsContext( + self.ui_components.reference_table, + ): + self.update_reference_point_table(self.selected_variables) # Only update plot if the selection has changed self.update_plots() - def update_reference_point_table(self, selected_variables: list[str]): + def update_reference_point_table(self, selected_variables: list[str]) -> None: """Disable and gray out reference points for selected variables.""" for i, var_name in enumerate(self.parameters["variables"]): @@ -387,14 +382,12 @@ def update_reference_point_table(self, selected_variables: list[str]): ref_item.setForeground(black) # Force the table to refresh and update its view - if self.ui_components.reference_table is not None: - viewport = self.ui_components.reference_table.viewport() - if viewport is not None: - viewport.update() + viewport = self.ui_components.reference_table.viewport() + viewport.update() def get_reference_points( self, ref_inputs: list[QTableWidgetItem], variable_names: list[str] - ): + ) -> dict[str, float]: reference_points: dict[str, float] = {} # Create a mapping from variable names to ref_inputs @@ -472,12 +465,11 @@ def update_plots( self.ui_components.update_variables(self.parameters) # Disable signals for the reference table to prevent updating the plot multiple times - if self.ui_components.reference_table is not None: - with BlockSignalsContext( - self.ui_components.reference_table, - ): - # Disable and gray out the reference points for selected variables - self.update_reference_point_table(selected_variables) + with BlockSignalsContext( + self.ui_components.reference_table, + ): + # Disable and gray out the reference points for selected variables + self.update_reference_point_table(selected_variables) # Get reference points for non-selected variables @@ -530,15 +522,15 @@ def update_routine(self, routine: Routine, generator_type: type[Generator]) -> N def set_best_reference_points( self, - ): + ) -> None: if self.generator.data is None: raise HandledException( ValueError, "No data available in generator for selecting best reference points", ) - index_arr, value_arr, input_params = self.routine.vocs.select_best( - self.generator.data + index_arr, value_arr, input_params = select_best( + self.routine.vocs, self.generator.data ) if not index_arr or not value_arr: diff --git a/src/badger/gui/components/bo_visualizer/plotting_area.py b/src/badger/gui/components/bo_visualizer/plotting_area.py index f93389e1..3e1bfe13 100644 --- a/src/badger/gui/components/bo_visualizer/plotting_area.py +++ b/src/badger/gui/components/bo_visualizer/plotting_area.py @@ -1,5 +1,18 @@ +import logging +import time from collections.abc import Callable from typing import Optional, cast + +from matplotlib.axes import Axes +from matplotlib.backends.backend_qt import NavigationToolbar2QT as NavigationToolbar +from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg as FigureCanvas +from matplotlib.figure import Figure +from PyQt5.QtWidgets import QVBoxLayout, QWidget +from xopt.generators.bayesian.bayesian_generator import BayesianGenerator +from xopt.generators.bayesian.visualize import ( + visualize_generator_model, +) + from badger.gui.components.bo_visualizer.types import ConfigurableOptions from badger.gui.components.extension_utilities import ( HandledException, @@ -7,29 +20,11 @@ clear_layout, requires_update, ) -from badger.routine import Routine -from matplotlib.backends.backend_qtagg import ( - FigureCanvasQTAgg as FigureCanvas, -) -from matplotlib.backends.backend_qt import ( - NavigationToolbar2QT as NavigationToolbar, -) -from matplotlib.figure import Figure -from matplotlib.axes import Axes -from PyQt5.QtWidgets import QVBoxLayout, QWidget -from badger.utils import BlockSignalsContext -from xopt.generators.bayesian.visualize import ( - visualize_generator_model, -) -from xopt.generators.bayesian.bayesian_generator import BayesianGenerator - -import time from badger.gui.components.plot_event_handlers import ( MatplotlibInteractionHandler, ) - - -import logging +from badger.routine import Routine +from badger.utils import BlockSignalsContext logger = logging.getLogger(__name__) @@ -59,7 +54,7 @@ def update_plot( n_grid: int, requires_rebuild: bool = False, interval: int = 500, # Interval in milliseconds - ): + ) -> None: logger.debug("Updating plot in PlottingArea") # Check if the plot was updated recently @@ -90,28 +85,29 @@ def update_plot( layout = self.layout() - if layout is not None: - with BlockSignalsContext(layout): - # Clear the existing layout (remove previous plot if any) - clear_layout(layout) - - with MatplotlibFigureContext(fig, ax) as (fig, ax): - # Create a new figure and canvas - canvas = FigureCanvas(fig) - toolbar = NavigationToolbar(canvas, self) - - variables = parameters["variables"] - - handler = MatplotlibInteractionHandler( - canvas, parameters, routine, variables, update_extension - ) - handler.connect_events() - - # Add the new canvas to the layout - layout.addWidget(canvas) - layout.addWidget(toolbar) - else: - raise HandledException(ValueError, "Layout is None and is not updated") + with BlockSignalsContext(layout): + # Clear the existing layout (remove previous plot if any) + clear_layout(layout) + + with MatplotlibFigureContext(fig, ax) as (fig, ax): + # Create a new figure and canvas + canvas = FigureCanvas(fig) + toolbar = NavigationToolbar(canvas, self) + + variables = parameters["variables"] + + handler = MatplotlibInteractionHandler( + canvas, + parameters, # pyright: ignore[reportArgumentType] + routine, + variables, + update_extension, + ) + handler.connect_events() + + # Add the new canvas to the layout + layout.addWidget(canvas) + layout.addWidget(toolbar) except HandledException as he: raise he except Exception as e: diff --git a/src/badger/gui/components/bo_visualizer/types.py b/src/badger/gui/components/bo_visualizer/types.py index 555833fb..dfee47ac 100644 --- a/src/badger/gui/components/bo_visualizer/types.py +++ b/src/badger/gui/components/bo_visualizer/types.py @@ -1,5 +1,7 @@ from typing import TypedDict +from gest_api.vocs import ContinuousVariable + class PlotOptions(TypedDict): n_grid: int @@ -16,5 +18,5 @@ class ConfigurableOptions(TypedDict): variable_2: int variables: list[str] reference_points: dict[str, float] - reference_points_range: dict[str, tuple[float, float]] + reference_points_range: dict[str, ContinuousVariable] include_variable_2: bool diff --git a/src/badger/gui/components/bo_visualizer/ui_components.py b/src/badger/gui/components/bo_visualizer/ui_components.py index b4b42500..c3bbb3c2 100644 --- a/src/badger/gui/components/bo_visualizer/ui_components.py +++ b/src/badger/gui/components/bo_visualizer/ui_components.py @@ -1,30 +1,28 @@ +import logging +from typing import cast + +from gest_api.vocs import BaseVariable, ContinuousVariable, VariableDict +from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( - QVBoxLayout, - QHBoxLayout, + QCheckBox, QComboBox, - QLabel, QGroupBox, + QHBoxLayout, + QHeaderView, + QLabel, + QPushButton, + QSpinBox, QTableWidget, QTableWidgetItem, - QSpinBox, - QPushButton, - QCheckBox, - QHeaderView, + QVBoxLayout, ) -from PyQt5.QtCore import Qt -from gest_api.vocs import BaseVariable from badger.gui.components.bo_visualizer.types import ConfigurableOptions from badger.gui.components.extension_utilities import ( - HandledException, to_precision_float, ) - -import logging - from badger.utils import BlockSignalsContext - logger = logging.getLogger(__name__) @@ -35,9 +33,9 @@ def __init__( self, default_parameters: ConfigurableOptions, ): - self.variable_checkboxes = {} + self.variable_checkboxes: dict[str, QCheckBox] = {} self.ref_inputs: list[QTableWidgetItem] = [] - self.reference_table = None # Will be initialized later + self.reference_table = QTableWidget() self.best_point_display = QLabel("") # Will be initialized later self.set_best_reference_point_button = QPushButton("Set Best Reference Point") @@ -74,7 +72,7 @@ def __init__( self.restrict_selection_variables(default_parameters) - def restrict_selection_variables(self, parameters: ConfigurableOptions): + def restrict_selection_variables(self, parameters: ConfigurableOptions) -> None: num_of_variables = len(parameters["variables"]) if num_of_variables < 2: parameters["include_variable_2"] = False @@ -89,13 +87,13 @@ def restrict_selection_variables(self, parameters: ConfigurableOptions): self.y_axis_checkbox.setEnabled(True) self.y_axis_combo.setEnabled(True) - def create_variable_checkboxes(self): + def create_variable_checkboxes(self) -> QGroupBox: group_box = QGroupBox("Select Variables") layout = self.variable_checkboxes_layout or QVBoxLayout() group_box.setLayout(layout) return group_box - def create_axis_layout(self): + def create_axis_layout(self) -> QVBoxLayout: layout = QVBoxLayout() x_layout = QHBoxLayout() @@ -124,7 +122,7 @@ def create_axis_layout(self): def initialize_ui_components( self, configurable_options: ConfigurableOptions, - ): + ) -> None: self.populate_reference_table( configurable_options["variables"], configurable_options["reference_points"], @@ -133,28 +131,39 @@ def initialize_ui_components( def initialize_variables( self, configurable_options: ConfigurableOptions, - vocs_variables: dict[str, BaseVariable], - ): + vocs_variables: VariableDict, + ) -> None: """Initialize the variable checkboxes with the provided variable names.""" # Initialize the parameters with the routine's variables configurable_options["reference_points_range"] = vocs_variables - configurable_options["reference_points"] = { - var: to_precision_float( - (vocs_variables[var].domain[1] - vocs_variables[var].domain[0]) / 2.0 + + reference_points: dict[str, float] = {} + + variables = cast(dict[str, BaseVariable], vocs_variables) + for var_name, variable in variables.items(): + if not isinstance(variable, ContinuousVariable): + raise ValueError( + f"Variable '{var_name}' is not continuous. Only continuous variables are supported for reference points." + ) + + domain = cast( + tuple[float, float], + variable.domain, # pyright: ignore[reportUnknownMemberType] + ) + reference_points[var_name] = to_precision_float( + (domain[1] - domain[0]) / 2.0 ) - for var in vocs_variables - } - def create_reference_inputs(self): + configurable_options["reference_points"] = reference_points + + def create_reference_inputs(self) -> QGroupBox: group_box = QGroupBox("Reference Points") layout = QVBoxLayout() - self.reference_table = QTableWidget() self.reference_table.setColumnCount(2) self.reference_table.setHorizontalHeaderLabels(["Variable", "Ref. Point"]) horizontal_header = self.reference_table.horizontalHeader() - if horizontal_header is not None: - horizontal_header.setSectionResizeMode(QHeaderView.ResizeMode.Stretch) + horizontal_header.setSectionResizeMode(QHeaderView.ResizeMode.Stretch) layout.addWidget(self.reference_table) layout.addWidget(self.set_best_reference_point_button) @@ -166,12 +175,10 @@ def populate_reference_table( self, variables: list[str], reference_points: dict[str, float], - ): + ) -> None: """Populate the reference table based on the current vocs variable names.""" logger.debug("Populating reference table") - if self.reference_table is None: - raise HandledException(ValueError, "Reference Table is None") with BlockSignalsContext(self.reference_table): self.reference_table.setRowCount(len(variables)) @@ -192,7 +199,7 @@ def populate_reference_table( self.ref_inputs.append(reference_point_item) self.reference_table.setItem(i, 1, reference_point_item) - def create_options_section(self): + def create_options_section(self) -> QGroupBox: group_box = QGroupBox("Plot Options") layout = QVBoxLayout() @@ -210,7 +217,7 @@ def create_options_section(self): group_box.setLayout(layout) return group_box - def create_buttons(self): + def create_buttons(self) -> QHBoxLayout: layout = QHBoxLayout() self.update_button = QPushButton("Update") @@ -223,7 +230,7 @@ def create_buttons(self): def update_variables( self, configurable_options: ConfigurableOptions, - ): + ) -> None: with BlockSignalsContext([self.x_axis_combo, self.y_axis_combo]): self.x_axis_combo.clear() self.y_axis_combo.clear() diff --git a/src/badger/gui/components/extensions_palette.py b/src/badger/gui/components/extensions_palette.py index 211e0ae1..5371d251 100644 --- a/src/badger/gui/components/extensions_palette.py +++ b/src/badger/gui/components/extensions_palette.py @@ -1,22 +1,28 @@ import traceback +from typing import TYPE_CHECKING +from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( + QLabel, QMainWindow, + QMessageBox, QPushButton, + QSizePolicy, QVBoxLayout, QWidget, - QLabel, - QMessageBox, - QSizePolicy, ) -from PyQt5.QtCore import Qt + from badger.gui.components.analysis_extensions import ( AnalysisExtension, - ParetoFrontViewer, + BaxVisualizer, BOVisualizer, + ParetoFrontViewer, ) from badger.gui.components.extension_utilities import HandledException +if TYPE_CHECKING: + from badger.gui.components.run_monitor import BadgerOptMonitor + class ExtensionsPalette(QMainWindow): """ @@ -24,7 +30,7 @@ class ExtensionsPalette(QMainWindow): Parameters ---------- - run_monitor : RunMonitor + run_monitor : BadgerOptMonitor The run monitor associated with the palette. Attributes @@ -49,13 +55,13 @@ class ExtensionsPalette(QMainWindow): """ - def __init__(self, run_monitor): + def __init__(self, run_monitor: "BadgerOptMonitor") -> None: """ Initialize the ExtensionsPalette. Parameters ---------- - run_monitor : RunMonitor + run_monitor : BadgerOptMonitor The run monitor associated with the palette. """ @@ -78,9 +84,11 @@ def __init__(self, run_monitor): self.btn_data_viewer = QPushButton("Pareto Front Viewer") self.btn_bo_visualizer = QPushButton("Bayesian Optimization Visualizer") + self.btn_bax_visualizer = QPushButton("Bax Visualizer") layout.addWidget(self.btn_data_viewer) layout.addWidget(self.btn_bo_visualizer) + layout.addWidget(self.btn_bax_visualizer) layout.addStretch() layout.addWidget(self.text_box) @@ -88,9 +96,10 @@ def __init__(self, run_monitor): self.btn_data_viewer.clicked.connect(self.add_pf_viewer) self.btn_bo_visualizer.clicked.connect(self.add_bo_visualizer) + self.btn_bax_visualizer.clicked.connect(self.add_bax_visualizer) @property - def n_active_extensions(self): + def n_active_extensions(self) -> int: """ Property to get the number of active extensions. @@ -102,26 +111,59 @@ def n_active_extensions(self): """ return len(self.run_monitor.active_extensions) - def update_palette(self): + def update_palette(self) -> None: self.text_box.setText(self.base_text + str(self.n_active_extensions)) - def add_pf_viewer(self): + def add_pf_viewer(self) -> None: """ Open the ParetoFrontViewer extension. """ + if self.run_monitor.routine is None: + QMessageBox.warning( + self, + "No Routine Error", + "Please start a routine before opening the Pareto Front Viewer.", + ) + return + self.add_child_window_to_monitor( ParetoFrontViewer(routine=self.run_monitor.routine) ) - def add_bo_visualizer(self): + def add_bo_visualizer(self) -> None: """ Open the BOVisualizer extension. """ + if self.run_monitor.routine is None: + QMessageBox.warning( + self, + "No Routine Error", + "Please start a routine before opening the BO Visualizer.", + ) + return + self.add_child_window_to_monitor(BOVisualizer(routine=self.run_monitor.routine)) - def add_child_window_to_monitor(self, child_window: AnalysisExtension): + def add_bax_visualizer(self) -> None: + """ + Open the BaxVisualizer extension. + + """ + if self.run_monitor.routine is None: + QMessageBox.warning( + self, + "No Routine Error", + "Please start a routine before opening the Bax Visualizer.", + ) + return + + self.add_child_window_to_monitor( + BaxVisualizer(routine=self.run_monitor.routine) + ) + + def add_child_window_to_monitor(self, child_window: AnalysisExtension) -> None: """ Add a child window to the run monitor. diff --git a/src/badger/gui/components/run_monitor.py b/src/badger/gui/components/run_monitor.py index 4c7d0f9c..952dc4d1 100644 --- a/src/badger/gui/components/run_monitor.py +++ b/src/badger/gui/components/run_monitor.py @@ -1,15 +1,13 @@ +import logging import os import traceback from importlib import resources -from typing import List +from typing import TYPE_CHECKING, List, Optional -from badger.gui.components.analysis_extensions import AnalysisExtension import numpy as np import pandas as pd import pyqtgraph as pg from PyQt5.QtCore import pyqtSignal -from pyqtgraph.Qt import QtGui, QtCore - from PyQt5.QtGui import QIcon from PyQt5.QtWidgets import ( QCheckBox, @@ -23,21 +21,23 @@ QVBoxLayout, QWidget, ) +from pyqtgraph.Qt import QtCore, QtGui from xopt.vocs import VOCS, normalize_inputs, select_best -from badger.archive import archive_run, BADGER_ARCHIVE_ROOT +from badger.archive import BADGER_ARCHIVE_ROOT, archive_run +from badger.gui.components.analysis_extensions import AnalysisExtension +from badger.gui.components.extensions_palette import ExtensionsPalette from badger.gui.components.pydantic_editor import BadgerPydanticEditor +from badger.gui.components.routine_runner import BadgerRoutineSubprocess +from badger.gui.windows.message_dialog import BadgerScrollableMessageBox # from ...utils import AURORA_PALETTE, FROST_PALETTE from badger.logbook import BADGER_LOGBOOK_ROOT, send_to_logbook from badger.routine import Routine from badger.tests.utils import get_current_vars -from badger.gui.windows.message_dialog import BadgerScrollableMessageBox - -from badger.gui.components.extensions_palette import ExtensionsPalette -from badger.gui.components.routine_runner import BadgerRoutineSubprocess -import logging +if TYPE_CHECKING: + from badger.gui.components.process_manager import ProcessManager logger = logging.getLogger(__name__) @@ -64,7 +64,7 @@ class BadgerOptMonitor(QWidget): sig_toggle_other = pyqtSignal(bool) sig_env_ready = pyqtSignal() - def __init__(self, process_manager=None): + def __init__(self, process_manager: "Optional[ProcessManager]" = None): super().__init__() # self.setAttribute(Qt.WA_DeleteOnClose, True) @@ -111,7 +111,7 @@ def vocs(self) -> VOCS: def states(self, new_states: dict) -> None: self._states = new_states - def init_ui(self): + def init_ui(self) -> None: # Load all icons icon_ref = resources.files(__package__) / "../images/play.png" with resources.as_file(icon_ref) as icon_path: @@ -189,7 +189,7 @@ def init_ui(self): vbox.addWidget(monitor) # noinspection PyUnresolvedReferences - def config_logic(self): + def config_logic(self) -> None: """ Configure the logic and connections for various interactive elements in the application. @@ -231,7 +231,7 @@ def config_logic(self): self.cb_plot_y.currentIndexChanged.connect(self.select_x_plot_y_axis) self.check_relative.stateChanged.connect(self.toggle_x_plot_y_axis_relative) - def init_plots(self, routine: Routine = None, run_filename: str = None): + def init_plots(self, routine: Routine = None, run_filename: str = None) -> None: """ Initialize and configure the plots and related components in the application. @@ -397,7 +397,7 @@ def init_plots(self, routine: Routine = None, run_filename: str = None): self.sig_toggle_other.emit(False) - def _configure_plot(self, plot_object, inspector, names): + def _configure_plot(self, plot_object, inspector, names: list[str]) -> dict: plot_object.clear() plot_object.addItem(inspector) curves = {} @@ -430,7 +430,7 @@ def _configure_plot(self, plot_object, inspector, names): return curves - def init_routine_runner(self): + def init_routine_runner(self) -> None: self.reset_routine_runner() self.routine_runner = routine_runner = BadgerRoutineSubprocess( @@ -451,7 +451,7 @@ def init_routine_runner(self): self.sig_pause.connect(routine_runner.ctrl_routine) self.sig_stop.connect(routine_runner.stop_routine) - def reset_routine_runner(self): + def reset_routine_runner(self) -> None: if self.routine_runner: self.sig_pause.disconnect() self.sig_stop.disconnect() @@ -462,7 +462,7 @@ def start( use_termination_condition: bool = False, run_data_flag: bool = False, init_points_flag: bool = True, - ): + ) -> None: self.sig_new_run.emit() self.sig_status.emit(f"Running routine {self.routine.name}...") if not run_data_flag: @@ -478,10 +478,10 @@ def start( self.sig_run_started.emit() self.sig_lock.emit(True) - def save_termination_condition(self, tc): + def save_termination_condition(self, tc) -> None: self.termination_condition = tc - def enable_auto_range(self): + def enable_auto_range(self) -> None: # Enable autorange self.plot_obj.enableAutoRange() self.plot_var.enableAutoRange() @@ -491,14 +491,14 @@ def enable_auto_range(self): if self.vocs.observable_names: self.plot_obs.enableAutoRange() - def open_extensions_palette(self): + def open_extensions_palette(self) -> None: self.extensions_palette.show() - def extension_window_closed(self, child_window: AnalysisExtension): + def extension_window_closed(self, child_window: AnalysisExtension) -> None: self.active_extensions.remove(child_window) self.extensions_palette.update_palette() - def extract_timestamp(self, data=None): + def extract_timestamp(self, data: pd.DataFrame | None = None) -> np.ndarray: if data is None: data = self.routine.sorted_data @@ -530,7 +530,7 @@ def update(self, results: pd.DataFrame) -> None: # Check critical condition self.check_critical() - def update_curves(self, results=None): + def update_curves(self, results: pd.DataFrame | None = None) -> None: use_time_axis = self.plot_x_axis == 1 norm_inputs = self.x_plot_y_axis == 1 @@ -694,7 +694,7 @@ def destroy_unused_env(self) -> None: except AttributeError: # env already destroyed pass - def on_error(self, error): + def on_error(self, error: Exception) -> None: details = error._details if hasattr(error, "_details") else None dialog = BadgerScrollableMessageBox( @@ -705,10 +705,10 @@ def on_error(self, error): dialog.exec_() # Do not show info -- too distracting - def on_info(self, msg): + def on_info(self, msg) -> None: pass - def logbook(self): + def logbook(self) -> None: try: send_to_logbook(self.routine, self.monitor) except Exception as e: @@ -721,39 +721,39 @@ def logbook(self): # QMessageBox.information( # self, 'Success!', f'') - def ctrl_routine(self, status): + def ctrl_routine(self, status) -> None: self.sig_pause.emit(status) - def ins_obj_dragged(self, ins_obj): + def ins_obj_dragged(self, ins_obj) -> None: self.inspector_variable.setValue(ins_obj.value()) if self.vocs.constraint_names: self.inspector_constraint.setValue(ins_obj.value()) if self.vocs.observable_names: self.inspector_state.setValue(ins_obj.value()) - def ins_con_dragged(self, ins_con): + def ins_con_dragged(self, ins_con) -> None: self.inspector_variable.setValue(ins_con.value()) self.inspector_objective.setValue(ins_con.value()) if self.vocs.observable_names: self.inspector_state.setValue(ins_con.value()) - def ins_sta_dragged(self, ins_sta): + def ins_sta_dragged(self, ins_sta) -> None: self.inspector_variable.setValue(ins_sta.value()) self.inspector_objective.setValue(ins_sta.value()) if self.vocs.constraint_names: self.inspector_constraint.setValue(ins_sta.value()) - def ins_var_dragged(self, ins_var): + def ins_var_dragged(self, ins_var) -> None: self.inspector_objective.setValue(ins_var.value()) if self.vocs.constraint_names: self.inspector_constraint.setValue(ins_var.value()) if self.vocs.observable_names: self.inspector_state.setValue(ins_var.value()) - def ins_drag_done(self, ins): + def ins_drag_done(self, ins) -> None: self.sync_ins(ins.value()) - def sync_ins(self, pos): + def sync_ins(self, pos) -> None: if self.plot_x_axis: # x-axis is time value, idx = self.closest_ts(pos) else: @@ -771,7 +771,7 @@ def sync_ins(self, pos): self.sig_inspect.emit(int(idx)) - def closest_ts(self, t): + def closest_ts(self, t: float) -> tuple[float, int]: # Get the closest timestamp in data regarding t ts = self.extract_timestamp() ts -= ts[0] @@ -779,7 +779,7 @@ def closest_ts(self, t): return ts[idx], idx - def reset_env(self): + def reset_env(self) -> None: reply = QMessageBox.question( self, "Reset Environment", @@ -813,7 +813,7 @@ def reset_env(self): # QMessageBox.information(self, 'Reset Environment', # f'Env vars {curr_vars} -> {self.init_vars}') - def get_checkpoint(self): + def get_checkpoint(self) -> Optional[dict[str, float]]: if not self.routine or not self.routine.environment or not self.routine.vocs: return None @@ -821,7 +821,7 @@ def get_checkpoint(self): [*self.routine.environment.variables, *self.routine.vocs.variables.keys()] ) - def save_checkpoint(self): + def save_checkpoint(self) -> None: checkpoint = self.get_checkpoint() if checkpoint is None: QMessageBox.critical( @@ -834,7 +834,7 @@ def save_checkpoint(self): f"Checkpoint saved with the following data: {self.checkpoint_data}" ) - def edit_checkpoint(self): + def edit_checkpoint(self) -> None: if self.checkpoint_data is None: QMessageBox.information( self, "Edit Checkpoint", "No checkpoint data has been saved yet." @@ -884,7 +884,7 @@ def popup_editor_accepted() -> None: popup.exec() - def load_checkpoint(self): + def load_checkpoint(self) -> None: if self.checkpoint_data is None: QMessageBox.information( self, "Load Checkpoint", "No checkpoint data has been saved yet." @@ -903,7 +903,7 @@ def load_checkpoint(self): f"Checkpoint loaded with the following data: {self.checkpoint_data}" ) - def jump_to_optimal(self): + def jump_to_optimal(self) -> None: try: best_idx, _, _ = select_best( self.routine.vocs, self.routine.sorted_data, n=1 @@ -920,7 +920,7 @@ def jump_to_optimal(self): "Jump to optimum is not supported for multi-objective optimization yet", ) - def jump_to_solution(self, idx): + def jump_to_solution(self, idx: int) -> None: if self.plot_x_axis: # x-axis is time ts = self.extract_timestamp() value = ts[idx] - ts[0] @@ -934,7 +934,7 @@ def jump_to_solution(self, idx): self.inspector_state.setValue(value) self.inspector_variable.setValue(value) - def set_vars(self): + def set_vars(self) -> None: df = self.routine.sorted_data if self.plot_x_axis: # x-axis is time pos, idx = self.closest_ts(self.inspector_objective.value()) @@ -980,7 +980,7 @@ def set_vars(self): # QMessageBox.information( # self, 'Set Environment', f'Env vars have been set to {solution}') - def select_x_axis(self, i): + def select_x_axis(self, i: int) -> None: self.plot_x_axis = i # Switch the x-axis labels @@ -1015,11 +1015,11 @@ def select_x_axis(self, i): self.update_curves() self.enable_auto_range() - def select_x_plot_y_axis(self, i): + def select_x_plot_y_axis(self, i: int) -> None: self.x_plot_y_axis = i self.update_curves() - def toggle_x_plot_y_axis_relative(self): + def toggle_x_plot_y_axis_relative(self) -> None: self.x_plot_relative = self.check_relative.isChecked() # Change axes labels depending on if relative is checked @@ -1032,7 +1032,7 @@ def toggle_x_plot_y_axis_relative(self): self.update_curves() - def on_mouse_click(self, event): + def on_mouse_click(self, event) -> None: # https://stackoverflow.com/a/64081483 coor_obj = self.plot_obj.vb.mapSceneToView(event._scenePos) if self.vocs and self.vocs.constraint_names: @@ -1052,18 +1052,18 @@ def on_mouse_click(self, event): if flag: self.sync_ins(coor_obj.x()) - def delete_run(self): + def delete_run(self) -> None: self.sig_del.emit() - def stop(self): + def stop(self) -> None: self.sig_stop.emit() self.sig_stop_run.emit() - def register_post_run_action(self, action): + def register_post_run_action(self, action) -> None: self.post_run_actions.append(action) -def add_axes(monitor, ylabel, title, cursor_line, **kwargs): +def add_axes(monitor, ylabel, title, cursor_line, **kwargs) -> pg.PlotItem: plot_obj = monitor.addPlot(title=title, **kwargs) plot_obj.setLabel("left", ylabel) plot_obj.setLabel("bottom", "iterations") @@ -1076,7 +1076,7 @@ def add_axes(monitor, ylabel, title, cursor_line, **kwargs): return plot_obj -def create_cursor_line(): +def create_cursor_line() -> pg.InfiniteLine: return pg.InfiniteLine( movable=True, angle=90, @@ -1090,7 +1090,7 @@ def create_cursor_line(): ) -def set_data(names: List[str], curves: dict, data: pd.DataFrame, ts=None): +def set_data(names: List[str], curves: dict, data: pd.DataFrame, ts=None) -> None: # Split data into live and not live live_mask = data["live"].astype(bool) live_data = data.loc[live_mask] diff --git a/src/badger/gui/pages/home_page.py b/src/badger/gui/pages/home_page.py index 127d48c2..31e40a35 100644 --- a/src/badger/gui/pages/home_page.py +++ b/src/badger/gui/pages/home_page.py @@ -2,52 +2,55 @@ import os import traceback from importlib import resources +from typing import TYPE_CHECKING, Optional import numpy as np from pandas import DataFrame -from PyQt5.QtCore import pyqtSignal, Qt, QModelIndex +from PyQt5.QtCore import QModelIndex, Qt, pyqtSignal from PyQt5.QtGui import QIcon, QKeySequence from PyQt5.QtWidgets import ( + QLabel, QMessageBox, QShortcut, QSplitter, + QTabWidget, QVBoxLayout, QWidget, - QLabel, - QTabWidget, ) from badger.archive import ( delete_run, get_base_run_filename, - load_run, get_runs, + load_run, save_tmp_run, ) +from badger.errors import BadgerRoutineError +from badger.gui.components.action_bar import BadgerActionBar +from badger.gui.components.data_panel import filter_metadata from badger.gui.components.data_table import ( add_row, data_table, reset_table, update_table, ) - -from badger.gui.components.navigators import HistoryNavigator -from badger.gui.components.navigators import TemplateNavigator +from badger.gui.components.navigators import HistoryNavigator, TemplateNavigator from badger.gui.components.routine_page import BadgerRoutinePage from badger.gui.components.run_monitor import BadgerOptMonitor from badger.gui.components.status_bar import BadgerStatusBar -from badger.gui.components.action_bar import BadgerActionBar -from badger.gui.components.data_panel import filter_metadata -from badger.utils import get_header -from badger.settings import init_settings +from badger.gui.utils import ModalOverlay # from PyQt5.QtGui import QBrush, QColor from badger.gui.windows.message_dialog import BadgerScrollableMessageBox from badger.gui.windows.terminition_condition_dialog import ( BadgerTerminationConditionDialog, ) -from badger.gui.utils import ModalOverlay -from badger.errors import BadgerRoutineError +from badger.settings import init_settings +from badger.utils import get_header + +if TYPE_CHECKING: + from badger.gui.components.process_manager import ProcessManager + from badger.routine import VOCS import logging @@ -73,7 +76,7 @@ class BadgerHomePage(QWidget): sig_routine_activated = pyqtSignal(bool) sig_routine_invalid = pyqtSignal() - def __init__(self, process_manager=None): + def __init__(self, process_manager: "Optional[ProcessManager]" = None): logger.info("Initializing BadgerHomePage.") super().__init__() @@ -88,7 +91,7 @@ def __init__(self, process_manager=None): self.load_all_runs() self.init_home_page() - def init_ui(self): + def init_ui(self) -> None: logger.info("Initializing UI for BadgerHomePage.") self.config_singleton = init_settings() icon_ref = resources.files(__package__) / "../images/add.png" @@ -141,13 +144,11 @@ def init_ui(self): vbox_table = QVBoxLayout(panel_table) vbox_table.setContentsMargins(0, 0, 0, 0) title_label = QLabel("Run Data") - title_label.setStyleSheet( - """ + title_label.setStyleSheet(""" background-color: #455364; font-weight: bold; padding: 4px; - """ - ) + """) title_label.setAlignment(Qt.AlignCenter) # Center-align the title vbox_table.addWidget(title_label, 0) self.run_table = run_table = data_table() @@ -209,7 +210,7 @@ def init_ui(self): status_bar.set_summary("Badger is ready!") vbox.addWidget(status_bar) - def config_logic(self): + def config_logic(self) -> None: logger.info("Configuring logic for BadgerHomePage.") self.colors = ["c", "g", "m", "y", "b", "r", "w"] self.symbols = ["o", "t", "t1", "s", "p", "h", "d"] @@ -278,21 +279,21 @@ def config_logic(self): self.shortcut_go_search = QShortcut(QKeySequence("Ctrl+L"), self) self.shortcut_go_search.activated.connect(self.go_search) - def go_search(self): + def go_search(self) -> None: logger.info("Activating search bar.") self.sbar.setFocus() - def load_all_runs(self): + def load_all_runs(self) -> None: logger.info("Loading all runs into history browser.") runs = get_runs() self.history_browser.updateItems(runs) - def init_home_page(self): + def init_home_page(self) -> None: logger.info("Initializing home page.") # Load the default generator self.routine_editor.generator_box.cb.setCurrentIndex(0) - def go_run(self, i: int = None): + def go_run(self, i: int = None) -> None: logger.info(f"Activating run: {i}") gc.collect() @@ -351,7 +352,7 @@ def go_run(self, i: int = None): self.run_monitor.update_analysis_extensions() - def go_template(self, index: QModelIndex): + def go_template(self, index: QModelIndex) -> None: path = self.template_browser.file_sys_model.filePath(index) # if directory, expand it. if os.path.isdir(path): @@ -364,16 +365,16 @@ def go_template(self, index: QModelIndex): self.status_bar.set_summary(f"Current template {path}") return - def inspect_solution(self, idx): + def inspect_solution(self, idx: int) -> None: logger.info(f"Inspecting solution at index: {idx}") self.run_table.selectRow(idx) self.run_table_2.selectRow(idx) - def solution_selected(self, r, c): + def solution_selected(self, r: int, c: int) -> None: logger.info(f"Solution selected at row {r}, column {c}") self.run_monitor.jump_to_solution(r) - def table_selection_changed(self): + def table_selection_changed(self) -> None: logger.info("Table selection changed.") indices = self.run_table.selectedIndexes() indices = self.run_table_2.selectedIndexes() @@ -397,7 +398,7 @@ def table_selection_changed(self): self.run_monitor.jump_to_solution(row) - def toggle_lock(self, lock, lock_tab=1): + def toggle_lock(self, lock: bool, lock_tab: int = 1) -> None: logger.info(f"Toggling lock: {lock}, tab: {lock_tab}") if lock: self.history_browser.setDisabled(True) @@ -406,7 +407,7 @@ def toggle_lock(self, lock, lock_tab=1): self.uncover_page() - def validate_loaded_data_keys(self, vocs): + def validate_loaded_data_keys(self, vocs: "VOCS") -> None: """ This function is called when adding historical data to a new routine. It makes sure that the keys of data to be loaded from data_panel match the @@ -458,7 +459,9 @@ def validate_loaded_data_keys(self, vocs): self.run_action_bar.routine_finished() # Reset action bar raise BadgerRoutineError("Routine initialization cancelled by user.") - def prepare_run(self, data=None, init_points_flag=True): + def prepare_run( + self, data: Optional["DataFrame"] = None, init_points_flag: bool = True + ) -> None: """ Prepares the run by composing the routine, validating data if present, saving created routine to a yaml file, and passing the routine to @@ -519,7 +522,7 @@ def prepare_run(self, data=None, init_points_flag=True): # Tell monitor to start the run self.run_monitor.init_plots(routine) - def start_run(self, use_termination_condition: bool = False): + def start_run(self, use_termination_condition: bool = False) -> None: """ Prepares and starts optimization run with provided options. - Termination Condition is provided when called via BadgerTerminationConditionDialog @@ -557,7 +560,7 @@ def start_run(self, use_termination_condition: bool = False): init_points_flag=init_points_flag, ) - def start_run_until(self): + def start_run_until(self) -> None: logger.info("Starting run until condition met.") dlg = BadgerTerminationConditionDialog( self, @@ -572,7 +575,7 @@ def start_run_until(self): self.tc_dialog = None # self.run_monitor.start_until() - def new_run(self): + def new_run(self) -> None: logger.info("Creating new run.") self.cover_page() @@ -582,17 +585,17 @@ def new_run(self): header = get_header(self.current_routine) reset_table(self.run_table, header) - def run_name(self, name): + def run_name(self, name: str) -> None: logger.info(f"Updating run name: {name}") runs = get_runs() self.history_browser.updateItems(runs) self.history_browser._selectItemByRun(name) - def update_status(self, info): + def update_status(self, info: str) -> None: logger.info(f"Updating status: {info}") self.status_bar.set_summary(info) - def progress(self, solution: DataFrame): + def progress(self, solution: DataFrame) -> None: vocs = self.current_routine.vocs vars = list(solution[vocs.variable_names].to_numpy()[0]) objs = list(solution[vocs.objective_names].to_numpy()[0]) @@ -601,7 +604,7 @@ def progress(self, solution: DataFrame): add_row(self.run_table, objs + cons + vars + stas) self.data_panel.add_live_data(solution) - def delete_run(self): + def delete_run(self) -> None: logger.info("Deleting run.") run_name = get_base_run_filename(self.history_browser.currentText()) @@ -622,7 +625,7 @@ def delete_run(self): self.history_browser.history_tree_widget.blockSignals(False) self.go_run(-1) - def cover_page(self): + def cover_page(self) -> None: logger.info("Covering page with overlay.") return # disable overlay for now @@ -637,7 +640,7 @@ def cover_page(self): self.overlay = ModalOverlay(main_window) self.overlay.show() - def uncover_page(self): + def uncover_page(self) -> None: logger.info("Uncovering page overlay.") return # disable overlay for now diff --git a/src/badger/utils.py b/src/badger/utils.py index cb58342c..40b88c90 100644 --- a/src/badger/utils.py +++ b/src/badger/utils.py @@ -1,17 +1,20 @@ -from importlib import metadata import json import logging import os -import sys import pathlib +import sys from datetime import datetime +from importlib import metadata from types import TracebackType -from typing import Iterable, Optional +from typing import TYPE_CHECKING, Iterable, Optional import yaml +from PyQt5.QtWidgets import QLayout, QWidget from badger.errors import BadgerLoadConfigError -from PyQt5.QtWidgets import QWidget, QLayout + +if TYPE_CHECKING: + from badger.routine import Routine logger = logging.getLogger(__name__) @@ -175,7 +178,7 @@ def curr_ts_to_str(format="lcls-log"): return ts_to_str(datetime.now(), format) -def create_archive_run_filename(routine, format: str = "lcls-fname") -> str: +def create_archive_run_filename(routine: "Routine", format: str = "lcls-fname") -> str: data = routine.sorted_data env_name = routine.environment.name data_dict = data.to_dict("list") From 716f14eb4221d0645b07f774536987e7e3b8e907 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Tue, 16 Jun 2026 16:55:03 -0700 Subject: [PATCH 02/18] Working to add bax algorithm --- src/badger/factory.py | 30 ++-- src/badger/gui/components/pydantic_editor.py | 173 +++++++++++-------- src/badger/gui/components/routine_page.py | 104 ++++++----- 3 files changed, 171 insertions(+), 136 deletions(-) diff --git a/src/badger/factory.py b/src/badger/factory.py index 30f7987d..27292eef 100644 --- a/src/badger/factory.py +++ b/src/badger/factory.py @@ -1,23 +1,23 @@ -from typing import Any, TypedDict, cast, TYPE_CHECKING -from badger.settings import init_settings -from badger.utils import get_value_or_none +import importlib +import logging +import os +import re +import sys +from pathlib import Path +from typing import TYPE_CHECKING, Any, TypedDict, cast + +import yaml +from xopt.generators import generators, get_generator_defaults + from badger.errors import ( BadgerConfigError, - BadgerInvalidPluginError, BadgerInvalidDocsError, + BadgerInvalidPluginError, BadgerPluginNotFoundError, ) - from badger.interface import Interface as BadgerInterface -import sys -import os -import importlib -import yaml -import re -from pathlib import Path -from xopt.generators import generators, get_generator_defaults - -import logging +from badger.settings import init_settings +from badger.utils import get_value_or_none if TYPE_CHECKING: from badger.environment import Environment as BadgerEnvironment @@ -33,7 +33,7 @@ "time_dependent_upper_confidence_bound", "multi_fidelity", "nsga2", - "BAX", + # "BAX", ] diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index 286a13bb..f7e07a83 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -1,50 +1,49 @@ +import ast +import logging +import re from dataclasses import dataclass +from inspect import isclass from types import NoneType from typing import ( Annotated, + Any, Callable, Optional, - Any, Sequence, TypeVar, - cast, - get_origin, Union, + cast, get_args, + get_origin, ) -from inspect import isclass -from pydantic_core import PydanticUndefined import yaml +from pydantic import BaseModel, Field, ValidationError, create_model +from pydantic.fields import FieldInfo +from pydantic_core import PydanticUndefined from PyQt5.QtCore import Qt, pyqtSignal from PyQt5.QtWidgets import ( - QTreeWidget, - QTreeWidgetItem, - QSpinBox, - QDoubleSpinBox, QCheckBox, - QLineEdit, + QComboBox, + QDoubleSpinBox, + QHBoxLayout, QLabel, - QWidget, + QLineEdit, QPushButton, - QVBoxLayout, QScrollArea, - QHBoxLayout, - QComboBox, QSizePolicy, + QSpinBox, + QTreeWidget, + QTreeWidgetItem, + QVBoxLayout, + QWidget, ) -from pydantic import BaseModel, Field, ValidationError, create_model -from pydantic.fields import FieldInfo from xopt.generators import get_generator +from xopt.generators.bayesian.bax.algorithms import Algorithm +from xopt.generators.bayesian.bax_generator import BaxGenerator +from xopt.generators.bayesian.bayesian_generator import BayesianGenerator from xopt.generators.bayesian.turbo import TurboController from xopt.numerical_optimizer import NumericalOptimizer -from xopt.generators.bayesian.bayesian_generator import BayesianGenerator - -import re -import ast - -import logging - from xopt.vocs import VOCS logger = logging.getLogger(__name__) @@ -57,8 +56,16 @@ class CustomSafeLoader(yaml.SafeLoader): - def tuple_constructor(self, node: yaml.ScalarNode | yaml.MappingNode): - value = self.construct_scalar(node) + def __init__(self, stream: Any) -> None: + super().__init__(stream) + self.add_constructor( + "tag:yaml.org,2002:null", type(self).handle_null_constructor + ) + self.add_constructor("tag:yaml.org,2002:str", type(self).tuple_constructor) + + @staticmethod + def tuple_constructor(loader: Any, node: yaml.ScalarNode) -> Any: + value = loader.construct_scalar(node) if TUPLE_PATTERN.match(value): try: return ast.literal_eval(value) @@ -66,10 +73,12 @@ def tuple_constructor(self, node: yaml.ScalarNode | yaml.MappingNode): pass return value - -CustomSafeLoader.add_constructor( - "tag:yaml.org,2002:str", CustomSafeLoader.tuple_constructor -) + @staticmethod + def handle_null_constructor(loader: Any, node: yaml.ScalarNode) -> Any: + value = loader.construct_scalar(node) + if value.lower() in ["null", "none", ""]: + return None + return value def convert_to_type(value: Any, type: Callable[[Any], T]) -> T: @@ -82,9 +91,9 @@ def convert_to_type(value: Any, type: Callable[[Any], T]) -> T: def _set_value_for_basic_widget( - widget: Any, + widget: QWidget, value: str | float | int | bool | None, -): +) -> None: if isinstance(widget, QLabel) or isinstance(widget, QLineEdit): widget.setText("null" if value is None else str(value)) elif isinstance(widget, QDoubleSpinBox): @@ -103,7 +112,7 @@ class BadgerResolvedType: @classmethod def find_primary( - cls, annotations: list[Any] | tuple[Any, ...] + cls, annotations: tuple[Any, ...] ) -> Optional["BadgerResolvedType"]: if len(annotations) == 0: return None @@ -147,11 +156,11 @@ def resolve( if origin == Union: if NoneType in args: origin = Optional - args = [arg for arg in args if arg != NoneType] + args = tuple(arg for arg in args if arg != NoneType) else: if len(args) == 1: origin = args[0] - args = [] + args = () nullable = True if origin is not None: @@ -168,7 +177,7 @@ def resolve( if origin == Optional: origin = args[0] - args = [] + args = () nullable = True if origin is not None: @@ -199,7 +208,7 @@ def resolve_qt( editor_info: tuple["BadgerPydanticEditor", QTreeWidgetItem] | None = None, ) -> QWidget | None: resolved_type = BadgerResolvedType.resolve(annotation) - widget = QLabel() + widget: QWidget | None = None if resolved_type.main is None: widget = QLineEdit() @@ -209,6 +218,8 @@ def resolve_qt( widget = QComboBox() elif issubclass(resolved_type.main, NumericalOptimizer): widget = QComboBox() + elif issubclass(resolved_type.main, Algorithm): + widget = QComboBox() else: return None if resolved_type.nullable: @@ -308,7 +319,7 @@ def resolve_qt( return widget -def handle_changed(editor_info: tuple["BadgerPydanticEditor", QTreeWidgetItem]): +def handle_changed(editor_info: tuple["BadgerPydanticEditor", QTreeWidgetItem]) -> None: tree_widget, _ = editor_info tree_widget.validate() @@ -427,7 +438,6 @@ def __init__(self, editor: "BadgerListEditor", parent: QWidget | None = None): QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Preferred ) if isinstance(self.parameter_value, QLineEdit): - self.parameter_value = cast(QLineEdit, self.parameter_value) self.parameter_value.editingFinished.connect( lambda: self.editor.listChanged.emit() ) @@ -444,7 +454,6 @@ def __init__(self, editor: "BadgerListEditor", parent: QWidget | None = None): ) if isinstance(self.parameter_value2, QDoubleSpinBox): - self.parameter_value2 = cast(QDoubleSpinBox, self.parameter_value2) self.parameter_value2.valueChanged.connect( lambda: self.editor.listChanged.emit() ) @@ -458,13 +467,13 @@ def __init__(self, editor: "BadgerListEditor", parent: QWidget | None = None): remove_button.clicked.connect(self.remove) layout.addWidget(remove_button, alignment=Qt.AlignmentFlag.AlignRight) - def parameter1(self): + def parameter1(self) -> QWidget | None: return self.parameter_value - def parameter2(self): + def parameter2(self) -> QWidget | None: return self.parameter_value2 - def remove(self): + def remove(self) -> None: self.setParent(None) self.deleteLater() self.editor.listChanged.emit() @@ -504,16 +513,16 @@ def __init__( layout.addLayout(button_layout) - def handle_button_click(self): + def handle_button_click(self) -> None: self.add_widget() - def add_widget(self): + def add_widget(self) -> BadgerListItem: widget = BadgerListItem(self) self.list_layout.addWidget(widget) self.listChanged.emit() return widget - def get_parameters_yaml(self): + def get_parameters_yaml(self) -> str | None: child_values: list[str | None] = [ _qt_widget_to_yaml_value(child.parameter_value) for child in self.list_container.children() @@ -539,7 +548,7 @@ def get_parameters_yaml(self): + "}" ) - def get_parameters_dict(self): + def get_parameters_dict(self) -> dict[str, Any] | None: child_values: list[str | None] = [ _qt_widget_to_value(child.parameter_value) for child in self.list_container.children() @@ -593,7 +602,7 @@ def _set_params_recurse( fields: dict[str, FieldInfo], defaults: dict[str, Any] | None, hidden: bool, - ): + ) -> None: for field_name, field_info in fields.items(): child = QTreeWidgetItem( [field_name if i == 0 else "" for i in range(0, self.value_col + 1)] @@ -603,7 +612,8 @@ def _set_params_recurse( self.addTopLevelItem(child) else: parent.addChild(child) - child.setToolTip(0, field_info.description) + if field_info.description: + child.setToolTip(0, field_info.description) widget = BadgerResolvedType.resolve_qt( annotation=field_info.annotation, @@ -645,27 +655,30 @@ def initialize_combo_widget( self, widget: QComboBox, selections: Sequence[type[BaseModel] | None], - ): + ) -> None: # Clear out existing children widget.clear() for selection in selections: if selection is None: widget.addItem("null", selection) else: + logger.debug( + f"Adding selection {selection} with name {selection.model_fields['name'].default} to combo box" + ) widget.addItem(selection.model_fields["name"].default, selection) - def set_params_from_class(self, pydantic_class: type[Any]): + def set_params_from_class(self, pydantic_class: type[Any]) -> None: self.clear() self.model_class = pydantic_class self._set_params_recurse(None, self.model_class.model_fields, None, False) self.set_params_post_setup({}) self.validate() - def set_params_from_dict(self, params: dict[str, Any]): + def set_params_from_dict(self, params: dict[str, Any]) -> None: self.clear() field_definitions: dict[str, tuple[type, FieldInfo]] = { k: (type(v), Field()) for k, v in params.items() - } # type: ignore + } if field_definitions == {}: logger.warning("No fields found in params dictionary") return @@ -684,7 +697,7 @@ def set_params_from_generator( defaults: dict[str, Any], vocs: VOCS | None = None, validate: bool = True, - ): + ) -> None: logger.debug(f"vocs: {vocs}") logger.debug(f"defaults: {defaults}") self.vocs = vocs or VOCS(variables={}) @@ -722,7 +735,7 @@ def set_params_from_generator( if validate: self.validate() - def set_params_post_setup(self, defaults: dict[str, Any]): + def set_params_post_setup(self, defaults: dict[str, Any]) -> None: if self.model_class is None: raise ValueError("Model class is not set.") @@ -733,7 +746,14 @@ def set_params_post_setup(self, defaults: dict[str, Any]): if self.model_class.model_fields.get("numerical_optimizer") is not None: self.initialize_special_field(defaults, "numerical_optimizer") - def initialize_special_field(self, defaults: dict[str, Any], field: str): + if issubclass(self.model_class, BaxGenerator): + logger.debug( + f"Generator is a {self.model_class.__name__} checking for algorithm field." + ) + if self.model_class.model_fields.get("algorithm") is not None: + self.initialize_special_field(defaults, "algorithm") + + def initialize_special_field(self, defaults: dict[str, Any], field: str) -> None: widget_items = self.findItems(field, Qt.MatchFlag.MatchExactly) if len(widget_items) == 0: @@ -748,7 +768,6 @@ def initialize_special_field(self, defaults: dict[str, Any], field: str): widget = self.itemWidget(special_item, self.value_col) if widget is None or not isinstance(widget, QComboBox): raise ValueError(f"{field} does not have a combo box widget.") - widget = cast(QComboBox, widget) selections = self.get_all_compatible_classes(field) @@ -778,10 +797,12 @@ def initialize_special_field(self, defaults: dict[str, Any], field: str): ) widget.currentIndexChanged.connect( - lambda: self.on_radio_changed(special_item, "turbo_controller") + lambda: self.on_radio_changed(special_item, field) ) - def get_all_compatible_classes(self, field_name: str): + def get_all_compatible_classes( + self, field_name: str + ) -> Sequence[type[BaseModel] | None]: if self.model_class is None: raise ValueError("Model class is not set.") @@ -795,6 +816,10 @@ def get_all_compatible_classes(self, field_name: str): if not issubclass(self.model_class, BayesianGenerator): raise ValueError("Generator does not support turbo controllers.") compatible_classes = self.model_class.get_compatible_turbo_controllers() + elif field_name == "algorithm": + if not issubclass(self.model_class, BaxGenerator): + raise ValueError("Generator does not support algorithms.") + compatible_classes = self.model_class.get_compatible_algorithms() else: raise ValueError(f"Field name {field_name} is not recognized.") @@ -816,7 +841,7 @@ def get_compatible_class(self, name: str, field_name: str) -> type[BaseModel]: if selected_class is None: raise ValueError( - f"Generator has numerical optimizer set but no compatible numerical optimizer with name {name} exists." + f"Generator has {field_name} set but no compatible {field_name} with name {name} exists." ) return selected_class @@ -825,7 +850,7 @@ def on_radio_changed( self, tree_widget_item: QTreeWidgetItem, field_name: str, - ): + ) -> None: # Clear out existing children for cc in tree_widget_item.takeChildren(): del cc @@ -833,7 +858,6 @@ def on_radio_changed( widget = self.itemWidget(tree_widget_item, self.value_col) if widget is None or not isinstance(widget, QComboBox): raise ValueError("tree widget item does not have a combo box widget.") - widget = cast(QComboBox, widget) name = widget.currentText() if name == "null": @@ -857,7 +881,7 @@ def update_params_from_generator_class( name: str, field_name: str, defaults: dict[str, Any], - ): + ) -> None: try: pydantic_class = self.get_compatible_class(name, field_name) except ValueError: @@ -920,7 +944,7 @@ def exclude_condition(k: str) -> bool: return filtered_class_fields, removed_class_fields @staticmethod - def get_defaults_from_type(pydantic_class: type[Any]): + def get_defaults_from_type(pydantic_class: type[Any]) -> dict[str, Any]: if not issubclass(pydantic_class, BaseModel): raise ValueError("Provided class is not a Pydantic model") defaults: dict[str, Any] = {} @@ -972,7 +996,7 @@ def find_widget_at_path( return None - def remove_style(self, item: QTreeWidgetItem | None): + def remove_style(self, item: QTreeWidgetItem | None) -> None: # Have to reset border styling in case some errors were fixed if item is None: return @@ -984,7 +1008,7 @@ def remove_style(self, item: QTreeWidgetItem | None): for j in range(item.childCount()): self.remove_style(item.child(j)) - def update_vocs(self, vocs: VOCS): + def update_vocs(self, vocs: VOCS) -> None: logger.debug(f"Updating VOCS in BadgerPydanticEditor: {vocs}") self.vocs = vocs @@ -996,7 +1020,7 @@ def update_vocs(self, vocs: VOCS): self.set_params_from_generator(self.generator_name, defaults, self.vocs) self.validate() - def update_after_validate(self, defaults: dict[str, Any]): + def update_after_validate(self, defaults: dict[str, Any]) -> None: model_class = self.model_class if model_class is None: return @@ -1029,9 +1053,9 @@ def update_after_validate(self, defaults: dict[str, Any]): if self.update_callback is not None: self.update_callback(self) - def validate(self): + def validate(self) -> bool: if self.model_class is None: - return False + raise ValueError("Model class is not set.") self.setStyleSheet("") for i in range(self.topLevelItemCount()): @@ -1041,7 +1065,7 @@ def validate(self): parameters = self.get_parameters_yaml() parameters_dict = yaml.load(parameters, Loader=CustomSafeLoader) - def convert_dict(val): + def convert_dict(val: Any) -> Any: # Convert str-encoded dicts (and lists) back into actual dict objects.""" if isinstance(val, str): stripped = val.strip() @@ -1068,10 +1092,7 @@ def convert_dict(val): # Not valid unlesss has at least one objective if "vocs" not in parameters_dict: - return False - obj_data = parameters_dict["vocs"].get("objectives") - if not obj_data or len(obj_data) == 0: - return False + raise KeyError("vocs field is required in parameters") model = self.model_class.model_validate(parameters_dict) @@ -1085,14 +1106,19 @@ def convert_dict(val): self.update_after_validate(defaults) return True + except KeyError as e: logger.error(e) + return False except ValidationError as e: logger.error(e) for error in e.errors(): loc = error["loc"] msg = error["msg"] + error_widget: ( + QTreeWidgetItem | QTreeWidget | "BadgerPydanticEditor" | None + ) = None if len(loc) > 0: error_widget = self.find_widget_at_path(loc) else: @@ -1107,7 +1133,6 @@ def convert_dict(val): '*[error="true"] { border: 2px dashed red }' ) if isinstance(widget, BadgerListEditor): - widget = cast(BadgerListEditor, widget) widget.list_container.setProperty("error", True) widget.list_container.setStyleSheet( '*[error="true"] { border: 2px dashed red }' diff --git a/src/badger/gui/components/routine_page.py b/src/badger/gui/components/routine_page.py index 0a5d1c35..12a19df6 100644 --- a/src/badger/gui/components/routine_page.py +++ b/src/badger/gui/components/routine_page.py @@ -1,38 +1,60 @@ -from typing import Any -import warnings -import traceback import copy -from functools import partial +import logging import os -import yaml +import traceback +import warnings +from datetime import datetime +from functools import partial +from typing import Any import numpy as np import pandas as pd -from PyQt5.QtCore import Qt, pyqtSignal, QTimer -from PyQt5.QtWidgets import QLineEdit, QLabel, QPushButton, QFileDialog -from PyQt5.QtWidgets import QMessageBox, QWidget, QTabWidget -from PyQt5.QtWidgets import QVBoxLayout, QHBoxLayout, QScrollArea -from PyQt5.QtWidgets import QTableWidgetItem, QPlainTextEdit +import yaml from coolname import generate_slug -from xopt import VOCS -from xopt.vocs import random_inputs -from xopt.generators import ( - get_generator_defaults, - all_generator_names, - get_generator_dynamic, -) -from xopt.utils import get_local_region from gest_api.vocs import ( BaseConstraint, BaseObjective, GreaterThanConstraint, LessThanConstraint, - MinimizeObjective, MaximizeObjective, + MinimizeObjective, ) from pydantic import ValidationError +from PyQt5.QtCore import Qt, QTimer, pyqtSignal +from PyQt5.QtWidgets import ( + QFileDialog, + QHBoxLayout, + QLabel, + QLineEdit, + QMessageBox, + QPlainTextEdit, + QPushButton, + QScrollArea, + QTableWidgetItem, + QTabWidget, + QVBoxLayout, + QWidget, +) +from xopt import VOCS +from xopt.generators import ( + all_generator_names, + get_generator_defaults, + get_generator_dynamic, +) +from xopt.utils import get_local_region +from xopt.vocs import random_inputs -from badger.gui.components.generator_cbox import BadgerAlgoBox +from badger.archive import update_run +from badger.environment import instantiate_env +from badger.errors import ( + BadgerEnvInstantiationError, + BadgerEnvNotFoundError, + BadgerEnvVarError, + BadgerRoutineError, + VariableRangeError, +) +from badger.factory import get_env, list_env, list_generators +from badger.gui.components.archive_search import ArchiveSearchWidget from badger.gui.components.data_panel import BadgerDataPanel from badger.gui.components.data_table import ( get_table_content_as_dict, @@ -41,42 +63,28 @@ ) from badger.gui.components.env_cbox import BadgerEnvBox from badger.gui.components.filter_cbox import BadgerFilterBox +from badger.gui.components.generator_cbox import BadgerAlgoBox +from badger.gui.utils import filter_generator_config +from badger.gui.windows.add_random_dialog import BadgerAddRandomDialog from badger.gui.windows.docs_window import BadgerDocsWindow from badger.gui.windows.edit_script_dialog import BadgerEditScriptDialog -from badger.gui.windows.lim_vrange_dialog import BadgerLimitVariableRangeDialog from badger.gui.windows.ind_lim_vrange_dialog import ( BadgerIndividualLimitVariableRangeDialog, ) -from badger.gui.windows.review_dialog import BadgerReviewDialog -from badger.gui.windows.add_random_dialog import BadgerAddRandomDialog +from badger.gui.windows.lim_vrange_dialog import BadgerLimitVariableRangeDialog from badger.gui.windows.message_dialog import BadgerScrollableMessageBox -from badger.gui.utils import filter_generator_config -from badger.gui.components.archive_search import ArchiveSearchWidget -from badger.archive import update_run -from badger.environment import instantiate_env -from badger.errors import ( - BadgerEnvNotFoundError, - BadgerRoutineError, - BadgerEnvVarError, - BadgerEnvInstantiationError, - VariableRangeError, -) -from badger.factory import list_generators, list_env, get_env +from badger.gui.windows.review_dialog import BadgerReviewDialog from badger.routine import Routine from badger.settings import init_settings -from datetime import datetime from badger.utils import ( BlockSignalsContext, - load_config, - strtobool, get_badger_version, get_xopt_version, + load_config, + strtobool, ts_float_to_str, ) - -import logging - logger = logging.getLogger(__name__) @@ -265,8 +273,7 @@ def init_ui(self): self.env_box = BadgerEnvBox(env_dict, None, self.envs) scroll_area = QScrollArea() scroll_area.setFrameShape(QScrollArea.NoFrame) - scroll_area.setStyleSheet( - """ + scroll_area.setStyleSheet(""" QScrollArea { border: none; /* Remove border */ margin: 0px; /* Remove margin */ @@ -275,8 +282,7 @@ def init_ui(self): QScrollArea > QWidget { margin: 0px; /* Remove margin inside */ } - """ - ) + """) scroll_content_env = QWidget() scroll_layout_env = QVBoxLayout(scroll_content_env) scroll_layout_env.setContentsMargins(0, 0, 15, 0) @@ -1790,9 +1796,13 @@ def _compose_routine(self) -> Routine: if not vocs.variables: logger.error("No variables selected.") raise BadgerRoutineError("no variables selected") + + NO_OBJECTIVE_GENERATORS = ["bax"] + if not vocs.objectives: - logger.error("No objectives selected.") - raise BadgerRoutineError("no objectives selected") + if generator_name not in NO_OBJECTIVE_GENERATORS: + logger.error("No objectives selected.") + raise BadgerRoutineError("no objectives selected") # Initial points init_points_df = pd.DataFrame.from_dict( From 91362752269578fa8b20b69a7396aa26e959ee13 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Tue, 16 Jun 2026 17:47:11 -0700 Subject: [PATCH 03/18] Fixed None issue and computed fields validation --- src/badger/gui/components/pydantic_editor.py | 52 +++++++++++++++++--- 1 file changed, 45 insertions(+), 7 deletions(-) diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index f7e07a83..54892d0d 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -338,14 +338,11 @@ def _qt_widget_to_yaml_value(widget: Any) -> str | None: return "null" return f'"{widget.currentText()}"' elif isinstance(widget, (QLabel, QLineEdit)): - if widget.text() == "null" or widget.text() == "None": - return '"null"' + text = widget.text() + if text == "null" or text == "None" or text == "": + return "null" else: - return ( - ('"' + widget.text() + '"') - if isinstance(widget, QLineEdit) - else widget.text() - ) + return ('"' + text + '"') if isinstance(widget, QLineEdit) else text return "null" @@ -1053,6 +1050,45 @@ def update_after_validate(self, defaults: dict[str, Any]) -> None: if self.update_callback is not None: self.update_callback(self) + def _inject_computed_fields( + self, + parameters_dict: Any, + model_class: type[BaseModel], + ) -> None: + # Pydantic @computed_field values aren't editable in the form, but some + # validators (e.g. Xopt's discriminated-union dispatch on `class_path`) + # require them in the input dict. Walk the dict, derive each level's + # computed fields from its pydantic class, and inject them. + if not isinstance(parameters_dict, dict): + return + + try: + instance = model_class.model_construct(**parameters_dict) + for cf_name in model_class.model_computed_fields: + if cf_name in parameters_dict: + continue + try: + parameters_dict[cf_name] = getattr(instance, cf_name) + except Exception as e: + logger.debug( + f"Could not compute {cf_name} on {model_class.__name__}: {e}" + ) + except Exception as e: + logger.debug( + f"Could not model_construct {model_class.__name__} for computed-field injection: {e}" + ) + + if isclass(model_class) and issubclass(model_class, BayesianGenerator): + for sub_field in ("algorithm", "numerical_optimizer", "turbo_controller"): + sub = parameters_dict.get(sub_field) + if not isinstance(sub, dict) or "name" not in sub: + continue + try: + sub_class = self.get_compatible_class(sub["name"], sub_field) + except (ValueError, AttributeError): + continue + self._inject_computed_fields(sub, sub_class) + def validate(self) -> bool: if self.model_class is None: raise ValueError("Model class is not set.") @@ -1094,6 +1130,8 @@ def convert_dict(val: Any) -> Any: if "vocs" not in parameters_dict: raise KeyError("vocs field is required in parameters") + self._inject_computed_fields(parameters_dict, self.model_class) + model = self.model_class.model_validate(parameters_dict) # After we validate the model, some fields may have been changed due to any validation logic present within the Pydantic model (i.e. field validators changing default values). We need to update the tree to reflect these changes. From 87fa2e43200c2f43006c457e460576318b1f73e3 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Thu, 18 Jun 2026 18:03:05 -0700 Subject: [PATCH 04/18] Fixed issue with pydantic editor to respect optional types --- src/badger/gui/components/pydantic_editor.py | 72 ++++++++++++++------ 1 file changed, 53 insertions(+), 19 deletions(-) diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index 54892d0d..b1386956 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -58,9 +58,6 @@ class CustomSafeLoader(yaml.SafeLoader): def __init__(self, stream: Any) -> None: super().__init__(stream) - self.add_constructor( - "tag:yaml.org,2002:null", type(self).handle_null_constructor - ) self.add_constructor("tag:yaml.org,2002:str", type(self).tuple_constructor) @staticmethod @@ -73,13 +70,6 @@ def tuple_constructor(loader: Any, node: yaml.ScalarNode) -> Any: pass return value - @staticmethod - def handle_null_constructor(loader: Any, node: yaml.ScalarNode) -> Any: - value = loader.construct_scalar(node) - if value.lower() in ["null", "none", ""]: - return None - return value - def convert_to_type(value: Any, type: Callable[[Any], T]) -> T: try: @@ -94,14 +84,24 @@ def _set_value_for_basic_widget( widget: QWidget, value: str | float | int | bool | None, ) -> None: + nullable = bool(widget.property("badger_nullable")) if isinstance(widget, QLabel) or isinstance(widget, QLineEdit): widget.setText("null" if value is None else str(value)) elif isinstance(widget, QDoubleSpinBox): - widget.setValue(convert_to_type(value, float)) + if value is None and nullable: + widget.setValue(widget.minimum()) + else: + widget.setValue(convert_to_type(value, float)) elif isinstance(widget, QSpinBox): - widget.setValue(convert_to_type(value, int)) + if value is None and nullable: + widget.setValue(widget.minimum()) + else: + widget.setValue(convert_to_type(value, int)) elif isinstance(widget, QCheckBox): - widget.setChecked(convert_to_type(value, bool)) + if value is None and nullable: + widget.setCheckState(Qt.CheckState.PartiallyChecked) + else: + widget.setChecked(convert_to_type(value, bool)) @dataclass @@ -285,23 +285,43 @@ def resolve_qt( widget = QDoubleSpinBox() widget.setRange(float("-inf"), float("inf")) widget.setDecimals(6) - value = convert_to_type(default, float) if default is not None else 0.0 - widget.setValue(value) + if resolved_type.nullable: + # The minimum value doubles as the "null" sentinel. + widget.setSpecialValueText("null") + if default is not None: + widget.setValue(convert_to_type(default, float)) + elif resolved_type.nullable: + widget.setValue(widget.minimum()) + else: + widget.setValue(0.0) if editor_info is not None: widget.valueChanged.connect(lambda: handle_changed(editor_info)) elif resolved_type.main is int: widget = QSpinBox() widget.setRange(-(2**31), 2**31 - 1) # int32 min/max - value = convert_to_type(default, int) if default is not None else 0 - widget.setValue(value) + if resolved_type.nullable: + # The minimum value doubles as the "null" sentinel. + widget.setSpecialValueText("null") + if default is not None: + widget.setValue(convert_to_type(default, int)) + elif resolved_type.nullable: + widget.setValue(widget.minimum()) + else: + widget.setValue(0) if editor_info is not None: widget.valueChanged.connect(lambda: handle_changed(editor_info)) elif resolved_type.main is bool: widget = QCheckBox() - value = convert_to_type(default, bool) if default is not None else False - widget.setChecked(value) + if resolved_type.nullable: + widget.setTristate(True) + if default is not None: + widget.setChecked(convert_to_type(default, bool)) + elif resolved_type.nullable: + widget.setCheckState(Qt.CheckState.PartiallyChecked) + else: + widget.setChecked(False) if editor_info is not None: widget.stateChanged.connect(lambda: handle_changed(editor_info)) @@ -330,8 +350,15 @@ def _qt_widget_to_yaml_value(widget: Any) -> str | None: elif isinstance(widget, BadgerListEditor): return widget.get_parameters_yaml() elif isinstance(widget, QSpinBox) or isinstance(widget, QDoubleSpinBox): + if widget.property("badger_nullable") and widget.value() == widget.minimum(): + return "null" return str(widget.value()) elif isinstance(widget, QCheckBox): + if ( + widget.property("badger_nullable") + and widget.checkState() == Qt.CheckState.PartiallyChecked + ): + return "null" return "true" if widget.isChecked() else "false" elif isinstance(widget, QComboBox): if widget.currentText() == "null": @@ -386,8 +413,15 @@ def _qt_widget_to_value(widget: Any) -> Any: elif isinstance(widget, BadgerListEditor): return widget.get_parameters_dict() elif isinstance(widget, QSpinBox) or isinstance(widget, QDoubleSpinBox): + if widget.property("badger_nullable") and widget.value() == widget.minimum(): + return None return widget.value() elif isinstance(widget, QCheckBox): + if ( + widget.property("badger_nullable") + and widget.checkState() == Qt.CheckState.PartiallyChecked + ): + return None return widget.isChecked() elif isinstance(widget, QComboBox): return widget.currentText() From 54f5a4b986aca72b31947c7f17ee383998dc294b Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Wed, 24 Jun 2026 21:33:16 -0700 Subject: [PATCH 05/18] Added both sets of plots to widget --- src/badger/factory.py | 2 +- .../gui/components/bax_visualizer/.gitignore | 1 + .../components/bax_visualizer/bax_widget.py | 91 +++++++++- .../gui/components/bax_visualizer/plotting.py | 161 ++++++++++++++++++ .../gui/components/bax_visualizer/ui.py | 39 +++++ .../gui/components/bo_visualizer/bo_widget.py | 2 +- .../gui/components/extension_utilities.py | 15 +- .../gui/components/pf_viewer/pf_widget.py | 73 ++++---- 8 files changed, 329 insertions(+), 55 deletions(-) create mode 100644 src/badger/gui/components/bax_visualizer/.gitignore create mode 100644 src/badger/gui/components/bax_visualizer/plotting.py create mode 100644 src/badger/gui/components/bax_visualizer/ui.py diff --git a/src/badger/factory.py b/src/badger/factory.py index a29aa6f9..956e4c17 100644 --- a/src/badger/factory.py +++ b/src/badger/factory.py @@ -44,7 +44,7 @@ "time_dependent_upper_confidence_bound", "multi_fidelity", "nsga2", - "bax", + # "bax", ] diff --git a/src/badger/gui/components/bax_visualizer/.gitignore b/src/badger/gui/components/bax_visualizer/.gitignore new file mode 100644 index 00000000..d2ff2d25 --- /dev/null +++ b/src/badger/gui/components/bax_visualizer/.gitignore @@ -0,0 +1 @@ +bax_algorithms/ diff --git a/src/badger/gui/components/bax_visualizer/bax_widget.py b/src/badger/gui/components/bax_visualizer/bax_widget.py index aca10fca..16070689 100644 --- a/src/badger/gui/components/bax_visualizer/bax_widget.py +++ b/src/badger/gui/components/bax_visualizer/bax_widget.py @@ -1,25 +1,101 @@ +"""Widget that hosts the BAX visualizer extension within the Badger GUI.""" + +import logging +import time +from dataclasses import dataclass from typing import Optional -from PyQt5.QtWidgets import QDialog +from PyQt5.QtWidgets import QWidget +from xopt.generators.bayesian.bax_generator import BaxGenerator from xopt.generators.bayesian.bayesian_generator import BayesianGenerator from badger.gui.components.analysis_widget import AnalysisWidget -from badger.gui.components.extension_utilities import HandledException +from badger.gui.components.bax_visualizer.ui import UI +from badger.gui.components.extension_utilities import HandledException, requires_update from badger.routine import Routine +from badger.utils import create_archive_run_filename + +logger = logging.getLogger(__name__) + + +@dataclass(frozen=True) +class Parameters: + n_grid: int = 50 + n_samples: int = 100 + fig_size: tuple[int, int] = (5, 5) + + +DEFAULT_PARAMETERS = Parameters(n_grid=50, n_samples=100, fig_size=(5, 5)) -class BaxWidget(AnalysisWidget): # type: ignore[misc] - def __init__(self, routine: Routine, parent: Optional[QDialog] = None): +class BaxWidget(AnalysisWidget): + parameters = DEFAULT_PARAMETERS + generator: BaxGenerator + + def __init__(self, routine: Routine, parent: Optional[QWidget] = None): + logger.debug("Initializing BaxWidget") super().__init__(routine=routine, parent=parent) + self.ui = UI(routine=self.routine, parameters=self.parameters, parent=self) + + self.setWindowTitle("BAX Visualizer") + self.setMinimumSize(800, 600) + def initialize_widget(self) -> None: pass def requires_reinitialization(self) -> bool: + # Check if the extension needs to be reinitialized + logger.debug("Checking if Bax Visualizer needs to be reinitialized") + + archive_name = create_archive_run_filename(self.routine) + + logger.debug(f"Archive name: {archive_name}") + + if not self.initialized: + logger.debug("Reset - Extension never initialized") + # Set up connections + logger.debug("Setting up connections") + self.setup_connections() + self.routine_identifier = archive_name + self.initialized = True + return True + + if self.routine_identifier != archive_name: + logger.debug("Reset - Routine name has changed") + self.routine_identifier = archive_name + self.reset_widget() + return True + + if self.routine.data is None: + logger.debug("Reset - No data available") + + return True + + previous_len = self.df_length + self.df_length = len(self.routine.data) + new_length = self.df_length + + if previous_len > new_length: + logger.debug("Reset - Data length is smaller") + self.df_length = float("inf") + return True + return False + def reset_widget(self) -> None: + logger.debug("Resetting BaxWidget") + self.initialized = False + self.routine_identifier = "" + self.df_length = float("inf") + def update_plots(self, requires_rebuild: bool, interval: int) -> None: - pass + if not requires_update(self.last_updated, interval, requires_rebuild): + return + + self.ui.plotting_area.update_tab_widget() + + self.last_updated = time.time() def setup_connections(self) -> None: pass @@ -29,3 +105,8 @@ def isValidRoutine(self, routine: Routine) -> None: raise HandledException( ValueError, "Bax Visualizer can only be used with a BayesianGenerator." ) + if len(routine.vocs.objective_names) > 0: + raise HandledException( + ValueError, + "BAX Visualizer uses observations to visualize the optimization process, and therefore cannot be used with routines that have objectives defined. Please remove the objectives from your routine and try again.", + ) diff --git a/src/badger/gui/components/bax_visualizer/plotting.py b/src/badger/gui/components/bax_visualizer/plotting.py new file mode 100644 index 00000000..50cc66df --- /dev/null +++ b/src/badger/gui/components/bax_visualizer/plotting.py @@ -0,0 +1,161 @@ +"""Matplotlib-based plotting widget for visualizing BAX virtual measurements.""" + +import logging +import os +import sys +from typing import TYPE_CHECKING, Optional + +from matplotlib.axes import Axes +from matplotlib.backends.backend_qt import NavigationToolbar2QT +from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg +from matplotlib.figure import Figure +from PyQt5.QtWidgets import QTabWidget, QVBoxLayout, QWidget +from xopt.generators.bayesian.bax_generator import BaxGenerator + +from badger.gui.components.extension_utilities import ( + MatplotlibFigureContext, + clear_tabs, +) +from badger.utils import BlockSignalsContext + +# Temporary: the vendored ``bax_algorithms`` package uses absolute imports +# rooted at the top-level name ``bax_algorithms`` (e.g. +# ``from bax_algorithms.utils import ...``). Add the directory that contains the +# package to ``sys.path`` so it is importable as a top-level package until it is +# properly published as part of xopt. +_BAX_ALGORITHMS_DIR = os.path.join(os.path.dirname(__file__), "bax_algorithms") +if _BAX_ALGORITHMS_DIR not in sys.path: + sys.path.insert(0, _BAX_ALGORITHMS_DIR) + +from bax_algorithms.visualize import ( # noqa: E402 + plot_bax_input_convergence, + plot_bax_objective_convergence, + visualize_virtual_measurement_result, +) + +if TYPE_CHECKING: + from badger.gui.components.bax_visualizer.bax_widget import Parameters + + +logger = logging.getLogger(__name__) + + +class PlottingWidget(QWidget): + def __init__( + self, + generator: BaxGenerator, + parameters: "Parameters", + parent: Optional[QWidget] = None, + ): + logger.debug("Initializing PlottingWidget") + super().__init__(parent=parent) + + self.generator = generator + self.parameters = parameters + self._initialize_plotting_area() + + def _initialize_plotting_area(self) -> None: + main_layout = QVBoxLayout() + self.plot_tab_widget = QTabWidget() + + main_layout.addWidget(self.plot_tab_widget) + + self.setLayout(main_layout) + + self.update_tab_widget() + + def create_first_plot(self) -> tuple[Figure, Axes]: + logger.debug("Creating first plot") + fig, ax = visualize_virtual_measurement_result( + self.generator, + variable_names=["x0", "x1"], + idx=0, + reference_point=None, # type: ignore[arg-type] + n_grid=self.parameters.n_grid, + n_samples=self.parameters.n_samples, + show_observations=True, + result_keys=["objective"], + ) + return fig, ax + + def create_second_plot(self) -> tuple[Figure, Axes]: + logger.debug("Creating second plot") + fig, ax = plot_bax_objective_convergence( + self.generator, + ) + return fig, ax + + def create_third_plot(self) -> tuple[Figure, Axes]: + logger.debug("Creating third plot") + fig, ax = plot_bax_input_convergence( + self.generator, + ) + return fig, ax + + def update_first_tab(self) -> None: + + with MatplotlibFigureContext(fig_size=self.parameters.fig_size) as (fig, ax): + try: + fig, ax = self.create_first_plot() + canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] + toolbar = NavigationToolbar2QT(canvas, self) # type: ignore[no-untyped-call] + + # handler = MatplotlibInteractionHandler(canvas, ) + + widget = QWidget() + layout = QVBoxLayout() + layout.addWidget(canvas) + layout.addWidget(toolbar) + widget.setLayout(layout) + self.plot_tab_widget.addTab(widget, "FirstPlot") + + except Exception as e: + logger.error(f"Error creating plot: {e}") + blank_canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] + self.plot_tab_widget.addTab(blank_canvas, "Error") + + def update_second_tab(self) -> None: + + widget = QWidget() + layout = QVBoxLayout() + + with MatplotlibFigureContext(fig_size=self.parameters.fig_size) as (fig, ax): + try: + fig, ax = self.create_second_plot() + canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] + toolbar = NavigationToolbar2QT(canvas, self) # type: ignore[no-untyped-call] + + # handler = MatplotlibInteractionHandler(canvas, ) + + layout.addWidget(canvas) + layout.addWidget(toolbar) + + except Exception as e: + logger.error(f"Error creating plot: {e}") + blank_canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] + self.plot_tab_widget.addTab(blank_canvas, "Error") + + with MatplotlibFigureContext(fig_size=self.parameters.fig_size) as (fig, ax): + try: + fig, ax = self.create_third_plot() + canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] + toolbar = NavigationToolbar2QT(canvas, self) # type: ignore[no-untyped-call] + + # handler = MatplotlibInteractionHandler(canvas, ) + + layout.addWidget(canvas) + layout.addWidget(toolbar) + + except Exception as e: + logger.error(f"Error creating plot: {e}") + blank_canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] + self.plot_tab_widget.addTab(blank_canvas, "Error") + + widget.setLayout(layout) + self.plot_tab_widget.addTab(widget, "SecondPlot") + + def update_tab_widget(self) -> None: + with BlockSignalsContext(self.plot_tab_widget): + clear_tabs(self.plot_tab_widget) + self.update_first_tab() + self.update_second_tab() diff --git a/src/badger/gui/components/bax_visualizer/ui.py b/src/badger/gui/components/bax_visualizer/ui.py new file mode 100644 index 00000000..645611b4 --- /dev/null +++ b/src/badger/gui/components/bax_visualizer/ui.py @@ -0,0 +1,39 @@ +"""UI layout definitions for the BAX visualizer widget.""" + +from typing import TYPE_CHECKING, Optional + +if TYPE_CHECKING: + from badger.gui.components.bax_visualizer.bax_widget import Parameters + +from PyQt5.QtWidgets import QVBoxLayout, QWidget + +from badger.gui.components.bax_visualizer.plotting import PlottingWidget +from badger.routine import Routine + + +class UI(QWidget): + def __init__( + self, + routine: Routine, + parameters: "Parameters", + parent: Optional[QWidget] = None, + ): + super().__init__(parent=parent) + + self.routine = routine + self.parameters = parameters + self._initialize_ui() + + def _initialize_ui(self) -> None: + main_layout = QVBoxLayout() + controls_layout = QVBoxLayout() + + main_layout.addLayout(controls_layout) + + self.plotting_area = PlottingWidget( + generator=self.routine.generator, parameters=self.parameters + ) + + main_layout.addWidget(self.plotting_area) + + self.setLayout(main_layout) diff --git a/src/badger/gui/components/bo_visualizer/bo_widget.py b/src/badger/gui/components/bo_visualizer/bo_widget.py index 38db0377..0241625a 100644 --- a/src/badger/gui/components/bo_visualizer/bo_widget.py +++ b/src/badger/gui/components/bo_visualizer/bo_widget.py @@ -56,7 +56,7 @@ } -class BOPlotWidget(AnalysisWidget): # type: ignore[misc] +class BOPlotWidget(AnalysisWidget): generator: BayesianGenerator # pyright: ignore[reportIncompatibleVariableOverride] parameters: ConfigurableOptions = DEFAULT_PARAMETERS.copy() df_length: float = float("inf") diff --git a/src/badger/gui/components/extension_utilities.py b/src/badger/gui/components/extension_utilities.py index be9e2015..b687506a 100644 --- a/src/badger/gui/components/extension_utilities.py +++ b/src/badger/gui/components/extension_utilities.py @@ -1,18 +1,17 @@ """Shared helpers for analysis extensions: matplotlib figure management, update throttling, error-handling decorators, and numeric formatting.""" +import logging import time -from functools import wraps import traceback -from typing import Any, Callable, Optional, ParamSpec +from functools import wraps from types import TracebackType -from PyQt5.QtWidgets import QLayout, QTabWidget +from typing import Any, Callable, Optional, ParamSpec +import matplotlib.pyplot as plt from matplotlib.axes import Axes from matplotlib.figure import Figure -import matplotlib.pyplot as plt - -import logging +from PyQt5.QtWidgets import QLayout, QTabWidget logger = logging.getLogger(__name__) @@ -124,7 +123,7 @@ def __init__( else: self.ax = ax - def __enter__(self): + def __enter__(self) -> tuple[Figure, Axes]: return self.fig, self.ax def __exit__( @@ -132,5 +131,5 @@ def __exit__( exc_type: Optional[type[BaseException]], exc_value: Optional[BaseException], exc_traceback: Optional[TracebackType], - ): + ) -> None: plt.close(self.fig) diff --git a/src/badger/gui/components/pf_viewer/pf_widget.py b/src/badger/gui/components/pf_viewer/pf_widget.py index f56cb2f1..20cd0202 100644 --- a/src/badger/gui/components/pf_viewer/pf_widget.py +++ b/src/badger/gui/components/pf_viewer/pf_widget.py @@ -2,62 +2,55 @@ with color-mapped iteration progress and optional Pareto-only filtering, updating live as new solutions come in.""" +import logging import time from typing import Optional +import matplotlib.pyplot as plt +import pandas as pd +from matplotlib.axes import Axes +from matplotlib.backends.backend_qt import NavigationToolbar2QT as NavigationToolbar +from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg as FigureCanvas +from matplotlib.colors import Normalize +from matplotlib.figure import Figure +from matplotlib.ticker import MaxNLocator +from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( - QWidget, - QSizePolicy, - QVBoxLayout, - QHBoxLayout, - QGridLayout, - QPushButton, - QComboBox, QCheckBox, - QLabel, + QComboBox, + QGridLayout, QGroupBox, + QHBoxLayout, + QLabel, + QPushButton, + QSizePolicy, QTabWidget, + QVBoxLayout, + QWidget, ) - -from PyQt5.QtCore import Qt -from badger.gui.components.plot_event_handlers import ( - MatplotlibInteractionHandler, -) -from badger.gui.components.pf_viewer.types import PFUI, ConfigurableOptions -from badger.routine import Routine - -from badger.utils import BlockSignalsContext, create_archive_run_filename -from matplotlib.axes import Axes -from matplotlib.figure import Figure -import matplotlib.pyplot as plt -from matplotlib.colors import Normalize -from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg as FigureCanvas -from matplotlib.ticker import MaxNLocator -from matplotlib.backends.backend_qt import ( - NavigationToolbar2QT as NavigationToolbar, -) -import pandas as pd from torch import Tensor - from xopt.generators.bayesian.mobo import MOBOGenerator -from badger.gui.components.pf_viewer.types import ( - PFUIWidgets, - PFUILayouts, -) - from badger.gui.components.analysis_widget import AnalysisWidget - from badger.gui.components.extension_utilities import ( HandledException, MatplotlibFigureContext, - signal_logger, - requires_update, - clear_tabs, clear_layout, + clear_tabs, + requires_update, + signal_logger, ) - -import logging +from badger.gui.components.pf_viewer.types import ( + PFUI, + ConfigurableOptions, + PFUILayouts, + PFUIWidgets, +) +from badger.gui.components.plot_event_handlers import ( + MatplotlibInteractionHandler, +) +from badger.routine import Routine +from badger.utils import BlockSignalsContext, create_archive_run_filename logger = logging.getLogger(__name__) @@ -168,7 +161,7 @@ def update_plots( # Update the last updated time self.last_updated = time.time() - def setup_connections(self): + def setup_connections(self) -> None: self.ui["components"]["update"].clicked.connect( lambda: signal_logger("Update button clicked")( lambda: self.on_button_click() From 785ad448dd3ec07bea7948aa705d8a5c0cbd38f1 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Thu, 25 Jun 2026 20:53:45 -0700 Subject: [PATCH 06/18] Added variable selection and connections --- .../gui/components/analysis_extensions.py | 14 +- .../components/bax_visualizer/bax_widget.py | 91 +++++++++++-- .../gui/components/bax_visualizer/controls.py | 118 ++++++++++++++++ .../gui/components/bax_visualizer/plotting.py | 126 ++++++++++-------- .../gui/components/bax_visualizer/ui.py | 32 ++++- .../gui/components/extensions_palette.py | 8 +- 6 files changed, 312 insertions(+), 77 deletions(-) create mode 100644 src/badger/gui/components/bax_visualizer/controls.py diff --git a/src/badger/gui/components/analysis_extensions.py b/src/badger/gui/components/analysis_extensions.py index f4e4cd98..3ffac040 100644 --- a/src/badger/gui/components/analysis_extensions.py +++ b/src/badger/gui/components/analysis_extensions.py @@ -4,7 +4,7 @@ import logging from typing import Optional, cast -from PyQt5.QtCore import pyqtSignal +from PyQt5.QtCore import Qt, pyqtSignal from PyQt5.QtGui import QCloseEvent from PyQt5.QtWidgets import QVBoxLayout, QWidget from xopt import Generator @@ -27,6 +27,10 @@ class AnalysisExtension(QWidget): def __init__(self, parent: Optional[QWidget] = None): super().__init__(parent=parent) + # A parented QWidget is a child widget by default. Setting the Window + # flag keeps the palette as the owner (for lifetime/stacking) while + # still displaying this as a separate top-level window. + self.setWindowFlag(Qt.WindowType.Window) def update_window(self, routine: Routine) -> None: try: @@ -101,7 +105,9 @@ def __init__( super().__init__(parent=parent) self.initialize_extension( - extension_widget=BOPlotWidget(routine=routine), + extension_widget=BOPlotWidget( + routine=routine, + ), extension_name="Bayesian Optimization Visualizer", generator_type=cast(type[Generator], BayesianGenerator), ) @@ -116,7 +122,9 @@ def __init__( super().__init__(parent=parent) self.initialize_extension( - extension_widget=BaxWidget(routine=routine), + extension_widget=BaxWidget( + routine=routine, + ), extension_name="Bax Visualizer", generator_type=cast(type[Generator], BayesianGenerator), ) diff --git a/src/badger/gui/components/bax_visualizer/bax_widget.py b/src/badger/gui/components/bax_visualizer/bax_widget.py index 16070689..77b05998 100644 --- a/src/badger/gui/components/bax_visualizer/bax_widget.py +++ b/src/badger/gui/components/bax_visualizer/bax_widget.py @@ -2,10 +2,10 @@ import logging import time -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Optional -from PyQt5.QtWidgets import QWidget +from PyQt5.QtWidgets import QSizePolicy, QVBoxLayout, QWidget from xopt.generators.bayesian.bax_generator import BaxGenerator from xopt.generators.bayesian.bayesian_generator import BayesianGenerator @@ -18,31 +18,63 @@ logger = logging.getLogger(__name__) -@dataclass(frozen=True) -class Parameters: +@dataclass() +class PlotParameters: n_grid: int = 50 n_samples: int = 100 - fig_size: tuple[int, int] = (5, 5) -DEFAULT_PARAMETERS = Parameters(n_grid=50, n_samples=100, fig_size=(5, 5)) +@dataclass() +class Parameters: + tab_1: PlotParameters = field(default_factory=PlotParameters) + tab_2: PlotParameters = field(default_factory=PlotParameters) + variables: list[str] = field(default_factory=list) + variable_idx_x: int = 0 + variable_idx_y: int = 1 + include_y: bool = True + + +DEFAULT_PARAMETERS = Parameters() class BaxWidget(AnalysisWidget): - parameters = DEFAULT_PARAMETERS generator: BaxGenerator + parameters: Parameters = DEFAULT_PARAMETERS def __init__(self, routine: Routine, parent: Optional[QWidget] = None): logger.debug("Initializing BaxWidget") super().__init__(routine=routine, parent=parent) - self.ui = UI(routine=self.routine, parameters=self.parameters, parent=self) + self.ui = UI(routine=self.routine, parameters=self.parameters) + + # The UI must live inside a layout, otherwise Qt never manages its + # geometry and the widget's size hints / minimum size are ignored. + layout = QVBoxLayout() + layout.setContentsMargins(0, 0, 0, 0) + layout.addWidget(self.ui) + self.setLayout(layout) - self.setWindowTitle("BAX Visualizer") + self.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding) self.setMinimumSize(800, 600) + self.initialize_widget() + def initialize_widget(self) -> None: - pass + logger.debug("Initializing BaxWidget") + + variable_names = list(self.routine.vocs.variable_names) + self.parameters.variables = variable_names + + temp_x = self.parameters.variable_idx_x + temp_y = self.parameters.variable_idx_y + if len(variable_names) < 2: + self.parameters.include_y = False + self.parameters.variable_idx_x = 0 + self.parameters.variable_idx_y = -1 + else: + self.parameters.include_y = True + self.parameters.variable_idx_x = min(temp_x, len(variable_names) - 1) + self.parameters.variable_idx_y = min(temp_y, len(variable_names) - 1) def requires_reinitialization(self) -> bool: # Check if the extension needs to be reinitialized @@ -85,7 +117,6 @@ def requires_reinitialization(self) -> bool: def reset_widget(self) -> None: logger.debug("Resetting BaxWidget") - self.initialized = False self.routine_identifier = "" self.df_length = float("inf") @@ -93,12 +124,48 @@ def update_plots(self, requires_rebuild: bool, interval: int) -> None: if not requires_update(self.last_updated, interval, requires_rebuild): return + self.ui.controls_area.update_controls() + self.ui.plotting_area.update_tab_widget() self.last_updated = time.time() def setup_connections(self) -> None: - pass + self.ui.controls_area.update_button.clicked.connect( + lambda: self.update_plots(requires_rebuild=True, interval=0) + ) + + self.ui.controls_area.x_axis_combo_box.currentIndexChanged.connect( + lambda: self.update_variables() + ) + self.ui.controls_area.y_axis_combo_box.currentIndexChanged.connect( + lambda: self.update_variables() + ) + self.ui.controls_area.y_axis_checkbox.stateChanged.connect( + lambda: self.update_y_axis_controls() + ) + + def update_variables(self) -> None: + + self.parameters.variable_idx_x = ( + self.ui.controls_area.x_axis_combo_box.currentIndex() + ) + self.parameters.variable_idx_y = ( + self.ui.controls_area.y_axis_combo_box.currentIndex() + ) + + self.update_plots(requires_rebuild=True, interval=0) + + def update_y_axis_controls(self) -> None: + + self.parameters.include_y = self.ui.controls_area.y_axis_checkbox.isChecked() + + if not self.parameters.include_y: + self.ui.controls_area.y_axis_combo_box.setEnabled(False) + else: + self.ui.controls_area.y_axis_combo_box.setEnabled(True) + + self.update_plots(requires_rebuild=True, interval=0) def isValidRoutine(self, routine: Routine) -> None: if not isinstance(routine.generator, BayesianGenerator): diff --git a/src/badger/gui/components/bax_visualizer/controls.py b/src/badger/gui/components/bax_visualizer/controls.py new file mode 100644 index 00000000..6354b53b --- /dev/null +++ b/src/badger/gui/components/bax_visualizer/controls.py @@ -0,0 +1,118 @@ +"""Controls widget for the BAX visualizer. + +This module provides the ControlsWidget class which manages the UI controls +for variable selection and visualization updates in the BAX visualizer. +""" + +from typing import TYPE_CHECKING, Optional + +from PyQt5.QtWidgets import ( + QCheckBox, + QComboBox, + QHBoxLayout, + QLabel, + QPushButton, + QVBoxLayout, + QWidget, +) + +from badger.routine import Routine +from badger.utils import BlockSignalsContext + +if TYPE_CHECKING: + from badger.gui.components.bax_visualizer.bax_widget import Parameters + + +class ControlsWidget(QWidget): + def __init__( + self, + routine: Routine, + parameters: "Parameters", + parent: Optional[QWidget] = None, + ) -> None: + super().__init__(parent=parent) + self.routine = routine + self.parameters = parameters + + self._initialize_ui() + + def _initialize_ui(self) -> None: + # Create the layout for the controls + controls_layout = QVBoxLayout() + + controls_layout.addLayout(self._create_variable_layout()) + + # Add the controls to the layout + + self.update_button = self._create_update_button() + + # Add the controls to the layout + + controls_layout.addWidget(self.update_button) + + self.setLayout(controls_layout) + + def update_controls(self) -> None: + self.update_variables() + with BlockSignalsContext((self.x_axis_combo_box, self.y_axis_combo_box)): + # Update the combo boxes and checkbox based on the current parameters + self.x_axis_combo_box.setCurrentIndex(self.parameters.variable_idx_x) + self.y_axis_combo_box.setCurrentIndex(self.parameters.variable_idx_y) + self.y_axis_checkbox.setChecked(self.parameters.include_y) + + def update_variables(self) -> None: + # Update the parameters with the current variable names + self.parameters.variables = self.routine.vocs.variable_names + + with BlockSignalsContext((self.x_axis_combo_box, self.y_axis_combo_box)): + # Update the combo boxes with the new variable names + self.x_axis_combo_box.clear() + self.x_axis_combo_box.addItems(self.parameters.variables) + + self.y_axis_combo_box.clear() + self.y_axis_combo_box.addItems(self.parameters.variables) + + def _create_variable_layout(self) -> QVBoxLayout: + layout = QVBoxLayout() + x_axis_combo_box, self.x_axis_combo_box = self._create_variable_combo_box( + is_x_axis=True + ) + y_axis_combo_box, self.y_axis_combo_box = self._create_variable_combo_box( + is_x_axis=False, disabled=not self.parameters.include_y + ) + self.y_axis_checkbox = self._create_include_y_checkbox() + + layout.addLayout(x_axis_combo_box) + layout.addLayout(y_axis_combo_box) + layout.addWidget(self.y_axis_checkbox) + return layout + + def _create_variable_combo_box( + self, is_x_axis: bool = True, disabled: bool = False + ) -> tuple[QHBoxLayout, QComboBox]: + layout = QHBoxLayout() + # Create a combo box for selecting variables + combo_box = QComboBox() + label = QLabel("Variable 1:" if is_x_axis else "Variable 2:") + combo_box.addItems(self.parameters.variables) + if is_x_axis: + combo_box.setCurrentIndex(self.parameters.variable_idx_x) + else: + combo_box.setCurrentIndex(self.parameters.variable_idx_y) + combo_box.setDisabled(disabled) + + layout.addWidget(label) + layout.addWidget(combo_box) + + return layout, combo_box + + def _create_include_y_checkbox(self) -> QCheckBox: + # Create a checkbox for including/excluding the second variable + checkbox = QCheckBox("Include Variable 2") + checkbox.setChecked(self.parameters.include_y) + return checkbox + + def _create_update_button(self) -> QPushButton: + # Create a button for updating the plots + button = QPushButton("Update") # Replace with actual button implementation + return button diff --git a/src/badger/gui/components/bax_visualizer/plotting.py b/src/badger/gui/components/bax_visualizer/plotting.py index 50cc66df..aa443986 100644 --- a/src/badger/gui/components/bax_visualizer/plotting.py +++ b/src/badger/gui/components/bax_visualizer/plotting.py @@ -5,15 +5,21 @@ import sys from typing import TYPE_CHECKING, Optional +import matplotlib.pyplot as plt from matplotlib.axes import Axes from matplotlib.backends.backend_qt import NavigationToolbar2QT from matplotlib.backends.backend_qtagg import FigureCanvasQTAgg from matplotlib.figure import Figure -from PyQt5.QtWidgets import QTabWidget, QVBoxLayout, QWidget +from PyQt5.QtWidgets import ( + QScrollArea, + QSizePolicy, + QTabWidget, + QVBoxLayout, + QWidget, +) from xopt.generators.bayesian.bax_generator import BaxGenerator from badger.gui.components.extension_utilities import ( - MatplotlibFigureContext, clear_tabs, ) from badger.utils import BlockSignalsContext @@ -66,13 +72,22 @@ def _initialize_plotting_area(self) -> None: def create_first_plot(self) -> tuple[Figure, Axes]: logger.debug("Creating first plot") + + selected_variable_names = [ + self.parameters.variables[self.parameters.variable_idx_x] + ] + if self.parameters.include_y: + selected_variable_names.append( + self.parameters.variables[self.parameters.variable_idx_y] + ) + fig, ax = visualize_virtual_measurement_result( self.generator, - variable_names=["x0", "x1"], + variable_names=selected_variable_names, idx=0, reference_point=None, # type: ignore[arg-type] - n_grid=self.parameters.n_grid, - n_samples=self.parameters.n_samples, + n_grid=self.parameters.tab_1.n_grid, + n_samples=self.parameters.tab_1.n_samples, show_observations=True, result_keys=["objective"], ) @@ -92,67 +107,68 @@ def create_third_plot(self) -> tuple[Figure, Axes]: ) return fig, ax - def update_first_tab(self) -> None: - - with MatplotlibFigureContext(fig_size=self.parameters.fig_size) as (fig, ax): - try: - fig, ax = self.create_first_plot() - canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] - toolbar = NavigationToolbar2QT(canvas, self) # type: ignore[no-untyped-call] + def _build_plot_widget(self, fig: Figure, add_stretch: bool = False) -> QWidget: + """Wrap a figure in a canvas + toolbar, constraining the canvas height + to the figure's natural pixel size. - # handler = MatplotlibInteractionHandler(canvas, ) + Matplotlib canvases default to an Expanding/Expanding size policy, which + is what causes a short single plot to be stretched to match a taller + tab. Fixing the vertical policy and a minimum height keeps each plot at + its intended aspect ratio. + """ + canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] + toolbar = NavigationToolbar2QT(canvas, self) # type: ignore[no-untyped-call] - widget = QWidget() - layout = QVBoxLayout() - layout.addWidget(canvas) - layout.addWidget(toolbar) - widget.setLayout(layout) - self.plot_tab_widget.addTab(widget, "FirstPlot") - - except Exception as e: - logger.error(f"Error creating plot: {e}") - blank_canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] - self.plot_tab_widget.addTab(blank_canvas, "Error") - - def update_second_tab(self) -> None: + width_inches, height_inches = fig.get_size_inches() + dpi = fig.get_dpi() + canvas.setMinimumSize(int(width_inches * dpi), int(height_inches * dpi)) + canvas.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Fixed) widget = QWidget() - layout = QVBoxLayout() + layout = QVBoxLayout(widget) + layout.addWidget(toolbar) + layout.addWidget(canvas) + if add_stretch: + # Absorb any extra vertical space so the canvas keeps its height + # instead of stretching to fill the tab. + layout.addStretch(1) + return widget + + @staticmethod + def _scrollable(content: QWidget) -> QScrollArea: + """Place tab content in a scroll area so plots taller than the visible + area scroll instead of forcing the whole extension window to grow.""" + scroll_area = QScrollArea() + scroll_area.setWidgetResizable(True) + scroll_area.setWidget(content) + return scroll_area - with MatplotlibFigureContext(fig_size=self.parameters.fig_size) as (fig, ax): - try: - fig, ax = self.create_second_plot() - canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] - toolbar = NavigationToolbar2QT(canvas, self) # type: ignore[no-untyped-call] - - # handler = MatplotlibInteractionHandler(canvas, ) - - layout.addWidget(canvas) - layout.addWidget(toolbar) + def update_first_tab(self) -> None: + try: + fig, _ = self.create_first_plot() + content = self._build_plot_widget(fig, add_stretch=True) + self.plot_tab_widget.addTab(self._scrollable(content), "FirstPlot") + plt.close(fig) + except Exception as e: + logger.error(f"Error creating plot: {e}") + self.plot_tab_widget.addTab(QWidget(), "Error") - except Exception as e: - logger.error(f"Error creating plot: {e}") - blank_canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] - self.plot_tab_widget.addTab(blank_canvas, "Error") + def update_second_tab(self) -> None: + content = QWidget() + layout = QVBoxLayout(content) - with MatplotlibFigureContext(fig_size=self.parameters.fig_size) as (fig, ax): + for create_plot in (self.create_second_plot, self.create_third_plot): try: - fig, ax = self.create_third_plot() - canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] - toolbar = NavigationToolbar2QT(canvas, self) # type: ignore[no-untyped-call] - - # handler = MatplotlibInteractionHandler(canvas, ) - - layout.addWidget(canvas) - layout.addWidget(toolbar) - + fig, _ = create_plot() + layout.addWidget(self._build_plot_widget(fig)) + plt.close(fig) except Exception as e: logger.error(f"Error creating plot: {e}") - blank_canvas = FigureCanvasQTAgg(fig) # type: ignore[no-untyped-call] - self.plot_tab_widget.addTab(blank_canvas, "Error") - widget.setLayout(layout) - self.plot_tab_widget.addTab(widget, "SecondPlot") + # Keep the plots top-aligned at their natural height; the scroll area + # handles the case where their combined height exceeds the viewport. + layout.addStretch(1) + self.plot_tab_widget.addTab(self._scrollable(content), "SecondPlot") def update_tab_widget(self) -> None: with BlockSignalsContext(self.plot_tab_widget): diff --git a/src/badger/gui/components/bax_visualizer/ui.py b/src/badger/gui/components/bax_visualizer/ui.py index 645611b4..68bb3852 100644 --- a/src/badger/gui/components/bax_visualizer/ui.py +++ b/src/badger/gui/components/bax_visualizer/ui.py @@ -2,10 +2,12 @@ from typing import TYPE_CHECKING, Optional +from badger.gui.components.bax_visualizer.controls import ControlsWidget + if TYPE_CHECKING: from badger.gui.components.bax_visualizer.bax_widget import Parameters -from PyQt5.QtWidgets import QVBoxLayout, QWidget +from PyQt5.QtWidgets import QHBoxLayout, QSizePolicy, QVBoxLayout, QWidget from badger.gui.components.bax_visualizer.plotting import PlottingWidget from badger.routine import Routine @@ -24,16 +26,38 @@ def __init__( self.parameters = parameters self._initialize_ui() + self.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding) + self.setMinimumSize(1250, 600) + def _initialize_ui(self) -> None: - main_layout = QVBoxLayout() + main_layout = QHBoxLayout() + + main_layout.setContentsMargins(0, 0, 0, 0) + main_layout.setSpacing(0) + controls_layout = QVBoxLayout() + controls_layout.setContentsMargins(0, 0, 0, 0) + controls_layout.setSpacing(0) + + self.controls_area = ControlsWidget( + routine=self.routine, parameters=self.parameters + ) + + self.controls_area.setSizePolicy( + QSizePolicy.Policy.Fixed, QSizePolicy.Policy.Expanding + ) + self.controls_area.setMinimumWidth(250) + + controls_layout.addWidget(self.controls_area, stretch=0) + main_layout.addLayout(controls_layout) self.plotting_area = PlottingWidget( - generator=self.routine.generator, parameters=self.parameters + generator=self.routine.generator, + parameters=self.parameters, ) - main_layout.addWidget(self.plotting_area) + main_layout.addWidget(self.plotting_area, stretch=1) self.setLayout(main_layout) diff --git a/src/badger/gui/components/extensions_palette.py b/src/badger/gui/components/extensions_palette.py index dc388bff..1a73d87d 100644 --- a/src/badger/gui/components/extensions_palette.py +++ b/src/badger/gui/components/extensions_palette.py @@ -131,7 +131,7 @@ def add_pf_viewer(self) -> None: return self.add_child_window_to_monitor( - ParetoFrontViewer(routine=self.run_monitor.routine) + ParetoFrontViewer(routine=self.run_monitor.routine, parent=self) ) def add_bo_visualizer(self) -> None: @@ -147,7 +147,9 @@ def add_bo_visualizer(self) -> None: ) return - self.add_child_window_to_monitor(BOVisualizer(routine=self.run_monitor.routine)) + self.add_child_window_to_monitor( + BOVisualizer(routine=self.run_monitor.routine, parent=self) + ) def add_bax_visualizer(self) -> None: """ @@ -163,7 +165,7 @@ def add_bax_visualizer(self) -> None: return self.add_child_window_to_monitor( - BaxVisualizer(routine=self.run_monitor.routine) + BaxVisualizer(routine=self.run_monitor.routine, parent=self) ) def add_child_window_to_monitor(self, child_window: AnalysisExtension) -> None: From 92bf81f38611269f16b27ef283827d4310bfd988 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Wed, 1 Jul 2026 19:05:35 -0700 Subject: [PATCH 07/18] Added BADGER_TEMP_DIRECTORY to settings for BAX generator visualization functions --- pyproject.toml | 2 + src/badger/actions/__init__.py | 5 +- .../gui/components/bax_visualizer/plotting.py | 12 --- src/badger/gui/components/pydantic_editor.py | 75 ++++++++++++++----- src/badger/gui/utils.py | 40 ++++++++-- src/badger/gui/windows/settings_dialog.py | 48 +++++++----- src/badger/settings.py | 69 ++++++++++++++++- src/badger/tests/conftest.py | 19 ++++- src/badger/tests/test_factory.py | 8 +- 9 files changed, 216 insertions(+), 62 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 779d725d..faa8c665 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -28,6 +28,8 @@ dependencies = [ "pillow", "requests", "xopt>=3.0.0", + "bax-algorithms", + ] dynamic = ["version"] [tool.setuptools_scm] diff --git a/src/badger/actions/__init__.py b/src/badger/actions/__init__.py index 85dea25b..973da226 100644 --- a/src/badger/actions/__init__.py +++ b/src/badger/actions/__init__.py @@ -4,9 +4,10 @@ import os from importlib import metadata + from badger.actions.doctor import check_n_config_paths +from badger.settings import get_user_config_folder, init_settings from badger.utils import yprint -from badger.settings import init_settings, get_user_config_folder def show_info(args): @@ -35,6 +36,7 @@ def show_info(args): BADGER_LOGBOOK_ROOT = config_singleton.read_value("BADGER_LOGBOOK_ROOT") BADGER_ARCHIVE_ROOT = config_singleton.read_value("BADGER_ARCHIVE_ROOT") BADGER_LOG_DIRECTORY = config_singleton.read_value("BADGER_LOG_DIRECTORY") + BADGER_TEMP_DIRECTORY = config_singleton.read_value("BADGER_TEMP_DIRECTORY") BADGER_LOG_LEVEL = config_singleton.read_value("BADGER_LOG_LEVEL") BADGER_TENSOR_STRATEGY = config_singleton.read_value( "BADGER_PYTORCH_TENSOR_SHARING_STRATEGY" @@ -51,6 +53,7 @@ def show_info(args): "logging directory": BADGER_LOG_DIRECTORY, "logging level": BADGER_LOG_LEVEL, "pytorch tensor sharing strategy": BADGER_TENSOR_STRATEGY, + "temporary directory": BADGER_TEMP_DIRECTORY, # 'plugin installation url': read_value('BADGER_PLUGINS_URL') } diff --git a/src/badger/gui/components/bax_visualizer/plotting.py b/src/badger/gui/components/bax_visualizer/plotting.py index aa443986..dbb34ad3 100644 --- a/src/badger/gui/components/bax_visualizer/plotting.py +++ b/src/badger/gui/components/bax_visualizer/plotting.py @@ -1,8 +1,6 @@ """Matplotlib-based plotting widget for visualizing BAX virtual measurements.""" import logging -import os -import sys from typing import TYPE_CHECKING, Optional import matplotlib.pyplot as plt @@ -23,16 +21,6 @@ clear_tabs, ) from badger.utils import BlockSignalsContext - -# Temporary: the vendored ``bax_algorithms`` package uses absolute imports -# rooted at the top-level name ``bax_algorithms`` (e.g. -# ``from bax_algorithms.utils import ...``). Add the directory that contains the -# package to ``sys.path`` so it is importable as a top-level package until it is -# properly published as part of xopt. -_BAX_ALGORITHMS_DIR = os.path.join(os.path.dirname(__file__), "bax_algorithms") -if _BAX_ALGORITHMS_DIR not in sys.path: - sys.path.insert(0, _BAX_ALGORITHMS_DIR) - from bax_algorithms.visualize import ( # noqa: E402 plot_bax_input_convergence, plot_bax_objective_convergence, diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index c15dbf4a..eb10cb57 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -22,7 +22,7 @@ import yaml from pydantic import BaseModel, Field, ValidationError, create_model from pydantic.fields import FieldInfo -from pydantic_core import PydanticUndefined +from pydantic_core import PydanticUndefined, PydanticUndefinedType from PyQt5.QtCore import Qt, pyqtSignal from PyQt5.QtWidgets import ( QCheckBox, @@ -49,6 +49,7 @@ from xopt.numerical_optimizer import NumericalOptimizer from xopt.vocs import VOCS + logger = logging.getLogger(__name__) @@ -212,6 +213,7 @@ def resolve_qt( ) -> QWidget | None: resolved_type = BadgerResolvedType.resolve(annotation) widget: QWidget | None = None + property_name = editor_info[1].text(0) if editor_info is not None else "" if resolved_type.main is None: widget = QLineEdit() @@ -239,15 +241,21 @@ def resolve_qt( elif resolved_type.main is dict: subtypes = resolved_type.subtype if subtypes is None: - raise ValueError("Dict type must have subtypes") + raise ValueError( + f"Property name {property_name}: Dict type must have subtypes" + ) if not isinstance(subtypes, list) or len(subtypes) != 2: - raise ValueError("Dict type must have two subtypes") + raise ValueError( + f"Property name {property_name}: Dict type must have two subtypes" + ) primary_type = subtypes[0] secondary_type = subtypes[1] if primary_type.main is None or secondary_type.main is None: - raise ValueError("Dict subtypes must be basic types") + raise ValueError( + f"Property name {property_name}: Dict subtypes must be basic types" + ) widget = BadgerListEditor(primary_type.main, secondary_type.main) @@ -261,7 +269,9 @@ def resolve_qt( widget.listChanged.connect(lambda: handle_changed(editor_info)) elif resolved_type.main is list: if resolved_type.subtype is None: - raise ValueError("List type must have a subtype") + raise ValueError( + f"Property name {property_name}: List type must have a subtype" + ) if isinstance(resolved_type.subtype, list): primary_type = resolved_type.subtype[0] secondary_type = ( @@ -272,7 +282,9 @@ def resolve_qt( primary_type = resolved_type.subtype secondary_type = None if primary_type.main is None: - raise ValueError("List subtype must be a basic type") + raise ValueError( + f"Property name {property_name}: List subtype must be a basic type" + ) widget = BadgerListEditor( primary_type.main, secondary_type.main if secondary_type else None ) @@ -288,43 +300,46 @@ def resolve_qt( widget = QDoubleSpinBox() widget.setRange(float("-inf"), float("inf")) widget.setDecimals(6) - if resolved_type.nullable: - # The minimum value doubles as the "null" sentinel. - widget.setSpecialValueText("null") - if default is not None: + if default is not None and not isinstance(default, PydanticUndefinedType): widget.setValue(convert_to_type(default, float)) elif resolved_type.nullable: + # The minimum value doubles as the "null" sentinel. + widget.setSpecialValueText("null") widget.setValue(widget.minimum()) else: - widget.setValue(0.0) + raise ValueError( + f"Property name {property_name}: Float type must have a default value" + ) if editor_info is not None: widget.valueChanged.connect(lambda: handle_changed(editor_info)) elif resolved_type.main is int: widget = QSpinBox() widget.setRange(-(2**31), 2**31 - 1) # int32 min/max - if resolved_type.nullable: - # The minimum value doubles as the "null" sentinel. - widget.setSpecialValueText("null") - if default is not None: + if default is not None and not isinstance(default, PydanticUndefinedType): widget.setValue(convert_to_type(default, int)) elif resolved_type.nullable: + # The minimum value doubles as the "null" sentinel. + widget.setSpecialValueText("null") widget.setValue(widget.minimum()) else: - widget.setValue(0) + raise ValueError( + f"Property name {property_name}: Int type must have a default value" + ) if editor_info is not None: widget.valueChanged.connect(lambda: handle_changed(editor_info)) elif resolved_type.main is bool: widget = QCheckBox() - if resolved_type.nullable: - widget.setTristate(True) - if default is not None: + if default is not None and not isinstance(default, PydanticUndefinedType): widget.setChecked(convert_to_type(default, bool)) elif resolved_type.nullable: + widget.setTristate(True) widget.setCheckState(Qt.CheckState.PartiallyChecked) else: - widget.setChecked(False) + raise ValueError( + f"Property name {property_name}: Bool type must have a default value" + ) if editor_info is not None: widget.stateChanged.connect(lambda: handle_changed(editor_info)) @@ -854,6 +869,11 @@ def get_all_compatible_classes( if not issubclass(self.model_class, BaxGenerator): raise ValueError("Generator does not support algorithms.") compatible_classes = self.model_class.get_compatible_algorithms() + # TODO: Add in additional from BAX algorithms. + # compatible_classes = list(compatible_classes) + [ + # EmittanceAlgorithm, + # PathwiseSolenoidAlignment, + # ] else: raise ValueError(f"Field name {field_name} is not recognized.") @@ -943,6 +963,21 @@ def update_params_from_generator_class( True, ) + # ``class_path`` is a Pydantic computed field (absent from ``model_fields``) + # so it never becomes a widget on its own. BaxGenerator.validate_algorithm + # relies on it to import algorithms that are not registered in the + # generator's compatible list (e.g. the vendored BAX algorithms). Render it + # as a hidden item so its value is carried through get_parameters_yaml()/ + # get_parameters_dict() into the final config. + if "class_path" in pydantic_class.model_computed_fields: + class_path_value = f"{pydantic_class.__module__}.{pydantic_class.__name__}" + self._set_params_recurse( + tree_widget_item, + {"class_path": FieldInfo(annotation=str, default=class_path_value)}, + {"class_path": class_path_value}, + True, + ) + self.expandItem(tree_widget_item) @staticmethod diff --git a/src/badger/gui/utils.py b/src/badger/gui/utils.py index ef5f03a9..a6ed5e19 100644 --- a/src/badger/gui/utils.py +++ b/src/badger/gui/utils.py @@ -2,13 +2,34 @@ scroll-wheel filters for spinboxes, custom combo boxes, and dialog utilities.""" +import copy +import logging +import os from importlib import resources from typing import Any -from PyQt5.QtWidgets import QAbstractSpinBox, QPushButton, QComboBox, QToolButton -from PyQt5.QtWidgets import QDialog, QVBoxLayout, QLabel -from PyQt5.QtCore import Qt, QObject, QEvent, QSize + +from PyQt5.QtCore import QEvent, QObject, QSize, Qt from PyQt5.QtGui import QIcon -import copy +from PyQt5.QtWidgets import ( + QAbstractSpinBox, + QComboBox, + QDialog, + QLabel, + QPushButton, + QToolButton, + QVBoxLayout, +) + +from badger.errors import BadgerConfigError +from badger.settings import init_settings + +logger = logging.getLogger(__name__) + +# Check badger optimization run archive root +config_singleton = init_settings() +BADGER_TEMP_DIRECTORY = config_singleton.read_value("BADGER_TEMP_DIRECTORY") +if BADGER_TEMP_DIRECTORY is None: + raise BadgerConfigError("Please set the BADGER_TEMP_DIRECTORY env var!") def preventAnnoyingSpinboxScrollBehaviour(self, control: QAbstractSpinBox) -> None: @@ -59,7 +80,10 @@ def create_button( return btn -def filter_generator_config(name: str, config: dict[str, Any]): +DEFAULT_ALGORITHM_RESULTS_FILE = "algorithm_results" + + +def filter_generator_config(name: str, config: dict[str, Any]) -> dict[str, Any]: filtered_config: dict[str, Any] = {} if name == "neldermead": filtered_config["adaptive"] = config["adaptive"] @@ -79,6 +103,12 @@ def filter_generator_config(name: str, config: dict[str, Any]): filtered_config["numerical_optimizer"] = config["numerical_optimizer"] filtered_config["max_travel_distances"] = config["max_travel_distances"] filtered_config["reference_point"] = config["reference_point"] + elif name == "bax": + filtered_config = config + filtered_config["algorithm_results_file"] = ( + f"{BADGER_TEMP_DIRECTORY}{os.sep}{DEFAULT_ALGORITHM_RESULTS_FILE}" + ) + else: filtered_config = config diff --git a/src/badger/gui/windows/settings_dialog.py b/src/badger/gui/windows/settings_dialog.py index b98a8ef0..43a055bb 100644 --- a/src/badger/gui/windows/settings_dialog.py +++ b/src/badger/gui/windows/settings_dialog.py @@ -5,26 +5,25 @@ import logging import os +from PyQt5.QtCore import Qt + # from PyQt5.QtCore import QRegExp # from PyQt5.QtGui import QRegExpValidator from PyQt5.QtWidgets import ( + QApplication, QComboBox, + QDialog, + QDialogButtonBox, QGridLayout, - QVBoxLayout, - QWidget, QLabel, QLineEdit, + QVBoxLayout, + QWidget, ) -from PyQt5.QtWidgets import ( - QDialog, - QDialogButtonBox, - QApplication, -) -from PyQt5.QtCore import Qt -from qdarkstyle import load_stylesheet, DarkPalette, LightPalette -from badger.settings import init_settings -from badger.log import get_logging_manager +from qdarkstyle import DarkPalette, LightPalette, load_stylesheet +from badger.log import get_logging_manager +from badger.settings import init_settings logger = logging.getLogger(__name__) @@ -106,12 +105,21 @@ def init_ui(self): grid.addWidget(archive_root, 4, 0) grid.addWidget(archive_root_path, 4, 1) - self.log_dir_label = QLabel("Log Directory") - self.log_dir_path = QLineEdit( + # Log directory setting + self.log_dir_label = log_dir_label = QLabel("Log Directory") + self.log_dir_path = log_dir_path = QLineEdit( self.config_singleton.read_value("BADGER_LOG_DIRECTORY") ) - grid.addWidget(self.log_dir_label, 5, 0) - grid.addWidget(self.log_dir_path, 5, 1) + grid.addWidget(log_dir_label, 5, 0) + grid.addWidget(log_dir_path, 5, 1) + + # Temporary directory setting + self.temp_dir_label = temp_dir_label = QLabel("Temporary Directory") + self.temp_dir_path = temp_dir_path = QLineEdit( + self.config_singleton.read_value("BADGER_TEMP_DIRECTORY") + ) + grid.addWidget(temp_dir_label, 6, 0) + grid.addWidget(temp_dir_path, 6, 1) # Log level setting self.logging_level = logging_level = QLabel("Logging Level") @@ -128,8 +136,8 @@ def init_ui(self): ]: self.logging_level_setting.setCurrentText(current_level) - grid.addWidget(logging_level, 6, 0) - grid.addWidget(logging_level_setting, 6, 1) + grid.addWidget(logging_level, 7, 0) + grid.addWidget(logging_level_setting, 7, 1) # Auto refresh # self.auto_refresh = auto_refresh = QLabel("Auto Refresh") @@ -225,6 +233,12 @@ def apply_settings(self): self.config_singleton.write_value( "BADGER_ARCHIVE_ROOT", self.archive_root_path.text() ) + self.config_singleton.write_value( + "BADGER_LOG_DIRECTORY", self.log_dir_path.text() + ) + self.config_singleton.write_value( + "BADGER_TEMP_DIRECTORY", self.temp_dir_path.text() + ) # self.config_singleton.write_value( # "AUTO_REFRESH", self.enable_auto_refresh.isChecked() # ) diff --git a/src/badger/settings.py b/src/badger/settings.py index f8642b20..29639390 100644 --- a/src/badger/settings.py +++ b/src/badger/settings.py @@ -8,16 +8,18 @@ Run `badger config` from the CLI to edit settings interactively. """ +import logging import os import platform -import yaml import shutil from importlib import resources -from badger.utils import get_datadir -from pydantic import BaseModel, Field, ValidationError from typing import Any, Dict, Optional, Union + +import yaml +from pydantic import BaseModel, Field, ValidationError + from badger.errors import BadgerLoadConfigError -import logging +from badger.utils import get_datadir logger = logging.getLogger(__name__) @@ -63,6 +65,8 @@ class BadgerConfig(BaseModel): Setting for the logging level. BADGER_LOG_DIRECTORY : Setting Setting for the location of logfile. + BADGER_TEMP_DIRECTORY : Setting + Setting for the location of temporary files. BADGER_DATA_DUMP_PERIOD : Setting Setting for the minimum time interval between data dumps (in seconds). BADGER_THEME : Setting @@ -109,6 +113,12 @@ class BadgerConfig(BaseModel): value="logs", is_path=True, ) + BADGER_TEMP_DIRECTORY: Setting = Setting( + display_name="temp directory", + description="Directory where temporary files will be stored", + value="temp", + is_path=True, + ) BADGER_DATA_DUMP_PERIOD: Setting = Setting( display_name="data dump period", description="Minimum time interval between data dumps, unit is second", @@ -442,9 +452,56 @@ def init_settings(config_arg: str = None) -> ConfigSingleton: user_flag = True config_singleton = ConfigSingleton(file_path, user_flag) + get_or_create_temp_directory(config_singleton) return config_singleton +def get_or_create_temp_directory(config_singleton: ConfigSingleton) -> str: + """Resolve BADGER_TEMP_DIRECTORY to an absolute path under the user config + folder and ensure the directory exists on disk. + + This migrates older configs that either lack the key or hold the relative + default ("temp"): they get rewritten to an absolute path anchored under + ``get_user_config_folder()`` so the temp location is OS-appropriate and + stable regardless of the current working directory. + + Parameters + ---------- + config_singleton: ConfigSingleton + The active configuration singleton. + + Returns + ------- + str + The absolute path to the temp directory that now exists on disk. + """ + try: + temp_dir = config_singleton.read_value("BADGER_TEMP_DIRECTORY") + except KeyError: + temp_dir = None + + # Migrate: unset or a relative path -> anchor under the user config folder + if not temp_dir or not os.path.isabs(os.path.expanduser(temp_dir)): + temp_dir = os.path.join(get_user_config_folder(), "temp") + config_singleton.write_value("BADGER_TEMP_DIRECTORY", temp_dir) + + temp_dir = os.path.expanduser(str(temp_dir)) + + # Ensure the directory exists, falling back to the config folder on failure + try: + os.makedirs(temp_dir, exist_ok=True) + except (PermissionError, FileExistsError): + logger.warning( + "Cannot use temp directory %s, falling back to the user config folder", + temp_dir, + ) + temp_dir = os.path.join(get_user_config_folder(), "temp") + config_singleton.write_value("BADGER_TEMP_DIRECTORY", temp_dir) + os.makedirs(temp_dir, exist_ok=True) + + return temp_dir + + def apply_pytorch_multiprocess_tensor_sharing_setting( config_singleton: ConfigSingleton, ) -> None: @@ -513,6 +570,10 @@ def mock_settings(): os.makedirs(templates_dir, exist_ok=True) config_singleton.write_value("BADGER_TEMPLATE_ROOT", templates_dir) + temp_dir = str(app_data_dir / "temp") + os.makedirs(temp_dir, exist_ok=True) + config_singleton.write_value("BADGER_TEMP_DIRECTORY", temp_dir) + # Set other settings to the default values for key in config_singleton.config.model_dump(by_alias=True).keys(): config_singleton.write_value( diff --git a/src/badger/tests/conftest.py b/src/badger/tests/conftest.py index 3d7961c7..0f21ea23 100644 --- a/src/badger/tests/conftest.py +++ b/src/badger/tests/conftest.py @@ -1,5 +1,6 @@ import os import shutil + import pytest @@ -19,6 +20,7 @@ def config_test_settings( mock_archive_root, mock_log_directory, mock_logging_level, + mock_temp_directory, ): from badger.settings import init_settings @@ -34,6 +36,7 @@ def config_test_settings( old_archived = config_singleton.read_value("BADGER_ARCHIVE_ROOT") old_log_directory = config_singleton.read_value("BADGER_LOG_DIRECTORY") old_logging_level = config_singleton.read_value("BADGER_LOG_LEVEL") + old_temp_directory = config_singleton.read_value("BADGER_TEMP_DIRECTORY") except KeyError: pass @@ -44,7 +47,7 @@ def config_test_settings( config_singleton.write_value("BADGER_ARCHIVE_ROOT", mock_archive_root) config_singleton.write_value("BADGER_LOG_DIRECTORY", mock_log_directory) config_singleton.write_value("BADGER_LOG_LEVEL", mock_logging_level) - + config_singleton.write_value("BADGER_TEMP_DIRECTORY", mock_temp_directory) yield # Restoring the original settings @@ -55,6 +58,7 @@ def config_test_settings( config_singleton.write_value("BADGER_ARCHIVE_ROOT", old_archived) config_singleton.write_value("BADGER_LOG_DIRECTORY", old_log_directory) config_singleton.write_value("BADGER_LOG_LEVEL", old_logging_level) + config_singleton.write_value("BADGER_TEMP_DIRECTORY", old_temp_directory) # check if any "old_..." vars didn't get created b/c any of the config values didn't exist in user's config. except NameError: pass @@ -62,13 +66,18 @@ def config_test_settings( @pytest.fixture(scope="module", autouse=True) def clean_up( - mock_template_root, mock_logbook_root, mock_archive_root, mock_log_directory + mock_template_root, + mock_logbook_root, + mock_archive_root, + mock_log_directory, + mock_temp_directory, ): # Clean before tests shutil.rmtree(mock_template_root, True) # ignore errors shutil.rmtree(mock_logbook_root, True) shutil.rmtree(mock_archive_root, True) shutil.rmtree(mock_log_directory, True) + shutil.rmtree(mock_temp_directory, True) yield @@ -77,6 +86,7 @@ def clean_up( shutil.rmtree(mock_logbook_root, True) shutil.rmtree(mock_archive_root, True) shutil.rmtree(mock_log_directory, True) + shutil.rmtree(mock_temp_directory, True) @pytest.fixture(scope="module") @@ -109,6 +119,11 @@ def mock_log_directory(mock_root): return os.path.join(mock_root, "logs") +@pytest.fixture(scope="module") +def mock_temp_directory(mock_root): + return os.path.join(mock_root, "temp") + + @pytest.fixture(scope="module") def mock_logging_level(mock_root): return "WARNING" diff --git a/src/badger/tests/test_factory.py b/src/badger/tests/test_factory.py index 74732d57..ab4d3e3a 100644 --- a/src/badger/tests/test_factory.py +++ b/src/badger/tests/test_factory.py @@ -1,9 +1,9 @@ import json from copy import deepcopy +from gest_api.vocs import ExploreObjective from xopt.generators import get_generator, get_generator_defaults from xopt.resources.testing import TEST_VOCS_BASE -from gest_api.vocs import ExploreObjective class TestFactory: @@ -42,6 +42,12 @@ def test_generator_generation(self): k: ExploreObjective() for k in test_vocs.objectives.keys() } gen_class(vocs=test_vocs, **gen_config) + elif name == "bax": + test_vocs = deepcopy(TEST_VOCS_BASE) + test_vocs.objectives = {} + test_vocs.observables = ["f"] + json.dumps(gen_config) + gen_class(vocs=test_vocs, **gen_config) else: json.dumps(gen_config) gen_class(vocs=TEST_VOCS_BASE, **gen_config) From cb564e1c1dc044e0c4b45683958844e541f204dd Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Wed, 1 Jul 2026 20:28:05 -0700 Subject: [PATCH 08/18] Consistency between switching tabs and excluding not supported pydantic editor fields --- .../components/bax_visualizer/bax_widget.py | 7 +++ .../gui/components/bax_visualizer/plotting.py | 1 + src/badger/gui/components/pydantic_editor.py | 55 +++++++++++++++++-- 3 files changed, 58 insertions(+), 5 deletions(-) diff --git a/src/badger/gui/components/bax_visualizer/bax_widget.py b/src/badger/gui/components/bax_visualizer/bax_widget.py index 77b05998..352b15c1 100644 --- a/src/badger/gui/components/bax_visualizer/bax_widget.py +++ b/src/badger/gui/components/bax_visualizer/bax_widget.py @@ -28,6 +28,7 @@ class PlotParameters: class Parameters: tab_1: PlotParameters = field(default_factory=PlotParameters) tab_2: PlotParameters = field(default_factory=PlotParameters) + active_tab: int = 0 variables: list[str] = field(default_factory=list) variable_idx_x: int = 0 variable_idx_y: int = 1 @@ -144,6 +145,12 @@ def setup_connections(self) -> None: self.ui.controls_area.y_axis_checkbox.stateChanged.connect( lambda: self.update_y_axis_controls() ) + self.ui.plotting_area.plot_tab_widget.currentChanged.connect( + lambda index: self.update_tab_index(index) + ) + + def update_tab_index(self, index: int) -> None: + self.parameters.active_tab = index def update_variables(self) -> None: diff --git a/src/badger/gui/components/bax_visualizer/plotting.py b/src/badger/gui/components/bax_visualizer/plotting.py index dbb34ad3..7e6ab6b3 100644 --- a/src/badger/gui/components/bax_visualizer/plotting.py +++ b/src/badger/gui/components/bax_visualizer/plotting.py @@ -163,3 +163,4 @@ def update_tab_widget(self) -> None: clear_tabs(self.plot_tab_widget) self.update_first_tab() self.update_second_tab() + self.plot_tab_widget.setCurrentIndex(self.parameters.active_tab) diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index eb10cb57..c67dedc1 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -49,7 +49,6 @@ from xopt.numerical_optimizer import NumericalOptimizer from xopt.vocs import VOCS - logger = logging.getLogger(__name__) @@ -623,6 +622,41 @@ class BadgerPydanticEditor(QTreeWidget): generator_name: str = "" model_class: type[BaseModel] | None = None + # Fields holding runtime state that has no editable widget representation + # (e.g. pandas DataFrames populated during/after optimization). These are + # dropped from the tree entirely so their stringified values never reach + # validation. + # + # ``COMMON_EXCLUDED_FIELDS`` applies to every generator. Add generator + # specific exclusions to ``GENERATOR_EXCLUDED_FIELDS`` keyed by the + # generator name (i.e. the value passed to ``set_params_from_generator`` / + # the generator's ``name`` field). The effective set is the union of both, + # resolved by ``get_excluded_fields``. + COMMON_EXCLUDED_FIELDS: frozenset[str] = frozenset({"computation_time"}) + GENERATOR_EXCLUDED_FIELDS: dict[str, frozenset[str]] = { + # "bax": frozenset({"algorithm_results"}), + } + + def get_excluded_fields(self) -> frozenset[str]: + """Return the set of fields to exclude from the tree for the current + generator: the common fields plus any generator-specific ones.""" + excluded: set[str] = set(self.COMMON_EXCLUDED_FIELDS) + + # Resolve the generator name from the loaded model class when available, + # falling back to the name provided to ``set_params_from_generator``. + names: set[str] = set() + if self.generator_name: + names.add(self.generator_name) + if self.model_class is not None: + name_field = self.model_class.model_fields.get("name") + if name_field is not None and isinstance(name_field.default, str): + names.add(name_field.default) + + for name in names: + excluded |= self.GENERATOR_EXCLUDED_FIELDS.get(name, frozenset()) + + return frozenset(excluded) + def __init__( self, parent: QTreeWidget | None = None, @@ -761,7 +795,11 @@ def set_params_from_generator( fields_to_remove = ["vocs"] filtered_class_fields, removed_class_fields = self.filter_class_fields( - self.model_class, fields_to_remove, defaults, include_defaults=True + self.model_class, + fields_to_remove, + defaults, + include_defaults=True, + excluded_fields=self.get_excluded_fields(), ) self._set_params_recurse( @@ -986,6 +1024,7 @@ def filter_class_fields( fields_to_remove: list[str] = [], defaults: dict[str, Any] = {}, include_defaults: bool = False, + excluded_fields: frozenset[str] = frozenset(), ) -> tuple[dict[str, FieldInfo], dict[str, FieldInfo]]: condition: Callable[[str], bool] @@ -1001,13 +1040,15 @@ def exclude_condition(k: str) -> bool: condition = exclude_condition filtered_class_fields = { - k: v for k, v in pydantic_class.model_fields.items() if condition(k) + k: v + for k, v in pydantic_class.model_fields.items() + if condition(k) and k not in excluded_fields } removed_class_fields = { k: v for k, v in pydantic_class.model_fields.items() - if k in fields_to_remove + if k in fields_to_remove and k not in excluded_fields } return filtered_class_fields, removed_class_fields @@ -1099,7 +1140,11 @@ def update_after_validate(self, defaults: dict[str, Any]) -> None: fields_to_remove = ["vocs"] filtered_class_fields, removed_class_fields = self.filter_class_fields( - model_class, fields_to_remove, defaults, include_defaults=True + model_class, + fields_to_remove, + defaults, + include_defaults=True, + excluded_fields=self.get_excluded_fields(), ) self._set_params_recurse( From abd359ca6d996495ffb718938adbda0b13da1838 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Wed, 1 Jul 2026 21:00:45 -0700 Subject: [PATCH 09/18] Added creation of temp folder per bax run --- src/badger/gui/components/pydantic_editor.py | 15 ++++++++++ src/badger/gui/pages/home_page.py | 9 +++++- src/badger/gui/utils.py | 29 ++++++++++++++++++++ 3 files changed, 52 insertions(+), 1 deletion(-) diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index c67dedc1..34c1755c 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -49,6 +49,8 @@ from xopt.numerical_optimizer import NumericalOptimizer from xopt.vocs import VOCS +from badger.gui.utils import build_bax_results_file, needs_new_bax_results_folder + logger = logging.getLogger(__name__) @@ -794,6 +796,14 @@ def set_params_from_generator( fields_to_remove = ["vocs"] + if issubclass(self.model_class, BaxGenerator): + # The results file is derived, not user-editable: give the run its + # own temp folder and hide the field from the tree. The physical + # directory is created at run start (see prepare_run), not here. + if needs_new_bax_results_folder(defaults.get("algorithm_results_file")): + defaults["algorithm_results_file"] = build_bax_results_file() + fields_to_remove.append("algorithm_results_file") + filtered_class_fields, removed_class_fields = self.filter_class_fields( self.model_class, fields_to_remove, @@ -1139,6 +1149,11 @@ def update_after_validate(self, defaults: dict[str, Any]) -> None: fields_to_remove = ["vocs"] + if isclass(model_class) and issubclass(model_class, BaxGenerator): + # Keep the derived results file hidden after re-rendering; its value + # is preserved from the (hidden) tree item via ``defaults``. + fields_to_remove.append("algorithm_results_file") + filtered_class_fields, removed_class_fields = self.filter_class_fields( model_class, fields_to_remove, diff --git a/src/badger/gui/pages/home_page.py b/src/badger/gui/pages/home_page.py index 574dac0f..447b0ab9 100644 --- a/src/badger/gui/pages/home_page.py +++ b/src/badger/gui/pages/home_page.py @@ -51,7 +51,7 @@ from badger.gui.components.routine_page import BadgerRoutinePage from badger.gui.components.run_monitor import BadgerOptMonitor from badger.gui.components.status_bar import BadgerStatusBar -from badger.gui.utils import ModalOverlay +from badger.gui.utils import ModalOverlay, build_bax_results_file # from PyQt5.QtGui import QBrush, QColor from badger.gui.windows.message_dialog import BadgerScrollableMessageBox @@ -493,6 +493,13 @@ def prepare_run( self.sig_routine_invalid.emit() raise e + # Give this run its own results folder so BAX pkl dumps don't reuse the + # folder set at config time. The folder is created here, at run start. + if routine.generator.name == "bax": + results_file = build_bax_results_file(create_dir=True) + routine.generator.algorithm_results_file = results_file + logger.debug(f"BAX results file set for run: {results_file}") + # Add data to routine before saving tmp file if data is not None: # Make sure selected generator is compatible with prior data diff --git a/src/badger/gui/utils.py b/src/badger/gui/utils.py index a6ed5e19..b18df80f 100644 --- a/src/badger/gui/utils.py +++ b/src/badger/gui/utils.py @@ -5,6 +5,7 @@ import copy import logging import os +from datetime import datetime from importlib import resources from typing import Any @@ -83,6 +84,34 @@ def create_button( DEFAULT_ALGORITHM_RESULTS_FILE = "algorithm_results" +def needs_new_bax_results_folder(value: str | None) -> bool: + """Return True when the BAX results path should get a fresh per-run folder. + + A value is considered "already scoped" when it lives in its own sub-folder + of the temp directory (``temp//...``). The shared default + (``temp/algorithm_results``) or an empty value triggers a new folder. + """ + if not value: + return True + parent = os.path.dirname(value) + return bool(os.path.normpath(parent) == os.path.normpath(BADGER_TEMP_DIRECTORY)) + + +def build_bax_results_file(create_dir: bool = False) -> str: + """Build a ``temp//algorithm_results`` prefix for BAX pkl dumps. + + The archive name isn't known until a run starts, so the folder id is a + timestamp. BaxGenerator does not create the parent directory and appends + ``_.pkl`` to the returned prefix, so ``create_dir`` should be True + whenever the path is going to be used for an actual run. + """ + folder_id = datetime.now().strftime("%Y%m%d-%H%M%S-%f") + results_dir = os.path.join(BADGER_TEMP_DIRECTORY, folder_id) + if create_dir: + os.makedirs(results_dir, exist_ok=True) + return os.path.join(results_dir, DEFAULT_ALGORITHM_RESULTS_FILE) + + def filter_generator_config(name: str, config: dict[str, Any]) -> dict[str, Any]: filtered_config: dict[str, Any] = {} if name == "neldermead": From 6a98675c690ffb224ad016a6ae7897252b4965ac Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Thu, 2 Jul 2026 16:25:56 -0700 Subject: [PATCH 10/18] Fixed issue with bax temp file path not being consistent with runs --- .../components/bax_visualizer/bax_widget.py | 5 ++++ .../gui/components/bax_visualizer/plotting.py | 3 +++ src/badger/gui/components/pydantic_editor.py | 10 +++---- src/badger/gui/pages/home_page.py | 9 ++++--- src/badger/gui/utils.py | 26 +++++-------------- 5 files changed, 23 insertions(+), 30 deletions(-) diff --git a/src/badger/gui/components/bax_visualizer/bax_widget.py b/src/badger/gui/components/bax_visualizer/bax_widget.py index 352b15c1..b88bda02 100644 --- a/src/badger/gui/components/bax_visualizer/bax_widget.py +++ b/src/badger/gui/components/bax_visualizer/bax_widget.py @@ -125,6 +125,11 @@ def update_plots(self, requires_rebuild: bool, interval: int) -> None: if not requires_update(self.last_updated, interval, requires_rebuild): return + # The plotting area caches the generator it was built with. Re-sync it + # with the current routine's generator so it reads this run's + # algorithm_results_file instead of a stale one from a previous run. + self.ui.plotting_area.generator = self.generator + self.ui.controls_area.update_controls() self.ui.plotting_area.update_tab_widget() diff --git a/src/badger/gui/components/bax_visualizer/plotting.py b/src/badger/gui/components/bax_visualizer/plotting.py index 7e6ab6b3..75e02cf6 100644 --- a/src/badger/gui/components/bax_visualizer/plotting.py +++ b/src/badger/gui/components/bax_visualizer/plotting.py @@ -69,6 +69,7 @@ def create_first_plot(self) -> tuple[Figure, Axes]: self.parameters.variables[self.parameters.variable_idx_y] ) + logger.debug(f"Results file: {self.generator.algorithm_results_file}") fig, ax = visualize_virtual_measurement_result( self.generator, variable_names=selected_variable_names, @@ -83,6 +84,7 @@ def create_first_plot(self) -> tuple[Figure, Axes]: def create_second_plot(self) -> tuple[Figure, Axes]: logger.debug("Creating second plot") + logger.debug(f"Results file: {self.generator.algorithm_results_file}") fig, ax = plot_bax_objective_convergence( self.generator, ) @@ -90,6 +92,7 @@ def create_second_plot(self) -> tuple[Figure, Axes]: def create_third_plot(self) -> tuple[Figure, Axes]: logger.debug("Creating third plot") + logger.debug(f"Results file: {self.generator.algorithm_results_file}") fig, ax = plot_bax_input_convergence( self.generator, ) diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index 34c1755c..c8d58535 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -49,8 +49,6 @@ from xopt.numerical_optimizer import NumericalOptimizer from xopt.vocs import VOCS -from badger.gui.utils import build_bax_results_file, needs_new_bax_results_folder - logger = logging.getLogger(__name__) @@ -797,11 +795,9 @@ def set_params_from_generator( fields_to_remove = ["vocs"] if issubclass(self.model_class, BaxGenerator): - # The results file is derived, not user-editable: give the run its - # own temp folder and hide the field from the tree. The physical - # directory is created at run start (see prepare_run), not here. - if needs_new_bax_results_folder(defaults.get("algorithm_results_file")): - defaults["algorithm_results_file"] = build_bax_results_file() + # The results file is derived, not user-editable: it is assigned at + # run start (see prepare_run) to match the run's archive name, so + # hide it from the tree rather than exposing a placeholder value. fields_to_remove.append("algorithm_results_file") filtered_class_fields, removed_class_fields = self.filter_class_fields( diff --git a/src/badger/gui/pages/home_page.py b/src/badger/gui/pages/home_page.py index 447b0ab9..a45da927 100644 --- a/src/badger/gui/pages/home_page.py +++ b/src/badger/gui/pages/home_page.py @@ -493,10 +493,13 @@ def prepare_run( self.sig_routine_invalid.emit() raise e - # Give this run its own results folder so BAX pkl dumps don't reuse the - # folder set at config time. The folder is created here, at run start. + # Give this run its own results folder, named after the run's archive + # name (-) so the folder used during the run matches + # the archived run and stays consistent with the visualizer plots. The + # folder is created here, at run start. if routine.generator.name == "bax": - results_file = build_bax_results_file(create_dir=True) + archive_name = f"{routine.environment.name}-{routine.creation_ts}" + results_file = build_bax_results_file(archive_name, create_dir=True) routine.generator.algorithm_results_file = results_file logger.debug(f"BAX results file set for run: {results_file}") diff --git a/src/badger/gui/utils.py b/src/badger/gui/utils.py index b18df80f..ee2030e1 100644 --- a/src/badger/gui/utils.py +++ b/src/badger/gui/utils.py @@ -5,7 +5,6 @@ import copy import logging import os -from datetime import datetime from importlib import resources from typing import Any @@ -84,28 +83,15 @@ def create_button( DEFAULT_ALGORITHM_RESULTS_FILE = "algorithm_results" -def needs_new_bax_results_folder(value: str | None) -> bool: - """Return True when the BAX results path should get a fresh per-run folder. - - A value is considered "already scoped" when it lives in its own sub-folder - of the temp directory (``temp//...``). The shared default - (``temp/algorithm_results``) or an empty value triggers a new folder. - """ - if not value: - return True - parent = os.path.dirname(value) - return bool(os.path.normpath(parent) == os.path.normpath(BADGER_TEMP_DIRECTORY)) - - -def build_bax_results_file(create_dir: bool = False) -> str: +def build_bax_results_file(folder_id: str, create_dir: bool = False) -> str: """Build a ``temp//algorithm_results`` prefix for BAX pkl dumps. - The archive name isn't known until a run starts, so the folder id is a - timestamp. BaxGenerator does not create the parent directory and appends - ``_.pkl`` to the returned prefix, so ``create_dir`` should be True - whenever the path is going to be used for an actual run. + ``folder_id`` should match the run's archive name (``-``) + so the results directory is traceable to the archived run and stays + consistent with what the visualizer plots. BaxGenerator does not create the + parent directory and appends ``_.pkl`` to the returned prefix, so + ``create_dir`` should be True whenever the path is used for an actual run. """ - folder_id = datetime.now().strftime("%Y%m%d-%H%M%S-%f") results_dir = os.path.join(BADGER_TEMP_DIRECTORY, folder_id) if create_dir: os.makedirs(results_dir, exist_ok=True) From b8b5bc93c32865aeb23780194a1e3299f792653a Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Thu, 2 Jul 2026 17:31:46 -0700 Subject: [PATCH 11/18] Added dynamic plot options dependent on which algorithm is being used. --- .../components/bax_visualizer/bax_widget.py | 151 ++++++++++++++++-- .../gui/components/bax_visualizer/controls.py | 71 +++++++- .../gui/components/bax_visualizer/plotting.py | 36 ++++- 3 files changed, 244 insertions(+), 14 deletions(-) diff --git a/src/badger/gui/components/bax_visualizer/bax_widget.py b/src/badger/gui/components/bax_visualizer/bax_widget.py index b88bda02..a9333188 100644 --- a/src/badger/gui/components/bax_visualizer/bax_widget.py +++ b/src/badger/gui/components/bax_visualizer/bax_widget.py @@ -13,21 +13,51 @@ from badger.gui.components.bax_visualizer.ui import UI from badger.gui.components.extension_utilities import HandledException, requires_update from badger.routine import Routine -from badger.utils import create_archive_run_filename +from badger.utils import BlockSignalsContext, create_archive_run_filename logger = logging.getLogger(__name__) +@dataclass +class GridOptimizePlots: + objective: bool = True + + +@dataclass +class EmittancePlots: + emittance_x: bool = True + emittance_y: bool = True + bmag_x: bool = True + bmag_y: bool = True + + +@dataclass +class PathwiseSolenoidAlignmentPlots: + misalignment_x: bool = True + misalignment_y: bool = True + + +@dataclass() +class Plot1Parameters: + n_grid: int = 50 + n_samples: int = 100 + grid_optimize: GridOptimizePlots = field(default_factory=GridOptimizePlots) + emittance: EmittancePlots = field(default_factory=EmittancePlots) + pathwise_solenoid_alignment: PathwiseSolenoidAlignmentPlots = field( + default_factory=PathwiseSolenoidAlignmentPlots + ) + + @dataclass() -class PlotParameters: +class Plot2Parameters: n_grid: int = 50 n_samples: int = 100 @dataclass() class Parameters: - tab_1: PlotParameters = field(default_factory=PlotParameters) - tab_2: PlotParameters = field(default_factory=PlotParameters) + tab_1: Plot1Parameters = field(default_factory=Plot1Parameters) + tab_2: Plot2Parameters = field(default_factory=Plot2Parameters) active_tab: int = 0 variables: list[str] = field(default_factory=list) variable_idx_x: int = 0 @@ -63,6 +93,8 @@ def __init__(self, routine: Routine, parent: Optional[QWidget] = None): def initialize_widget(self) -> None: logger.debug("Initializing BaxWidget") + self.parameters = DEFAULT_PARAMETERS + variable_names = list(self.routine.vocs.variable_names) self.parameters.variables = variable_names @@ -77,6 +109,24 @@ def initialize_widget(self) -> None: self.parameters.variable_idx_x = min(temp_x, len(variable_names) - 1) self.parameters.variable_idx_y = min(temp_y, len(variable_names) - 1) + # Hide plotting options that are not relevant to the current algorithm + algorithm_type = self.generator.algorithm.name + if algorithm_type == "grid_optimize": + self.ui.controls_area.emittance_x_checkbox.setVisible(False) + self.ui.controls_area.emittance_y_checkbox.setVisible(False) + self.ui.controls_area.bmag_x_checkbox.setVisible(False) + self.ui.controls_area.bmag_y_checkbox.setVisible(False) + self.ui.controls_area.alignment_x_checkbox.setVisible(False) + self.ui.controls_area.alignment_y_checkbox.setVisible(False) + elif algorithm_type == "emittance": + self.ui.controls_area.grid_optimize_checkbox.setVisible(False) + self.ui.controls_area.alignment_x_checkbox.setVisible(False) + self.ui.controls_area.alignment_y_checkbox.setVisible(False) + elif algorithm_type == "pathwise_solenoid_alignment": + self.ui.controls_area.grid_optimize_checkbox.setVisible(False) + self.ui.controls_area.emittance_x_checkbox.setVisible(False) + self.ui.controls_area.emittance_y_checkbox.setVisible(False) + def requires_reinitialization(self) -> bool: # Check if the extension needs to be reinitialized logger.debug("Checking if Bax Visualizer needs to be reinitialized") @@ -154,17 +204,98 @@ def setup_connections(self) -> None: lambda index: self.update_tab_index(index) ) + self.ui.controls_area.n_grid_spin_box.valueChanged.connect( + lambda value: self.update_n_grid(value) + ) + self.ui.controls_area.n_samples_spin_box.valueChanged.connect( + lambda value: self.update_n_samples(value) + ) + + # Plotting options checkboxes + for label, value in [ + ("Grid Optimize", self.parameters.tab_1.grid_optimize.objective), + ("Emittance X", self.parameters.tab_1.emittance.emittance_x), + ("Emittance Y", self.parameters.tab_1.emittance.emittance_y), + ("Bmag X", self.parameters.tab_1.emittance.bmag_x), + ("Bmag Y", self.parameters.tab_1.emittance.bmag_y), + ( + "Alignment X", + self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_x, + ), + ( + "Alignment Y", + self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_y, + ), + ]: + checkbox = getattr( + self.ui.controls_area, label.replace(" ", "_").lower() + "_checkbox" + ) + checkbox.stateChanged.connect( + lambda _, lbl=label: self.update_plot_option(lbl) + ) + + def update_n_grid(self, value: int) -> None: + self.parameters.tab_1.n_grid = value + self.update_plots(requires_rebuild=True, interval=0) + + def update_n_samples(self, value: int) -> None: + self.parameters.tab_1.n_samples = value + self.update_plots(requires_rebuild=True, interval=0) + + def update_plot_option(self, label: str) -> None: + checkbox = getattr( + self.ui.controls_area, label.replace(" ", "_").lower() + "_checkbox" + ) + is_checked = checkbox.isChecked() + + if label == "Grid Optimize": + self.parameters.tab_1.grid_optimize.objective = is_checked + elif label == "Emittance X": + self.parameters.tab_1.emittance.emittance_x = is_checked + elif label == "Emittance Y": + self.parameters.tab_1.emittance.emittance_y = is_checked + elif label == "Bmag X": + self.parameters.tab_1.emittance.bmag_x = is_checked + elif label == "Bmag Y": + self.parameters.tab_1.emittance.bmag_y = is_checked + elif label == "Alignment X": + self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_x = ( + is_checked + ) + elif label == "Alignment Y": + self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_y = ( + is_checked + ) + + self.update_plots(requires_rebuild=True, interval=0) + def update_tab_index(self, index: int) -> None: self.parameters.active_tab = index def update_variables(self) -> None: - self.parameters.variable_idx_x = ( - self.ui.controls_area.x_axis_combo_box.currentIndex() - ) - self.parameters.variable_idx_y = ( - self.ui.controls_area.y_axis_combo_box.currentIndex() - ) + previous_x_index = self.parameters.variable_idx_x + previous_y_index = self.parameters.variable_idx_y + + current_x_index = self.ui.controls_area.x_axis_combo_box.currentIndex() + current_y_index = self.ui.controls_area.y_axis_combo_box.currentIndex() + + if current_x_index == current_y_index: + with BlockSignalsContext( + ( + self.ui.controls_area.x_axis_combo_box, + self.ui.controls_area.y_axis_combo_box, + ) + ): + # If the user selects the same variable for both axes, we can either swap the previous indices or reset to defaults. Here, we choose to swap. + self.ui.controls_area.x_axis_combo_box.setCurrentIndex(previous_y_index) + self.ui.controls_area.y_axis_combo_box.setCurrentIndex(previous_x_index) + # If the user selects the same variable for both axes, we can either swap the previous indices or reset to defaults. Here, we choose to swap. + current_x_index = self.parameters.variable_idx_y + current_y_index = self.parameters.variable_idx_x + + self.parameters.variable_idx_x = current_x_index + self.parameters.variable_idx_y = current_y_index self.update_plots(requires_rebuild=True, interval=0) diff --git a/src/badger/gui/components/bax_visualizer/controls.py b/src/badger/gui/components/bax_visualizer/controls.py index 6354b53b..2e69ce67 100644 --- a/src/badger/gui/components/bax_visualizer/controls.py +++ b/src/badger/gui/components/bax_visualizer/controls.py @@ -9,9 +9,11 @@ from PyQt5.QtWidgets import ( QCheckBox, QComboBox, + QGroupBox, QHBoxLayout, QLabel, QPushButton, + QSpinBox, QVBoxLayout, QWidget, ) @@ -40,7 +42,9 @@ def _initialize_ui(self) -> None: # Create the layout for the controls controls_layout = QVBoxLayout() - controls_layout.addLayout(self._create_variable_layout()) + controls_layout.addWidget(self._create_variable_group()) + controls_layout.addWidget(self._create_plot_options()) + controls_layout.addStretch() # Add stretch to push controls to the top # Add the controls to the layout @@ -52,6 +56,65 @@ def _initialize_ui(self) -> None: self.setLayout(controls_layout) + def _create_plot_options(self) -> QGroupBox: + layout = QVBoxLayout() + group_widget = QGroupBox("Optional Plots") + + n_grid_label = QLabel("Number of Grid Points:") + self.n_grid_spin_box = QSpinBox() + self.n_grid_spin_box.setRange(10, 100) + self.n_grid_spin_box.setSingleStep(10) + self.n_grid_spin_box.setValue(self.parameters.tab_1.n_grid) + + n_samples_label = QLabel("Number of Samples:") + self.n_samples_spin_box = QSpinBox() + self.n_samples_spin_box.setRange(10, 100) + self.n_samples_spin_box.setSingleStep(10) + self.n_samples_spin_box.setValue(self.parameters.tab_1.n_samples) + + layout.addWidget(n_grid_label) + layout.addWidget(self.n_grid_spin_box) + layout.addWidget(n_samples_label) + layout.addWidget(self.n_samples_spin_box) + + # Create checkboxes for optional plots based on the parameters + + self.grid_optimize_checkbox = QCheckBox("Grid Optimize") + self.grid_optimize_checkbox.setChecked( + self.parameters.tab_1.grid_optimize.objective + ) + self.emittance_x_checkbox = QCheckBox("Emittance X") + self.emittance_x_checkbox.setChecked( + self.parameters.tab_1.emittance.emittance_x + ) + self.emittance_y_checkbox = QCheckBox("Emittance Y") + self.emittance_y_checkbox.setChecked( + self.parameters.tab_1.emittance.emittance_y + ) + self.bmag_x_checkbox = QCheckBox("Bmag X") + self.bmag_x_checkbox.setChecked(self.parameters.tab_1.emittance.bmag_x) + self.bmag_y_checkbox = QCheckBox("Bmag Y") + self.bmag_y_checkbox.setChecked(self.parameters.tab_1.emittance.bmag_y) + self.alignment_x_checkbox = QCheckBox("Alignment X") + self.alignment_x_checkbox.setChecked( + self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_x + ) + self.alignment_y_checkbox = QCheckBox("Alignment Y") + self.alignment_y_checkbox.setChecked( + self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_y + ) + + layout.addWidget(self.grid_optimize_checkbox) + layout.addWidget(self.emittance_x_checkbox) + layout.addWidget(self.emittance_y_checkbox) + layout.addWidget(self.bmag_x_checkbox) + layout.addWidget(self.bmag_y_checkbox) + layout.addWidget(self.alignment_x_checkbox) + layout.addWidget(self.alignment_y_checkbox) + + group_widget.setLayout(layout) + return group_widget + def update_controls(self) -> None: self.update_variables() with BlockSignalsContext((self.x_axis_combo_box, self.y_axis_combo_box)): @@ -72,7 +135,8 @@ def update_variables(self) -> None: self.y_axis_combo_box.clear() self.y_axis_combo_box.addItems(self.parameters.variables) - def _create_variable_layout(self) -> QVBoxLayout: + def _create_variable_group(self) -> QGroupBox: + group_box = QGroupBox("Variable Selection") layout = QVBoxLayout() x_axis_combo_box, self.x_axis_combo_box = self._create_variable_combo_box( is_x_axis=True @@ -85,7 +149,8 @@ def _create_variable_layout(self) -> QVBoxLayout: layout.addLayout(x_axis_combo_box) layout.addLayout(y_axis_combo_box) layout.addWidget(self.y_axis_checkbox) - return layout + group_box.setLayout(layout) + return group_box def _create_variable_combo_box( self, is_x_axis: bool = True, disabled: bool = False diff --git a/src/badger/gui/components/bax_visualizer/plotting.py b/src/badger/gui/components/bax_visualizer/plotting.py index 75e02cf6..fed7aae7 100644 --- a/src/badger/gui/components/bax_visualizer/plotting.py +++ b/src/badger/gui/components/bax_visualizer/plotting.py @@ -58,6 +58,37 @@ def _initialize_plotting_area(self) -> None: self.update_tab_widget() + def get_plot_results_keys(self) -> list[str]: + algorithm_type = self.generator.algorithm.name + if algorithm_type == "grid_optimize": + plot_options_dict = { + "objective": self.parameters.tab_1.grid_optimize.objective + } + + elif algorithm_type == "emittance": + plot_options_dict = { + "emittance_x": self.parameters.tab_1.emittance.emittance_x, + "emittance_y": self.parameters.tab_1.emittance.emittance_y, + "bmag_x": self.parameters.tab_1.emittance.bmag_x, + "bmag_y": self.parameters.tab_1.emittance.bmag_y, + } + elif algorithm_type == "pathwise_solenoid_alignment": + plot_options_dict = { + "misalignment_x": self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_x, + "misalignment_y": self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_y, + } + else: + raise ValueError(f"Unsupported algorithm type: {algorithm_type}") + + result_keys = [key for key, enabled in plot_options_dict.items() if enabled] + + if len(result_keys) == 0: + logger.warning( + "No results keys selected for plotting. Please enable at least one plot option." + ) + + return result_keys + def create_first_plot(self) -> tuple[Figure, Axes]: logger.debug("Creating first plot") @@ -69,6 +100,9 @@ def create_first_plot(self) -> tuple[Figure, Axes]: self.parameters.variables[self.parameters.variable_idx_y] ) + results_keys = self.get_plot_results_keys() + logger.debug(f"Results keys for plotting: {results_keys}") + logger.debug(f"Results file: {self.generator.algorithm_results_file}") fig, ax = visualize_virtual_measurement_result( self.generator, @@ -78,7 +112,7 @@ def create_first_plot(self) -> tuple[Figure, Axes]: n_grid=self.parameters.tab_1.n_grid, n_samples=self.parameters.tab_1.n_samples, show_observations=True, - result_keys=["objective"], + result_keys=results_keys, ) return fig, ax From 4f4e6198941472ccee3e0241937516f4ebd1d6c6 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Wed, 15 Jul 2026 17:37:02 -0700 Subject: [PATCH 12/18] Added reference point to controls area --- .../components/bax_visualizer/bax_widget.py | 87 ++++++++++- .../gui/components/bax_visualizer/controls.py | 141 +++++++++++++++++- .../gui/components/bax_visualizer/plotting.py | 9 +- .../gui/components/bax_visualizer/ui.py | 25 ++++ .../components/bo_visualizer/ui_components.py | 2 +- src/badger/gui/components/pydantic_editor.py | 11 +- 6 files changed, 253 insertions(+), 22 deletions(-) diff --git a/src/badger/gui/components/bax_visualizer/bax_widget.py b/src/badger/gui/components/bax_visualizer/bax_widget.py index a9333188..9d48b554 100644 --- a/src/badger/gui/components/bax_visualizer/bax_widget.py +++ b/src/badger/gui/components/bax_visualizer/bax_widget.py @@ -3,7 +3,7 @@ import logging import time from dataclasses import dataclass, field -from typing import Optional +from typing import Optional, cast from PyQt5.QtWidgets import QSizePolicy, QVBoxLayout, QWidget from xopt.generators.bayesian.bax_generator import BaxGenerator @@ -11,7 +11,11 @@ from badger.gui.components.analysis_widget import AnalysisWidget from badger.gui.components.bax_visualizer.ui import UI -from badger.gui.components.extension_utilities import HandledException, requires_update +from badger.gui.components.extension_utilities import ( + HandledException, + requires_update, + to_precision_float, +) from badger.routine import Routine from badger.utils import BlockSignalsContext, create_archive_run_filename @@ -41,6 +45,8 @@ class PathwiseSolenoidAlignmentPlots: class Plot1Parameters: n_grid: int = 50 n_samples: int = 100 + use_reference_point: bool = False + reference_points: dict[str, float] = field(default_factory=dict) grid_optimize: GridOptimizePlots = field(default_factory=GridOptimizePlots) emittance: EmittancePlots = field(default_factory=EmittancePlots) pathwise_solenoid_alignment: PathwiseSolenoidAlignmentPlots = field( @@ -48,16 +54,14 @@ class Plot1Parameters: ) -@dataclass() -class Plot2Parameters: - n_grid: int = 50 - n_samples: int = 100 +# @dataclass() +# class Plot2Parameters: @dataclass() class Parameters: tab_1: Plot1Parameters = field(default_factory=Plot1Parameters) - tab_2: Plot2Parameters = field(default_factory=Plot2Parameters) + # tab_2: Plot2Parameters = field(default_factory=Plot2Parameters) active_tab: int = 0 variables: list[str] = field(default_factory=list) variable_idx_x: int = 0 @@ -127,6 +131,8 @@ def initialize_widget(self) -> None: self.ui.controls_area.emittance_x_checkbox.setVisible(False) self.ui.controls_area.emittance_y_checkbox.setVisible(False) + self.ui.initialize_reference_table() + def requires_reinitialization(self) -> bool: # Check if the extension needs to be reinitialized logger.debug("Checking if Bax Visualizer needs to be reinitialized") @@ -211,6 +217,18 @@ def setup_connections(self) -> None: lambda value: self.update_n_samples(value) ) + self.ui.controls_area.reference_point_checkbox.stateChanged.connect( + lambda: self.update_use_reference_point() + ) + + self.ui.controls_area.reference_table.cellChanged.connect( + lambda: self.update_reference_point() + ) + + self.ui.controls_area.select_best_reference_point_button.clicked.connect( + lambda: self.set_best_reference_points() + ) + # Plotting options checkboxes for label, value in [ ("Grid Optimize", self.parameters.tab_1.grid_optimize.objective), @@ -234,6 +252,55 @@ def setup_connections(self) -> None: lambda _, lbl=label: self.update_plot_option(lbl) ) + def set_best_reference_points( + self, + ) -> None: + if self.generator.data is None: + raise HandledException( + ValueError, + "No data available in generator for selecting best reference points", + ) + + input_params = ( + # -1 index is used to select the last row of the DataFrame, which corresponds to the best reference points + self.generator.data[self.routine.vocs.variable_names].iloc[-1].to_dict() + ) + + logger.debug(f"Best reference points: {input_params}") + + # Update the reference table with the best reference points + self.parameters.tab_1.reference_points = cast( + dict[str, float], + {var: to_precision_float(input_params[var]) for var in input_params}, + ) + self.ui.controls_area.best_point_display.setText( + f"Best Reference Points: {', '.join(f'{k}: {v}' for k, v in input_params.items())}" + ) + + self.ui.controls_area.populate_reference_table() + + self.update_plots(requires_rebuild=True, interval=0) + + def update_use_reference_point(self) -> None: + self.parameters.tab_1.use_reference_point = ( + self.ui.controls_area.reference_point_checkbox.isChecked() + ) + + if self.parameters.tab_1.use_reference_point: + self.ui.controls_area.reference_table.setEnabled(True) + self.ui.controls_area.select_best_reference_point_button.setEnabled(True) + else: + self.ui.controls_area.reference_table.setEnabled(False) + self.ui.controls_area.select_best_reference_point_button.setEnabled(False) + + self.update_plots(requires_rebuild=True, interval=0) + + def update_reference_point(self) -> None: + self.parameters.tab_1.reference_points = ( + self.ui.controls_area.get_reference_points(self.parameters.variables) + ) + self.update_plots(requires_rebuild=True, interval=0) + def update_n_grid(self, value: int) -> None: self.parameters.tab_1.n_grid = value self.update_plots(requires_rebuild=True, interval=0) @@ -297,6 +364,9 @@ def update_variables(self) -> None: self.parameters.variable_idx_x = current_x_index self.parameters.variable_idx_y = current_y_index + with BlockSignalsContext(self.ui.controls_area.reference_table): + self.ui.controls_area.update_reference_point_table_editability() + self.update_plots(requires_rebuild=True, interval=0) def update_y_axis_controls(self) -> None: @@ -308,6 +378,9 @@ def update_y_axis_controls(self) -> None: else: self.ui.controls_area.y_axis_combo_box.setEnabled(True) + with BlockSignalsContext(self.ui.controls_area.reference_table): + self.ui.controls_area.update_reference_point_table_editability() + self.update_plots(requires_rebuild=True, interval=0) def isValidRoutine(self, routine: Routine) -> None: diff --git a/src/badger/gui/components/bax_visualizer/controls.py b/src/badger/gui/components/bax_visualizer/controls.py index 2e69ce67..5a2cfe6c 100644 --- a/src/badger/gui/components/bax_visualizer/controls.py +++ b/src/badger/gui/components/bax_visualizer/controls.py @@ -6,14 +6,18 @@ from typing import TYPE_CHECKING, Optional +from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( QCheckBox, QComboBox, QGroupBox, QHBoxLayout, + QHeaderView, QLabel, QPushButton, QSpinBox, + QTableWidget, + QTableWidgetItem, QVBoxLayout, QWidget, ) @@ -24,8 +28,14 @@ if TYPE_CHECKING: from badger.gui.components.bax_visualizer.bax_widget import Parameters +import logging + +logger = logging.getLogger(__name__) + class ControlsWidget(QWidget): + ref_inputs: list[QTableWidgetItem] = [] + def __init__( self, routine: Routine, @@ -43,6 +53,7 @@ def _initialize_ui(self) -> None: controls_layout = QVBoxLayout() controls_layout.addWidget(self._create_variable_group()) + controls_layout.addWidget(self._create_reference_point_group()) controls_layout.addWidget(self._create_plot_options()) controls_layout.addStretch() # Add stretch to push controls to the top @@ -56,6 +67,120 @@ def _initialize_ui(self) -> None: self.setLayout(controls_layout) + # Initialize the reference table based on the current vocs variable names + if self.parameters.tab_1.use_reference_point: + self.reference_table.setEnabled(True) + self.select_best_reference_point_button.setEnabled(True) + else: + self.reference_table.setEnabled(False) + self.select_best_reference_point_button.setEnabled(False) + + def _create_reference_point_group(self) -> QGroupBox: + layout = QVBoxLayout() + group_widget = QGroupBox("Reference Point") + + self.reference_point_checkbox = QCheckBox("Use Reference Point") + self.reference_point_checkbox.setChecked( + self.parameters.tab_1.use_reference_point + ) + + self.reference_table = QTableWidget() + self.reference_table.setColumnCount(2) + self.reference_table.setHorizontalHeaderLabels(["Variable", "Value"]) + horizontal_header = self.reference_table.horizontalHeader() + horizontal_header.setSectionResizeMode(QHeaderView.ResizeMode.Stretch) + + self.select_best_reference_point_button = QPushButton( + "Set Best Reference Point" + ) + self.best_point_display = QLabel("Best Reference Point: N/A") + + layout.addWidget(self.reference_point_checkbox) + layout.addWidget(self.reference_table) + layout.addWidget(self.select_best_reference_point_button) + layout.addWidget(self.best_point_display) + + group_widget.setLayout(layout) + return group_widget + + def populate_reference_table( + self, + ) -> None: + """Populate the reference table based on the current vocs variable names.""" + + logger.debug("Populating reference table") + + with BlockSignalsContext(self.reference_table): + self.reference_table.setRowCount(len(self.parameters.variables)) + self.ref_inputs: list[QTableWidgetItem] = [] + + for i, var_name in enumerate(self.parameters.variables): + variable_item = QTableWidgetItem(var_name) + itemIsEditable = Qt.ItemFlag.ItemIsEditable + + variable_item.setFlags( + variable_item.flags() & ~Qt.ItemFlags(itemIsEditable) + ) + self.reference_table.setItem(i, 0, variable_item) + + value = self.parameters.tab_1.reference_points[var_name] + + reference_point_item = QTableWidgetItem(str(value)) + self.ref_inputs.append(reference_point_item) + self.reference_table.setItem(i, 1, reference_point_item) + + self.update_reference_point_table_editability() + + def get_reference_points(self, variable_names: list[str]) -> dict[str, float]: + reference_points: dict[str, float] = {} + + # Create a mapping from variable names to ref_inputs + ref_inputs_dict = dict(zip(self.parameters.variables, self.ref_inputs)) + for var in self.parameters.variables: + if var in variable_names: + ref_value = float(ref_inputs_dict[var].text()) + reference_points[var] = ref_value + return reference_points + + def update_reference_point_table_editability(self) -> None: + """Disable and gray out reference points for selected variables.""" + + selected_variables = self.get_selected_variables() + + white = Qt.GlobalColor.white + lightGray = Qt.GlobalColor.lightGray + black = Qt.GlobalColor.black + + itemIsEditable = Qt.ItemFlag.ItemIsEditable + + for i, var_name in enumerate(self.parameters.variables): + # Get the reference point item from the table + ref_item = self.ref_inputs[i] + + if var_name in selected_variables: + # Disable editing and gray out the background + ref_item.setFlags(ref_item.flags() & ~Qt.ItemFlags(itemIsEditable)) + ref_item.setBackground(lightGray) + ref_item.setForeground(white) + else: + # Re-enable editing and set background to white + ref_item.setFlags(ref_item.flags() | Qt.ItemFlags(itemIsEditable)) + ref_item.setBackground(white) + ref_item.setForeground(black) + + # Force the table to refresh and update its view + viewport = self.reference_table.viewport() + viewport.update() + + def get_selected_variables(self) -> list[str]: + """Get the currently selected variables from the combo boxes.""" + selected_variables = [self.parameters.variables[self.parameters.variable_idx_x]] + if self.parameters.include_y: + selected_variables.append( + self.parameters.variables[self.parameters.variable_idx_y] + ) + return selected_variables + def _create_plot_options(self) -> QGroupBox: layout = QVBoxLayout() group_widget = QGroupBox("Optional Plots") @@ -79,27 +204,27 @@ def _create_plot_options(self) -> QGroupBox: # Create checkboxes for optional plots based on the parameters - self.grid_optimize_checkbox = QCheckBox("Grid Optimize") + self.grid_optimize_checkbox = QCheckBox("Show Objective") self.grid_optimize_checkbox.setChecked( self.parameters.tab_1.grid_optimize.objective ) - self.emittance_x_checkbox = QCheckBox("Emittance X") + self.emittance_x_checkbox = QCheckBox("Show Emittance X") self.emittance_x_checkbox.setChecked( self.parameters.tab_1.emittance.emittance_x ) - self.emittance_y_checkbox = QCheckBox("Emittance Y") + self.emittance_y_checkbox = QCheckBox("Show Emittance Y") self.emittance_y_checkbox.setChecked( self.parameters.tab_1.emittance.emittance_y ) - self.bmag_x_checkbox = QCheckBox("Bmag X") + self.bmag_x_checkbox = QCheckBox("Show Bmag X") self.bmag_x_checkbox.setChecked(self.parameters.tab_1.emittance.bmag_x) - self.bmag_y_checkbox = QCheckBox("Bmag Y") + self.bmag_y_checkbox = QCheckBox("Show Bmag Y") self.bmag_y_checkbox.setChecked(self.parameters.tab_1.emittance.bmag_y) - self.alignment_x_checkbox = QCheckBox("Alignment X") + self.alignment_x_checkbox = QCheckBox("Show Alignment X") self.alignment_x_checkbox.setChecked( self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_x ) - self.alignment_y_checkbox = QCheckBox("Alignment Y") + self.alignment_y_checkbox = QCheckBox("Show Alignment Y") self.alignment_y_checkbox.setChecked( self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_y ) @@ -179,5 +304,5 @@ def _create_include_y_checkbox(self) -> QCheckBox: def _create_update_button(self) -> QPushButton: # Create a button for updating the plots - button = QPushButton("Update") # Replace with actual button implementation + button = QPushButton("Update") return button diff --git a/src/badger/gui/components/bax_visualizer/plotting.py b/src/badger/gui/components/bax_visualizer/plotting.py index fed7aae7..fdfdaf2e 100644 --- a/src/badger/gui/components/bax_visualizer/plotting.py +++ b/src/badger/gui/components/bax_visualizer/plotting.py @@ -104,11 +104,16 @@ def create_first_plot(self) -> tuple[Figure, Axes]: logger.debug(f"Results keys for plotting: {results_keys}") logger.debug(f"Results file: {self.generator.algorithm_results_file}") + + ref_point = None + if self.parameters.tab_1.use_reference_point: + ref_point = self.parameters.tab_1.reference_points + fig, ax = visualize_virtual_measurement_result( self.generator, variable_names=selected_variable_names, - idx=0, - reference_point=None, # type: ignore[arg-type] + idx=-1, + reference_point=ref_point, n_grid=self.parameters.tab_1.n_grid, n_samples=self.parameters.tab_1.n_samples, show_observations=True, diff --git a/src/badger/gui/components/bax_visualizer/ui.py b/src/badger/gui/components/bax_visualizer/ui.py index 68bb3852..d1fd5d62 100644 --- a/src/badger/gui/components/bax_visualizer/ui.py +++ b/src/badger/gui/components/bax_visualizer/ui.py @@ -3,10 +3,12 @@ from typing import TYPE_CHECKING, Optional from badger.gui.components.bax_visualizer.controls import ControlsWidget +from badger.gui.components.extension_utilities import to_precision_float if TYPE_CHECKING: from badger.gui.components.bax_visualizer.bax_widget import Parameters +from gest_api.vocs import ContinuousVariable from PyQt5.QtWidgets import QHBoxLayout, QSizePolicy, QVBoxLayout, QWidget from badger.gui.components.bax_visualizer.plotting import PlottingWidget @@ -61,3 +63,26 @@ def _initialize_ui(self) -> None: main_layout.addWidget(self.plotting_area, stretch=1) self.setLayout(main_layout) + + def initialize_reference_table(self) -> None: + """Initialize the reference table in the controls area.""" + + reference_points: dict[str, float] = {} + + vocs_variables = self.routine.vocs.variables + + for var_name, variable in vocs_variables.items(): + if not isinstance(variable, ContinuousVariable): + raise ValueError( + f"Variable '{var_name}' is not continuous. Only continuous variables are supported for reference points." + ) + + domain_range = variable.domain[1] - variable.domain[0] + + reference_points[var_name] = to_precision_float( + variable.domain[0] + (domain_range / 2.0) + ) + + self.parameters.tab_1.reference_points = reference_points + + self.controls_area.populate_reference_table() diff --git a/src/badger/gui/components/bo_visualizer/ui_components.py b/src/badger/gui/components/bo_visualizer/ui_components.py index e6e13dde..82843c32 100644 --- a/src/badger/gui/components/bo_visualizer/ui_components.py +++ b/src/badger/gui/components/bo_visualizer/ui_components.py @@ -154,7 +154,7 @@ def initialize_variables( variable.domain, # pyright: ignore[reportUnknownMemberType] ) reference_points[var_name] = to_precision_float( - (domain[1] - domain[0]) / 2.0 + domain[0] + ((domain[1] - domain[0]) / 2.0) ) configurable_options["reference_points"] = reference_points diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index c8d58535..17d0e25a 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -49,6 +49,9 @@ from xopt.numerical_optimizer import NumericalOptimizer from xopt.vocs import VOCS +from bax_algorithms.emittance import EmittanceAlgorithm +from bax_algorithms.solenoid_alignment import PathwiseSolenoidAlignment + logger = logging.getLogger(__name__) @@ -914,10 +917,10 @@ def get_all_compatible_classes( raise ValueError("Generator does not support algorithms.") compatible_classes = self.model_class.get_compatible_algorithms() # TODO: Add in additional from BAX algorithms. - # compatible_classes = list(compatible_classes) + [ - # EmittanceAlgorithm, - # PathwiseSolenoidAlignment, - # ] + compatible_classes = list(compatible_classes) + [ + EmittanceAlgorithm, + PathwiseSolenoidAlignment, + ] else: raise ValueError(f"Field name {field_name} is not recognized.") From bdb6750e51a6b0fa6abedd5f922014a28725391b Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Thu, 16 Jul 2026 19:18:13 -0700 Subject: [PATCH 13/18] Added reference point changes to BAX and BO extensions --- src/badger/gui/components/analysis_widget.py | 71 ++++++++-- .../components/bax_visualizer/bax_widget.py | 133 +++++------------- .../gui/components/bax_visualizer/controls.py | 111 ++++++++++++--- .../gui/components/bax_visualizer/plotting.py | 4 +- .../gui/components/bax_visualizer/ui.py | 55 +++++--- .../gui/components/bo_visualizer/bo_widget.py | 112 ++++++++------- .../gui/components/bo_visualizer/types.py | 3 - .../components/bo_visualizer/ui_components.py | 37 ++--- .../gui/components/extension_utilities.py | 13 ++ src/badger/gui/components/pydantic_editor.py | 3 - 10 files changed, 315 insertions(+), 227 deletions(-) diff --git a/src/badger/gui/components/analysis_widget.py b/src/badger/gui/components/analysis_widget.py index 68502fb0..ee4c3b56 100644 --- a/src/badger/gui/components/analysis_widget.py +++ b/src/badger/gui/components/analysis_widget.py @@ -11,6 +11,7 @@ from badger.gui.components.extension_utilities import HandledException from badger.routine import Routine +from badger.utils import create_archive_run_filename logger = logging.getLogger(__name__) @@ -43,14 +44,6 @@ def initialize_widget(self) -> None: """ raise NotImplementedError("initialize_widget method not implemented") - @abstractmethod - def requires_reinitialization(self) -> bool: - """ - Check if the widget requires reinitialization. - This is used to determine if the widget needs to be reset or updated. - """ - raise NotImplementedError("requires_reinitialization method not implemented") - @abstractmethod def update_plots(self, requires_rebuild: bool, interval: int) -> None: """ @@ -75,6 +68,68 @@ def isValidRoutine(self, routine: Routine) -> None: """ raise NotImplementedError("isValidRoutine method not implemented") + @abstractmethod + def reset_widget(self) -> None: + """ + Reset the widget to its initial state. + This method should be implemented to clear the current state and prepare the widget for a new routine. + """ + raise NotImplementedError("reset_widget method not implemented") + + def requires_reinitialization(self) -> bool: + # Check if the extension needs to be reinitialized + logger.debug("Checking if AnalysisWidget needs to be reinitialized") + + archive_name = create_archive_run_filename(self.routine) + + logger.debug(f"Archive name: {archive_name}") + + if not self.initialized: + logger.debug("Reset - Extension never initialized") + # Set up connections + logger.debug("Setting up connections") + self.setup_connections() + self.routine_identifier = archive_name + self.initialized = True + # Track the current data length so the growth check does not treat + # the first post-init update as a shrink and reinitialize again. + if self.routine.data is not None: + self.df_length = len(self.routine.data) + return True + + if self.routine_identifier != archive_name: + logger.debug("Reset - Routine name has changed") + # Reset first: reset_widget() clears routine_identifier, so the new + # identifier must be assigned afterwards. Assigning before the reset + # would be clobbered back to "" and force a reinitialization on every + # subsequent update during the same run. + self.reset_widget() + self.routine_identifier = archive_name + # Sync the tracked data length to the new routine so the growth + # check below does not immediately treat the next update as a + # shrink (df_length is left at inf by reset_widget()). + if self.routine.data is not None: + self.df_length = len(self.routine.data) + return True + + if self.routine.data is None: + logger.debug("Reset - No data available") + + return True + + previous_len = self.df_length + self.df_length = len(self.routine.data) + new_length = self.df_length + + if previous_len > new_length: + logger.debug("Reset - Data length is smaller") + # Keep df_length at the current (smaller) length rather than resetting + # it to inf. Leaving it at inf would make every subsequent update look + # like a shrink and reinitialize the widget on a loop. + return True + + return False + def update_routine(self, routine: Routine, generator_type: type[Generator]) -> None: self.routine = routine diff --git a/src/badger/gui/components/bax_visualizer/bax_widget.py b/src/badger/gui/components/bax_visualizer/bax_widget.py index 9d48b554..5e800728 100644 --- a/src/badger/gui/components/bax_visualizer/bax_widget.py +++ b/src/badger/gui/components/bax_visualizer/bax_widget.py @@ -3,7 +3,7 @@ import logging import time from dataclasses import dataclass, field -from typing import Optional, cast +from typing import Optional from PyQt5.QtWidgets import QSizePolicy, QVBoxLayout, QWidget from xopt.generators.bayesian.bax_generator import BaxGenerator @@ -13,11 +13,11 @@ from badger.gui.components.bax_visualizer.ui import UI from badger.gui.components.extension_utilities import ( HandledException, + get_latest_reference_points, requires_update, - to_precision_float, ) from badger.routine import Routine -from badger.utils import BlockSignalsContext, create_archive_run_filename +from badger.utils import BlockSignalsContext logger = logging.getLogger(__name__) @@ -45,7 +45,6 @@ class PathwiseSolenoidAlignmentPlots: class Plot1Parameters: n_grid: int = 50 n_samples: int = 100 - use_reference_point: bool = False reference_points: dict[str, float] = field(default_factory=dict) grid_optimize: GridOptimizePlots = field(default_factory=GridOptimizePlots) emittance: EmittancePlots = field(default_factory=EmittancePlots) @@ -69,17 +68,18 @@ class Parameters: include_y: bool = True -DEFAULT_PARAMETERS = Parameters() - - class BaxWidget(AnalysisWidget): generator: BaxGenerator - parameters: Parameters = DEFAULT_PARAMETERS + parameters: Parameters # type: ignore[assignment] def __init__(self, routine: Routine, parent: Optional[QWidget] = None): logger.debug("Initializing BaxWidget") super().__init__(routine=routine, parent=parent) + # Instance-level parameters. Never use a class-level default here: it + # would be shared (and mutated) across every BaxWidget instance. + self.parameters = Parameters() + self.ui = UI(routine=self.routine, parameters=self.parameters) # The UI must live inside a layout, otherwise Qt never manages its @@ -97,7 +97,17 @@ def __init__(self, routine: Routine, parent: Optional[QWidget] = None): def initialize_widget(self) -> None: logger.debug("Initializing BaxWidget") - self.parameters = DEFAULT_PARAMETERS + # Start from a clean parameter state and propagate that single object to + # every child widget so the whole UI subtree resets in lock-step. Simply + # rebinding ``self.parameters`` would leave the UI holding the previous + # run's object and cause state to desync. + self.parameters = Parameters() + self.ui.set_parameters(self.parameters) + + # The child widgets also cache the routine/generator they were built + # with. Refresh them too, otherwise the controls read the previous + # routine's variables (leaving stale reference-point keys behind). + self.ui.set_routine(self.routine) variable_names = list(self.routine.vocs.variable_names) self.parameters.variables = variable_names @@ -112,70 +122,15 @@ def initialize_widget(self) -> None: self.parameters.include_y = True self.parameters.variable_idx_x = min(temp_x, len(variable_names) - 1) self.parameters.variable_idx_y = min(temp_y, len(variable_names) - 1) - - # Hide plotting options that are not relevant to the current algorithm - algorithm_type = self.generator.algorithm.name - if algorithm_type == "grid_optimize": - self.ui.controls_area.emittance_x_checkbox.setVisible(False) - self.ui.controls_area.emittance_y_checkbox.setVisible(False) - self.ui.controls_area.bmag_x_checkbox.setVisible(False) - self.ui.controls_area.bmag_y_checkbox.setVisible(False) - self.ui.controls_area.alignment_x_checkbox.setVisible(False) - self.ui.controls_area.alignment_y_checkbox.setVisible(False) - elif algorithm_type == "emittance": - self.ui.controls_area.grid_optimize_checkbox.setVisible(False) - self.ui.controls_area.alignment_x_checkbox.setVisible(False) - self.ui.controls_area.alignment_y_checkbox.setVisible(False) - elif algorithm_type == "pathwise_solenoid_alignment": - self.ui.controls_area.grid_optimize_checkbox.setVisible(False) - self.ui.controls_area.emittance_x_checkbox.setVisible(False) - self.ui.controls_area.emittance_y_checkbox.setVisible(False) + self.ui.reset_ui() self.ui.initialize_reference_table() - def requires_reinitialization(self) -> bool: - # Check if the extension needs to be reinitialized - logger.debug("Checking if Bax Visualizer needs to be reinitialized") - - archive_name = create_archive_run_filename(self.routine) - - logger.debug(f"Archive name: {archive_name}") - - if not self.initialized: - logger.debug("Reset - Extension never initialized") - # Set up connections - logger.debug("Setting up connections") - self.setup_connections() - self.routine_identifier = archive_name - self.initialized = True - return True - - if self.routine_identifier != archive_name: - logger.debug("Reset - Routine name has changed") - self.routine_identifier = archive_name - self.reset_widget() - return True - - if self.routine.data is None: - logger.debug("Reset - No data available") - - return True - - previous_len = self.df_length - self.df_length = len(self.routine.data) - new_length = self.df_length - - if previous_len > new_length: - logger.debug("Reset - Data length is smaller") - self.df_length = float("inf") - return True - - return False - def reset_widget(self) -> None: logger.debug("Resetting BaxWidget") self.routine_identifier = "" self.df_length = float("inf") + self.ui.reset_ui() def update_plots(self, requires_rebuild: bool, interval: int) -> None: if not requires_update(self.last_updated, interval, requires_rebuild): @@ -186,8 +141,6 @@ def update_plots(self, requires_rebuild: bool, interval: int) -> None: # algorithm_results_file instead of a stale one from a previous run. self.ui.plotting_area.generator = self.generator - self.ui.controls_area.update_controls() - self.ui.plotting_area.update_tab_widget() self.last_updated = time.time() @@ -217,16 +170,12 @@ def setup_connections(self) -> None: lambda value: self.update_n_samples(value) ) - self.ui.controls_area.reference_point_checkbox.stateChanged.connect( - lambda: self.update_use_reference_point() - ) - self.ui.controls_area.reference_table.cellChanged.connect( lambda: self.update_reference_point() ) - self.ui.controls_area.select_best_reference_point_button.clicked.connect( - lambda: self.set_best_reference_points() + self.ui.controls_area.select_latest_reference_point_button.clicked.connect( + lambda: self.set_latest_reference_points() ) # Plotting options checkboxes @@ -252,49 +201,31 @@ def setup_connections(self) -> None: lambda _, lbl=label: self.update_plot_option(lbl) ) - def set_best_reference_points( + def set_latest_reference_points( self, ) -> None: if self.generator.data is None: raise HandledException( ValueError, - "No data available in generator for selecting best reference points", + "No data available in generator for selecting latest reference points", ) - input_params = ( - # -1 index is used to select the last row of the DataFrame, which corresponds to the best reference points - self.generator.data[self.routine.vocs.variable_names].iloc[-1].to_dict() + reference_points = get_latest_reference_points( + self.generator.data, self.routine.vocs.variable_names ) - logger.debug(f"Best reference points: {input_params}") + logger.debug(f"Latest reference points: {reference_points}") - # Update the reference table with the best reference points - self.parameters.tab_1.reference_points = cast( - dict[str, float], - {var: to_precision_float(input_params[var]) for var in input_params}, - ) - self.ui.controls_area.best_point_display.setText( - f"Best Reference Points: {', '.join(f'{k}: {v}' for k, v in input_params.items())}" + # Update the reference table with the latest reference points + self.parameters.tab_1.reference_points = reference_points + self.ui.controls_area.reference_point_display.setText( + f"Latest Reference Points: {', '.join(f'{k}: {v}' for k, v in reference_points.items())}" ) self.ui.controls_area.populate_reference_table() self.update_plots(requires_rebuild=True, interval=0) - def update_use_reference_point(self) -> None: - self.parameters.tab_1.use_reference_point = ( - self.ui.controls_area.reference_point_checkbox.isChecked() - ) - - if self.parameters.tab_1.use_reference_point: - self.ui.controls_area.reference_table.setEnabled(True) - self.ui.controls_area.select_best_reference_point_button.setEnabled(True) - else: - self.ui.controls_area.reference_table.setEnabled(False) - self.ui.controls_area.select_best_reference_point_button.setEnabled(False) - - self.update_plots(requires_rebuild=True, interval=0) - def update_reference_point(self) -> None: self.parameters.tab_1.reference_points = ( self.ui.controls_area.get_reference_points(self.parameters.variables) diff --git a/src/badger/gui/components/bax_visualizer/controls.py b/src/badger/gui/components/bax_visualizer/controls.py index 5a2cfe6c..74efd59c 100644 --- a/src/badger/gui/components/bax_visualizer/controls.py +++ b/src/badger/gui/components/bax_visualizer/controls.py @@ -67,38 +67,111 @@ def _initialize_ui(self) -> None: self.setLayout(controls_layout) - # Initialize the reference table based on the current vocs variable names - if self.parameters.tab_1.use_reference_point: - self.reference_table.setEnabled(True) - self.select_best_reference_point_button.setEnabled(True) - else: - self.reference_table.setEnabled(False) - self.select_best_reference_point_button.setEnabled(False) + def reset_controls_widget(self) -> None: + """Reset the controls to their initial state.""" + + with BlockSignalsContext( + [ + self.emittance_x_checkbox, + self.emittance_y_checkbox, + self.bmag_x_checkbox, + self.bmag_y_checkbox, + self.alignment_x_checkbox, + self.alignment_y_checkbox, + self.grid_optimize_checkbox, + self.n_grid_spin_box, + self.n_samples_spin_box, + self.y_axis_checkbox, + self.reference_table, + self.x_axis_combo_box, + self.y_axis_combo_box, + ] + ): + # Start from every option visible, then hide the ones that are not + # relevant to the current algorithm. Resetting visibility first + # ensures a checkbox hidden by a previous run's algorithm is shown + # again when the routine changes. + all_option_checkboxes = [ + self.grid_optimize_checkbox, + self.emittance_x_checkbox, + self.emittance_y_checkbox, + self.bmag_x_checkbox, + self.bmag_y_checkbox, + self.alignment_x_checkbox, + self.alignment_y_checkbox, + ] + for checkbox in all_option_checkboxes: + checkbox.setVisible(True) + + # Hide plotting options that are not relevant to the current algorithm + algorithm_type = self.routine.generator.algorithm.name + if algorithm_type == "grid_optimize": + self.emittance_x_checkbox.setVisible(False) + self.emittance_y_checkbox.setVisible(False) + self.bmag_x_checkbox.setVisible(False) + self.bmag_y_checkbox.setVisible(False) + self.alignment_x_checkbox.setVisible(False) + self.alignment_y_checkbox.setVisible(False) + elif algorithm_type == "emittance": + self.grid_optimize_checkbox.setVisible(False) + self.alignment_x_checkbox.setVisible(False) + self.alignment_y_checkbox.setVisible(False) + elif algorithm_type == "pathwise_solenoid_alignment": + self.grid_optimize_checkbox.setVisible(False) + self.emittance_x_checkbox.setVisible(False) + self.emittance_y_checkbox.setVisible(False) + else: + raise ValueError(f"Unsupported algorithm type: {algorithm_type}") + + # Push the (freshly reset) parameter values back into every widget so + # no toggle, spin value or selection lingers from the previous run. + self._sync_plot_options_from_parameters() + + self.reference_table.clearContents() + + self.x_axis_combo_box.clear() + self.y_axis_combo_box.clear() + + def _sync_plot_options_from_parameters(self) -> None: + """Set every plot-option widget to match the current parameters. + + Callers are responsible for blocking signals; this only mutates widget + state to mirror ``self.parameters``. + """ + tab = self.parameters.tab_1 + + self.n_grid_spin_box.setValue(tab.n_grid) + self.n_samples_spin_box.setValue(tab.n_samples) + + self.grid_optimize_checkbox.setChecked(tab.grid_optimize.objective) + self.emittance_x_checkbox.setChecked(tab.emittance.emittance_x) + self.emittance_y_checkbox.setChecked(tab.emittance.emittance_y) + self.bmag_x_checkbox.setChecked(tab.emittance.bmag_x) + self.bmag_y_checkbox.setChecked(tab.emittance.bmag_y) + self.alignment_x_checkbox.setChecked( + tab.pathwise_solenoid_alignment.misalignment_x + ) + self.alignment_y_checkbox.setChecked( + tab.pathwise_solenoid_alignment.misalignment_y + ) + self.y_axis_checkbox.setChecked(self.parameters.include_y) def _create_reference_point_group(self) -> QGroupBox: layout = QVBoxLayout() group_widget = QGroupBox("Reference Point") - self.reference_point_checkbox = QCheckBox("Use Reference Point") - self.reference_point_checkbox.setChecked( - self.parameters.tab_1.use_reference_point - ) - self.reference_table = QTableWidget() self.reference_table.setColumnCount(2) self.reference_table.setHorizontalHeaderLabels(["Variable", "Value"]) horizontal_header = self.reference_table.horizontalHeader() horizontal_header.setSectionResizeMode(QHeaderView.ResizeMode.Stretch) - self.select_best_reference_point_button = QPushButton( - "Set Best Reference Point" - ) - self.best_point_display = QLabel("Best Reference Point: N/A") + self.select_latest_reference_point_button = QPushButton("Set Latest") + self.reference_point_display = QLabel("") - layout.addWidget(self.reference_point_checkbox) layout.addWidget(self.reference_table) - layout.addWidget(self.select_best_reference_point_button) - layout.addWidget(self.best_point_display) + layout.addWidget(self.select_latest_reference_point_button) + layout.addWidget(self.reference_point_display) group_widget.setLayout(layout) return group_widget diff --git a/src/badger/gui/components/bax_visualizer/plotting.py b/src/badger/gui/components/bax_visualizer/plotting.py index fdfdaf2e..301c68c8 100644 --- a/src/badger/gui/components/bax_visualizer/plotting.py +++ b/src/badger/gui/components/bax_visualizer/plotting.py @@ -105,9 +105,7 @@ def create_first_plot(self) -> tuple[Figure, Axes]: logger.debug(f"Results file: {self.generator.algorithm_results_file}") - ref_point = None - if self.parameters.tab_1.use_reference_point: - ref_point = self.parameters.tab_1.reference_points + ref_point = self.parameters.tab_1.reference_points fig, ax = visualize_virtual_measurement_result( self.generator, diff --git a/src/badger/gui/components/bax_visualizer/ui.py b/src/badger/gui/components/bax_visualizer/ui.py index d1fd5d62..e7b1d75f 100644 --- a/src/badger/gui/components/bax_visualizer/ui.py +++ b/src/badger/gui/components/bax_visualizer/ui.py @@ -2,13 +2,15 @@ from typing import TYPE_CHECKING, Optional + from badger.gui.components.bax_visualizer.controls import ControlsWidget -from badger.gui.components.extension_utilities import to_precision_float +from badger.gui.components.extension_utilities import ( + get_latest_reference_points, +) if TYPE_CHECKING: from badger.gui.components.bax_visualizer.bax_widget import Parameters -from gest_api.vocs import ContinuousVariable from PyQt5.QtWidgets import QHBoxLayout, QSizePolicy, QVBoxLayout, QWidget from badger.gui.components.bax_visualizer.plotting import PlottingWidget @@ -26,12 +28,12 @@ def __init__( self.routine = routine self.parameters = parameters - self._initialize_ui() + self.initialize_ui() self.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding) self.setMinimumSize(1250, 600) - def _initialize_ui(self) -> None: + def initialize_ui(self) -> None: main_layout = QHBoxLayout() main_layout.setContentsMargins(0, 0, 0, 0) @@ -64,24 +66,41 @@ def _initialize_ui(self) -> None: self.setLayout(main_layout) - def initialize_reference_table(self) -> None: - """Initialize the reference table in the controls area.""" + def set_parameters(self, parameters: "Parameters") -> None: + """Point this widget and every child widget at the same parameters object. - reference_points: dict[str, float] = {} - - vocs_variables = self.routine.vocs.variables + This must be called whenever the top-level parameters object is + replaced (e.g. on reinitialization) so the controls and plotting areas + do not keep reading/writing a stale object from a previous run. + """ + self.parameters = parameters + self.controls_area.parameters = parameters + self.plotting_area.parameters = parameters + + def set_routine(self, routine: Routine) -> None: + """Point this widget and every child widget at the current routine. + + The child widgets cache the routine (and its generator) they were built + with. When the user switches routines, those caches must be refreshed or + the controls/reference table will read the previous run's variables and + data (e.g. leftover reference-point keys that no longer exist in the new + routine's vocs). + """ + self.routine = routine + self.controls_area.routine = routine + self.plotting_area.generator = routine.generator - for var_name, variable in vocs_variables.items(): - if not isinstance(variable, ContinuousVariable): - raise ValueError( - f"Variable '{var_name}' is not continuous. Only continuous variables are supported for reference points." - ) + def reset_ui(self) -> None: + """Reset the UI to its initial state.""" + self.controls_area.reset_controls_widget() + self.controls_area.update_controls() - domain_range = variable.domain[1] - variable.domain[0] + def initialize_reference_table(self) -> None: + """Initialize the reference table in the controls area.""" - reference_points[var_name] = to_precision_float( - variable.domain[0] + (domain_range / 2.0) - ) + reference_points = get_latest_reference_points( + self.routine.generator.data, self.parameters.variables + ) self.parameters.tab_1.reference_points = reference_points diff --git a/src/badger/gui/components/bo_visualizer/bo_widget.py b/src/badger/gui/components/bo_visualizer/bo_widget.py index 0241625a..3e4c827a 100644 --- a/src/badger/gui/components/bo_visualizer/bo_widget.py +++ b/src/badger/gui/components/bo_visualizer/bo_widget.py @@ -29,11 +29,12 @@ from badger.gui.components.bo_visualizer.ui_components import UIComponents from badger.gui.components.extension_utilities import ( HandledException, + get_latest_reference_points, signal_logger, to_precision_float, ) from badger.routine import Routine -from badger.utils import BlockSignalsContext, create_archive_run_filename +from badger.utils import BlockSignalsContext logger = logging.getLogger(__name__) @@ -51,7 +52,6 @@ "variable_2": 1, "variables": [], "reference_points": {}, - "reference_points_range": {}, "include_variable_2": True, } @@ -126,9 +126,11 @@ def initialize_widget(self) -> None: self.parameters["include_variable_2"] = False self.parameters["variable_2"] = -1 - vocs_variables = self.routine.vocs.variables + vocs_variables = self.routine.vocs.variable_names - self.ui_components.initialize_variables(self.parameters, vocs_variables) + self.ui_components.initialize_variables( + self.routine.generator.data, self.parameters, vocs_variables + ) self.ui_components.update_variables(self.parameters) @@ -205,6 +207,12 @@ def setup_connections(self) -> None: )() ) + self.ui_components.set_latest_reference_points_button.clicked.connect( + lambda: signal_logger("Set latest reference points clicked")( + lambda: self.on_set_latest_reference_points_clicked() + )() + ) + def on_button_clicked(self) -> None: self.update_extension(self.routine, True) @@ -221,6 +229,19 @@ def on_set_best_reference_point_clicked(self) -> None: ) self.update_plots(requires_rebuild=True) + def on_set_latest_reference_points_clicked(self) -> None: + logger.debug("Setting latest reference points") + try: + self.set_latest_reference_points() + except Exception as e: + logger.error(f"Error getting latest reference points: {e}") + QMessageBox.critical( + self, + "Error", + f"Error getting latest reference points: {e}", + ) + self.update_plots(requires_rebuild=True) + def on_plot_options_changed(self) -> None: self.update_plots(requires_rebuild=True) @@ -245,45 +266,6 @@ def reset_widget(self) -> None: DEFAULT_PARAMETERS.copy() ) - def requires_reinitialization(self) -> bool: - # Check if the extension needs to be reinitialized - logger.debug("Checking if BO Visualizer needs to be reinitialized") - - archive_name = create_archive_run_filename(self.routine) - - logger.debug(f"Archive name: {archive_name}") - - if not self.initialized: - logger.debug("Reset - Extension never initialized") - # Set up connections - logger.debug("Setting up connections") - self.setup_connections() - self.routine_identifier = archive_name - self.initialized = True - return True - - if self.routine_identifier != archive_name: - logger.debug("Reset - Routine name has changed") - self.routine_identifier = archive_name - self.reset_widget() - return True - - if self.routine.data is None: - logger.debug("Reset - No data available") - - return True - - previous_len = self.df_length - self.df_length = len(self.routine.data) - new_length = self.df_length - - if previous_len > new_length: - logger.debug("Reset - Data length is the same or smaller") - self.df_length = float("inf") - return True - - return False - def on_axis_selection_changed(self) -> None: logger.debug("Axis selection changed") @@ -482,12 +464,8 @@ def update_plots( # Get reference points for non-selected variables - non_selected_variables = [ - var for var in self.parameters["variables"] if var not in selected_variables - ] - reference_point = self.get_reference_points( - self.ui_components.ref_inputs, non_selected_variables + self.ui_components.ref_inputs, self.parameters["variables"] ) logger.debug("Updating plot with selected variables and reference points") @@ -562,3 +540,41 @@ def set_best_reference_points( self.ui_components.best_point_display.setText( f"Best Point Index: {index}\nValue: {to_precision_float(value)}" ) + self.ui_components.populate_reference_table( + self.parameters["variables"], + self.parameters["reference_points"], + ) + + def set_latest_reference_points( + self, + ) -> None: + if self.generator.data is None: + raise HandledException( + ValueError, + "No data available in generator for selecting latest reference points", + ) + + reference_points = get_latest_reference_points( + self.generator.data, self.routine.vocs.variable_names + ) + + if not reference_points: + raise HandledException(ValueError, "No latest reference points found") + + logger.debug(f"Latest reference points: {reference_points}") + + # Update the reference table with the latest reference points + self.parameters["reference_points"] = cast( + dict[str, float], + { + var: to_precision_float(reference_points[var]) + for var in reference_points + }, + ) + + self.ui_components.best_point_display.setText("Latest Reference Points Set") + + self.ui_components.populate_reference_table( + self.parameters["variables"], + self.parameters["reference_points"], + ) diff --git a/src/badger/gui/components/bo_visualizer/types.py b/src/badger/gui/components/bo_visualizer/types.py index 54533e0d..eaf18732 100644 --- a/src/badger/gui/components/bo_visualizer/types.py +++ b/src/badger/gui/components/bo_visualizer/types.py @@ -3,8 +3,6 @@ from typing import TypedDict -from gest_api.vocs import ContinuousVariable - class PlotOptions(TypedDict): n_grid: int @@ -21,5 +19,4 @@ class ConfigurableOptions(TypedDict): variable_2: int variables: list[str] reference_points: dict[str, float] - reference_points_range: dict[str, ContinuousVariable] include_variable_2: bool diff --git a/src/badger/gui/components/bo_visualizer/ui_components.py b/src/badger/gui/components/bo_visualizer/ui_components.py index 82843c32..e8f9ea57 100644 --- a/src/badger/gui/components/bo_visualizer/ui_components.py +++ b/src/badger/gui/components/bo_visualizer/ui_components.py @@ -2,9 +2,8 @@ point table, grid resolution, and plot option checkboxes.""" import logging -from typing import cast -from gest_api.vocs import BaseVariable, ContinuousVariable, VariableDict +import pandas as pd from PyQt5.QtCore import Qt from PyQt5.QtWidgets import ( QCheckBox, @@ -22,7 +21,7 @@ from badger.gui.components.bo_visualizer.types import ConfigurableOptions from badger.gui.components.extension_utilities import ( - to_precision_float, + get_latest_reference_points, ) from badger.utils import BlockSignalsContext @@ -40,7 +39,8 @@ def __init__( self.ref_inputs: list[QTableWidgetItem] = [] self.reference_table = QTableWidget() self.best_point_display = QLabel("") # Will be initialized later - self.set_best_reference_point_button = QPushButton("Set Best Reference Point") + self.set_best_reference_point_button = QPushButton("Set Best") + self.set_latest_reference_points_button = QPushButton("Set Latest") # Initialize other UI components self.update_button = QPushButton("Update") @@ -133,29 +133,13 @@ def initialize_ui_components( def initialize_variables( self, + data: pd.DataFrame | None, configurable_options: ConfigurableOptions, - vocs_variables: VariableDict, + vocs_variables: list[str], ) -> None: """Initialize the variable checkboxes with the provided variable names.""" - # Initialize the parameters with the routine's variables - configurable_options["reference_points_range"] = vocs_variables - reference_points: dict[str, float] = {} - - variables = cast(dict[str, BaseVariable], vocs_variables) - for var_name, variable in variables.items(): - if not isinstance(variable, ContinuousVariable): - raise ValueError( - f"Variable '{var_name}' is not continuous. Only continuous variables are supported for reference points." - ) - - domain = cast( - tuple[float, float], - variable.domain, # pyright: ignore[reportUnknownMemberType] - ) - reference_points[var_name] = to_precision_float( - domain[0] + ((domain[1] - domain[0]) / 2.0) - ) + reference_points = get_latest_reference_points(data, vocs_variables) configurable_options["reference_points"] = reference_points @@ -169,7 +153,12 @@ def create_reference_inputs(self) -> QGroupBox: horizontal_header.setSectionResizeMode(QHeaderView.ResizeMode.Stretch) layout.addWidget(self.reference_table) - layout.addWidget(self.set_best_reference_point_button) + + btn_group = QHBoxLayout() + btn_group.addWidget(self.set_latest_reference_points_button) + btn_group.addWidget(self.set_best_reference_point_button) + + layout.addLayout(btn_group) layout.addWidget(self.best_point_display) group_box.setLayout(layout) return group_box diff --git a/src/badger/gui/components/extension_utilities.py b/src/badger/gui/components/extension_utilities.py index b687506a..ad54f3c3 100644 --- a/src/badger/gui/components/extension_utilities.py +++ b/src/badger/gui/components/extension_utilities.py @@ -9,6 +9,7 @@ from typing import Any, Callable, Optional, ParamSpec import matplotlib.pyplot as plt +import pandas as pd from matplotlib.axes import Axes from matplotlib.figure import Figure from PyQt5.QtWidgets import QLayout, QTabWidget @@ -79,6 +80,18 @@ def to_precision_float(value: Any, precision: int = 4) -> float: ) +def get_latest_reference_points( + data: pd.DataFrame | None, variable_names: list[str] +) -> dict[str, float]: + + if data is None or data.empty: + raise ValueError("No data available to extract the latest reference point.") + + reference_points = data[variable_names].iloc[-1].to_dict() + + return {str(k): to_precision_float(v) for k, v in reference_points.items()} + + class HandledException(Exception): """ Custom exception class to handle exceptions in a way that can be caught and logged. diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index 17d0e25a..97b97001 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -748,9 +748,6 @@ def initialize_combo_widget( if selection is None: widget.addItem("null", selection) else: - logger.debug( - f"Adding selection {selection} with name {selection.model_fields['name'].default} to combo box" - ) widget.addItem(selection.model_fields["name"].default, selection) def set_params_from_class(self, pydantic_class: type[Any]) -> None: From 8a0f8a5285e79ba2f48b6038e696a384d109f094 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Thu, 16 Jul 2026 19:47:09 -0700 Subject: [PATCH 14/18] Reworked Pareto Front extension to be more in line with how ui components are referenced in the other extensions --- .../gui/components/pf_viewer/pf_widget.py | 215 ++++++------------ src/badger/gui/components/pf_viewer/types.py | 53 +---- src/badger/gui/components/pydantic_editor.py | 22 +- 3 files changed, 90 insertions(+), 200 deletions(-) diff --git a/src/badger/gui/components/pf_viewer/pf_widget.py b/src/badger/gui/components/pf_viewer/pf_widget.py index 20cd0202..5a97af69 100644 --- a/src/badger/gui/components/pf_viewer/pf_widget.py +++ b/src/badger/gui/components/pf_viewer/pf_widget.py @@ -41,16 +41,13 @@ signal_logger, ) from badger.gui.components.pf_viewer.types import ( - PFUI, ConfigurableOptions, - PFUILayouts, - PFUIWidgets, ) from badger.gui.components.plot_event_handlers import ( MatplotlibInteractionHandler, ) from badger.routine import Routine -from badger.utils import BlockSignalsContext, create_archive_run_filename +from badger.utils import BlockSignalsContext logger = logging.getLogger(__name__) @@ -79,6 +76,14 @@ class ParetoFrontWidget(AnalysisWidget): pf_mask: Optional[Tensor] = None plot_size: tuple[float, float] = (8, 6) + # UI component references + update_button: QPushButton + variable_1_combo: QComboBox + variable_2_combo: QComboBox + show_only_pareto_front_checkbox: QCheckBox + pareto_tab_widget: QTabWidget + hypervolume_layout: QVBoxLayout + def __init__( self, routine: Routine, @@ -108,43 +113,6 @@ def reset_widget(self) -> None: self.pf_2 = None self.pf_mask = None - def requires_reinitialization(self) -> bool: - # Check if the extension needs to be reinitialized - logger.debug("Checking if extension needs to be reinitialized") - - archive_name = create_archive_run_filename(self.routine) - - if not self.initialized: - logger.debug("Reset - Extension never initialized") - self.initialized = True - self.routine_identifier = archive_name - self.reset_widget() - self.setup_connections() - return True - - if self.routine_identifier != archive_name: - logger.debug("Reset - Routine name has changed") - self.routine_identifier = archive_name - self.reset_widget() - return True - - if self.routine.data is None: - logger.debug("Reset - No data available") - self.reset_widget() - return True - - previous_len = self.df_length - self.df_length = len(self.routine.data) - new_length = self.df_length - - if previous_len > new_length: - logger.debug("Reset - Data length is smaller") - self.reset_widget() - self.df_length = float("inf") - return True - - return False - def update_plots( self, requires_rebuild: bool = False, @@ -162,85 +130,59 @@ def update_plots( self.last_updated = time.time() def setup_connections(self) -> None: - self.ui["components"]["update"].clicked.connect( + self.update_button.clicked.connect( lambda: signal_logger("Update button clicked")( lambda: self.on_button_click() )() ) - self.ui["components"]["variables"]["variable_1"].currentIndexChanged.connect( + self.variable_1_combo.currentIndexChanged.connect( lambda: signal_logger("Variable 1 has changed")( lambda: self.on_variable_change() )() ) - self.ui["components"]["variables"]["variable_2"].currentIndexChanged.connect( + self.variable_2_combo.currentIndexChanged.connect( lambda: signal_logger("Variable 2 has changed")( lambda: self.on_variable_change() )() ) - self.ui["components"]["plot"]["pareto"].currentChanged.connect( + self.pareto_tab_widget.currentChanged.connect( lambda: signal_logger("Tab changed")(lambda: self.on_tab_change())() ) - self.ui["components"]["options"]["show_only_pareto_front"].clicked.connect( + self.show_only_pareto_front_checkbox.clicked.connect( lambda: signal_logger("Sample checkbox changed")(lambda: self.update_ui())() ) - def create_ui(self): - update_button = QPushButton("Update") - variable_1_combo = QComboBox() - variable_1_combo.setMinimumWidth(100) - variable_2_combo = QComboBox() - variable_2_combo.setMinimumWidth(100) - show_only_pareto_front = QCheckBox("Show only Pareto Front") - - components: PFUIWidgets = { - "variables": { - "variable_1": variable_1_combo, - "variable_2": variable_2_combo, - }, - "options": { - "show_only_pareto_front": show_only_pareto_front, - }, - "update": update_button, - "plot": { - "pareto": QTabWidget(), - "hypervolume": QVBoxLayout(), - }, - } - - layouts: PFUILayouts = { - "main": QHBoxLayout(), - "settings": QVBoxLayout(), - "plot": QGridLayout(), - "options": QVBoxLayout(), - "variables": QVBoxLayout(), - "update": QVBoxLayout(), - } - - self.ui: PFUI = {"components": components, "layouts": layouts} - - main_layout = self.ui["layouts"]["main"] + def create_ui(self) -> None: + self.update_button = QPushButton("Update") + self.variable_1_combo = QComboBox() + self.variable_1_combo.setMinimumWidth(100) + self.variable_2_combo = QComboBox() + self.variable_2_combo.setMinimumWidth(100) + self.show_only_pareto_front_checkbox = QCheckBox("Show only Pareto Front") + self.pareto_tab_widget = QTabWidget() + self.hypervolume_layout = QVBoxLayout() - # Left side of the layout - settings_layout = self.ui["layouts"]["settings"] + main_layout = QHBoxLayout() + # Left side of the layout + settings_layout = QVBoxLayout() settings_layout.setAlignment(Qt.AlignmentFlag.AlignTop) # Variables layout - - variables_layout = self.ui["layouts"]["variables"] + variables_layout = QVBoxLayout() variable_1_layout = QHBoxLayout() variable_1_layout.addWidget(QLabel("X Axis")) - variable_1_layout.addWidget(variable_1_combo) + variable_1_layout.addWidget(self.variable_1_combo) variable_2_layout = QHBoxLayout() variable_2_layout.addWidget( QLabel("Y Axis"), alignment=Qt.AlignmentFlag.AlignLeft ) - variable_2_layout.addWidget(variable_2_combo) + variable_2_layout.addWidget(self.variable_2_combo) variables_layout.addLayout(variable_1_layout) variables_layout.addLayout(variable_2_layout) @@ -251,25 +193,19 @@ def create_ui(self): settings_layout.addWidget(variables_group) # Options layout - options_layout = self.ui["layouts"]["options"] + options_layout = QVBoxLayout() - show_only_pareto_front = self.ui["components"]["options"][ - "show_only_pareto_front" - ] - show_only_pareto_front.setChecked( + self.show_only_pareto_front_checkbox.setChecked( self.parameters["plot_options"]["show_only_pareto_front"] ) - options_layout.addWidget(show_only_pareto_front) + options_layout.addWidget(self.show_only_pareto_front_checkbox) settings_layout.addLayout(options_layout) - # Update layout - - update_button = self.ui["components"]["update"] - + # Update button settings_layout.addStretch(1) - settings_layout.addWidget(update_button) + settings_layout.addWidget(self.update_button) settings_widget = QWidget() settings_widget.setLayout(settings_layout) settings_widget.setSizePolicy( @@ -278,19 +214,19 @@ def create_ui(self): main_layout.addWidget(settings_widget) # Right side of the layout - plot_layout = self.ui["layouts"]["plot"] + plot_layout = QGridLayout() plot_layout.setAlignment(Qt.AlignmentFlag.AlignCenter) - plot_tab_widget = self.ui["components"]["plot"]["pareto"] - plot_tab_widget.setCurrentIndex(self.parameters["plot_tab"]) - plot_tab_widget.setMinimumWidth(400) + self.pareto_tab_widget.setCurrentIndex(self.parameters["plot_tab"]) + self.pareto_tab_widget.setMinimumWidth(400) plot_hypervolume_widget = QWidget() - plot_hypervolume = self.ui["components"]["plot"]["hypervolume"] - plot_hypervolume_widget.setLayout(plot_hypervolume) + plot_hypervolume_widget.setLayout(self.hypervolume_layout) plot_hypervolume_widget.setMinimumWidth(400) plot_hypervolume_widget.setMaximumWidth(600) - plot_layout.addWidget(plot_tab_widget, 0, 0, Qt.AlignmentFlag.AlignCenter) + plot_layout.addWidget( + self.pareto_tab_widget, 0, 0, Qt.AlignmentFlag.AlignCenter + ) plot_layout.addWidget( plot_hypervolume_widget, 0, 1, Qt.AlignmentFlag.AlignCenter ) @@ -307,27 +243,22 @@ def initialize_widget(self) -> None: self.parameters["variables"] = variable_names self.parameters["objectives"] = objective_names - variable_1_combo = self.ui["components"]["variables"]["variable_1"] - variable_2_combo = self.ui["components"]["variables"]["variable_2"] - - with BlockSignalsContext([variable_1_combo, variable_2_combo]): - variable_1_combo.clear() - variable_2_combo.clear() + with BlockSignalsContext([self.variable_1_combo, self.variable_2_combo]): + self.variable_1_combo.clear() + self.variable_2_combo.clear() for variable_name in variable_names: - variable_1_combo.addItem(variable_name) - variable_2_combo.addItem(variable_name) + self.variable_1_combo.addItem(variable_name) + self.variable_2_combo.addItem(variable_name) - variable_1_combo.setCurrentIndex(self.parameters["variable_1"]) - variable_2_combo.setCurrentIndex(self.parameters["variable_2"]) + self.variable_1_combo.setCurrentIndex(self.parameters["variable_1"]) + self.variable_2_combo.setCurrentIndex(self.parameters["variable_2"]) - def on_tab_change(self): - self.parameters["plot_tab"] = self.ui["components"]["plot"][ - "pareto" - ].currentIndex() + def on_tab_change(self) -> None: + self.parameters["plot_tab"] = self.pareto_tab_widget.currentIndex() # change x and y axis options - x_combo = self.ui["components"]["variables"]["variable_1"] - y_combo = self.ui["components"]["variables"]["variable_2"] + x_combo = self.variable_1_combo + y_combo = self.variable_2_combo plot_tab = self.parameters["plot_tab"] @@ -355,39 +286,31 @@ def on_tab_change(self): # Update the plot self.update_pareto_front_plot() - def on_variable_change(self): + def on_variable_change(self) -> None: plot_tab = self.parameters["plot_tab"] if plot_tab == 0: - self.parameters["variable_1"] = self.ui["components"]["variables"][ - "variable_1" - ].currentIndex() - self.parameters["variable_2"] = self.ui["components"]["variables"][ - "variable_2" - ].currentIndex() + self.parameters["variable_1"] = self.variable_1_combo.currentIndex() + self.parameters["variable_2"] = self.variable_2_combo.currentIndex() elif plot_tab == 1: - self.parameters["objective_1"] = self.ui["components"]["variables"][ - "variable_1" - ].currentIndex() - self.parameters["objective_2"] = self.ui["components"]["variables"][ - "variable_2" - ].currentIndex() + self.parameters["objective_1"] = self.variable_1_combo.currentIndex() + self.parameters["objective_2"] = self.variable_2_combo.currentIndex() else: raise HandledException(ValueError, "Invalid plot tab") self.update_pareto_front_plot() - def on_button_click(self): + def on_button_click(self) -> None: self.update_extension(self.routine, True) - def update_ui(self): + def update_ui(self) -> None: self.update_plots(requires_rebuild=True) def update_pareto_front_plot( self, - ): + ) -> None: self.update_pareto_front() - plot_tab_widget = self.ui["components"]["plot"]["pareto"] + plot_tab_widget = self.pareto_tab_widget variables = self.parameters["variables"] with BlockSignalsContext(plot_tab_widget): @@ -464,10 +387,10 @@ def update_pareto_front_plot( def update_hypervolume_plot( self, - ): + ) -> None: self.update_hypervolume() - plot_hypervolume = self.ui["components"]["plot"]["hypervolume"] + plot_hypervolume = self.hypervolume_layout with BlockSignalsContext(plot_hypervolume): clear_layout(plot_hypervolume) @@ -482,11 +405,9 @@ def update_hypervolume_plot( blank_canvas = FigureCanvas(fig) plot_hypervolume.addWidget(blank_canvas) - def create_pareto_plot(self, fig: Figure, ax: Axes): + def create_pareto_plot(self, fig: Figure, ax: Axes) -> tuple[Figure, Axes]: current_tab = self.parameters["plot_tab"] - show_only_pareto_front = self.ui["components"]["options"][ - "show_only_pareto_front" - ].isChecked() + show_only_pareto_front = self.show_only_pareto_front_checkbox.isChecked() if current_tab == 0: x_axis = self.parameters["variable_1"] @@ -587,7 +508,7 @@ def create_pareto_plot(self, fig: Figure, ax: Axes): return fig, ax - def create_hypervolume_plot(self, fig: Figure, ax: Axes): + def create_hypervolume_plot(self, fig: Figure, ax: Axes) -> tuple[Figure, Axes]: data_points = self.hypervolume_history if len(data_points) == 0: raise HandledException( @@ -613,7 +534,7 @@ def create_hypervolume_plot(self, fig: Figure, ax: Axes): return fig, ax - def update_hypervolume(self): + def update_hypervolume(self) -> None: # Get the hypervolume from the generator self.generator.update_pareto_front_history() pareto_front_history_df = self.generator.pareto_front_history @@ -623,7 +544,7 @@ def update_hypervolume(self): self.hypervolume_history = pareto_front_history_df - def update_pareto_front(self): + def update_pareto_front(self) -> None: pf_1, pf_2, pf_mask, _ = self.generator.get_pareto_front_and_hypervolume() if pf_mask is None or pf_1 is None or pf_2 is None: diff --git a/src/badger/gui/components/pf_viewer/types.py b/src/badger/gui/components/pf_viewer/types.py index aaa80299..94f127c5 100644 --- a/src/badger/gui/components/pf_viewer/types.py +++ b/src/badger/gui/components/pf_viewer/types.py @@ -1,17 +1,8 @@ -"""TypedDict definitions for the Pareto front viewer — plot options, -objective/variable selection, and internal UI widget references.""" +"""TypedDict definitions for the Pareto front viewer — plot options and +objective/variable selection.""" from typing import TypedDict -from PyQt5.QtWidgets import ( - QRadioButton, - QComboBox, - QVBoxLayout, - QHBoxLayout, - QGridLayout, - QTabWidget, -) - class PlotOptions(TypedDict): show_only_pareto_front: bool @@ -26,43 +17,3 @@ class ConfigurableOptions(TypedDict): objective_1: int objective_2: int plot_tab: int - - -class PFOptionsUIWidgets(TypedDict): - show_only_pareto_front: QRadioButton - - -class PFVariablesUIWidgets(TypedDict): - variable_1: QComboBox - variable_2: QComboBox - - -class PFPlotUIWidgets(TypedDict): - pareto: QTabWidget - hypervolume: QVBoxLayout - - -class PFUIWidgets(TypedDict): - variables: PFVariablesUIWidgets - options: PFOptionsUIWidgets - update: QRadioButton - plot: PFPlotUIWidgets - - -class PFVariablesLayouts(TypedDict): - variable_1: QVBoxLayout - variable_2: QVBoxLayout - - -class PFUILayouts(TypedDict): - main: QHBoxLayout - settings: QVBoxLayout - plot: QGridLayout - options: QVBoxLayout - variables: QVBoxLayout - update: QVBoxLayout - - -class PFUI(TypedDict): - components: PFUIWidgets - layouts: PFUILayouts diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index 97b97001..1d89f517 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -475,6 +475,24 @@ def _qt_widgets_to_values_recurse( return out +def _default_for_new_row( + widget_type: type[Any] | None, +) -> float | int | bool | None: + """Provide a sensible default for a freshly-added list/dict row. + + Numeric and boolean widgets raise if resolved without a default value, so a + new (empty) row must supply one. Other types (e.g. str) already handle a + missing default gracefully, so ``None`` is returned for them. + """ + if widget_type is float: + return 0.0 + if widget_type is int: + return 0 + if widget_type is bool: + return False + return None + + class BadgerListItem(QWidget): def __init__(self, editor: "BadgerListEditor", parent: QWidget | None = None): super().__init__(parent) @@ -482,7 +500,7 @@ def __init__(self, editor: "BadgerListEditor", parent: QWidget | None = None): layout = QHBoxLayout(self) layout.setContentsMargins(0, 0, 0, 0) self.parameter_value = BadgerResolvedType.resolve_qt( - editor.widget_type, default=None + editor.widget_type, default=_default_for_new_row(editor.widget_type) ) if self.parameter_value: self.parameter_value.setSizePolicy( @@ -497,7 +515,7 @@ def __init__(self, editor: "BadgerListEditor", parent: QWidget | None = None): self.parameter_value2 = None if editor.widget_type2 is not None: self.parameter_value2 = BadgerResolvedType.resolve_qt( - editor.widget_type2, default=None + editor.widget_type2, default=_default_for_new_row(editor.widget_type2) ) if self.parameter_value2: self.parameter_value2.setSizePolicy( From 3a6e1adce32d64b95db4326b0ad6c488d0247509 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Mon, 27 Jul 2026 13:03:05 -0700 Subject: [PATCH 15/18] Fixed incorrect import from merge --- src/badger/gui/components/routine_page.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/src/badger/gui/components/routine_page.py b/src/badger/gui/components/routine_page.py index eed3651b..a3937bf0 100644 --- a/src/badger/gui/components/routine_page.py +++ b/src/badger/gui/components/routine_page.py @@ -68,8 +68,7 @@ get_generator_defaults, get_generator_dynamic, ) -from xopt.utils import get_local_region -from xopt.vocs import random_inputs +from xopt.vocs import get_local_region, random_inputs from badger.archive import update_run from badger.environment import instantiate_env From 586f4a6c05c57c55c558d4f9fd0947f9f415e62b Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Mon, 27 Jul 2026 13:40:30 -0700 Subject: [PATCH 16/18] Updated PathwiseMinimizeEmittance import --- src/badger/gui/components/pydantic_editor.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index e0015fe4..aee55f8d 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -56,7 +56,7 @@ from xopt.numerical_optimizer import NumericalOptimizer from xopt.vocs import VOCS -from bax_algorithms.emittance import EmittanceAlgorithm +from bax_algorithms.emittance import PathwiseMinimizeEmittance from bax_algorithms.solenoid_alignment import PathwiseSolenoidAlignment logger = logging.getLogger(__name__) @@ -940,7 +940,7 @@ def get_all_compatible_classes( compatible_classes = self.model_class.get_compatible_algorithms() # TODO: Add in additional from BAX algorithms. compatible_classes = list(compatible_classes) + [ - EmittanceAlgorithm, + PathwiseMinimizeEmittance, PathwiseSolenoidAlignment, ] else: From 72f8330f440a38a4dcbf134a5cadf6545ffb0c50 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Mon, 27 Jul 2026 15:46:27 -0700 Subject: [PATCH 17/18] Fixed issues with bax widget and added tensor parsing to str for pydanticeditor --- .../components/bax_visualizer/bax_widget.py | 38 ++++++++++------- .../gui/components/bax_visualizer/controls.py | 34 +++++++++------ .../gui/components/bax_visualizer/plotting.py | 16 +++---- .../gui/components/bo_visualizer/bo_widget.py | 15 ++++--- src/badger/gui/components/pydantic_editor.py | 42 +++++++++++++++++-- 5 files changed, 99 insertions(+), 46 deletions(-) diff --git a/src/badger/gui/components/bax_visualizer/bax_widget.py b/src/badger/gui/components/bax_visualizer/bax_widget.py index 5e800728..c080216e 100644 --- a/src/badger/gui/components/bax_visualizer/bax_widget.py +++ b/src/badger/gui/components/bax_visualizer/bax_widget.py @@ -22,9 +22,8 @@ logger = logging.getLogger(__name__) -@dataclass -class GridOptimizePlots: - objective: bool = True +# @dataclass +# class GridOptimizePlots: @dataclass @@ -43,11 +42,12 @@ class PathwiseSolenoidAlignmentPlots: @dataclass() class Plot1Parameters: - n_grid: int = 50 + n_grid: int = 20 n_samples: int = 100 reference_points: dict[str, float] = field(default_factory=dict) - grid_optimize: GridOptimizePlots = field(default_factory=GridOptimizePlots) - emittance: EmittancePlots = field(default_factory=EmittancePlots) + objective: bool = True + # grid_optimize: GridOptimizePlots = field(default_factory=GridOptimizePlots) + pathwise_minimize_emittance: EmittancePlots = field(default_factory=EmittancePlots) pathwise_solenoid_alignment: PathwiseSolenoidAlignmentPlots = field( default_factory=PathwiseSolenoidAlignmentPlots ) @@ -180,11 +180,17 @@ def setup_connections(self) -> None: # Plotting options checkboxes for label, value in [ - ("Grid Optimize", self.parameters.tab_1.grid_optimize.objective), - ("Emittance X", self.parameters.tab_1.emittance.emittance_x), - ("Emittance Y", self.parameters.tab_1.emittance.emittance_y), - ("Bmag X", self.parameters.tab_1.emittance.bmag_x), - ("Bmag Y", self.parameters.tab_1.emittance.bmag_y), + ("Grid Optimize", self.parameters.tab_1.objective), + ( + "Emittance X", + self.parameters.tab_1.pathwise_minimize_emittance.emittance_x, + ), + ( + "Emittance Y", + self.parameters.tab_1.pathwise_minimize_emittance.emittance_y, + ), + ("Bmag X", self.parameters.tab_1.pathwise_minimize_emittance.bmag_x), + ("Bmag Y", self.parameters.tab_1.pathwise_minimize_emittance.bmag_y), ( "Alignment X", self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_x, @@ -247,15 +253,15 @@ def update_plot_option(self, label: str) -> None: is_checked = checkbox.isChecked() if label == "Grid Optimize": - self.parameters.tab_1.grid_optimize.objective = is_checked + self.parameters.tab_1.objective = is_checked elif label == "Emittance X": - self.parameters.tab_1.emittance.emittance_x = is_checked + self.parameters.tab_1.pathwise_minimize_emittance.emittance_x = is_checked elif label == "Emittance Y": - self.parameters.tab_1.emittance.emittance_y = is_checked + self.parameters.tab_1.pathwise_minimize_emittance.emittance_y = is_checked elif label == "Bmag X": - self.parameters.tab_1.emittance.bmag_x = is_checked + self.parameters.tab_1.pathwise_minimize_emittance.bmag_x = is_checked elif label == "Bmag Y": - self.parameters.tab_1.emittance.bmag_y = is_checked + self.parameters.tab_1.pathwise_minimize_emittance.bmag_y = is_checked elif label == "Alignment X": self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_x = ( is_checked diff --git a/src/badger/gui/components/bax_visualizer/controls.py b/src/badger/gui/components/bax_visualizer/controls.py index 74efd59c..b39bac41 100644 --- a/src/badger/gui/components/bax_visualizer/controls.py +++ b/src/badger/gui/components/bax_visualizer/controls.py @@ -120,6 +120,8 @@ def reset_controls_widget(self) -> None: self.grid_optimize_checkbox.setVisible(False) self.emittance_x_checkbox.setVisible(False) self.emittance_y_checkbox.setVisible(False) + self.bmag_x_checkbox.setVisible(False) + self.bmag_y_checkbox.setVisible(False) else: raise ValueError(f"Unsupported algorithm type: {algorithm_type}") @@ -143,11 +145,15 @@ def _sync_plot_options_from_parameters(self) -> None: self.n_grid_spin_box.setValue(tab.n_grid) self.n_samples_spin_box.setValue(tab.n_samples) - self.grid_optimize_checkbox.setChecked(tab.grid_optimize.objective) - self.emittance_x_checkbox.setChecked(tab.emittance.emittance_x) - self.emittance_y_checkbox.setChecked(tab.emittance.emittance_y) - self.bmag_x_checkbox.setChecked(tab.emittance.bmag_x) - self.bmag_y_checkbox.setChecked(tab.emittance.bmag_y) + self.grid_optimize_checkbox.setChecked(tab.objective) + self.emittance_x_checkbox.setChecked( + tab.pathwise_minimize_emittance.emittance_x + ) + self.emittance_y_checkbox.setChecked( + tab.pathwise_minimize_emittance.emittance_y + ) + self.bmag_x_checkbox.setChecked(tab.pathwise_minimize_emittance.bmag_x) + self.bmag_y_checkbox.setChecked(tab.pathwise_minimize_emittance.bmag_y) self.alignment_x_checkbox.setChecked( tab.pathwise_solenoid_alignment.misalignment_x ) @@ -266,7 +272,7 @@ def _create_plot_options(self) -> QGroupBox: n_samples_label = QLabel("Number of Samples:") self.n_samples_spin_box = QSpinBox() - self.n_samples_spin_box.setRange(10, 100) + self.n_samples_spin_box.setRange(10, 1000) self.n_samples_spin_box.setSingleStep(10) self.n_samples_spin_box.setValue(self.parameters.tab_1.n_samples) @@ -278,21 +284,23 @@ def _create_plot_options(self) -> QGroupBox: # Create checkboxes for optional plots based on the parameters self.grid_optimize_checkbox = QCheckBox("Show Objective") - self.grid_optimize_checkbox.setChecked( - self.parameters.tab_1.grid_optimize.objective - ) + self.grid_optimize_checkbox.setChecked(self.parameters.tab_1.objective) self.emittance_x_checkbox = QCheckBox("Show Emittance X") self.emittance_x_checkbox.setChecked( - self.parameters.tab_1.emittance.emittance_x + self.parameters.tab_1.pathwise_minimize_emittance.emittance_x ) self.emittance_y_checkbox = QCheckBox("Show Emittance Y") self.emittance_y_checkbox.setChecked( - self.parameters.tab_1.emittance.emittance_y + self.parameters.tab_1.pathwise_minimize_emittance.emittance_y ) self.bmag_x_checkbox = QCheckBox("Show Bmag X") - self.bmag_x_checkbox.setChecked(self.parameters.tab_1.emittance.bmag_x) + self.bmag_x_checkbox.setChecked( + self.parameters.tab_1.pathwise_minimize_emittance.bmag_x + ) self.bmag_y_checkbox = QCheckBox("Show Bmag Y") - self.bmag_y_checkbox.setChecked(self.parameters.tab_1.emittance.bmag_y) + self.bmag_y_checkbox.setChecked( + self.parameters.tab_1.pathwise_minimize_emittance.bmag_y + ) self.alignment_x_checkbox = QCheckBox("Show Alignment X") self.alignment_x_checkbox.setChecked( self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_x diff --git a/src/badger/gui/components/bax_visualizer/plotting.py b/src/badger/gui/components/bax_visualizer/plotting.py index 301c68c8..13f9c225 100644 --- a/src/badger/gui/components/bax_visualizer/plotting.py +++ b/src/badger/gui/components/bax_visualizer/plotting.py @@ -61,19 +61,19 @@ def _initialize_plotting_area(self) -> None: def get_plot_results_keys(self) -> list[str]: algorithm_type = self.generator.algorithm.name if algorithm_type == "grid_optimize": - plot_options_dict = { - "objective": self.parameters.tab_1.grid_optimize.objective - } + plot_options_dict = {"objective": self.parameters.tab_1.objective} - elif algorithm_type == "emittance": + elif algorithm_type == "pathwise_minimize_emittance": plot_options_dict = { - "emittance_x": self.parameters.tab_1.emittance.emittance_x, - "emittance_y": self.parameters.tab_1.emittance.emittance_y, - "bmag_x": self.parameters.tab_1.emittance.bmag_x, - "bmag_y": self.parameters.tab_1.emittance.bmag_y, + "objective": self.parameters.tab_1.objective, + "emittance_x": self.parameters.tab_1.pathwise_minimize_emittance.emittance_x, + "emittance_y": self.parameters.tab_1.pathwise_minimize_emittance.emittance_y, + "bmag_x": self.parameters.tab_1.pathwise_minimize_emittance.bmag_x, + "bmag_y": self.parameters.tab_1.pathwise_minimize_emittance.bmag_y, } elif algorithm_type == "pathwise_solenoid_alignment": plot_options_dict = { + "objective": self.parameters.tab_1.objective, "misalignment_x": self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_x, "misalignment_y": self.parameters.tab_1.pathwise_solenoid_alignment.misalignment_y, } diff --git a/src/badger/gui/components/bo_visualizer/bo_widget.py b/src/badger/gui/components/bo_visualizer/bo_widget.py index 3e4c827a..2538d4dc 100644 --- a/src/badger/gui/components/bo_visualizer/bo_widget.py +++ b/src/badger/gui/components/bo_visualizer/bo_widget.py @@ -20,6 +20,7 @@ QWidget, ) from xopt.generator import Generator +from xopt.generators.bayesian.bax_generator import BaxGenerator from xopt.generators.bayesian.bayesian_generator import BayesianGenerator from xopt.vocs import select_best @@ -84,11 +85,7 @@ def isValidRoutine(self, routine: Routine) -> None: ValueError, "BO Visualizer requires at least one variable in the VOCS", ) - if len(routine.vocs.objective_names) < 1: - raise HandledException( - ValueError, - "BO Visualizer requires at least one objective in the VOCS", - ) + if not isinstance(routine.generator, BayesianGenerator): raise HandledException( ValueError, @@ -217,6 +214,7 @@ def on_button_clicked(self) -> None: self.update_extension(self.routine, True) def on_set_best_reference_point_clicked(self) -> None: + logger.debug("Setting best reference points") try: self.set_best_reference_points() @@ -490,6 +488,13 @@ def update_plots( def update_routine(self, routine: Routine, generator_type: type[Generator]) -> None: super().update_routine(routine, generator_type) + # The BAX generator has no objective, so "Set Best" (which relies on + # select_best over an objective) would fail. Disable the button for + # BAX routines and re-enable it for regular Bayesian ones. + self.ui_components.set_best_reference_point_button.setEnabled( + not isinstance(self.generator, BaxGenerator) + ) + # Handle the edge case where the extension has been opened after an optimization has already finished. if self.generator.model is None: logger.warning("Model not found in generator") diff --git a/src/badger/gui/components/pydantic_editor.py b/src/badger/gui/components/pydantic_editor.py index aee55f8d..b4732915 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -30,7 +30,7 @@ from pydantic import BaseModel, Field, ValidationError, create_model from pydantic.fields import FieldInfo from pydantic_core import PydanticUndefined, PydanticUndefinedType -from PyQt5.QtCore import Qt, pyqtSignal +from PyQt5.QtCore import Qt, QTimer, pyqtSignal from PyQt5.QtWidgets import ( QCheckBox, QComboBox, @@ -53,6 +53,7 @@ from xopt.generators.bayesian.bax_generator import BaxGenerator from xopt.generators.bayesian.bayesian_generator import BayesianGenerator from xopt.generators.bayesian.turbo import TurboController +from torch import Tensor from xopt.numerical_optimizer import NumericalOptimizer from xopt.vocs import VOCS @@ -321,7 +322,10 @@ def resolve_qt( ) if editor_info is not None: - widget.valueChanged.connect(lambda: handle_changed(editor_info)) + # Validate on ``editingFinished`` (focus loss / Enter) rather than + # ``valueChanged``: rebuilding the tree on every value change would + # destroy this spinbox mid-edit and drop the cursor. + widget.editingFinished.connect(lambda: handle_changed(editor_info)) elif resolved_type.main is int: widget = QSpinBox() widget.setRange(-(2**31), 2**31 - 1) # int32 min/max @@ -337,7 +341,10 @@ def resolve_qt( ) if editor_info is not None: - widget.valueChanged.connect(lambda: handle_changed(editor_info)) + # Validate on ``editingFinished`` (focus loss / Enter) rather than + # ``valueChanged``: rebuilding the tree on every value change would + # destroy this spinbox mid-edit and drop the cursor. + widget.editingFinished.connect(lambda: handle_changed(editor_info)) elif resolved_type.main is bool: widget = QCheckBox() if default is not None and not isinstance(default, PydanticUndefinedType): @@ -356,11 +363,23 @@ def resolve_qt( widget = QLineEdit() if default is None: widget.setText("null") + elif isinstance(default, Tensor): + # Tensor-typed fields (e.g. ``Tensor | None``) resolve to a bare + # union here, so they land in this catch-all. Render them as a plain + # nested list string (e.g. "[[1.0, 1.0], [0.0, 1.0]]") rather than + # the "tensor(...)" repr, so the value round-trips cleanly through + # the model's field validator. + widget.setText(str(default.tolist())) else: widget.setText(str(default)) if editor_info is not None: - widget.textChanged.connect(lambda: handle_changed(editor_info)) + # Validate on ``editingFinished`` (focus loss / Enter) rather than + # ``textChanged``. ``handle_changed`` rebuilds the whole tree, which + # destroys and recreates this very QLineEdit; doing that on every + # keystroke kills the text cursor and makes the view jump. Waiting + # until the user is done editing keeps the cursor active while typing. + widget.editingFinished.connect(lambda: handle_changed(editor_info)) widget.setProperty("badger_nullable", resolved_type.nullable) return widget @@ -1166,6 +1185,12 @@ def update_after_validate(self, defaults: dict[str, Any]) -> None: if model_class is None: return + # Rebuilding the tree resets the scrollbars, making the view jump back + # to the top on every edit. Capture the current scroll positions so we + # can restore them once the tree has been repopulated. + h_scroll = self.horizontalScrollBar().value() + v_scroll = self.verticalScrollBar().value() + self.clear() fields_to_remove = ["vocs"] @@ -1200,6 +1225,15 @@ def update_after_validate(self, defaults: dict[str, Any]) -> None: # Update parameters with defaults from generator class self.set_params_post_setup(defaults) + # Restore the scroll positions captured before the rebuild. Defer to the + # next event-loop iteration so the restore runs after the tree has laid + # out its (re)created items and updated the scrollbar ranges. + def restore_scroll() -> None: + self.horizontalScrollBar().setValue(h_scroll) + self.verticalScrollBar().setValue(v_scroll) + + QTimer.singleShot(0, restore_scroll) + if self.update_callback is not None: self.update_callback(self) From 9c81d0d92c5d481a463e4703682090c0849c8190 Mon Sep 17 00:00:00 2001 From: Mitchell Victoriano Date: Mon, 27 Jul 2026 17:51:28 -0700 Subject: [PATCH 18/18] Fixing validation issues due to new BADGER_TEMP_DIR config variable and new BAX algorithm implementation --- src/badger/tests/test_cli_basic.py | 2 +- src/badger/tests/test_gui_basic.py | 6 +++++- src/badger/tests/test_settings.py | 8 +++++++- 3 files changed, 13 insertions(+), 3 deletions(-) diff --git a/src/badger/tests/test_cli_basic.py b/src/badger/tests/test_cli_basic.py index 42a8efb2..1ceaf1bc 100644 --- a/src/badger/tests/test_cli_basic.py +++ b/src/badger/tests/test_cli_basic.py @@ -18,7 +18,7 @@ def test_cli_main(): # Check output lines outlines = out.splitlines() - assert len(outlines) == 11 + assert len(outlines) == 12 # Check name assert outlines[0] == "name: Badger the optimizer" diff --git a/src/badger/tests/test_gui_basic.py b/src/badger/tests/test_gui_basic.py index 09ddb919..d094ab1d 100644 --- a/src/badger/tests/test_gui_basic.py +++ b/src/badger/tests/test_gui_basic.py @@ -223,7 +223,11 @@ def test_default_low_noise_prior_in_bo(qtbot, init_multiprocessing): params_dict = yaml.safe_load(params) if "gp_constructor" in params_dict: - assert not params_dict["gp_constructor"]["use_low_noise_prior"] + # use_low_noise_prior may not be exposed in the GUI for every + # generator; default to False so a hidden key doesn't error. + assert not params_dict["gp_constructor"].get( + "use_low_noise_prior", False + ) else: # that part of params is hidden so we need to dig deeper pass diff --git a/src/badger/tests/test_settings.py b/src/badger/tests/test_settings.py index d9a24a29..55978f06 100644 --- a/src/badger/tests/test_settings.py +++ b/src/badger/tests/test_settings.py @@ -107,10 +107,16 @@ def test_init_settings(self): with patch( "badger.settings.ConfigSingleton", return_value=mock_config_singleton ) as mock_config_cls: - config_singleton = init_settings() + with patch( + "badger.settings.get_or_create_temp_directory" + ) as mock_get_or_create_temp_directory: + config_singleton = init_settings() mock_config_cls.assert_called_once_with( "/mock/config/folder/config.yaml", False ) + mock_get_or_create_temp_directory.assert_called_once_with( + mock_config_singleton + ) assert config_singleton == mock_config_singleton def test_config_singleton_initialization(self):