diff --git a/deeplabcut/gui/tabs/create_videos.py b/deeplabcut/gui/tabs/create_videos.py index dcbec6a232..0fba518a59 100644 --- a/deeplabcut/gui/tabs/create_videos.py +++ b/deeplabcut/gui/tabs/create_videos.py @@ -27,6 +27,8 @@ class CreateVideos(DefaultTab): def __init__(self, root, parent, h1_description): super().__init__(root, parent, h1_description) + self.skeleton_builder = None + self.bodyparts_to_use = self.root.all_bodyparts self._set_page() @@ -71,6 +73,9 @@ def _set_page(self): self.help_button.clicked.connect(self.show_help_dialog) self.main_layout.addWidget(self.help_button, alignment=Qt.AlignLeft) + def _on_skeleton_builder_destroyed(self): + self.skeleton_builder = None + def show_help_dialog(self): dialog = QtWidgets.QDialog(self) layout = QtWidgets.QVBoxLayout() @@ -287,4 +292,10 @@ def create_videos(self): def build_skeleton(self, *args): from deeplabcut.gui.widgets import SkeletonBuilder - SkeletonBuilder(self.root.config) + if self.skeleton_builder is None: + self.skeleton_builder = SkeletonBuilder( + config_path=self.root.config, + parent=self.root, + ) + self.skeleton_builder.destroyed.connect(self._on_skeleton_builder_destroyed) + self.skeleton_builder.show() diff --git a/deeplabcut/gui/tabs/label_frames.py b/deeplabcut/gui/tabs/label_frames.py index 509b04fcd3..661875f2ca 100644 --- a/deeplabcut/gui/tabs/label_frames.py +++ b/deeplabcut/gui/tabs/label_frames.py @@ -18,8 +18,7 @@ from deeplabcut.generate_training_dataset import check_labels from deeplabcut.gui.components import DefaultTab -from deeplabcut.gui.widgets import launch_napari -from deeplabcut.utils.skeleton import SkeletonBuilder +from deeplabcut.gui.widgets import SkeletonBuilder, launch_napari def label_frames(config_path: str | Path | None = None, image_folder: str | None = None): @@ -103,6 +102,7 @@ def __init__(self, root, parent, h1_description): super().__init__(root, parent, h1_description) self._set_page() + self.skeleton_builder = None def _set_page(self): self.label_frames_btn = QtWidgets.QPushButton("Label Frames") @@ -139,5 +139,14 @@ def check_labels(self): labeled_images = (Path(self.root.config).parent / "labeled-data").rglob("*_labeled/*.png") _ = launch_napari(labeled_images, plugin="napari", stack=True) + def _on_skeleton_builder_destroyed(self): + self.skeleton_builder = None + def build_skeleton(self, *args): - SkeletonBuilder(self.root.config) + if self.skeleton_builder is None: + self.skeleton_builder = SkeletonBuilder( + config_path=self.root.config, + parent=self.root, + ) + self.skeleton_builder.show() + self.skeleton_builder.destroyed.connect(self._on_skeleton_builder_destroyed) diff --git a/deeplabcut/gui/widgets.py b/deeplabcut/gui/widgets.py index f8a1e4c932..138787883c 100644 --- a/deeplabcut/gui/widgets.py +++ b/deeplabcut/gui/widgets.py @@ -19,7 +19,7 @@ from matplotlib.figure import Figure from matplotlib.widgets import Button, LassoSelector, RectangleSelector from PySide6 import QtCore, QtWidgets -from PySide6.QtGui import QAction, QCursor, QStandardItem, QStandardItemModel +from PySide6.QtGui import QAction, QCursor, QStandardItem, QStandardItemModel, Qt from deeplabcut.utils import auxiliaryfunctions from deeplabcut.utils.auxfun_videos import VideoWriter @@ -522,30 +522,32 @@ def display_help(self, *args): ) -class SkeletonBuilder(QtWidgets.QDialog, BaseSkeletonBuilder): - def __init__(self, config_path, parent=None): +class SkeletonBuilder(BaseSkeletonBuilder, QtWidgets.QDialog): + def __init__(self, config_path, *, parent=None): QtWidgets.QDialog.__init__(self, parent) self._parent = parent self.setWindowTitle("Skeleton Builder") + self.setAttribute(Qt.WA_DeleteOnClose, True) BaseSkeletonBuilder.__init__(self, config_path) def build_ui(self): self.fig = Figure() - self.ax = self.fig.add_subplot(111) - self.ax.axis("off") + self.canvas = FigureCanvas(self.fig) + + self._ax = self.fig.add_subplot(111) + self._ax.axis("off") - ax_clear = self.fig.add_axes([0.85, 0.55, 0.1, 0.1]) - ax_export = self.fig.add_axes([0.85, 0.45, 0.1, 0.1]) + ax_clear = self.fig.add_axes(self.clear_button_axes) + ax_export = self.fig.add_axes(self.export_button_axes) - self.clear_button = Button(ax_clear, "Clear") + self.clear_button = Button(ax_clear, self.clear_button_text) self.clear_button.on_clicked(self.clear) - self.export_button = Button(ax_export, "Export") + self.export_button = Button(ax_export, self.export_button_text) self.export_button.on_clicked(self.export) self.fig.canvas.mpl_connect("pick_event", self.on_pick) - self.canvas = FigureCanvas(self.fig) layout = QtWidgets.QVBoxLayout(self) layout.addWidget(self.canvas) self.setLayout(layout) @@ -554,18 +556,17 @@ def build_ui(self): hi = np.nanmax(self.xy, axis=0) center = (hi + lo) / 2 w, h = hi - lo - ampl = 1.3 - w *= ampl - h *= ampl - - self.ax.set_xlim(center[0] - w / 2, center[0] + w / 2) - self.ax.set_ylim(center[1] - h / 2, center[1] + h / 2) - self.ax.imshow(self.image) - self.ax.scatter(*self.xy.T, s=self.cfg["dotsize"] ** 2) - self.ax.add_collection(self.lines) - self.ax.invert_yaxis() - - self.lasso = LassoSelector(self.ax, onselect=self.on_select) + w *= self.ampl + h *= self.ampl + + self._ax.set_xlim(center[0] - w / 2, center[0] + w / 2) + self._ax.set_ylim(center[1] - h / 2, center[1] + h / 2) + self._ax.imshow(self.image) + self._ax.scatter(*self.xy.T, s=self.cfg["dotsize"] ** 2) + self._ax.add_collection(self.lines) + self._ax.invert_yaxis() + + self.lasso = LassoSelector(self._ax, onselect=self.on_select) self.canvas.draw_idle() def read_config(self, config_path): @@ -581,3 +582,12 @@ def write_config(self, config_path, cfg): def display(self): # No-op, the dialog is shown/exec'd by the caller pass + + def export(self, *args): + success = super().export(*args) + if success: + self._parent.logger.info("Skeleton exported successfully.") + self._parent.status_bar.showMessage("Skeleton exported successfully.", 5000) + else: + self._parent.logger.warning("Failed to export skeleton.") + self._parent.status_bar.showMessage("Failed to export skeleton.", 5000) diff --git a/deeplabcut/utils/skeleton.py b/deeplabcut/utils/skeleton.py index 738c9b064e..a16a31760b 100644 --- a/deeplabcut/utils/skeleton.py +++ b/deeplabcut/utils/skeleton.py @@ -18,6 +18,7 @@ Licensed under GNU Lesser General Public License v3.0 """ +import logging import os import warnings @@ -33,12 +34,23 @@ from deeplabcut.core.config import read_config_as_dict, write_config from deeplabcut.generate_training_dataset.trainingsetmanipulation import drop_likelihood_columns +logger = logging.getLogger(__name__) + class SkeletonBuilder: + ### Usage parameters + lasso_select_size = 10 + clear_button_axes = [0.85, 0.55, 0.1, 0.1] + clear_button_text = "Clear" + export_button_axes = [0.85, 0.45, 0.1, 0.1] + export_button_text = "Save" + ampl = 1.3 # Amplification factor for the zoomed-in view of the animal + def __init__(self, config_path): self.config_path = config_path self.cfg = read_config_as_dict(config_path) # Find uncropped labeled data + self._ax = None self.df = None found = False root = os.path.join(self.cfg["project_path"], "labeled-data") @@ -74,8 +86,9 @@ def __init__(self, config_path): self.inds = set() self.segs = set() # Draw the skeleton if already existent - if self.cfg["skeleton"]: - for bone in self.cfg["skeleton"]: + skeleton = self.cfg.get("skeleton", []) + if skeleton: + for bone in skeleton: pair = np.flatnonzero(self.bpts.isin(bone)) if len(pair) != 2: continue @@ -89,28 +102,27 @@ def __init__(self, config_path): def build_ui(self): self.fig = plt.figure() - ax = self.fig.add_subplot(111) - ax.axis("off") + self._ax = self.fig.add_subplot(111) + self._ax.axis("off") lo = np.nanmin(self.xy, axis=0) hi = np.nanmax(self.xy, axis=0) center = (hi + lo) / 2 w, h = hi - lo - ampl = 1.3 - w *= ampl - h *= ampl - ax.set_xlim(center[0] - w / 2, center[0] + w / 2) - ax.set_ylim(center[1] - h / 2, center[1] + h / 2) - ax.imshow(self.image) - ax.scatter(*self.xy.T, s=self.cfg["dotsize"] ** 2) - ax.add_collection(self.lines) - ax.invert_yaxis() - - self.lasso = LassoSelector(ax, onselect=self.on_select) - ax_clear = self.fig.add_axes([0.85, 0.55, 0.1, 0.1]) - ax_export = self.fig.add_axes([0.85, 0.45, 0.1, 0.1]) - self.clear_button = Button(ax_clear, "Clear") + w *= self.ampl + h *= self.ampl + self._ax.set_xlim(center[0] - w / 2, center[0] + w / 2) + self._ax.set_ylim(center[1] - h / 2, center[1] + h / 2) + self._ax.imshow(self.image) + self._ax.scatter(*self.xy.T, s=self.cfg["dotsize"] ** 2) + self._ax.add_collection(self.lines) + self._ax.invert_yaxis() + + self.lasso = LassoSelector(self._ax, onselect=self.on_select) + ax_clear = self.fig.add_axes(self.clear_button_axes) + ax_export = self.fig.add_axes(self.export_button_axes) + self.clear_button = Button(ax_clear, self.clear_button_text) self.clear_button.on_clicked(self.clear) - self.export_button = Button(ax_export, "Export") + self.export_button = Button(ax_export, self.export_button_text) self.export_button.on_clicked(self.export) self.fig.canvas.mpl_connect("pick_event", self.on_pick) @@ -145,17 +157,53 @@ def read_config(self, config_path): def write_config(self, config_path, cfg): write_config(config_path, cfg) - def export(self, *args): - inds_flat = set(ind for pair in self.inds for ind in pair) - unconnected = [i for i in range(len(self.xy)) if i not in inds_flat] - if len(unconnected): - warnings.warn( - "You didn't connect all the bodyparts (which is fine!). This is just a note to let you know.", - stacklevel=2, - ) - # sort to ensure consistent order in config.yaml - self.cfg["skeleton"] = [tuple(self.bpts[list(pair)]) for pair in sorted(self.inds)] - self.write_config(self.config_path, self.cfg) + def _show_export_feedback(self): + if not hasattr(self, "export_button"): + return + + button = self.export_button + canvas = self.fig.canvas + + original_text = button.label.get_text() + original_color = button.ax.get_facecolor() + + n_edges = len(self.cfg.get("skeleton") or []) + button.label.set_text(f"Saved {n_edges}") + button.ax.set_facecolor("#c8e6c9") # light green + canvas.draw_idle() + + def reset_button(): + button.label.set_text(original_text) + button.ax.set_facecolor(original_color) + canvas.draw_idle() + return False # stop Matplotlib timer + + timer = canvas.new_timer(interval=1200) + timer.add_callback(reset_button) + + # Keep a reference so the timer is not garbage-collected. + self._export_feedback_timer = timer + timer.start() + + def export(self, *args) -> bool: + try: + inds_flat = set(ind for pair in self.inds for ind in pair) + unconnected = [i for i in range(len(self.xy)) if i not in inds_flat] + # if empty, mention we are saving an empty skeleton + if not self.inds: + logger.warning("No bodyparts are connected. Saving an empty skeleton.") + elif len(unconnected): + logger.warning( + "Not all bodyparts are connected. Note that connecting all bodyparts is not necessary.", + ) + # sort to ensure consistent order in config.yaml + self.cfg["skeleton"] = [tuple(self.bpts[list(pair)]) for pair in sorted(self.inds)] + self.write_config(self.config_path, self.cfg) + self._show_export_feedback() + return True + except Exception as e: + logger.warning(f"Failed to export skeleton: {e}", stacklevel=2) + return False def on_pick(self, event): if event.mouseevent.button == 3: @@ -169,16 +217,27 @@ def on_pick(self, event): self.fig.canvas.draw_idle() def on_select(self, verts): - # self.path = Path(verts) - # self.verts = verts - inds = self.tree.query_ball_point(verts, 5) + # Transform keypoints and lasso vertices from image/data coordinates + # into display coordinates. This makes the grab radius independent of + # the image resolution and current zoom level. + xy_display = self._ax.transData.transform(self.xy) + verts_display = self._ax.transData.transform(np.asarray(verts)) + + tree_display = KDTree(xy_display) + inds = tree_display.query_ball_point( + verts_display, + self.lasso_select_size, + ) + inds_unique = [] for lst in inds: if len(lst) and lst[0] not in inds_unique: inds_unique.append(lst[0]) + for pair in zip(inds_unique, inds_unique[1:], strict=False): pair_sorted = tuple(sorted(pair)) self.inds.add(pair_sorted) self.segs.add(tuple(map(tuple, self.xy[pair_sorted, :]))) - self.lines.set_segments(self.segs) + + self.lines.set_segments(list(self.segs)) self.fig.canvas.draw_idle() diff --git a/tests/utils/test_skeleton.py b/tests/utils/test_skeleton.py index c83d5060d1..deba4d3c07 100644 --- a/tests/utils/test_skeleton.py +++ b/tests/utils/test_skeleton.py @@ -41,6 +41,9 @@ def make_test_builder(): def attach_fake_canvas(builder): builder.fig = Figure() + builder._ax = builder.fig.add_subplot(111) + builder._ax.set_xlim(-5, 25) + builder._ax.set_ylim(-5, 5) builder.fig.canvas.draw_idle = lambda: None @@ -136,7 +139,7 @@ def test_clear_resets_indices_segments_and_linecollection(): # --------------------------------------------------------------------- -def test_export_sorts_pairs_and_warns_for_unconnected(monkeypatch): +def test_export_sorts_pairs_and_warns_for_unconnected(monkeypatch, caplog): builder = make_test_builder() builder.config_path = "dummy_config.yaml" builder.xy = np.array( @@ -159,8 +162,9 @@ def fake_write_config(path, cfg): monkeypatch.setattr(skeleton_mod, "write_config", fake_write_config) - with pytest.warns(UserWarning, match="didn't connect all the bodyparts"): + with caplog.at_level("INFO"): builder.export() + assert "Not all bodyparts are connected" in caplog.text assert captured["path"] == "dummy_config.yaml" assert captured["cfg"]["skeleton"] == [