diff --git a/loopstructural/gui/visualisation/cross_section_utils.py b/loopstructural/gui/visualisation/cross_section_utils.py index 68dd5aa..43af389 100644 --- a/loopstructural/gui/visualisation/cross_section_utils.py +++ b/loopstructural/gui/visualisation/cross_section_utils.py @@ -69,3 +69,30 @@ def build_line_extrusion_mesh( yy = np.tile(xy[:, 1][:, None], (1, z_resolution)) zz = np.tile(z[None, :], (resolution, 1)) return pv.StructuredGrid(xx, yy, zz) + + +def build_block_model_mesh(origin, maximum, ncells) -> pv.ImageData: + """Build a regular voxel grid (block model) that fills a bounding box. + + Parameters + ---------- + origin, maximum : array_like + (3,) opposite corners of the box, in the model's coordinate system. + ncells : array_like + (3,) number of cells (blocks) along x, y and z. + + Returns + ------- + pv.ImageData + A grid with `prod(ncells)` cells. Callers evaluate the model at + `mesh.cell_centers().points` and store the result as cell data. + """ + origin = np.asarray(origin, dtype=float) + maximum = np.asarray(maximum, dtype=float) + ncells = np.maximum(np.asarray(ncells, dtype=int), 1) + spacing = (maximum - origin) / ncells + return pv.ImageData( + dimensions=tuple(int(n) for n in ncells + 1), + spacing=tuple(float(s) for s in spacing), + origin=tuple(float(o) for o in origin), + ) diff --git a/loopstructural/gui/visualisation/feature_list_widget.py b/loopstructural/gui/visualisation/feature_list_widget.py index b31d024..5bbba5f 100644 --- a/loopstructural/gui/visualisation/feature_list_widget.py +++ b/loopstructural/gui/visualisation/feature_list_widget.py @@ -38,7 +38,11 @@ from ..background_task import finish_background_task, start_background_task from ..compatibility import configure_layer_combo -from .cross_section_utils import build_line_extrusion_mesh, build_plane_mesh +from .cross_section_utils import ( + build_block_model_mesh, + build_line_extrusion_mesh, + build_plane_mesh, +) from .mesh_scalar_utils import stratigraphic_ids_to_rgb logger = logging.getLogger(__name__) @@ -88,6 +92,7 @@ def __init__(self, parent=None, *, model_manager=None, viewer=None, data_manager self._topography_progress = None self._build_cross_section_controls() + self._build_block_model_controls() # A single row of icon-only actions, in workflow order, replaces the # previous stack of full-width text buttons. @@ -97,16 +102,37 @@ def __init__(self, parent=None, *, model_manager=None, viewer=None, data_manager actionsRow.addWidget(self.addStratigraphicSurfacesButton) actionsRow.addWidget(self.addTopographyButton) actionsRow.addWidget(self.crossSectionButton) + actionsRow.addWidget(self.blockModelButton) actionsRow.addStretch(1) self.mainLayout.addLayout(actionsRow) self.mainLayout.addWidget(self.colourTopographyByStratigraphyCheckBox) + # Objects in the viewer are not rebuilt automatically when the model + # changes (that can re-solve the whole model after each small edit). + # They are marked out of date, and this button rebuilds them. + self.updateObjectsButton = QPushButton(self) + self.updateObjectsButton.setIcon(QgsApplication.getThemeIcon("mActionRefresh.svg")) + self.updateObjectsButton.clicked.connect(self.update_out_of_date_objects) + self.mainLayout.addWidget(self.updateObjectsButton) + self._update_objects_thread = None + self._update_objects_worker = None + self._update_objects_progress = None + if self.viewer is not None: + self.viewer.outOfDateChanged.connect(self._refresh_update_objects_button) + self._refresh_update_objects_button() + # background task handles shared by the plane and line cross-section # actions (only one can run at a time) self._cross_section_thread = None self._cross_section_worker = None self._cross_section_progress = None self._pending_cross_section_name = None + + # background task handles for the block model + self._block_model_thread = None + self._block_model_worker = None + self._block_model_progress = None + self._pending_block_model_name = None # Whether the user has hand-edited the plane's origin/normal/size -- # while False, `update_feature_list` keeps re-syncing those fields to # the model's current bounding box (see @@ -128,8 +154,8 @@ def __init__(self, parent=None, *, model_manager=None, viewer=None, data_manager self._disp_update = self.model_manager.attach( lambda _obs, _event, *a, **k: self.update_feature_list(), 'model_updated' ) - # also listen for model and feature updates so visualisation can refresh - # forward event and args into the handler so it can act on specific surfaces + # also listen for model and feature updates so the viewer + # objects built from the model can be marked out of date self._disp_feature = self.model_manager.attach( lambda _obs, _event, *a, **k: self._on_model_update(_event, *a), 'model_updated' ) @@ -281,6 +307,79 @@ def make_coord_spinbox(default=0.0): closeButton.clicked.connect(self.crossSectionDialog.close) dialogLayout.addWidget(closeButton) + def _build_block_model_controls(self): + """Build the "Block Model" icon button (added to the shared actions + row in __init__), which opens a dialog to set the number of blocks + along each axis. The block model fills the model bounding box and is + coloured by the stratigraphic column, like the cross sections. + """ + self.blockModelButton = self._make_custom_icon_tool_button( + "block_model.svg", "Block Model..." + ) + self.blockModelButton.clicked.connect(self._show_block_model_dialog) + self._block_model_resolution_initialised = False + + self.blockModelDialog = QDialog(self) + self.blockModelDialog.setWindowTitle("Block Model") + dialogLayout = QVBoxLayout(self.blockModelDialog) + dialogLayout.addWidget(QLabel("Blocks fill the model bounding box")) + form = QFormLayout() + + def make_cells_spinbox(): + box = QSpinBox(self) + box.setRange(1, 1000) + box.setValue(50) + return box + + self.blockModelNxSpinBox = make_cells_spinbox() + self.blockModelNySpinBox = make_cells_spinbox() + self.blockModelNzSpinBox = make_cells_spinbox() + form.addRow( + "Blocks (x, y, z)", + self._hbox( + self.blockModelNxSpinBox, self.blockModelNySpinBox, self.blockModelNzSpinBox + ), + ) + self.blockModelUseModelResolutionButton = QPushButton("Use Model Resolution", self) + self.blockModelUseModelResolutionButton.setToolTip( + "Set the number of blocks to the model's interpolation grid" + ) + self.blockModelUseModelResolutionButton.clicked.connect(self._reset_block_model_resolution) + form.addRow("", self.blockModelUseModelResolutionButton) + dialogLayout.addLayout(form) + + self.addBlockModelButton = QPushButton("Add Block Model", self) + self.addBlockModelButton.clicked.connect(self.add_block_model) + dialogLayout.addWidget(self.addBlockModelButton) + + closeButton = QPushButton("Close", self.blockModelDialog) + closeButton.clicked.connect(self.blockModelDialog.close) + dialogLayout.addWidget(closeButton) + + def _show_block_model_dialog(self): + # Start from the model's own resolution the first time only, so the + # user's block counts are kept between openings of the dialog. + if not self._block_model_resolution_initialised: + self._reset_block_model_resolution() + self._block_model_resolution_initialised = True + self.blockModelDialog.show() + self.blockModelDialog.raise_() + self.blockModelDialog.activateWindow() + + def _reset_block_model_resolution(self): + if not self.model_manager or self.model_manager.model is None: + return + try: + nsteps = np.asarray(self.model_manager.model.bounding_box.nsteps, dtype=int) + except Exception: + logger.info("Model bounding box has no resolution.") + return + for box, value in zip( + (self.blockModelNxSpinBox, self.blockModelNySpinBox, self.blockModelNzSpinBox), + nsteps, + ): + box.setValue(int(value)) + def _show_cross_section_dialog(self): self.crossSectionDialog.show() self.crossSectionDialog.raise_() @@ -373,10 +472,12 @@ def contextMenuEvent(self, event): elif action == add_data_action: self.add_data(feature_name) + def _build_scalar_field(self, feature_name): + return self.model_manager.model[feature_name].scalar_field().vtk() + def add_scalar_field(self, feature_name): - scalar_field = self.model_manager.model[feature_name].scalar_field() self.viewer.add_mesh_object( - scalar_field.vtk(), + self._build_scalar_field(feature_name), name=f'{feature_name}_scalar_field', source_feature=feature_name, source_type='feature_scalar', @@ -410,11 +511,20 @@ def add_surface(self, feature_name): isovalue=isovalue, ) - def add_vector_field(self, feature_name): + def _build_feature_surface(self, feature_name, isovalue): + feature = self.model_manager.model[feature_name] + surfaces = feature.surfaces(isovalue) if isovalue is not None else feature.surfaces() + if not surfaces: + raise ValueError(f"Feature '{feature_name}' has no surface at {isovalue}") + return surfaces[0].vtk() + + def _build_vector_field(self, feature_name): vector_field = self.model_manager.model[feature_name].vector_field() - scale = self._get_vector_scale() + return vector_field.vtk(scale=self._get_vector_scale()) + + def add_vector_field(self, feature_name): self.viewer.add_mesh_object( - vector_field.vtk(scale=scale), + self._build_vector_field(feature_name), name=f'{feature_name}_vector_field', source_feature=feature_name, source_type='feature_vector', @@ -445,23 +555,18 @@ def _get_fold(self, feature_name): fold = getattr(getattr(feature, 'builder', None), 'fold', None) return fold - def add_fold_constraint(self, feature_name, constraint): - """Add the vectors of one fold constraint of a folded feature to the viewer. + def _build_fold_constraint(self, feature_name, constraint): + """Return the vectors of one fold constraint of a folded feature as + a mesh, or None if the feature is not folded or the vectors are not + defined. The vectors are evaluated on the model grid, in the same way as the interpolator evaluates them on the element barycentres. - - Parameters - ---------- - feature_name : str - Name of the folded feature. - constraint : str - One of 'direction', 'axis' or 'norm'. """ fold = self._get_fold(feature_name) if fold is None: logger.info(f"Feature {feature_name} is not folded") - return + return None feature = self.model_manager.model[feature_name] # make sure the fold rotation angles are fitted feature.builder.up_to_date() @@ -474,53 +579,75 @@ def add_fold_constraint(self, feature_name, constraint): logger.warning( f"Fold {constraint} vectors do not match the grid points ({vectors.shape} != {points.shape})" ) - return + return None length = np.linalg.norm(vectors, axis=1) mask = np.all(np.isfinite(vectors), axis=1) & (length > 0) if not np.any(mask): logger.warning(f"Fold {constraint} vectors for {feature_name} are not defined") - return + return None vectors = vectors[mask] / length[mask, None] vector_points = VectorPoints(points[mask], vectors, f'{feature_name}_fold_{constraint}') + return vector_points.vtk(scale=self._get_vector_scale()) + + def add_fold_constraint(self, feature_name, constraint): + """Add the vectors of one fold constraint of a folded feature to the viewer. + + Parameters + ---------- + feature_name : str + Name of the folded feature. + constraint : str + One of 'direction', 'axis' or 'norm'. + """ + mesh = self._build_fold_constraint(feature_name, constraint) + if mesh is None: + return self.viewer.add_mesh_object( - vector_points.vtk(scale=self._get_vector_scale()), + mesh, name=f'{feature_name}_fold_{constraint}', color=self.FOLD_CONSTRAINT_COLOURS[constraint], source_feature=feature_name, source_type=f'fold_constraint_{constraint}', ) - def add_data(self, feature_name): - data = self.model_manager.model[feature_name].get_data() - for d in data: + def _build_data_meshes(self, feature_name): + """Return (name, mesh, source_type) for each data set of a feature.""" + meshes = [] + for d in self.model_manager.model[feature_name].get_data(): d.locations = self.model_manager.model.rescale(d.locations) if issubclass(type(d), VectorPoints): - scale = self._get_vector_scale() # tolerance is None means all points are shown - self.viewer.add_mesh_object( - d.vtk(scale=scale, tolerance=None), - name=f'{feature_name}_{d.name}_points', - source_feature=feature_name, - source_type='feature_points', + meshes.append( + ( + f'{feature_name}_{d.name}_points', + d.vtk(scale=self._get_vector_scale(), tolerance=None), + 'feature_points', + ) ) else: - self.viewer.add_mesh_object( - d.vtk(), - name=f'{feature_name}_{d.name}', - source_feature=feature_name, - source_type='feature_data', - ) + meshes.append((f'{feature_name}_{d.name}', d.vtk(), 'feature_data')) + return meshes + + def add_data(self, feature_name): + for name, mesh, source_type in self._build_data_meshes(feature_name): + self.viewer.add_mesh_object( + mesh, name=name, source_feature=feature_name, source_type=source_type + ) logger.info(f"Adding data to feature: {feature_name}") + def _build_bounding_box(self): + return self.model_manager.model.bounding_box.vtk().outline() + def add_model_bounding_box(self): if not self.model_manager: logger.info("Model manager is not set.") return - bb = self.model_manager.model.bounding_box.vtk().outline() self.viewer.add_mesh_object( - bb, name='model_bounding_box', source_feature='__model__', source_type='bounding_box' + self._build_bounding_box(), + name='model_bounding_box', + source_feature='__model__', + source_type='bounding_box', ) - # Logic for adding model bounding box logger.info("Adding model bounding box...") def add_fault_surfaces(self): @@ -539,6 +666,18 @@ def add_fault_surfaces(self): ) logger.info("Adding fault surfaces...") + def _find_fault_surface(self, name): + for surface in self.model_manager.model.get_fault_surfaces(): + if str(surface.name) == str(name): + return surface + raise ValueError(f"Fault surface '{name}' is not in the model") + + def _find_stratigraphic_surface(self, name): + for surface in self.model_manager.model.get_stratigraphic_surfaces(): + if str(surface.name) == str(name): + return surface + raise ValueError(f"Stratigraphic surface '{name}' is not in the model") + def add_stratigraphic_surfaces(self): if not self.model_manager: logger.info("Model manager is not set.") @@ -619,6 +758,7 @@ def _on_topography_grid_finished(self, result): cmap='terrain', show_scalar_bar=True, source_type='topography_surface', + metadata={'coloured': False}, ) self.colourTopographyByStratigraphyCheckBox.setEnabled(True) if self.colourTopographyByStratigraphyCheckBox.isChecked(): @@ -649,6 +789,7 @@ def _on_colour_topography_toggled(self, checked): cmap='terrain', show_scalar_bar=True, source_type='topography_surface', + metadata={'coloured': False}, ) def _colour_topography_surface(self): @@ -704,11 +845,12 @@ def _on_topography_colour_finished(self, result): ids, colours = result mesh = self.viewer.meshes['topography_surface']['mesh'] - rgb = stratigraphic_ids_to_rgb(ids, colours) + self._set_stratigraphy_arrays(mesh.point_data, ids, colours) + metadata = {'coloured': True, **self._stratigraphy_metadata(colours)} self.viewer.add_mesh_object( mesh, name='topography_surface', - scalars=rgb, + scalars='colour', rgb=True, show_scalar_bar=False, # Directional lighting shades the same colour differently depending @@ -718,6 +860,7 @@ def _on_topography_colour_finished(self, result): # keeps it a true, direct colour-for-colour match. lighting=False, source_type='topography_surface', + metadata=metadata, ) logger.info("Coloured topography surface by stratigraphic column.") @@ -913,11 +1056,11 @@ def _on_cross_section_progress(self, message): def _add_cross_section_mesh(self, result, source_type): mesh, ids, colours = result - rgb = stratigraphic_ids_to_rgb(ids, colours) + self._set_stratigraphy_arrays(mesh.point_data, ids, colours) self.viewer.add_mesh_object( mesh, name=self._pending_cross_section_name, - scalars=rgb, + scalars='colour', rgb=True, show_scalar_bar=False, # see `_on_topography_colour_finished` for why cross sections use @@ -925,6 +1068,7 @@ def _add_cross_section_mesh(self, result, source_type): # of how the section plane/line happens to be oriented. lighting=False, source_type=source_type, + metadata=self._stratigraphy_metadata(colours), ) logger.info(f"Added cross section '{self._pending_cross_section_name}'.") @@ -950,244 +1094,389 @@ def _on_cross_section_error(self, traceback_text): self.addLineCrossSectionButton.setEnabled(True) logger.error(f"Failed to build cross section: {traceback_text}") - def _on_model_update(self, event: str, *args): - """Called when the underlying model_manager notifies observers. - - We remove any meshes that were created from model features and re-add - them from the current model so visualisation follows model changes. + def add_block_model(self): + """Fill the model bounding box with blocks and colour each block by + the stratigraphic unit at its centre. - If the notification is for a specific feature (event == 'feature_updated') - and an isovalue is provided (either as second arg or stored in viewer - metadata), only the matching surface will be re-added. For generic - 'model_updated' notifications the previous behaviour (re-add all - affected feature representations) is preserved. + Evaluating the model at every block centre can be slow, so this runs + on a background thread (see `add_topography_surface`). """ + if not self.model_manager: + logger.info("Model manager is not set.") + return + if self.model_manager.model is None: + logger.info("No model available to build a block model.") + return + missing = self.model_manager.get_units_without_colour() + if missing: + QMessageBox.warning( + self, + "Missing unit colour", + "Cannot colour the block model by stratigraphy. These units have no " + "valid colour in the stratigraphic column:\n\n" + + "\n".join(missing) + + "\n\nSet a colour for each unit and try again.", + ) + return - # Prefer the DebugManager for logging when available (it forwards to - # the plugin/toolbelt logger and handles debug mode). Fall back to the - # module logger if no debug manager is present. - def _log(msg, level=0): - try: - dbg = None - if getattr(self, 'model_manager', None) is not None: - dbg = getattr(self.model_manager, '_debug_manager', None) - if dbg is not None and hasattr(dbg, 'log'): - # DebugManager.log expects message and log_level keyword - dbg.log(str(msg), log_level=level) - else: - logger.info(str(msg)) - except Exception: - try: - logger.info(str(msg)) - except Exception: - pass + bb = self.model_manager.model.bounding_box + origin = np.asarray(bb.origin, dtype=float) + maximum = np.asarray(bb.maximum, dtype=float) + ncells = ( + self.blockModelNxSpinBox.value(), + self.blockModelNySpinBox.value(), + self.blockModelNzSpinBox.value(), + ) + + def target(progress_callback): + progress_callback("Building block model grid...") + mesh = build_block_model_mesh(origin, maximum, ncells) + progress_callback("Evaluating stratigraphy on block model...") + ids = self.model_manager.evaluate_stratigraphy_on_points(mesh.cell_centers().points) + colours = self.model_manager.get_stratigraphic_column_colours() + return mesh, ids, colours + + self._pending_block_model_name = self._unique_cross_section_name('block_model') + self.addBlockModelButton.setEnabled(False) + self._block_model_thread, self._block_model_worker, self._block_model_progress = ( + start_background_task( + self, + target, + title="Block Model", + initial_label="Building block model grid...", + on_progress=self._on_block_model_progress, + on_finished=self._on_block_model_finished, + on_error=self._on_block_model_error, + ) + ) - _log(f"Model update event received: {event} with args: {args}") + def _on_block_model_progress(self, message): try: - _log([f"Mesh: {name}, Meta: {meta}" for name, meta in self.viewer.meshes.items()]) + self._block_model_progress.setLabelText(message) except Exception: - _log("Model update: failed to enumerate viewer meshes") + pass + + def _on_block_model_finished(self, result): + finish_background_task( + self._block_model_thread, self._block_model_worker, self._block_model_progress + ) + self.addBlockModelButton.setEnabled(True) - if not self.model_manager or not self.viewer: + mesh, ids, colours = result + self._set_stratigraphy_arrays(mesh.cell_data, ids, colours) + self.viewer.add_mesh_object( + mesh, + name=self._pending_block_model_name, + scalars='colour', + rgb=True, + show_scalar_bar=False, + show_edges=False, + source_type='block_model', + metadata={ + 'ncells': [int(n) for n in np.asarray(mesh.dimensions) - 1], + **self._stratigraphy_metadata(colours), + }, + ) + logger.info(f"Added block model '{self._pending_block_model_name}'.") + + def _on_block_model_error(self, traceback_text): + finish_background_task( + self._block_model_thread, self._block_model_worker, self._block_model_progress + ) + self.addBlockModelButton.setEnabled(True) + logger.error(f"Failed to build block model: {traceback_text}") + + @staticmethod + def _set_stratigraphy_arrays(data, ids, colours): + """Store the unit ids and their colours as the 'stratigraphy' and + 'colour' arrays of `data` (a mesh's point_data or cell_data). + + Named arrays (not an RGB array passed straight to the viewer) let the + object properties panel use the unit ids, and let + `update_out_of_date_objects` add the object again with the same + viewer settings. + """ + data['stratigraphy'] = np.asarray(ids) + data['colour'] = stratigraphic_ids_to_rgb(ids, colours) + + def _stratigraphy_metadata(self, colours): + """Unit names and colours, indexed by unit id, for the unit check + boxes of the object properties panel.""" + return { + 'unit_names': list(self.model_manager.get_stratigraphic_unit_names()), + 'unit_colours': list(colours), + } + + def _colour_by_stratigraphy(self, points, data, metadata): + """Colour a mesh by the stratigraphic unit at `points`, and store + the unit names and colours in `metadata`.""" + ids = self.model_manager.evaluate_stratigraphy_on_points(points) + colours = self.model_manager.get_stratigraphic_column_colours() + self._set_stratigraphy_arrays(data, ids, colours) + metadata.update(self._stratigraphy_metadata(colours)) + + # Viewer objects that `_rebuild_object` can build again from the model, + # by `source_type` (fold constraints use the 'fold_constraint_' prefix). + REBUILDABLE_SOURCE_TYPES = { + 'feature_scalar', + 'feature_surface', + 'feature_vector', + 'feature_vectors', + 'feature_points', + 'feature_data', + 'bounding_box', + 'fault_surface', + 'stratigraphic_surface', + 'cross_section_plane', + 'cross_section_line', + 'block_model', + 'topography_surface', + } + # Source types that are coloured by the stratigraphic column + STRATIGRAPHY_COLOURED_SOURCE_TYPES = { + 'cross_section_plane', + 'cross_section_line', + 'block_model', + } + # Source types that are built from one feature of the model + FEATURE_SOURCE_TYPES = { + 'feature_scalar', + 'feature_surface', + 'feature_vector', + 'feature_vectors', + 'feature_points', + 'feature_data', + } + + def _is_rebuildable(self, meta) -> bool: + source_type = meta.get('source_type') or '' + return source_type in self.REBUILDABLE_SOURCE_TYPES or source_type.startswith( + 'fold_constraint_' + ) + + def _uses_stratigraphy_colours(self, spec) -> bool: + if spec['source_type'] in self.STRATIGRAPHY_COLOURED_SOURCE_TYPES: + return True + return spec['source_type'] == 'topography_surface' and bool( + spec['metadata'].get('coloured') + ) + + def _on_model_update(self, event: str, *args): + """Mark the viewer objects built from the model as out of date. + + The objects are not rebuilt here: rebuilding can solve the model + again, which is slow after each small edit. The user rebuilds them + with the update button (see `update_out_of_date_objects`). + """ + if not self.viewer: return if event not in ('model_updated', 'feature_updated'): return - feature_name = None - if event == 'feature_updated' and len(args) >= 1: - feature_name = args[0] + names = [ + name for name, meta in list(self.viewer.meshes.items()) if self._is_rebuildable(meta) + ] + self.viewer.set_out_of_date(names, True) + + def _refresh_update_objects_button(self): + count = len(self.viewer.out_of_date_objects()) if self.viewer is not None else 0 + if self._update_objects_thread is not None: + self.updateObjectsButton.setText("Updating Viewer Objects...") + self.updateObjectsButton.setEnabled(False) + elif count: + noun = "Object" if count == 1 else "Objects" + self.updateObjectsButton.setText(f"Update {count} Out-of-Date {noun}") + self.updateObjectsButton.setToolTip( + "The model changed after these objects were added to the viewer. " + "Build them again from the current model." + ) + self.updateObjectsButton.setEnabled(True) + else: + self.updateObjectsButton.setText("Viewer Objects Up to Date") + self.updateObjectsButton.setToolTip("") + self.updateObjectsButton.setEnabled(False) - # If the model was reset (None) or features referenced by viewer meshes - # no longer exist in the current model, remove the linkage from those - # meshes so they are not treated as feature-driven on subsequent updates. - try: - try: - current_features = {f.name for f in self.model_manager.features()} - except Exception: - current_features = set() - - # If the model is None or a feature referenced by a mesh is missing, - # decouple that mesh from the feature so it remains visible but won't - # be auto-updated or re-added when the model changes. - for mesh_name, meta in list(self.viewer.meshes.items()): - sf = meta.get('source_feature', None) - if sf is None: - continue - if self.model_manager.model is None or sf not in current_features: - _log(f"Decoupling mesh '{mesh_name}' from missing feature '{sf}'") - meta.pop('source_feature', None) - meta.pop('source_type', None) - meta.pop('isovalue', None) - # mark as decoupled so other logic can detect it if needed - meta['decoupled_from_feature'] = True - except Exception: - _log('Failed while decoupling meshes from features') - - # Build a set of features that currently have viewer meshes - affected_features = set() - for _, meta in list(self.viewer.meshes.items()): - if feature_name is not None: - if meta.get('source_feature', None) == feature_name: - affected_features.add(feature_name) - _log(f"Updating visualisation for feature: {feature_name}") - continue - - sf = meta.get('source_feature', None) - - if sf is not None: - affected_features.add(sf) - _log(f"Affected features to update: {affected_features}") - # For each affected feature, only update existing meshes tied to that feature - for feature_name in affected_features: - # collect mesh names that belong to this feature (snapshot to avoid mutation while iterating) - meshes_for_feature = [ - name - for name, meta in list(self.viewer.meshes.items()) - if meta.get('source_feature') == feature_name - ] - _log(f"Re-adding meshes for feature: {feature_name}: {meshes_for_feature}") - - for mesh_name in meshes_for_feature: - meta = self.viewer.meshes.get(mesh_name, {}) - source_type = meta.get('source_type') - kwargs = meta.get('kwargs', {}) or {} - isovalue = meta.get('isovalue', None) - - # remove existing actor/entry so add_mesh_object can recreate with same name - try: - self.viewer.remove_object(mesh_name) - _log(f"Removed existing mesh: {mesh_name}") - except Exception: - _log(f"Failed to remove existing mesh: {mesh_name}") + def update_out_of_date_objects(self): + """Build all out-of-date viewer objects again from the current model. - try: - # Surfaces associated with individual features - if source_type == 'feature_surface': - surfaces = [] - try: - if isovalue is not None: - surfaces = self.model_manager.model[feature_name].surfaces(isovalue) - else: - surfaces = self.model_manager.model[feature_name].surfaces() - - if surfaces: - add_name = mesh_name - _log( - f"Re-adding surface for feature: {feature_name} with isovalue: {isovalue} and {kwargs}" - ) - kwargs['isovalue'] = isovalue - - self.viewer.add_mesh_object( - surfaces[0].vtk(), - name=add_name, - source_feature=feature_name, - source_type='feature_surface', - isovalue=isovalue, - **kwargs, - ) - continue - except Exception as e: - _log( - f"Failed to find matching surface for feature: {feature_name} with isovalue: {isovalue}, trying all surfaces. Error: {e}" - ) - - # Fault surfaces (added via add_fault_surfaces) - if source_type == 'fault_surface': - try: - fault_surfaces = self.model_manager.model.get_fault_surfaces() - match = next( - (s for s in fault_surfaces if str(s.name) == str(feature_name)), - None, - ) - if match is not None: - _log(f"Re-adding fault surface for: {feature_name}") - self.viewer.add_mesh_object( - match.vtk(), - name=mesh_name, - source_feature=feature_name, - source_type='fault_surface', - isovalue=meta.get('isovalue', 0.0), - **kwargs, - ) - continue - except Exception as e: - _log(f"Failed to re-add fault surface for {feature_name}: {e}") - - # Stratigraphic surfaces (added via add_stratigraphic_surfaces) - if source_type == 'stratigraphic_surface': - try: - strat_surfaces = self.model_manager.model.get_stratigraphic_surfaces() - match = next( - (s for s in strat_surfaces if str(s.name) == str(feature_name)), - None, - ) - if match is not None: - _log(f"Re-adding stratigraphic surface for: {feature_name}") - kwargs['color'] = getattr(match, 'colour', None) - - self.viewer.add_mesh_object( - match.vtk(), - name=mesh_name, - source_feature=feature_name, - source_type='stratigraphic_surface', - **kwargs, - ) - continue - except Exception as e: - _log(f"Failed to re-add stratigraphic surface for {feature_name}: {e}") - - # Vectors, points, scalar fields and other feature related objects - if source_type == 'feature_vector' or source_type == 'feature_vectors': - try: - self.add_vector_field(feature_name) - continue - except Exception as e: - _log(f"Failed to re-add vector field for {feature_name}: {e}") - - if source_type and source_type.startswith('fold_constraint_'): - try: - self.add_fold_constraint( - feature_name, source_type[len('fold_constraint_') :] - ) - continue - except Exception as e: - _log(f"Failed to re-add fold constraint for {feature_name}: {e}") - - if source_type in ('feature_points', 'feature_data'): - try: - self.add_data(feature_name) - continue - except Exception as e: - _log(f"Failed to re-add data for {feature_name}: {e}") - - if source_type == 'feature_scalar': - try: - self.add_scalar_field(feature_name) - continue - except Exception as e: - _log(f"Failed to re-add scalar field for {feature_name}: {e}") - - if source_type == 'bounding_box' or mesh_name == 'model_bounding_box': - try: - self.add_model_bounding_box() - continue - except Exception as e: - _log(f"Failed to re-add bounding box: {e}") - - # Fallback: if nothing matched, attempt to re-add by using viewer metadata - # Many viewer entries store the vtk source under meta['vtk'] or similar; try best-effort - try: - vtk_src = meta.get('vtk') - if vtk_src is not None: - _log(f"Fallback re-add for mesh {mesh_name}") - self.viewer.add_mesh_object(vtk_src, name=mesh_name, **kwargs) - except Exception: - pass + The objects are built on a background thread (this can solve the + model), then added to the viewer again on the GUI thread with the + same name and viewer settings (see `_on_update_objects_finished`). + """ + if not self.model_manager or self.viewer is None: + return + if self.model_manager.model is None: + logger.info("No model available to update the viewer objects.") + return + names = self.viewer.out_of_date_objects() + if not names: + return + specs = [] + for name in names: + meta = self.viewer.meshes[name] + spec = { + 'name': name, + 'source_type': meta.get('source_type') or '', + 'source_feature': meta.get('source_feature'), + 'isovalue': meta.get('isovalue'), + 'metadata': dict(meta.get('metadata') or {}), + } + if spec['source_type'] in ('cross_section_plane', 'cross_section_line'): + # the section geometry does not change; copy it so the + # background thread does not change the mesh on screen + spec['mesh'] = meta['mesh'].copy() + specs.append(spec) + + if any(self._uses_stratigraphy_colours(spec) for spec in specs): + missing = self.model_manager.get_units_without_colour() + if missing: + QMessageBox.warning( + self, + "Missing unit colour", + "Cannot update the objects coloured by stratigraphy. These units " + "have no valid colour in the stratigraphic column:\n\n" + + "\n".join(missing) + + "\n\nSet a colour for each unit and try again.", + ) + return + def target(progress_callback): + results = [] + for i, spec in enumerate(specs): + progress_callback(f"Updating {spec['name']} ({i + 1} of {len(specs)})...") + try: + mesh, overrides = self._rebuild_object(spec) + results.append((spec['name'], mesh, overrides, spec['metadata'], None)) except Exception as e: - _log(f"Failed to update visualisation for feature: {feature_name}. Error: {e}") + results.append((spec['name'], None, {}, None, str(e))) + return results + + self._update_objects_thread, self._update_objects_worker, self._update_objects_progress = ( + start_background_task( + self, + target, + title="Update Viewer Objects", + initial_label="Updating viewer objects...", + on_progress=self._on_update_objects_progress, + on_finished=self._on_update_objects_finished, + on_error=self._on_update_objects_error, + ) + ) + self._refresh_update_objects_button() + + def _rebuild_object(self, spec): + """Build one viewer object again from the current model. - # Refresh the viewer + Runs on a background thread, so it must not touch the viewer. + Returns (mesh, overrides), where overrides are viewer settings that + come from the model (e.g. a unit colour). `spec['metadata']` is + updated with new values from the model (e.g. the unit names). Raises + if the object cannot be built. + """ + source_type = spec['source_type'] + feature_name = spec['source_feature'] + metadata = spec['metadata'] + model = self.model_manager.model + + if source_type in self.FEATURE_SOURCE_TYPES or source_type.startswith('fold_constraint_'): + if feature_name is None or model.get_feature_by_name(feature_name) is None: + raise ValueError(f"Feature '{feature_name}' is not in the model") + + overrides = {} + if source_type == 'feature_scalar': + mesh = self._build_scalar_field(feature_name) + elif source_type == 'feature_surface': + mesh = self._build_feature_surface(feature_name, spec['isovalue']) + elif source_type in ('feature_vector', 'feature_vectors'): + mesh = self._build_vector_field(feature_name) + elif source_type.startswith('fold_constraint_'): + constraint = source_type[len('fold_constraint_') :] + mesh = self._build_fold_constraint(feature_name, constraint) + if mesh is None: + raise ValueError("The fold constraint vectors are not defined") + elif source_type in ('feature_points', 'feature_data'): + meshes = {name: m for name, m, _ in self._build_data_meshes(feature_name)} + if spec['name'] not in meshes: + raise ValueError(f"Feature '{feature_name}' has no data for this object") + mesh = meshes[spec['name']] + elif source_type == 'bounding_box': + mesh = self._build_bounding_box() + elif source_type == 'fault_surface': + mesh = self._find_fault_surface(feature_name).vtk() + elif source_type == 'stratigraphic_surface': + surface = self._find_stratigraphic_surface(feature_name) + mesh = surface.vtk() + overrides['color'] = surface.colour + elif source_type in ('cross_section_plane', 'cross_section_line'): + mesh = spec['mesh'] + self._colour_by_stratigraphy(mesh.points, mesh.point_data, metadata) + elif source_type == 'block_model': + bb = model.bounding_box + mesh = build_block_model_mesh(bb.origin, bb.maximum, metadata['ncells']) + self._colour_by_stratigraphy(mesh.cell_centers().points, mesh.cell_data, metadata) + elif source_type == 'topography_surface': + xx, yy, zz = self.model_manager.sample_dem_grid() + mesh = pv.StructuredGrid(xx, yy, zz) + mesh['Elevation'] = mesh.points[:, 2] + if metadata.get('coloured'): + self._colour_by_stratigraphy(mesh.points, mesh.point_data, metadata) + else: + raise ValueError(f"Cannot update objects of type '{source_type}'") + + if getattr(mesh, 'n_points', 1) == 0: + raise ValueError("The object has no geometry in the current model") + return mesh, overrides + + def _on_update_objects_progress(self, message): try: - self.viewer.update() + self._update_objects_progress.setLabelText(message) except Exception: pass + + def _finish_update_objects_task(self): + finish_background_task( + self._update_objects_thread, self._update_objects_worker, self._update_objects_progress + ) + self._update_objects_thread = None + self._update_objects_worker = None + self._update_objects_progress = None + + def _on_update_objects_finished(self, results): + self._finish_update_objects_task() + failed = [] + for name, mesh, overrides, metadata, error in results: + entry = self.viewer.meshes.get(name) + if entry is None: + # removed from the viewer while the update ran + continue + if error is not None: + failed.append(f"{name}: {error}") + continue + try: + self.viewer.replace_mesh_object( + name, mesh, overrides, out_of_date=False, metadata=metadata + ) + except Exception as e: + logger.exception(f"Cannot add updated object '{name}' to the viewer") + failed.append(f"{name}: {e}") + try: + self.viewer.render() + except Exception: + pass + self._refresh_update_objects_button() + if failed: + logger.warning("Cannot update viewer objects:\n" + "\n".join(failed)) + QMessageBox.warning( + self, + "Update Viewer Objects", + "These objects were not updated and are still out of date:\n\n" + "\n".join(failed), + ) + + def _on_update_objects_error(self, traceback_text): + self._finish_update_objects_task() + self._refresh_update_objects_button() + logger.error(f"Failed to update viewer objects: {traceback_text}") + QMessageBox.warning( + self, + "Update Viewer Objects", + "Cannot update the viewer objects. See the log for details.", + ) diff --git a/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py b/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py index b4d338c..f0c6dfa 100644 --- a/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py +++ b/loopstructural/gui/visualisation/loop_pyvistaqt_wrapper.py @@ -3,9 +3,13 @@ from pyvistaqt import QtInteractor from qgis.PyQt.QtCore import pyqtSignal +from .mesh_scalar_utils import threshold_mesh + class LoopPyVistaQTPlotter(QtInteractor): objectAdded = pyqtSignal(QtInteractor) # Signal to request deletion + # emitted when objects are marked out of date, or brought up to date + outOfDateChanged = pyqtSignal() def __init__(self, parent): super().__init__(parent=parent) @@ -42,6 +46,9 @@ def add_mesh_object( source_feature: Optional[str] = None, source_type: Optional[str] = None, isovalue: Optional[float] = None, + metadata: Optional[Dict[str, Any]] = None, + out_of_date: bool = False, + threshold: Optional[Dict[str, Any]] = None, **kwargs, ) -> None: """Add a mesh to the plotter. @@ -73,6 +80,16 @@ def add_mesh_object( source_type : Optional[str] A short tag describing the kind of source (e.g. 'feature_surface', 'fault_surface', 'bounding_box'). + metadata : Optional[dict] + Extra values needed to build the mesh again from the model (for + example the number of blocks of a block model). + out_of_date : bool + True if the mesh no longer matches the model. + threshold : Optional[dict] + Show only the part of the mesh whose values are in a range (see + `mesh_scalar_utils.threshold_mesh`). The full mesh is still + stored, so the filter can be changed or removed later. Raises + ValueError if no cells are in the range. Returns ------- @@ -113,8 +130,14 @@ def add_mesh_object( # merge any extra kwargs (allow caller to override default choices) add_kwargs.update(kwargs) + display_mesh = mesh + if threshold: + display_mesh = threshold_mesh(mesh, threshold) + if display_mesh.n_cells == 0: + raise ValueError("No cells are in the filter range") + # attempt to add to the underlying pyvista plotter - actor = self.add_mesh(mesh, name=name, **add_kwargs) + actor = self.add_mesh(display_mesh, name=name, **add_kwargs) # store the mesh, actor and kwargs for future re-adds # persist source metadata so callers can find meshes created from model features @@ -125,9 +148,101 @@ def add_mesh_object( 'source_feature': source_feature, 'source_type': source_type, 'isovalue': isovalue, + 'metadata': dict(metadata or {}), + 'out_of_date': out_of_date, + 'threshold': dict(threshold) if threshold else None, + # the mesh shown in the viewer (the filtered part of `mesh`) + 'display_mesh': display_mesh, } self.objectAdded.emit(self) + def get_source_metadata(self, name: str) -> Dict[str, Any]: + """Return the source values of an object as keyword arguments for + `add_mesh_object`, so that an object removed and added again (for + example to change its colour map) can still be built again from the + model. + """ + entry = self.meshes.get(name) + if not entry: + return {} + return { + 'source_feature': entry.get('source_feature'), + 'source_type': entry.get('source_type'), + 'isovalue': entry.get('isovalue'), + 'metadata': entry.get('metadata'), + 'out_of_date': bool(entry.get('out_of_date', False)), + 'threshold': entry.get('threshold'), + } + + def replace_mesh_object(self, name: str, mesh=None, overrides=None, **source_updates) -> None: + """Add the object `name` again, and keep its source values, viewer + settings (colour map, opacity, colour picked by the user, ...) and + visibility. + + Parameters + ---------- + name : str + Name of an object in the viewer. + mesh : optional + A new mesh for the object (e.g. built again from the model). If + None, the current mesh is used. + overrides : Optional[dict] + Viewer settings that replace the stored ones (e.g. a new unit + colour). A colour picked by the user still has priority. + **source_updates + Source values to change, e.g. `threshold=...` or + `out_of_date=False`. + + pyvista replaces the actor that has the same name, so if the new + object cannot be added, the old object stays and the error is raised. + """ + entry = self.meshes[name] + if mesh is None: + mesh = entry['mesh'] + source = self.get_source_metadata(name) + source.update(source_updates) + kwargs = { + key: value + for key, value in (entry.get('kwargs') or {}).items() + if key not in source and key != 'name' + } + kwargs.update(overrides or {}) + user_colour = entry.get('color') + if user_colour is not None: + kwargs['color'] = user_colour + actor = entry.get('actor') + visible = bool(getattr(actor, 'visibility', True)) + + try: + self.add_mesh_object(mesh, name=name, **source, **kwargs) + except Exception: + # e.g. a scalar array selected in the properties panel that the + # new mesh does not have; add it with the default colouring + for key in ('scalars', 'cmap', 'clim', 'rgb'): + kwargs.pop(key, None) + self.add_mesh_object(mesh, name=name, **source, **kwargs) + + new_entry = self.meshes[name] + if user_colour is not None: + new_entry['color'] = user_colour + if not visible and new_entry.get('actor') is not None: + new_entry['actor'].visibility = False + + def set_out_of_date(self, names, out_of_date: bool = True) -> None: + """Mark the named objects as out of date (or up to date).""" + changed = False + for name in names: + entry = self.meshes.get(name) + if entry is not None and bool(entry.get('out_of_date')) != out_of_date: + entry['out_of_date'] = out_of_date + changed = True + if changed: + self.outOfDateChanged.emit() + + def out_of_date_objects(self): + """Return the names of the objects that are out of date.""" + return [name for name, entry in self.meshes.items() if entry.get('out_of_date')] + def remove_object(self, name: str) -> None: """Remove an object by name and clean up stored metadata. @@ -157,6 +272,8 @@ def remove_object(self, name: str) -> None: del self.meshes[name] except Exception: pass + if entry.get('out_of_date'): + self.outOfDateChanged.emit() def set_object_visibility(self, name: str, visibility): """Change the visibility of an object.""" diff --git a/loopstructural/gui/visualisation/mesh_scalar_utils.py b/loopstructural/gui/visualisation/mesh_scalar_utils.py index 2af621e..be84b74 100644 --- a/loopstructural/gui/visualisation/mesh_scalar_utils.py +++ b/loopstructural/gui/visualisation/mesh_scalar_utils.py @@ -176,3 +176,68 @@ def apply_colormap_lut(mapper, cmap, clim=None, nan_color=(0.6, 0.6, 0.6, 1.0)): pass except Exception: pass + + +def filter_array_names(mesh): + """Return the names of the arrays of `mesh` that a threshold filter can + use: the single-component point arrays, and the single-component cell + arrays as `cell:` (the same naming as `get_scalar_values`). + """ + names = [] + for prefix, data in (('', 'point_data'), ('cell:', 'cell_data')): + arrays = getattr(mesh, data, None) or {} + for key in sorted(arrays.keys()): + values = np.asarray(arrays[key]) + if values.ndim == 1 and np.issubdtype(values.dtype, np.number): + names.append(f"{prefix}{key}") + return names + + +def threshold_mesh(mesh, threshold): + """Return the part of `mesh` whose values pass a filter. + + `threshold` is a dict with `scalars` (the array name, `cell:` for + a cell array) and one of: + + - `min`, `max`: keep the values in this range (the limits are + included); with `invert`, keep the values outside the range + - `values`: keep only these values (e.g. the ids of some units) + + For a point array, a cell is kept only if all its points pass. Raises + KeyError if the array is not on the mesh. + """ + name = threshold['scalars'] + preference = 'point' + if name.startswith('cell:'): + name = name.split(':', 1)[1] + preference = 'cell' + data = mesh.cell_data if preference == 'cell' else mesh.point_data + if name not in data: + raise KeyError(f"The object has no {preference} array '{name}'") + + if 'values' in threshold: + # threshold a 0/1 mask on a shallow copy, so the result has the same + # type as a range filter and the mesh of the caller does not change + mask = np.isin(np.asarray(data[name]), list(threshold['values'])).astype(float) + masked = mesh.copy(deep=False) + mask_data = masked.cell_data if preference == 'cell' else masked.point_data + mask_data['_filter_mask'] = mask + result = masked.threshold( + value=(0.5, 1.5), + scalars='_filter_mask', + preference=preference, + all_scalars=preference == 'point', + ) + for arrays in (result.cell_data, result.point_data): + if '_filter_mask' in arrays: + del arrays['_filter_mask'] + return result + + low, high = sorted((float(threshold['min']), float(threshold['max']))) + return mesh.threshold( + value=(low, high), + scalars=name, + preference=preference, + invert=bool(threshold.get('invert', False)), + all_scalars=preference == 'point', + ) diff --git a/loopstructural/gui/visualisation/object_list_widget.py b/loopstructural/gui/visualisation/object_list_widget.py index 8d19021..9910852 100644 --- a/loopstructural/gui/visualisation/object_list_widget.py +++ b/loopstructural/gui/visualisation/object_list_widget.py @@ -35,6 +35,7 @@ def __init__(self, parent=None, *, viewer=None, properties_widget=None): self.setLayout(self.mainLayout) self.viewer = viewer self.viewer.objectAdded.connect(self.update_object_list) + self.viewer.outOfDateChanged.connect(self._on_out_of_date_changed) self.treeWidget.installEventFilter(self) self.treeWidget.itemSelectionChanged.connect(self.on_object_selected) self.treeWidget.itemDoubleClicked.connect(self.onDoubleClick) @@ -94,6 +95,9 @@ def update_object_list(self, new_object): mesh = meshes[mesh_name] self.add_mesh_item(mesh_name, mesh) + def _on_out_of_date_changed(self): + self.update_object_list(None) + def add_mesh_item(self, mesh_name, mesh): """Add a top-level tree item for a mesh and populate children for point/cell data arrays. @@ -139,7 +143,13 @@ def _on_vis(state, name=mesh_name, m=mesh): itemLayout = QHBoxLayout(itemWidget) itemLayout.setContentsMargins(0, 0, 0, 0) itemLayout.addWidget(visibilityCheckbox) - itemLayout.addWidget(QLabel(mesh_name)) + nameLabel = QLabel(mesh_name) + if isinstance(mesh, dict) and mesh.get('out_of_date'): + # the label text is the object name used elsewhere, so show the + # state with the style and tooltip only + nameLabel.setStyleSheet("color: gray; font-style: italic;") + nameLabel.setToolTip("Out of date: the model changed after this object was added") + itemLayout.addWidget(nameLabel) itemWidget.setLayout(itemLayout) self.treeWidget.setItemWidget(top, 0, itemWidget) diff --git a/loopstructural/gui/visualisation/object_properties_widget.py b/loopstructural/gui/visualisation/object_properties_widget.py index 562446c..597a136 100644 --- a/loopstructural/gui/visualisation/object_properties_widget.py +++ b/loopstructural/gui/visualisation/object_properties_widget.py @@ -1,15 +1,22 @@ import matplotlib.pyplot as plt +import numpy as np # Add plotting imports for scalar histogram from matplotlib.backends.backend_qt5agg import FigureCanvasQTAgg as FigureCanvas from qgis.PyQt.QtCore import Qt +from qgis.PyQt.QtGui import QColor, QIcon, QPixmap from qgis.PyQt.QtWidgets import ( QCheckBox, QColorDialog, QComboBox, + QDoubleSpinBox, + QGroupBox, QHBoxLayout, QLabel, QLineEdit, + QListWidget, + QListWidgetItem, + QMessageBox, QPushButton, QSizePolicy, QSlider, @@ -17,7 +24,12 @@ QWidget, ) -from .mesh_scalar_utils import apply_colormap_lut, get_scalar_values, render_histogram +from .mesh_scalar_utils import ( + apply_colormap_lut, + filter_array_names, + get_scalar_values, + render_histogram, +) class ObjectPropertiesWidget(QWidget): @@ -110,6 +122,64 @@ def __init__(self, parent=None, *, viewer=None): self.hist_canvas.setSizePolicy(QSizePolicy.Expanding, QSizePolicy.Fixed) layout.addWidget(self.hist_canvas) + # Filter: show only the part of the object whose values are in a + # range, or (for the unit ids of a block model, cross section or + # topography) only the units that are checked + self.filter_group = QGroupBox("Filter (Threshold)") + filter_layout = QVBoxLayout(self.filter_group) + filter_layout.addWidget(QLabel("Array:")) + self.filter_array_combo = QComboBox() + self.filter_array_combo.setSizePolicy(QSizePolicy.Expanding, QSizePolicy.Fixed) + self.filter_array_combo.currentTextChanged.connect(self._on_filter_array_changed) + filter_layout.addWidget(self.filter_array_combo) + self.filter_range_widget = QWidget() + filter_range_widget_layout = QVBoxLayout(self.filter_range_widget) + filter_range_widget_layout.setContentsMargins(0, 0, 0, 0) + filter_range_layout = QHBoxLayout() + filter_range_layout.setSpacing(6) + filter_range_layout.addWidget(QLabel("Range:")) + self.filter_min = QDoubleSpinBox() + self.filter_max = QDoubleSpinBox() + for box in (self.filter_min, self.filter_max): + box.setRange(-1e12, 1e12) + box.setSizePolicy(QSizePolicy.Expanding, QSizePolicy.Fixed) + filter_range_layout.addWidget(box) + filter_range_widget_layout.addLayout(filter_range_layout) + self.filter_invert_checkbox = QCheckBox("Invert (show values outside the range)") + filter_range_widget_layout.addWidget(self.filter_invert_checkbox) + filter_layout.addWidget(self.filter_range_widget) + + self.filter_units_widget = QWidget() + filter_units_layout = QVBoxLayout(self.filter_units_widget) + filter_units_layout.setContentsMargins(0, 0, 0, 0) + filter_units_layout.addWidget(QLabel("Units to show:")) + self.filter_units_list = QListWidget() + self.filter_units_list.setMaximumHeight(160) + self.filter_units_list.itemChanged.connect(self._on_unit_check_changed) + filter_units_layout.addWidget(self.filter_units_list) + filter_units_buttons_layout = QHBoxLayout() + check_all_button = QPushButton("Check All") + check_all_button.clicked.connect(lambda: self._set_all_units_checked(True)) + uncheck_all_button = QPushButton("Uncheck All") + uncheck_all_button.clicked.connect(lambda: self._set_all_units_checked(False)) + filter_units_buttons_layout.addWidget(check_all_button) + filter_units_buttons_layout.addWidget(uncheck_all_button) + filter_units_layout.addLayout(filter_units_buttons_layout) + self.filter_units_widget.setVisible(False) + filter_layout.addWidget(self.filter_units_widget) + filter_buttons_layout = QHBoxLayout() + self.filter_apply_button = QPushButton("Apply Filter") + self.filter_apply_button.clicked.connect(self.apply_filter) + self.filter_clear_button = QPushButton("Clear Filter") + self.filter_clear_button.clicked.connect(self.clear_filter) + filter_buttons_layout.addWidget(self.filter_apply_button) + filter_buttons_layout.addWidget(self.filter_clear_button) + filter_layout.addLayout(filter_buttons_layout) + self.filter_status_label = QLabel("") + filter_layout.addWidget(self.filter_status_label) + self.filter_group.setEnabled(False) + layout.addWidget(self.filter_group) + # Surface Color surface_color_layout = QHBoxLayout() surface_color_layout.setSpacing(6) @@ -184,13 +254,8 @@ def set_opacity(self, value: float): pass # store in metadata if self.current_object_name in getattr(self.viewer, 'meshes', {}): - self.viewer.meshes[self.current_object_name].set( - 'kwargs', - { - **self.viewer.meshes[self.current_object_name].get('kwargs', {}), - 'opacity': value, - }, - ) + entry = self.viewer.meshes[self.current_object_name] + entry['kwargs'] = {**(entry.get('kwargs') or {}), 'opacity': value} except Exception: pass @@ -457,6 +522,194 @@ def setCurrentObject(self, object_name: str): except Exception: self._update_histogram(None) + self._load_filter_state(mesh_entry) + + def _load_filter_state(self, mesh_entry): + """Show the filter of the current object, or the full value range + of an array if the object has no filter.""" + names = filter_array_names(self.current_mesh) if self.current_mesh is not None else [] + threshold = mesh_entry.get('threshold') if isinstance(mesh_entry, dict) else None + self.filter_array_combo.blockSignals(True) + self.filter_array_combo.clear() + self.filter_array_combo.addItems(names) + self.filter_array_combo.blockSignals(False) + self.filter_group.setEnabled(bool(names)) + if threshold and threshold.get('scalars') in names: + self.filter_array_combo.blockSignals(True) + self.filter_array_combo.setCurrentText(threshold['scalars']) + self.filter_array_combo.blockSignals(False) + self._set_filter_range_to_data(threshold['scalars']) + if 'min' in threshold: + self.filter_min.setValue(float(threshold['min'])) + self.filter_max.setValue(float(threshold['max'])) + self.filter_invert_checkbox.setChecked(bool(threshold.get('invert', False))) + self._update_filter_mode(threshold['scalars'], threshold) + elif names: + # the unit ids of a block model or cross section are the most + # likely array to filter on + if 'cell:stratigraphy' in names: + self.filter_array_combo.setCurrentText('cell:stratigraphy') + elif 'stratigraphy' in names: + self.filter_array_combo.setCurrentText('stratigraphy') + self._set_filter_range_to_data(self.filter_array_combo.currentText()) + self.filter_invert_checkbox.setChecked(False) + self._update_filter_mode(self.filter_array_combo.currentText(), None) + self._update_filter_status() + + def _on_filter_array_changed(self, array_name: str): + if array_name: + self._set_filter_range_to_data(array_name) + entry = self.viewer.meshes.get(self.current_object_name) if self.viewer else None + threshold = entry.get('threshold') if entry else None + self._update_filter_mode(array_name, threshold) + + @staticmethod + def _is_unit_array(array_name: str) -> bool: + return array_name.split(':', 1)[-1] == 'stratigraphy' + + def _update_filter_mode(self, array_name: str, threshold): + """Show unit check boxes for a unit id array, and the range controls + for any other array.""" + unit_mode = self._is_unit_array(array_name) + self.filter_range_widget.setVisible(not unit_mode) + self.filter_apply_button.setVisible(not unit_mode) + self.filter_units_widget.setVisible(unit_mode) + if unit_mode: + self._populate_unit_list(array_name, threshold) + + def _populate_unit_list(self, array_name: str, threshold): + """Fill the unit check boxes with the units that are in the array. + + The unit names and colours come from the object's metadata (stored + when the object was coloured by the stratigraphic column), indexed + by unit id. + """ + from matplotlib.colors import to_hex + + entry = self.viewer.meshes.get(self.current_object_name, {}) if self.viewer else {} + metadata = entry.get('metadata') or {} + names = metadata.get('unit_names') or [] + colours = metadata.get('unit_colours') or [] + values = self._get_scalar_values(array_name) + ids = [int(v) for v in np.unique(values)] if values is not None else [] + checked = None + if threshold and threshold.get('scalars') == array_name and 'values' in threshold: + checked = {int(v) for v in threshold['values']} + + self.filter_units_list.blockSignals(True) + self.filter_units_list.clear() + for unit_id in ids: + if 0 <= unit_id < len(names): + label = str(names[unit_id]) + elif unit_id < 0: + label = "Outside all units" + else: + label = f"Unit {unit_id}" + item = QListWidgetItem(label) + item.setData(Qt.UserRole, unit_id) + item.setFlags(item.flags() | Qt.ItemIsUserCheckable) + item.setCheckState( + Qt.Checked if checked is None or unit_id in checked else Qt.Unchecked + ) + if 0 <= unit_id < len(colours): + try: + swatch = QPixmap(12, 12) + swatch.fill(QColor(to_hex(colours[unit_id]))) + item.setIcon(QIcon(swatch)) + except Exception: + pass + self.filter_units_list.addItem(item) + self.filter_units_list.blockSignals(False) + + def _checked_unit_ids(self): + ids, total = [], self.filter_units_list.count() + for i in range(total): + item = self.filter_units_list.item(i) + if item.checkState() == Qt.Checked: + ids.append(int(item.data(Qt.UserRole))) + return ids, total + + def _on_unit_check_changed(self, _item=None): + ids, total = self._checked_unit_ids() + if not ids: + # an object with no cells cannot be shown; keep the last filter + self.filter_status_label.setText("Check at least one unit to show") + return + if len(ids) == total: + self._set_threshold(None) + else: + self._set_threshold({'scalars': self.filter_array_combo.currentText(), 'values': ids}) + + def _set_all_units_checked(self, checked: bool): + self.filter_units_list.blockSignals(True) + for i in range(self.filter_units_list.count()): + self.filter_units_list.item(i).setCheckState(Qt.Checked if checked else Qt.Unchecked) + self.filter_units_list.blockSignals(False) + self._on_unit_check_changed() + + def _set_filter_range_to_data(self, array_name: str): + """Set the filter range to the full range of values of the array.""" + values = self._get_scalar_values(array_name) + if values is None: + return + values = np.asarray(values) + finite = values[np.isfinite(values)] if values.dtype.kind == 'f' else values + if finite.size == 0: + return + low, high = float(finite.min()), float(finite.max()) + # show integer arrays (e.g. unit ids) without decimals + integer = values.dtype.kind in 'iub' + for box in (self.filter_min, self.filter_max): + box.setDecimals(0 if integer else 4) + box.setSingleStep(1.0 if integer else max((high - low) / 100.0, 1e-4)) + self.filter_min.setValue(low) + self.filter_max.setValue(high) + + def apply_filter(self): + array_name = self.filter_array_combo.currentText() + if not array_name: + return + self._set_threshold( + { + 'scalars': array_name, + 'min': self.filter_min.value(), + 'max': self.filter_max.value(), + 'invert': self.filter_invert_checkbox.isChecked(), + } + ) + + def clear_filter(self): + self._set_threshold(None) + if self.filter_units_list.count(): + self.filter_units_list.blockSignals(True) + for i in range(self.filter_units_list.count()): + self.filter_units_list.item(i).setCheckState(Qt.Checked) + self.filter_units_list.blockSignals(False) + + def _set_threshold(self, threshold): + name = self.current_object_name + if not name or self.viewer is None or name not in self.viewer.meshes: + return + try: + self.viewer.replace_mesh_object(name, threshold=threshold) + except Exception as e: + QMessageBox.warning(self, "Filter", f"Cannot apply the filter:\n{e}") + return + try: + self.viewer.render() + except Exception: + pass + self._update_filter_status() + + def _update_filter_status(self): + entry = self.viewer.meshes.get(self.current_object_name) if self.viewer else None + if not entry or not entry.get('threshold'): + self.filter_status_label.setText("No filter") + return + total = getattr(entry.get('mesh'), 'n_cells', 0) + shown = getattr(entry.get('display_mesh'), 'n_cells', 0) + self.filter_status_label.setText(f"Showing {shown} of {total} cells") + def _on_scalar_changed(self, scalar_name: str): # update histogram preview immediately try: @@ -514,6 +767,7 @@ def _on_scalar_changed(self, scalar_name: str): opacity = old_kwargs.get('opacity', None) show_scalar_bar = self.scalar_bar_checkbox.isChecked() + source = self.viewer.get_source_metadata(self.current_object_name) try: self.viewer.remove_object(self.current_object_name) except Exception: @@ -528,11 +782,12 @@ def _on_scalar_changed(self, scalar_name: str): clim=clim, opacity=opacity, show_scalar_bar=show_scalar_bar, + **source, ) self.current_mesh = self.viewer.meshes.get(self.current_object_name, {}).get('mesh') except Exception: try: - self.viewer.add_mesh_object(mesh, name=self.current_object_name) + self.viewer.add_mesh_object(mesh, name=self.current_object_name, **source) self.current_mesh = self.viewer.meshes.get(self.current_object_name, {}).get('mesh') except Exception: pass @@ -645,6 +900,7 @@ def _on_colormap_changed(self, cmap: str): opacity = old_kwargs.get('opacity', None) show_scalar_bar = self.scalar_bar_checkbox.isChecked() + source = self.viewer.get_source_metadata(self.current_object_name) try: self.viewer.remove_object(self.current_object_name) except Exception: @@ -659,11 +915,12 @@ def _on_colormap_changed(self, cmap: str): clim=clim, opacity=opacity, show_scalar_bar=show_scalar_bar, + **source, ) self.current_mesh = self.viewer.meshes.get(self.current_object_name, {}).get('mesh') except Exception: try: - self.viewer.add_mesh_object(mesh, name=self.current_object_name) + self.viewer.add_mesh_object(mesh, name=self.current_object_name, **source) self.current_mesh = self.viewer.meshes.get(self.current_object_name, {}).get( 'mesh' ) diff --git a/loopstructural/main/model_manager.py b/loopstructural/main/model_manager.py index c3aa7be..fb573fc 100644 --- a/loopstructural/main/model_manager.py +++ b/loopstructural/main/model_manager.py @@ -2000,6 +2000,19 @@ def get_stratigraphic_column_colours(self) -> list: colours.extend(unit.colour for unit in group.units) return colours + def get_stratigraphic_unit_names(self) -> list: + """Return unit names ordered to line up with `evaluate_model`'s ids. + + `names[i]` is the name of whichever unit `evaluate_model` labels `i` + (see `get_stratigraphic_column_colours`). + """ + if self.model is None or self.model.stratigraphic_column is None: + return [] + names = [] + for group in reversed(self.model.stratigraphic_column.get_groups()): + names.extend(unit.name for unit in group.units) + return names + def get_units_without_colour(self) -> list: """Return the names of units whose colour is missing or invalid. diff --git a/loopstructural/resources/images/block_model.svg b/loopstructural/resources/images/block_model.svg new file mode 100644 index 0000000..14dfd21 --- /dev/null +++ b/loopstructural/resources/images/block_model.svg @@ -0,0 +1,13 @@ + + + + + + + + + + + + + diff --git a/tests/unit/test_block_model_mesh.py b/tests/unit/test_block_model_mesh.py new file mode 100644 index 0000000..e7ea9a6 --- /dev/null +++ b/tests/unit/test_block_model_mesh.py @@ -0,0 +1,27 @@ +"""Pytest tests for `build_block_model_mesh` (see +loopstructural/gui/visualisation/cross_section_utils.py). + +Only pyvista/numpy are needed, so these run in the fast tests/unit/ job. +""" + +import numpy as np + +from loopstructural.gui.visualisation.cross_section_utils import build_block_model_mesh + + +def test_block_model_fills_bounding_box(): + mesh = build_block_model_mesh((0, 10, -100), (100, 60, 0), (10, 5, 4)) + assert mesh.n_cells == 10 * 5 * 4 + np.testing.assert_allclose(mesh.bounds, (0, 100, 10, 60, -100, 0)) + + +def test_block_model_cell_centres_are_inside_box(): + mesh = build_block_model_mesh((0, 0, 0), (10, 10, 10), (2, 2, 2)) + centres = mesh.cell_centers().points + assert centres.shape == (8, 3) + np.testing.assert_allclose(np.unique(centres[:, 0]), [2.5, 7.5]) + + +def test_block_model_needs_at_least_one_cell_per_axis(): + mesh = build_block_model_mesh((0, 0, 0), (1, 1, 1), (0, 3, 3)) + assert mesh.n_cells == 1 * 3 * 3 diff --git a/tests/unit/test_threshold_mesh.py b/tests/unit/test_threshold_mesh.py new file mode 100644 index 0000000..6b463ed --- /dev/null +++ b/tests/unit/test_threshold_mesh.py @@ -0,0 +1,74 @@ +"""Pytest tests for the threshold filter helpers in +loopstructural/gui/visualisation/mesh_scalar_utils.py. +""" + +import numpy as np +import pytest +import pyvista as pv + +from loopstructural.gui.visualisation.mesh_scalar_utils import ( + filter_array_names, + threshold_mesh, +) + + +@pytest.fixture +def block_model(): + grid = pv.ImageData(dimensions=(5, 5, 5)) + grid.cell_data['stratigraphy'] = np.arange(grid.n_cells) % 4 + grid.cell_data['colour'] = np.zeros((grid.n_cells, 3), dtype=np.uint8) + grid.point_data['z'] = grid.points[:, 2] + return grid + + +def test_filter_array_names_lists_single_component_arrays(block_model): + assert filter_array_names(block_model) == ['z', 'cell:stratigraphy'] + + +def test_threshold_keeps_cells_in_range(block_model): + result = threshold_mesh(block_model, {'scalars': 'cell:stratigraphy', 'min': 1, 'max': 2}) + assert set(np.unique(result.cell_data['stratigraphy'])) == {1, 2} + + +def test_threshold_invert_keeps_cells_outside_range(block_model): + result = threshold_mesh( + block_model, {'scalars': 'cell:stratigraphy', 'min': 1, 'max': 2, 'invert': True} + ) + assert set(np.unique(result.cell_data['stratigraphy'])) == {0, 3} + + +def test_threshold_accepts_limits_in_either_order(block_model): + result = threshold_mesh(block_model, {'scalars': 'cell:stratigraphy', 'min': 2, 'max': 1}) + assert set(np.unique(result.cell_data['stratigraphy'])) == {1, 2} + + +def test_threshold_on_point_array(block_model): + result = threshold_mesh(block_model, {'scalars': 'z', 'min': 0, 'max': 2}) + assert result.n_cells > 0 + assert result.points[:, 2].max() <= 2 + + +def test_threshold_missing_array_raises(block_model): + with pytest.raises(KeyError): + threshold_mesh(block_model, {'scalars': 'cell:missing', 'min': 0, 'max': 1}) + + +def test_threshold_values_keeps_only_those_cells(block_model): + result = threshold_mesh(block_model, {'scalars': 'cell:stratigraphy', 'values': [0, 3]}) + assert set(np.unique(result.cell_data['stratigraphy'])) == {0, 3} + assert '_filter_mask' not in result.cell_data + assert '_filter_mask' not in block_model.cell_data + + +def test_threshold_values_on_point_array_keeps_arrays(): + grid = pv.ImageData(dimensions=(5, 5, 5)) + grid.point_data['stratigraphy'] = (grid.points[:, 2] >= 2).astype(int) + result = threshold_mesh(grid, {'scalars': 'stratigraphy', 'values': [1]}) + assert result.n_cells > 0 + assert set(np.unique(result.point_data['stratigraphy'])) == {1} + assert '_filter_mask' not in grid.point_data + + +def test_threshold_no_values_gives_empty_mesh(block_model): + result = threshold_mesh(block_model, {'scalars': 'cell:stratigraphy', 'values': []}) + assert result.n_cells == 0