Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 12 additions & 1 deletion deeplabcut/gui/tabs/create_videos.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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()
15 changes: 12 additions & 3 deletions deeplabcut/gui/tabs/label_frames.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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)
54 changes: 32 additions & 22 deletions deeplabcut/gui/widgets.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Comment thread
C-Achard marked this conversation as resolved.
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)
Expand All @@ -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):
Expand All @@ -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)
127 changes: 93 additions & 34 deletions deeplabcut/utils/skeleton.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
Licensed under GNU Lesser General Public License v3.0
"""

import logging
import os
import warnings

Expand All @@ -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")
Expand Down Expand Up @@ -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
Expand All @@ -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)

Expand Down Expand Up @@ -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.",
)
Comment thread
C-Achard marked this conversation as resolved.
# 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:
Expand All @@ -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()
8 changes: 6 additions & 2 deletions tests/utils/test_skeleton.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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(
Expand All @@ -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"] == [
Expand Down
Loading