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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion src/ert/config/analysis_module.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,7 +86,7 @@ def correlation_threshold(self, ensemble_size: int) -> float:
Section 2.3 - Localization in the CHOP problem
"""
if self.localization_correlation_threshold is None:
return 3 / math.sqrt(ensemble_size)
return min(1.0, 3 / math.sqrt(ensemble_size))
return self.localization_correlation_threshold


Expand Down
96 changes: 78 additions & 18 deletions src/ert/gui/ertwidgets/analysismoduleedit.py
Original file line number Diff line number Diff line change
@@ -1,49 +1,109 @@
from __future__ import annotations

from typing import TYPE_CHECKING
from collections import defaultdict

from PyQt6.QtCore import QMargins, Qt
from PyQt6.QtWidgets import QHBoxLayout, QPushButton, QWidget
from PyQt6.QtWidgets import (
QDialog,
QDialogButtonBox,
QHBoxLayout,
QPushButton,
QVBoxLayout,
QWidget,
)

from ert.config import ESSettings, LocalizationType, ParameterConfig
from ert.gui.icon_utils import load_icon

from .analysismodulevariablespanel import AnalysisModuleVariablesPanel
from .closabledialog import ClosableDialog

if TYPE_CHECKING:
from ert.config import AnalysisModule


class AnalysisModuleEdit(QWidget):
def __init__(
self,
analysis_module: AnalysisModule,
es_settings: ESSettings,
parameter_config: list[ParameterConfig],
ensemble_size: int,
) -> None:
self.analysis_module = analysis_module
self.ensemble_size = ensemble_size
QWidget.__init__(self)

self._es_settings: ESSettings = es_settings
self._parameter_config: list[ParameterConfig] = parameter_config
self._ensemble_size: int = ensemble_size

layout = QHBoxLayout()

variables_popup_button = QPushButton("Edit")
variables_popup_button.setObjectName("analysis_variables_popup_button")
variables_popup_button.setIcon(load_icon("edit.svg"))
variables_popup_button.clicked.connect(self.showVariablesPopup)
variables_popup_button.clicked.connect(self._show_update_settings_dialog)

layout.addWidget(variables_popup_button, 0, Qt.AlignmentFlag.AlignLeft)
layout.setContentsMargins(QMargins(0, 0, 0, 0))
layout.addStretch()

self.setLayout(layout)

def showVariablesPopup(self) -> None:
variable_dialog = AnalysisModuleVariablesPanel(
self.analysis_module, self.ensemble_size
def _show_update_settings_dialog(self) -> None:
dialog = QDialog(self.parent()) # type: ignore
dialog.setWindowTitle("Update settings")
dialog.setModal(True)
dialog.setWindowFlag(Qt.WindowType.CustomizeWindowHint, True)
dialog.setWindowFlag(Qt.WindowType.WindowContextHelpButtonHint, False)
dialog.setWindowFlag(Qt.WindowType.WindowCloseButtonHint, False)

layout = QVBoxLayout()

update_strategies: dict[str, LocalizationType] = defaultdict(
lambda: LocalizationType.GLOBAL
)
for parameter_config in self._parameter_config:
if parameter_config.update_strategy:
update_strategies[parameter_config.type.upper()] = (
parameter_config.update_strategy
)

correlation_threshold = 1.0
if self._ensemble_size != 0:
correlation_threshold = self._es_settings.correlation_threshold(
self._ensemble_size
)

update_settings_dialog = AnalysisModuleVariablesPanel(
update_strategies=update_strategies,
correlation_threshold=correlation_threshold,
enkf_truncation=self._es_settings.enkf_truncation,
)
Comment thread
frode-aarstad marked this conversation as resolved.
dialog = ClosableDialog(
"Edit variables",
variable_dialog,
self.parent(), # type: ignore

layout.addWidget(update_settings_dialog, stretch=1)

button_box = QDialogButtonBox(
QDialogButtonBox.StandardButton.Save
| QDialogButtonBox.StandardButton.Cancel
)
dialog.exec()
save_button = button_box.button(QDialogButtonBox.StandardButton.Save)
assert save_button is not None
save_button.setAutoDefault(False)
cancel_button = button_box.button(QDialogButtonBox.StandardButton.Cancel)
assert cancel_button is not None
cancel_button.setAutoDefault(False)
button_box.accepted.connect(dialog.accept)
button_box.rejected.connect(dialog.reject)

button_layout = QHBoxLayout()
button_layout.addStretch()
button_layout.addWidget(button_box)
layout.addLayout(button_layout)

dialog.setLayout(layout)
dialog.setFixedSize(450, 300)

if dialog.exec() == QDialog.DialogCode.Accepted:
self._es_settings.localization_correlation_threshold = (
update_settings_dialog.correlation_threshold
)
self._es_settings.enkf_truncation = update_settings_dialog.enkf_truncation
for name, strategy in update_settings_dialog.update_strategies.items():
for parameter_config in self._parameter_config:
if parameter_config.type.upper() == name:
parameter_config.update_strategy = strategy
Comment thread
frode-aarstad marked this conversation as resolved.
132 changes: 114 additions & 18 deletions src/ert/gui/ertwidgets/analysismodulevariablespanel.py
Original file line number Diff line number Diff line change
@@ -1,55 +1,157 @@
from __future__ import annotations

from functools import partial
from typing import cast

from annotated_types import Ge, Gt, Le
from PyQt6.QtCore import Qt
from PyQt6.QtGui import QStandardItem, QStandardItemModel
from PyQt6.QtWidgets import (
QComboBox,
QDoubleSpinBox,
QFormLayout,
QLabel,
QWidget,
)

from ert.config import AnalysisModule
from ert.config import AnalysisModule, LocalizationType


class _LocalizationTypeModel(QStandardItemModel):
def __init__(self, exclude: set[LocalizationType] | None = None) -> None:
super().__init__()

type_set = set(LocalizationType)
if exclude is not None:
type_set -= exclude

for localization_type in sorted(type_set, key=lambda lt: lt.name):
item = QStandardItem(localization_type.name)
item.setData(localization_type, Qt.ItemDataRole.UserRole)
self.appendRow(item)
Comment thread
frode-aarstad marked this conversation as resolved.


class AnalysisModuleVariablesPanel(QWidget):
def __init__(self, analysis_module: AnalysisModule, ensemble_size: int) -> None:
def __init__(
self,
update_strategies: dict[str, LocalizationType],
correlation_threshold: float,
enkf_truncation: float,
) -> None:
QWidget.__init__(self)
self.analysis_module = analysis_module

layout = QFormLayout()
self._update_strategies = update_strategies
self._correlation_threshold = correlation_threshold
self._enkf_truncation = enkf_truncation

layout = QFormLayout()
self.blockSignals(True)

layout.addRow(
QLabel("<b>Select the localization method for each parameter type</b>")
)

gen_kw_combobox = QComboBox(self)
gen_kw_combobox.setModel(
_LocalizationTypeModel(exclude={LocalizationType.DISTANCE})
)
Comment thread
frode-aarstad marked this conversation as resolved.
gen_kw_combobox.setCurrentIndex(
self._find_correct_index(gen_kw_combobox, "GEN_KW")
)
gen_kw_combobox.currentIndexChanged.connect(
lambda index: self._update_strategies.__setitem__(
"GEN_KW",
gen_kw_combobox.itemData(index, Qt.ItemDataRole.UserRole),
)
)
layout.addRow("GEN_KW", gen_kw_combobox)

field_combobox = QComboBox(self)
field_combobox.setModel(_LocalizationTypeModel())
field_combobox.setCurrentIndex(
self._find_correct_index(field_combobox, "FIELD")
)
field_combobox.currentIndexChanged.connect(
lambda index: self._update_strategies.__setitem__(
"FIELD",
field_combobox.itemData(index, Qt.ItemDataRole.UserRole),
)
)
layout.addRow("FIELD", field_combobox)

surface_combobox = QComboBox(self)
surface_combobox.setModel(_LocalizationTypeModel())
surface_combobox.setCurrentIndex(
self._find_correct_index(surface_combobox, "SURFACE")
)
surface_combobox.currentIndexChanged.connect(
lambda index: self._update_strategies.__setitem__(
"SURFACE",
surface_combobox.itemData(index, Qt.ItemDataRole.UserRole),
)
)
layout.addRow("SURFACE", surface_combobox)

layout.addRow(QLabel("<b>General settings</b>"))

var_name = "enkf_truncation"
metadata = AnalysisModule.model_fields[var_name]
self.truncation_spinner = self.createDoubleSpinBox(
self.truncation_spinner = self._create_double_spinbox(
var_name,
analysis_module.enkf_truncation,
self._enkf_truncation,
cast(float, next(v for v in metadata.metadata if isinstance(v, Gt)).gt)
+ 0.001,
cast(float, next(v for v in metadata.metadata if isinstance(v, Le)).le),
0.01,
)
self.truncation_spinner.valueChanged.connect(
lambda value: setattr(self, "_enkf_truncation", value)
)

layout.addRow("Singular value truncation", self.truncation_spinner)

var_name = "localization_correlation_threshold"
metadata = AnalysisModule.model_fields[var_name]
self.local_spinner = self.createDoubleSpinBox(
self.threshold_spinner = self._create_double_spinbox(
var_name,
analysis_module.correlation_threshold(ensemble_size),
self._correlation_threshold,
cast(float, next(v for v in metadata.metadata if isinstance(v, Ge)).ge),
cast(float, next(v for v in metadata.metadata if isinstance(v, Le)).le),
0.1,
)
self.local_spinner.setObjectName("localization_correlation_threshold")
layout.addRow("Adaptive localization correlation threshold", self.local_spinner)
self.threshold_spinner.setObjectName("localization_correlation_threshold")
self.threshold_spinner.valueChanged.connect(
lambda value: setattr(self, "_correlation_threshold", value)
)

layout.addRow(
"Adaptive localization correlation threshold", self.threshold_spinner
)

self.setLayout(layout)
self.blockSignals(False)

def createDoubleSpinBox(
@property
def update_strategies(self) -> dict[str, LocalizationType]:
return self._update_strategies

@property
def correlation_threshold(self) -> float:
return self._correlation_threshold

@property
def enkf_truncation(self) -> float:
return self._enkf_truncation

def _find_correct_index(self, combobox: QComboBox, type_name: str) -> int:
if type_name in self._update_strategies:
localization_type = self._update_strategies[type_name]
if (
index := combobox.findData(localization_type, Qt.ItemDataRole.UserRole)
) != -1:
return index
return combobox.findData(LocalizationType.GLOBAL, Qt.ItemDataRole.UserRole)

def _create_double_spinbox(
self,
variable_name: str,
variable_value: float,
Expand All @@ -61,16 +163,10 @@ def createDoubleSpinBox(
spinner.setDecimals(6)
spinner.setFixedWidth(180)
spinner.setObjectName(variable_name)

spinner.setRange(
min_value,
max_value,
)

spinner.setSingleStep(step_length)
spinner.setValue(variable_value)
spinner.valueChanged.connect(partial(self.valueChangedSpinner, variable_name))
return spinner

def valueChangedSpinner(self, name: str, value: float) -> None:
setattr(self.analysis_module, name, value)
7 changes: 4 additions & 3 deletions src/ert/gui/experiments/ensemble_smoother_panel.py
Original file line number Diff line number Diff line change
Expand Up @@ -95,14 +95,15 @@ def __init__(
layout.addRow("Ensemble format:", self._ensemble_format_field)

self._analysis_module_edit = AnalysisModuleEdit(
analysis_config.es_settings,
sum(
es_settings=analysis_config.es_settings,
parameter_config=parameter_configuration,
ensemble_size=sum(
active_realizations
), # only use active realizations for setting threshold
)
self._analysis_module_edit.setObjectName("ensemble_smoother_edit")
layout.addRow("Analysis module:", self._analysis_module_edit)

layout.addRow("Update settings:", self._analysis_module_edit)
self._active_realizations_field = StringBox(
ActiveRealizationsModel(len(active_realizations)), # type: ignore
"config/experiment/active_realizations",
Expand Down
7 changes: 6 additions & 1 deletion src/ert/gui/experiments/experiment_panel.py
Original file line number Diff line number Diff line change
Expand Up @@ -242,7 +242,12 @@ def __init__(
experiment_type_valid,
)
self.addExperimentConfigPanel(
ManualUpdatePanel(run_path, notifier, analysis_config),
ManualUpdatePanel(
run_path,
notifier,
analysis_config,
config.ensemble_config.parameter_configuration,
),
experiment_type_valid,
)

Expand Down
Loading