From 323d628e699db9e1cf2e91ae07228b2b4cbe7a72 Mon Sep 17 00:00:00 2001 From: Marvin Albert Date: Fri, 10 Apr 2026 01:12:10 +0200 Subject: [PATCH 1/2] Add beyond translation registration methods --- src/napari_stitcher/_stitcher_widget.py | 45 +++++++++++++++++++++++-- 1 file changed, 43 insertions(+), 2 deletions(-) diff --git a/src/napari_stitcher/_stitcher_widget.py b/src/napari_stitcher/_stitcher_widget.py index 3452a82..f83d013 100644 --- a/src/napari_stitcher/_stitcher_widget.py +++ b/src/napari_stitcher/_stitcher_widget.py @@ -103,6 +103,23 @@ def __init__(self, napari_viewer): 'Alternating pattern': 'alternating_pattern', } + self.reg_method = widgets.ComboBox( + choices=['Phase Correlation', 'ITKElastix'], + value='Phase Correlation', + label='Registration method:', + tooltip='Choose the pairwise registration method.\n' + '"Phase Correlation" is fast and works well for translation.\n' + '"ITKElastix" supports more transform types but requires the itk-elastix package.') + + self.antspy_transform_types = widgets.Select( + choices=['Translation', 'Rigid', 'Affine'], + value=['Translation', 'Rigid'], + label='Transform types:', + tooltip='Sequence of transform types applied in order. The last selected type is also used for global optimization.') + + self.reg_method.changed.connect(self._on_reg_method_changed) + self._on_reg_method_changed() + self.button_stitch = widgets.Button(text='Register', tooltip='Use the overlaps between tiles to determine their relative positions.') @@ -140,7 +157,12 @@ def __init__(self, napari_viewer): self.pair_pruning_method, ] - self.reg_config_widgets = self.reg_config_widgets_basic + self.reg_config_widgets_advanced + self.reg_config_widgets_method = [ + self.reg_method, + self.antspy_transform_types, + ] + + self.reg_config_widgets = self.reg_config_widgets_basic + self.reg_config_widgets_advanced + self.reg_config_widgets_method # Initialize tab screen self.reg_config_widgets_tabs = QTabWidget() @@ -150,7 +172,9 @@ def __init__(self, napari_viewer): self.reg_config_widgets_tabs.addTab( widgets.VBox(widgets=self.reg_config_widgets_basic).native, "Basic") self.reg_config_widgets_tabs.addTab( - widgets.VBox(widgets=self.reg_config_widgets_advanced).native, "More") + widgets.VBox(widgets=self.reg_config_widgets_advanced).native, "More") + self.reg_config_widgets_tabs.addTab( + widgets.VBox(widgets=self.reg_config_widgets_method).native, "Method") self.visualization_widgets = [ self.visualization_type_rbuttons, @@ -216,6 +240,10 @@ def __init__(self, napari_viewer): self.button_load_layers_sel.clicked.connect(self.load_layers_sel) + def _on_reg_method_changed(self, event=None): + """Show/hide transform type widget based on the selected method.""" + self.antspy_transform_types.visible = self.reg_method.value == 'ITKElastix' + def update_viewer_transformations(self, event=None): """ set transformations @@ -328,9 +356,22 @@ def run_registration(self): else: registration_binning = None + if self.reg_method.value == 'ITKElastix': + pairwise_reg_func = registration.registration_ITKElastix + transform_types = list(self.antspy_transform_types.value) + pairwise_reg_func_kwargs = {'transform_types': transform_types} + groupwise_resolution_kwargs = {'transform': transform_types[-1].lower()} + else: + pairwise_reg_func = registration.phase_correlation_registration + pairwise_reg_func_kwargs = None + groupwise_resolution_kwargs = None + params = registration.register( msims, registration_binning=registration_binning, + pairwise_reg_func=pairwise_reg_func, + pairwise_reg_func_kwargs=pairwise_reg_func_kwargs, + groupwise_resolution_kwargs=groupwise_resolution_kwargs, pre_registration_pruning_method=self.pair_pruning_method_mapping[self.pair_pruning_method.value], post_registration_do_quality_filter=self.do_quality_filter.value, post_registration_quality_threshold=self.quality_threshold.value, From 6d7bc555ae51a068012c0341628bfdabf870deba Mon Sep 17 00:00:00 2001 From: Marvin Albert Date: Sun, 12 Apr 2026 20:39:46 +0200 Subject: [PATCH 2/2] Add possibility to correct parameters after registration, including for time lapses. Also, registration runs will take into account currently visible parameters (both original and registered) without the need to reload layers --- src/napari_stitcher/_stitcher_widget.py | 228 +++++++++++++----- .../_tests/test_stitcher_widget.py | 158 +++++++++++- 2 files changed, 327 insertions(+), 59 deletions(-) diff --git a/src/napari_stitcher/_stitcher_widget.py b/src/napari_stitcher/_stitcher_widget.py index f83d013..d3ed385 100644 --- a/src/napari_stitcher/_stitcher_widget.py +++ b/src/napari_stitcher/_stitcher_widget.py @@ -23,6 +23,7 @@ fusion, spatial_image_utils, msi_utils, + param_utils, ) from napari.layers import Image, Labels @@ -226,11 +227,16 @@ def __init__(self, napari_viewer): self.fused_layers = [] self.params = dict() + # flag to suppress watch_layer_changes during programmatic affine updates + self._updating_viewer = False + # last timepoint that was applied to the viewer; used to skip no-op current_step events + self._last_applied_tp = None + # create temporary directory for storing dask arrays self.tmpdir = tempfile.TemporaryDirectory() self.visualization_type_rbuttons.changed.connect(self.update_viewer_transformations) - self.viewer.dims.events.connect(self.update_viewer_transformations) + self.viewer.dims.events.current_step.connect(self.update_viewer_transformations) self.button_stitch.clicked.connect(self.run_registration) # self.button_stabilize.clicked.connect(self.run_stabilization) @@ -249,15 +255,36 @@ def update_viewer_transformations(self, event=None): set transformations - for current timepoint - for each (compatible) layer loaded in viewer + + Called from exactly two sources: + 1. viewer.dims.events.current_step (timepoint scroll) + 2. visualization_type_rbuttons.changed (Show toggle) """ - try: - # events are constantly triggered by viewer.dims.events, - # but we only want to update if current_step changes - if hasattr(event, 'type') and\ - event.type != 'current_step': return - except AttributeError: - pass + # When called from a current_step event: + # - only proceed after registration has been performed + # - only proceed when the timepoint actually changed (transform-mode + # interactions also fire current_step without changing the tp) + if hasattr(event, 'type'): + if not self.visualization_type_rbuttons.enabled: + return + # Compute the candidate tp now so we can compare + # (replicated from the block below; simims may not be loaded yet) + if not len(self.msims): + return + _sims_check = [msi_utils.get_sim_from_msim(self.msims[l.name]) + for l in self.viewer.layers if l.name in self.msims] + if not _sims_check: + return + _highest_sdim = max( + len(spatial_image_utils.get_spatial_dims_from_sim(s)) + for s in _sims_check) + _candidate_tp = ( + self.viewer.dims.current_step[-_highest_sdim - 1] + if len(self.viewer.dims.current_step) > _highest_sdim + else 0) + if _candidate_tp == self._last_applied_tp: + return if not len(self.msims): return @@ -290,45 +317,121 @@ def update_viewer_transformations(self, event=None): else: transform_key = 'affine_registered' - for il, l in enumerate(compatible_layers): + self._last_applied_tp = curr_tp + self._updating_viewer = True + try: + for il, l in enumerate(compatible_layers): - try: - params = spatial_image_utils.get_affine_from_sim( - sims[il], transform_key=transform_key - ) - except: - # notifications.notification_manager.receive_info( - # 'Update transform: %s not available in %s' %(transform_key, l.name)) - continue + try: + params = spatial_image_utils.get_affine_from_sim( + sims[il], transform_key=transform_key + ) + except: + continue - try: - p = np.array(params.sel(t=sims[il].coords['t'][curr_tp])).squeeze() - if np.isnan(p).any(): - raise(Exception()) - except: - notifications.notification_manager.receive_info( - 'Timepoint %s: no parameters available, register first.' % curr_tp) - continue + try: + p = np.array(params.sel(t=sims[il].coords['t'][curr_tp])).squeeze() + if np.isnan(p).any(): + raise(Exception()) + except: + notifications.notification_manager.receive_info( + 'Timepoint %s: no parameters available, register first.' % curr_tp) + continue - # # if curr_tp not available, use nearest available parameter - # notifications.notification_manager.receive_info( - # 'Timepoint %s: no parameters available, taking nearest available one.' % curr_tp) - # p = np.array(params.sel(t=layer_sim.coords['t'][curr_tp], method='nearest')).squeeze() + ndim_layer_data = l.ndim - ndim_layer_data = l.ndim + # if stitcher sim has more dimensions than layer data (i.e. time) + vis_p = p[-(ndim_layer_data + 1):, -(ndim_layer_data + 1):] - # if stitcher sim has more dimensions than layer data (i.e. time) - vis_p = p[-(ndim_layer_data + 1):, -(ndim_layer_data + 1):] + # if layer data has more dimensions than stitcher sim + full_vis_p = np.eye(ndim_layer_data + 1) + full_vis_p[-len(vis_p):, -len(vis_p):] = vis_p - # if layer data has more dimensions than stitcher sim - full_vis_p = np.eye(ndim_layer_data + 1) - full_vis_p[-len(vis_p):, -len(vis_p):] = vis_p + l.affine = full_vis_p + finally: + self._updating_viewer = False - l.affine = full_vis_p + def _capture_layer_transforms_to_msims(self): + """ + Capture the current layer affines into affine_metadata in msims. + Called before registration/fusion when showing Original transforms, so + any manual layer adjustments made by the user are used as the starting point. + """ + for l in self.input_layers: + if l.name not in self.msims: + continue + msim = self.msims[l.name] + sim = msi_utils.get_sim_from_msim(msim) + ndim = spatial_image_utils.get_ndim_from_sim(sim) + affine = np.array(l.affine.affine_matrix)[-(ndim + 1):, -(ndim + 1):] + t_coords = sim.coords['t'] if 't' in sim.dims else None + affine_xr = param_utils.affine_to_xaffine(affine, t_coords=t_coords) + msi_utils.set_affine_transform(msim, affine_xr, transform_key='affine_metadata') + + def _update_registered_param_for_current_tp(self, l): + """ + Update affine_registered for the current timepoint from l.affine. + Only the current timepoint is modified; all others remain unchanged. + Called live when the user manually transforms a layer while showing Registered. + """ + if l.name not in self.msims: + return + msim = self.msims[l.name] + sim = msi_utils.get_sim_from_msim(msim) + ndim = spatial_image_utils.get_ndim_from_sim(sim) + + # Determine current timepoint (mirrors the logic in update_viewer_transformations) + sdims = spatial_image_utils.get_spatial_dims_from_sim(sim) + if len(self.viewer.dims.current_step) > len(sdims): + curr_tp = self.viewer.dims.current_step[-len(sdims) - 1] + else: + curr_tp = 0 + + curr_affine = np.array(l.affine.affine_matrix)[-(ndim + 1):, -(ndim + 1):] + + try: + existing_params = spatial_image_utils.get_affine_from_sim( + sim, transform_key='affine_registered').copy() + except Exception: + return # not registered yet + + if 't' in existing_params.dims: + t_val = sim.coords['t'][curr_tp] + existing_params.loc[{'t': t_val}] = curr_affine + else: + existing_params = param_utils.affine_to_xaffine(curr_affine, t_coords=None) + + msi_utils.set_affine_transform(msim, existing_params, transform_key='affine_registered') + + def _promote_registered_to_metadata(self): + """ + Copy affine_registered → affine_metadata for all msims so that a + subsequent registration uses the manually corrected positions as its + starting point rather than the original metadata positions. + """ + for l_name, msim in self.msims.items(): + sim = msi_utils.get_sim_from_msim(msim) + try: + registered_params = spatial_image_utils.get_affine_from_sim( + sim, transform_key='affine_registered') + except Exception: + continue + msi_utils.set_affine_transform( + msim, registered_params.copy(), transform_key='affine_metadata') def run_registration(self): + # Promote the current starting-point transforms into affine_metadata so + # that registration always uses the most up-to-date positions: + # - Showing Original (or pre-registration): capture layer affines → affine_metadata + # - Showing Registered: copy affine_registered → affine_metadata + if (self.visualization_type_rbuttons.enabled and + self.visualization_type_rbuttons.value == CHOICE_REGISTERED): + self._promote_registered_to_metadata() + else: + self._capture_layer_transforms_to_msims() + # select layers corresponding to the chosen registration channel msims_dict = {_utils.get_str_unique_to_view_from_layer_name(lname): msim for lname, msim in self.msims.items() @@ -395,6 +498,10 @@ def run_registration(self): self.visualization_type_rbuttons.enabled = True self.visualization_type_rbuttons.value = CHOICE_REGISTERED + # Always refresh the viewer after registration, even if already showing + # CHOICE_REGISTERED (setting the same value doesn't fire a changed event). + self._last_applied_tp = None + self.update_viewer_transformations() def run_fusion(self): @@ -403,6 +510,11 @@ def run_fusion(self): Split layers into channel groups and fuse each group separately. """ + # Capture manual layer adjustments if fusing with original transforms + if not (self.visualization_type_rbuttons.enabled and + self.visualization_type_rbuttons.value == CHOICE_REGISTERED): + self._capture_layer_transforms_to_msims() + channels = self.reg_ch_picker.choices for _, ch in enumerate(channels): @@ -463,6 +575,7 @@ def reset(self): self.times_slider.value = (-1, 0) self.input_layers = [] self.fused_layers = [] + self._last_applied_tp = None def load_metadata(self): @@ -563,28 +676,27 @@ def load_layers(self, layers): def watch_layer_changes(self, event): """ - Watch changes in layers and warn user or update msims accordingly. - I.e. changes in transformations. + Watch user-initiated layer transform changes and update stored parameters. + + - Pre-registration or showing Original: do nothing here; transforms are + captured at registration/fusion time via _capture_layer_transforms_to_msims. + - Post-registration, showing Registered: live-update affine_registered for + the current timepoint only, leaving other timepoints unchanged. """ - if event.type in ['affine', 'scale', 'translate']: - if not self.visualization_type_rbuttons.enabled: - # assume user is modifying transforms before registration - # reload layer into stitching widget - l = event.source - self.msims[l.name] = msi_utils.ensure_dim( - viewer_utils.image_layer_to_msim(l, self.viewer), - 't', - ) - # else: - # # inform user about the consequences of modifying transforms - # if self.visualization_type_rbuttons.value == CHOICE_METADATA: - # notifications.notification_manager.receive_info( - # 'Please reload the layers for a new registration.' - # ) - # elif self.visualization_type_rbuttons.value == CHOICE_REGISTERED: - # notifications.notification_manager.receive_info( - # 'Manual corrections of transforms will be supported soon!' - # ) + if self._updating_viewer: + return + + if event.type not in ['affine', 'scale', 'translate']: + return + + l = event.source + if l.name not in self.msims: + return + + # Post-registration + showing Registered: live update current tp + if (self.visualization_type_rbuttons.enabled and + self.visualization_type_rbuttons.value == CHOICE_REGISTERED): + self._update_registered_param_for_current_tp(l) def link_channel_layers(self, layers, attributes=('contrast_limits', 'visible')): @@ -636,7 +748,7 @@ def __del__(self): print('Deleting napari-stitcher widget') # clean up callbacks - self.viewer.dims.events.disconnect(self.update_viewer_transformations) + self.viewer.dims.events.current_step.disconnect(self.update_viewer_transformations) for l in self.viewer.layers: if l.name in self.layers_selection.choices: diff --git a/src/napari_stitcher/_tests/test_stitcher_widget.py b/src/napari_stitcher/_tests/test_stitcher_widget.py index 5973657..ac52249 100644 --- a/src/napari_stitcher/_tests/test_stitcher_widget.py +++ b/src/napari_stitcher/_tests/test_stitcher_widget.py @@ -10,7 +10,7 @@ viewer_utils, ) -from multiview_stitcher import msi_utils, registration, mv_graph +from multiview_stitcher import msi_utils, registration, mv_graph, spatial_image_utils from multiview_stitcher.io import METADATA_TRANSFORM_KEY from multiview_stitcher.sample_data import ( get_mosaic_sample_data_path, generate_tiled_dataset) @@ -278,3 +278,159 @@ def test_fuse_register_buttons_not_grayed_out(make_napari_viewer): # Check if the buttons are not grayed out assert wdg.button_stitch.enabled assert wdg.button_fuse.enabled + + +def test_manual_transform_pre_registration(make_napari_viewer): + """ + Pre-registration: manual layer adjustments should be captured into + affine_metadata when _capture_layer_transforms_to_msims is called + (as happens at the start of run_registration / run_fusion). + """ + viewer = make_napari_viewer() + wdg = StitcherQWidget(viewer) + viewer.window.add_dock_widget(wdg) + + ndim = 2 + sims = generate_tiled_dataset( + ndim=ndim, N_t=1, N_c=1, + tile_size=30, tiles_x=2, tiles_y=1, tiles_z=1, + overlap=5, zoom=10, dtype=np.uint8) + msims = [msi_utils.get_msim_from_sim(sim, scale_factors=[]) for sim in sims] + layer_tuples = viewer_utils.create_image_layer_tuples_from_msims( + msims, transform_key=METADATA_TRANSFORM_KEY) + for lt in layer_tuples: + viewer.add_image(lt[0], **lt[1]) + + wdg.button_load_layers_all.clicked() + + # Manually move the first tile in the viewer (simulating the transform tool) + l = wdg.input_layers[0] + original_affine = np.array(l.affine.affine_matrix).copy() + custom_affine = original_affine.copy() + custom_affine[-(ndim), -1] += 25 # shift in last spatial dim + + # Setting affine while _updating_viewer=False simulates a user drag + wdg._updating_viewer = False + l.affine = custom_affine + + # Directly call _capture (run_registration calls this when not in CHOICE_REGISTERED) + wdg._capture_layer_transforms_to_msims() + + sim = msi_utils.get_sim_from_msim(wdg.msims[l.name]) + stored = np.array( + spatial_image_utils.get_affine_from_sim(sim, 'affine_metadata').isel(t=0) + ).squeeze() + expected = custom_affine[-(ndim + 1):, -(ndim + 1):] + assert np.allclose(stored, expected), \ + f"affine_metadata should reflect the manual adjustment.\nExpected:\n{expected}\nGot:\n{stored}" + + +@pytest.mark.parametrize("ndim", [2, 3]) +def test_manual_registered_transform(ndim, make_napari_viewer): + """ + Post-registration with Show=Registered: manually moving a layer should + update affine_registered for the CURRENT timepoint only; all other + timepoints must remain unchanged. + """ + N_t = 3 + viewer = make_napari_viewer() + wdg = StitcherQWidget(viewer) + viewer.window.add_dock_widget(wdg) + + sims = generate_tiled_dataset( + ndim=ndim, N_t=N_t, N_c=1, + tile_size=30, tiles_x=2, tiles_y=1, tiles_z=1, + overlap=5, zoom=10, dtype=np.uint8) + msims = [msi_utils.get_msim_from_sim(sim, scale_factors=[]) for sim in sims] + layer_tuples = viewer_utils.create_image_layer_tuples_from_msims( + msims, transform_key=METADATA_TRANSFORM_KEY) + for lt in layer_tuples: + viewer.add_image(lt[0], **lt[1]) + + wdg.button_load_layers_all.clicked() + wdg.run_registration() + wdg.visualization_type_rbuttons.value = _stitcher_widget.CHOICE_REGISTERED + + # Use the second tile (index 1) – it typically has a non-identity registered transform + l = wdg.input_layers[1] + sim = msi_utils.get_sim_from_msim(wdg.msims[l.name]) + original_params = spatial_image_utils.get_affine_from_sim( + sim, 'affine_registered').copy() + + # Navigate to tp=1 + current_step = list(viewer.dims.current_step) + current_step[-ndim - 1] = 1 + viewer.dims.current_step = tuple(current_step) + + # Build a custom affine that is different from the current one at tp=1 + custom_affine = np.array(l.affine.affine_matrix).copy() + custom_affine[-(ndim), -1] += 99 # big shift in last spatial dim + + # Simulate user drag (flag must be False) + wdg._updating_viewer = False + l.affine = custom_affine + + # Re-fetch updated params + updated_sim = msi_utils.get_sim_from_msim(wdg.msims[l.name]) + updated_params = spatial_image_utils.get_affine_from_sim( + updated_sim, 'affine_registered') + + # tp=1 should now contain the custom affine (lower-right (ndim+1) block) + p_modified = np.array(updated_params.isel(t=1)).squeeze() + assert np.allclose( + p_modified, custom_affine[-(ndim + 1):, -(ndim + 1):] + ), "affine_registered for tp=1 should equal the custom affine" + + # tp=0 and tp=2 must be unchanged + for tp_idx in [0, 2]: + p_orig = np.array(original_params.isel(t=tp_idx)).squeeze() + p_new = np.array(updated_params.isel(t=tp_idx)).squeeze() + assert np.allclose(p_orig, p_new), \ + f"affine_registered for tp={tp_idx} should be unchanged" + + +def test_manual_transform_show_original_no_msim_update(make_napari_viewer): + """ + Post-registration with Show=Original: manually moving a layer must NOT + update affine_metadata in the msim (it will be captured at registration + time, not live). + """ + viewer = make_napari_viewer() + wdg = StitcherQWidget(viewer) + viewer.window.add_dock_widget(wdg) + + ndim = 2 + sims = generate_tiled_dataset( + ndim=ndim, N_t=2, N_c=1, + tile_size=30, tiles_x=2, tiles_y=1, tiles_z=1, + overlap=5, zoom=10, dtype=np.uint8) + msims = [msi_utils.get_msim_from_sim(sim, scale_factors=[]) for sim in sims] + layer_tuples = viewer_utils.create_image_layer_tuples_from_msims( + msims, transform_key=METADATA_TRANSFORM_KEY) + for lt in layer_tuples: + viewer.add_image(lt[0], **lt[1]) + + wdg.button_load_layers_all.clicked() + wdg.run_registration() + + # Switch to Original + wdg.visualization_type_rbuttons.value = _stitcher_widget.CHOICE_METADATA + + l = wdg.input_layers[0] + sim = msi_utils.get_sim_from_msim(wdg.msims[l.name]) + original_metadata = spatial_image_utils.get_affine_from_sim( + sim, 'affine_metadata').copy() + + # Simulate user drag while showing Original + wdg._updating_viewer = False + custom_affine = np.array(l.affine.affine_matrix).copy() + custom_affine[-(ndim), -1] += 77 + l.affine = custom_affine + + # affine_metadata in the msim must remain unchanged + updated_sim = msi_utils.get_sim_from_msim(wdg.msims[l.name]) + updated_metadata = spatial_image_utils.get_affine_from_sim( + updated_sim, 'affine_metadata') + assert np.allclose( + np.array(updated_metadata), np.array(original_metadata) + ), "affine_metadata should NOT be updated live when showing Original"