diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 86272cd4..f5ce9069 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 @@ -26,8 +26,13 @@ repos: exclude: ^src/badger/tests/ - 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..faa8c665 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", @@ -30,18 +28,15 @@ dependencies = [ "pillow", "requests", "xopt>=3.0.0", + "bax-algorithms", + ] dynamic = ["version"] [tool.setuptools_scm] 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 +55,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/actions/__init__.py b/src/badger/actions/__init__.py index 3e260442..c9226afa 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): @@ -43,6 +44,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" @@ -59,6 +61,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/factory.py b/src/badger/factory.py index cbebcc94..956e4c17 100644 --- a/src/badger/factory.py +++ b/src/badger/factory.py @@ -9,26 +9,26 @@ Also handles loading Markdown docs for the built-in documentation browser. """ -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 @@ -44,7 +44,7 @@ "time_dependent_upper_confidence_bound", "multi_fidelity", "nsga2", - "bax", + # "bax", ] diff --git a/src/badger/gui/components/analysis_extensions.py b/src/badger/gui/components/analysis_extensions.py index 4a4cbc14..3ffac040 100644 --- a/src/badger/gui/components/analysis_extensions.py +++ b/src/badger/gui/components/analysis_extensions.py @@ -1,31 +1,36 @@ """Dialog wrappers for analysis extensions (BO visualizer, Pareto front viewer). Each dialog receives live data updates from the run monitor.""" -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.QtCore import Qt, pyqtSignal 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) + # 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: @@ -62,8 +67,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) @@ -73,7 +76,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) @@ -82,7 +85,7 @@ class ParetoFrontViewer(AnalysisExtension): def __init__( self, routine: Routine, - parent: Optional[QDialog] = None, + parent: Optional[QWidget] = None, ): super().__init__(parent=parent) @@ -97,12 +100,31 @@ 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_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 ecf71238..ee4c3b56 100644 --- a/src/badger/gui/components/analysis_widget.py +++ b/src/badger/gui/components/analysis_widget.py @@ -1,22 +1,22 @@ """Base class that all analysis extension widgets must implement. Defines the interface for receiving routine updates and rendering plots.""" +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 +from badger.utils import create_archive_run_filename logger = logging.getLogger(__name__) -class AnalysisWidget(QDialog): +class AnalysisWidget(QWidget): routine: Routine generator: Generator parameters: dict[str, Any] = {} @@ -30,7 +30,7 @@ class AnalysisWidget(QDialog): def __init__( self, routine: Routine, - parent: Optional[QDialog] = None, + parent: Optional[QWidget] = None, ): super().__init__(parent=parent) self.routine = routine @@ -44,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: """ @@ -76,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/.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/__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..c080216e --- /dev/null +++ b/src/badger/gui/components/bax_visualizer/bax_widget.py @@ -0,0 +1,332 @@ +"""Widget that hosts the BAX visualizer extension within the Badger GUI.""" + +import logging +import time +from dataclasses import dataclass, field +from typing import Optional + +from PyQt5.QtWidgets import QSizePolicy, QVBoxLayout, 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.bax_visualizer.ui import UI +from badger.gui.components.extension_utilities import ( + HandledException, + get_latest_reference_points, + requires_update, +) +from badger.routine import Routine +from badger.utils import BlockSignalsContext + +logger = logging.getLogger(__name__) + + +# @dataclass +# class GridOptimizePlots: + + +@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 = 20 + n_samples: int = 100 + reference_points: dict[str, float] = field(default_factory=dict) + 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 + ) + + +# @dataclass() +# class Plot2Parameters: + + +@dataclass() +class Parameters: + 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 + variable_idx_y: int = 1 + include_y: bool = True + + +class BaxWidget(AnalysisWidget): + generator: BaxGenerator + 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 + # 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.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding) + self.setMinimumSize(800, 600) + + self.initialize_widget() + + def initialize_widget(self) -> None: + logger.debug("Initializing BaxWidget") + + # 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 + + 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) + self.ui.reset_ui() + + self.ui.initialize_reference_table() + + 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): + 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.plotting_area.update_tab_widget() + + self.last_updated = time.time() + + def setup_connections(self) -> None: + 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() + ) + self.ui.plotting_area.plot_tab_widget.currentChanged.connect( + 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) + ) + + self.ui.controls_area.reference_table.cellChanged.connect( + lambda: self.update_reference_point() + ) + + self.ui.controls_area.select_latest_reference_point_button.clicked.connect( + lambda: self.set_latest_reference_points() + ) + + # Plotting options checkboxes + for label, value in [ + ("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, + ), + ( + "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 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 + ) + + logger.debug(f"Latest reference points: {reference_points}") + + # 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_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) + + 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.objective = is_checked + elif label == "Emittance X": + self.parameters.tab_1.pathwise_minimize_emittance.emittance_x = is_checked + elif label == "Emittance Y": + self.parameters.tab_1.pathwise_minimize_emittance.emittance_y = is_checked + elif label == "Bmag X": + self.parameters.tab_1.pathwise_minimize_emittance.bmag_x = is_checked + elif label == "Bmag Y": + 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 + ) + 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: + + 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 + + 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: + + 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) + + 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: + if not isinstance(routine.generator, BayesianGenerator): + 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/controls.py b/src/badger/gui/components/bax_visualizer/controls.py new file mode 100644 index 00000000..b39bac41 --- /dev/null +++ b/src/badger/gui/components/bax_visualizer/controls.py @@ -0,0 +1,389 @@ +"""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.QtCore import Qt +from PyQt5.QtWidgets import ( + QCheckBox, + QComboBox, + QGroupBox, + QHBoxLayout, + QHeaderView, + QLabel, + QPushButton, + QSpinBox, + QTableWidget, + QTableWidgetItem, + 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 + +import logging + +logger = logging.getLogger(__name__) + + +class ControlsWidget(QWidget): + ref_inputs: list[QTableWidgetItem] = [] + + 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.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 + + # 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 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) + self.bmag_x_checkbox.setVisible(False) + self.bmag_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.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 + ) + 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_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_latest_reference_point_button = QPushButton("Set Latest") + self.reference_point_display = QLabel("") + + layout.addWidget(self.reference_table) + layout.addWidget(self.select_latest_reference_point_button) + layout.addWidget(self.reference_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") + + 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, 1000) + 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("Show 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.pathwise_minimize_emittance.emittance_x + ) + self.emittance_y_checkbox = QCheckBox("Show Emittance Y") + self.emittance_y_checkbox.setChecked( + 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.pathwise_minimize_emittance.bmag_x + ) + self.bmag_y_checkbox = QCheckBox("Show 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 + ) + self.alignment_y_checkbox = QCheckBox("Show 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)): + # 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_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 + ) + 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) + group_box.setLayout(layout) + return group_box + + 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") + return button 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..13f9c225 --- /dev/null +++ b/src/badger/gui/components/bax_visualizer/plotting.py @@ -0,0 +1,206 @@ +"""Matplotlib-based plotting widget for visualizing BAX virtual measurements.""" + +import logging +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 ( + QScrollArea, + QSizePolicy, + QTabWidget, + QVBoxLayout, + QWidget, +) +from xopt.generators.bayesian.bax_generator import BaxGenerator + +from badger.gui.components.extension_utilities import ( + clear_tabs, +) +from badger.utils import BlockSignalsContext +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 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.objective} + + elif algorithm_type == "pathwise_minimize_emittance": + plot_options_dict = { + "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, + } + 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") + + 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] + ) + + 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}") + + ref_point = self.parameters.tab_1.reference_points + + fig, ax = visualize_virtual_measurement_result( + self.generator, + variable_names=selected_variable_names, + 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, + result_keys=results_keys, + ) + return fig, ax + + 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, + ) + return fig, ax + + 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, + ) + return fig, ax + + 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. + + 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] + + 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(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 + + 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") + + def update_second_tab(self) -> None: + content = QWidget() + layout = QVBoxLayout(content) + + for create_plot in (self.create_second_plot, self.create_third_plot): + try: + fig, _ = create_plot() + layout.addWidget(self._build_plot_widget(fig)) + plt.close(fig) + except Exception as e: + logger.error(f"Error creating plot: {e}") + + # 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): + 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/bax_visualizer/ui.py b/src/badger/gui/components/bax_visualizer/ui.py new file mode 100644 index 00000000..e7b1d75f --- /dev/null +++ b/src/badger/gui/components/bax_visualizer/ui.py @@ -0,0 +1,107 @@ +"""UI layout definitions for the BAX visualizer widget.""" + +from typing import TYPE_CHECKING, Optional + + +from badger.gui.components.bax_visualizer.controls import ControlsWidget +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 PyQt5.QtWidgets import QHBoxLayout, QSizePolicy, 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() + + self.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding) + self.setMinimumSize(1250, 600) + + def initialize_ui(self) -> None: + 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, + ) + + main_layout.addWidget(self.plotting_area, stretch=1) + + self.setLayout(main_layout) + + def set_parameters(self, parameters: "Parameters") -> None: + """Point this widget and every child widget at the same parameters object. + + 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 + + def reset_ui(self) -> None: + """Reset the UI to its initial state.""" + self.controls_area.reset_controls_widget() + self.controls_area.update_controls() + + def initialize_reference_table(self) -> None: + """Initialize the reference table in the controls area.""" + + reference_points = get_latest_reference_points( + self.routine.generator.data, self.parameters.variables + ) + + self.parameters.tab_1.reference_points = reference_points + + self.controls_area.populate_reference_table() diff --git a/src/badger/gui/components/bo_visualizer/bo_widget.py b/src/badger/gui/components/bo_visualizer/bo_widget.py index 44502ddd..2538d4dc 100644 --- a/src/badger/gui/components/bo_visualizer/bo_widget.py +++ b/src/badger/gui/components/bo_visualizer/bo_widget.py @@ -7,34 +7,35 @@ routines using a BayesianGenerator. """ +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.bax_generator import BaxGenerator +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, + get_latest_reference_points, signal_logger, to_precision_float, ) 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 xopt.vocs import select_best -from badger.gui.components.analysis_widget import AnalysisWidget - -import logging +from badger.utils import BlockSignalsContext logger = logging.getLogger(__name__) @@ -52,14 +53,15 @@ "variable_2": 1, "variables": [], "reference_points": {}, - "reference_points_range": {}, "include_variable_2": True, } class BOPlotWidget(AnalysisWidget): - generator: BayesianGenerator # type: ignore - parameters: ConfigurableOptions = DEFAULT_PARAMETERS.copy() # type: ignore + generator: BayesianGenerator # pyright: ignore[reportIncompatibleVariableOverride] + parameters: ConfigurableOptions = DEFAULT_PARAMETERS.copy() + df_length: float = float("inf") + initialized: bool = False def __init__( self, @@ -83,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, @@ -125,19 +123,14 @@ 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.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) - 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, @@ -199,12 +192,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")( @@ -212,10 +204,17 @@ 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) def on_set_best_reference_point_clicked(self) -> None: + logger.debug("Setting best reference points") try: self.set_best_reference_points() @@ -228,6 +227,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) @@ -248,48 +260,11 @@ 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 - - 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 + self.parameters = ( # pyright: ignore[reportIncompatibleVariableOverride] + DEFAULT_PARAMETERS.copy() + ) - def on_axis_selection_changed(self): + def on_axis_selection_changed(self) -> None: logger.debug("Axis selection changed") selected_variables: list[str] = [] @@ -364,15 +339,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"]): @@ -397,14 +371,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 @@ -482,21 +454,16 @@ 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 - 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") @@ -521,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") @@ -540,7 +514,7 @@ 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, @@ -571,3 +545,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/plotting_area.py b/src/badger/gui/components/bo_visualizer/plotting_area.py index bd569f04..47a50e3b 100644 --- a/src/badger/gui/components/bo_visualizer/plotting_area.py +++ b/src/badger/gui/components/bo_visualizer/plotting_area.py @@ -1,8 +1,21 @@ """Matplotlib canvas for the BO visualizer. Renders surrogate model plots via Xopt's visualize_generator_model and handles mouse interaction.""" +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, @@ -10,29 +23,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__) @@ -62,7 +57,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 @@ -93,28 +88,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 10daade5..eaf18732 100644 --- a/src/badger/gui/components/bo_visualizer/types.py +++ b/src/badger/gui/components/bo_visualizer/types.py @@ -19,5 +19,4 @@ class ConfigurableOptions(TypedDict): variable_2: int variables: list[str] reference_points: dict[str, float] - reference_points_range: dict[str, tuple[float, float]] 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 1ff15f9f..e8f9ea57 100644 --- a/src/badger/gui/components/bo_visualizer/ui_components.py +++ b/src/badger/gui/components/bo_visualizer/ui_components.py @@ -1,33 +1,30 @@ """Control panel for the BO visualizer — variable selectors, reference point table, grid resolution, and plot option checkboxes.""" +import logging + +import pandas as pd +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, + get_latest_reference_points, ) - -import logging - from badger.utils import BlockSignalsContext - logger = logging.getLogger(__name__) @@ -38,11 +35,12 @@ 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") + 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") @@ -77,7 +75,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 @@ -92,13 +90,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() @@ -127,7 +125,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"], @@ -135,32 +133,32 @@ def initialize_ui_components( def initialize_variables( self, + data: pd.DataFrame | None, configurable_options: ConfigurableOptions, - vocs_variables: dict[str, BaseVariable], - ): + 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 - configurable_options["reference_points"] = { - var: to_precision_float( - (vocs_variables[var].domain[1] - vocs_variables[var].domain[0]) / 2.0 - ) - for var in vocs_variables - } - - def create_reference_inputs(self): + + reference_points = get_latest_reference_points(data, vocs_variables) + + 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) + + 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 @@ -169,12 +167,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)) @@ -195,7 +191,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() @@ -213,7 +209,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") @@ -226,7 +222,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/extension_utilities.py b/src/badger/gui/components/extension_utilities.py index be9e2015..ad54f3c3 100644 --- a/src/badger/gui/components/extension_utilities.py +++ b/src/badger/gui/components/extension_utilities.py @@ -1,18 +1,18 @@ """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 +import pandas as pd 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__) @@ -80,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. @@ -124,7 +136,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 +144,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/extensions_palette.py b/src/badger/gui/components/extensions_palette.py index 96b23a58..1a73d87d 100644 --- a/src/badger/gui/components/extensions_palette.py +++ b/src/badger/gui/components/extensions_palette.py @@ -2,24 +2,30 @@ Pareto Front Viewer) with buttons to open each in its own window.""" 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): """ @@ -27,7 +33,7 @@ class ExtensionsPalette(QMainWindow): Parameters ---------- - run_monitor : RunMonitor + run_monitor : BadgerOptMonitor The run monitor associated with the palette. Attributes @@ -52,13 +58,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. """ @@ -81,9 +87,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) @@ -91,9 +99,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. @@ -105,26 +114,61 @@ 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) + ParetoFrontViewer(routine=self.run_monitor.routine, parent=self) ) - def add_bo_visualizer(self): + def add_bo_visualizer(self) -> None: """ Open the BOVisualizer extension. """ - self.add_child_window_to_monitor(BOVisualizer(routine=self.run_monitor.routine)) + 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, parent=self) + ) + + 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, parent=self) + ) - def add_child_window_to_monitor(self, child_window: AnalysisExtension): + 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/pf_viewer/pf_widget.py b/src/badger/gui/components/pf_viewer/pf_widget.py index f56cb2f1..5a97af69 100644 --- a/src/badger/gui/components/pf_viewer/pf_widget.py +++ b/src/badger/gui/components/pf_viewer/pf_widget.py @@ -2,62 +2,52 @@ 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 ( + ConfigurableOptions, +) +from badger.gui.components.plot_event_handlers import ( + MatplotlibInteractionHandler, +) +from badger.routine import Routine +from badger.utils import BlockSignalsContext logger = logging.getLogger(__name__) @@ -86,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, @@ -115,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, @@ -168,86 +129,60 @@ def update_plots( # Update the last updated time self.last_updated = time.time() - def setup_connections(self): - self.ui["components"]["update"].clicked.connect( + def setup_connections(self) -> None: + 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) @@ -258,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( @@ -285,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 ) @@ -314,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"] @@ -362,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): @@ -471,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) @@ -489,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"] @@ -594,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( @@ -620,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 @@ -630,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 93310373..b4732915 100644 --- a/src/badger/gui/components/pydantic_editor.py +++ b/src/badger/gui/components/pydantic_editor.py @@ -29,8 +29,8 @@ 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 pydantic_core import PydanticUndefined, PydanticUndefinedType +from PyQt5.QtCore import Qt, QTimer, pyqtSignal from PyQt5.QtWidgets import ( QCheckBox, QComboBox, @@ -53,9 +53,13 @@ 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 +from bax_algorithms.emittance import PathwiseMinimizeEmittance +from bax_algorithms.solenoid_alignment import PathwiseSolenoidAlignment + logger = logging.getLogger(__name__) @@ -219,6 +223,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() @@ -246,15 +251,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) @@ -268,7 +279,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 = ( @@ -279,7 +292,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 ) @@ -295,43 +310,52 @@ 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)) + # 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 - 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)) + # 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 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)) @@ -339,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 @@ -465,6 +501,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) @@ -472,7 +526,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( @@ -487,7 +541,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( @@ -615,6 +669,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, @@ -703,9 +792,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: @@ -752,8 +838,18 @@ def set_params_from_generator( fields_to_remove = ["vocs"] + if issubclass(self.model_class, BaxGenerator): + # 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( - 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( @@ -861,6 +957,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) + [ + PathwiseMinimizeEmittance, + PathwiseSolenoidAlignment, + ] else: raise ValueError(f"Field name {field_name} is not recognized.") @@ -950,6 +1051,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 @@ -958,6 +1074,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] @@ -973,13 +1090,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 @@ -1066,12 +1185,27 @@ 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"] + 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, defaults, include_defaults=True + model_class, + fields_to_remove, + defaults, + include_defaults=True, + excluded_fields=self.get_excluded_fields(), ) self._set_params_recurse( @@ -1091,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) diff --git a/src/badger/gui/components/routine_page.py b/src/badger/gui/components/routine_page.py index c7215d0c..a3937bf0 100644 --- a/src/badger/gui/components/routine_page.py +++ b/src/badger/gui/components/routine_page.py @@ -25,41 +25,62 @@ buttons around it. """ -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.vocs 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.vocs import get_local_region, 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, @@ -68,42 +89,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__) @@ -292,8 +299,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 */ @@ -302,8 +308,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) @@ -1817,9 +1822,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( diff --git a/src/badger/gui/components/run_monitor.py b/src/badger/gui/components/run_monitor.py index b5e2e7bf..cb5a281a 100644 --- a/src/badger/gui/components/run_monitor.py +++ b/src/badger/gui/components/run_monitor.py @@ -7,18 +7,16 @@ (BO visualizer, Pareto front viewer) if the user has them open. """ +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, @@ -32,21 +30,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__) @@ -73,7 +73,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) @@ -119,7 +119,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: @@ -197,7 +197,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. @@ -239,7 +239,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. @@ -404,7 +404,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 = {} @@ -437,7 +437,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( @@ -458,7 +458,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() @@ -469,7 +469,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: @@ -485,10 +485,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() @@ -498,14 +498,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 @@ -532,7 +532,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 @@ -696,7 +696,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( @@ -707,10 +707,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: @@ -723,39 +723,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: @@ -773,7 +773,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] @@ -781,7 +781,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", @@ -815,7 +815,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 @@ -823,7 +823,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( @@ -836,7 +836,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." @@ -886,7 +886,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." @@ -905,7 +905,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 @@ -922,7 +922,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] @@ -936,7 +936,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()) @@ -982,7 +982,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 @@ -1017,11 +1017,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 @@ -1034,7 +1034,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: @@ -1054,18 +1054,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") @@ -1078,7 +1078,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, @@ -1092,7 +1092,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 7dfe8f6f..a45da927 100644 --- a/src/badger/gui/pages/home_page.py +++ b/src/badger/gui/pages/home_page.py @@ -15,6 +15,7 @@ import os import traceback from importlib import resources +from typing import TYPE_CHECKING, Optional import numpy as np from pandas import DataFrame @@ -50,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 @@ -60,6 +61,11 @@ 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 + + logger = logging.getLogger(__name__) stylesheet = """ @@ -82,7 +88,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__() @@ -97,7 +103,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" @@ -216,7 +222,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"] @@ -285,21 +291,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() @@ -358,7 +364,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): @@ -371,16 +377,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() @@ -404,7 +410,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) @@ -413,7 +419,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 @@ -465,7 +471,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 @@ -485,6 +493,16 @@ def prepare_run(self, data=None, init_points_flag=True): self.sig_routine_invalid.emit() raise e + # 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": + 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}") + # Add data to routine before saving tmp file if data is not None: # Make sure selected generator is compatible with prior data @@ -526,7 +544,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 @@ -564,7 +582,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, @@ -579,7 +597,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() @@ -589,17 +607,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]) @@ -608,7 +626,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()) @@ -629,7 +647,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 @@ -644,7 +662,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/gui/utils.py b/src/badger/gui/utils.py index ef5f03a9..ee2030e1 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,25 @@ def create_button( return btn -def filter_generator_config(name: str, config: dict[str, Any]): +DEFAULT_ALGORITHM_RESULTS_FILE = "algorithm_results" + + +def build_bax_results_file(folder_id: str, create_dir: bool = False) -> str: + """Build a ``temp//algorithm_results`` prefix for BAX pkl dumps. + + ``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. + """ + 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": filtered_config["adaptive"] = config["adaptive"] @@ -79,6 +118,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_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_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) 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): diff --git a/src/badger/utils.py b/src/badger/utils.py index 29f08387..4cd1362f 100644 --- a/src/badger/utils.py +++ b/src/badger/utils.py @@ -2,22 +2,25 @@ timestamp formatting, value normalization, run filename generation, and platform-specific data directory resolution.""" -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, Any +from typing import TYPE_CHECKING, Any, Iterable, Optional import yaml +from PyQt5.QtWidgets import QLayout, QWidget from badger.errors import BadgerLoadConfigError -from PyQt5.QtWidgets import QWidget, QLayout -from decimal import Decimal, ROUND_CEILING, ROUND_FLOOR +if TYPE_CHECKING: + from badger.routine import Routine + +from decimal import ROUND_CEILING, ROUND_FLOOR, Decimal from gest_api.vocs import ContinuousVariable @@ -183,7 +186,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")