Source code for spacr.qt.screens.train_cellpose

"""Combined interface for fine-tuning and applying Cellpose models.

The workbench embeds the ``train_cellpose`` and ``cellpose_masks`` workflows
as Train and Apply tabs, each backed by its own
:class:`~spacr.qt.screens.app_screen.AppScreen`. The Train tab reads
``src`` and ``mask_src`` (default ``<src>/masks``) and writes model checkpoints
beneath ``<src>/models``. The Apply tab reads ``<src>/*.tif`` from its
independent source directory and writes masks to ``<src>/masks``.

When the active tab changes, settings named by the ``cellpose_masks``
propagation map in :mod:`spacr.qt.preview_registry` are copied to the
destination tab when that tab exposes the corresponding setting. ``src`` and
``model_name`` are not copied. On entry to Apply, the workbench searches
``<train-src>/models/cellpose_model/models`` for checkpoints whose names begin
with the Train tab's ``model_name``. It prefers completed checkpoints over
periodic ``_epoch_`` checkpoints and selects the most recently modified
candidate from the preferred group. If found, that path is assigned to
``custom_model``, which
:func:`spacr.spacr_cellpose.identify_masks_finetune` resolves before
``model_name``; otherwise, Apply retains its current model selection.

Each tab retains its own settings model, console, Run action, drop handling,
and any registered preview. Search, recipe, and preview integrations are
installed directly because embedded tabs do not become the current page in
the main-window stack.

The combined screen changes only GUI registration. The ``train_cellpose`` and
``cellpose_masks`` pipeline keys, command-line entry points, and settings-file
formats remain available.
"""
from __future__ import annotations

import logging
import os
from typing import Dict, List, Optional, Tuple

from PySide6.QtCore import Qt, Signal
from PySide6.QtWidgets import QLabel, QTabWidget, QVBoxLayout, QWidget

from ..i18n import tr
from ..theme import SPACING
from .app_screen import AppScreen, ModuleHeader

LOG = logging.getLogger("spacr.qt.screens.train_cellpose")

#: The registry key this screen answers to, and the training half's key.
TRAIN_KEY = "train_cellpose"

#: The applying half. Still a CLI module and still a validate entry; it is
#: the registry ROW that this screen replaced, not the module.
APPLY_KEY = "cellpose_masks"

#: The name the whole loop goes by. Named for the loop rather than for the
#: training half so somebody looking for "segment this folder with a stock
#: model" finds it.
WORKBENCH_TITLE = "Cellpose Workbench"

#: The one line under the name in the registry and beside it on the page.
WORKBENCH_INTRO = (
    "Fine-tune a Cellpose model on your own labelled fields, then segment "
    "a folder of images with it or with a stock model"
)

#: tab index -> (module key, the sentence that says what ``src`` means there).
#: The sentence is the whole reason there are two path fields, so it is shown
#: rather than left for the user to infer from the folder they picked.
TABS: Tuple[Tuple[str, str, str], ...] = (
    (TRAIN_KEY, "Train",
     "Choose an image folder and optional mask folder (default: images/masks). Train writes the "
     "checkpoint under <src>/models."),
    (APPLY_KEY, "Apply",
     "Apply reads every .tif in <src> and writes one mask per image into "
     "<src>/masks."),
)

#: How ``spacr.submodules.train_cellpose`` lays out what it saves, relative
#: to the training ``src``: it hands Cellpose ``<src>/models/cellpose_model``
#: as a save path, and Cellpose puts the checkpoint in a ``models`` folder
#: inside that.
CHECKPOINT_DIR = ("models", "cellpose_model", "models")

#: Infix Cellpose stamps onto the periodic saves it makes during a run. The
#: final save has no infix, so a finished run is preferred over a mid-run
#: snapshot of itself.
EPOCH_INFIX = "_epoch_"


[docs] def carried_setting_keys() -> Tuple[str, ...]: """The knobs copied from one tab to the other when the tab changes. Read from the propagation map :mod:`spacr.qt.preview_registry` declares for ``cellpose_masks`` — the map already answers "which settings is a Cellpose run judged by", and a second copy of that list here would be one that could disagree with it. ``model_name`` is excluded: it is the one shared name whose MEANING differs between the two halves, and it crosses as a checkpoint path instead (see :func:`CellposeWorkbenchScreen.trained_checkpoint`). ``src`` is not in the map at all, which is what keeps the two path fields independent. """ try: from ..preview_registry import PREVIEWS spec = PREVIEWS.get(APPLY_KEY) except Exception: # noqa: BLE001 LOG.debug("could not read the preview propagation map", exc_info=True) return () if spec is None: return () return tuple(name for name in spec.propagation.values() if name not in ("model_name", "src"))
def _install_screen_seams(screen: AppScreen) -> None: """Give an embedded module the strip, the recipes button and the preview. A module page normally collects these when the window's stack switches to it. These two are tabs rather than stack pages, so the watchers never see them and the seams are installed directly. Order matters: the search strip is what the other two hang their buttons on. Every step is guarded on its own — a missing preview toggle must not cost anyone the screen it would have sat above. """ try: from ..settings_search import install as install_search install_search(screen) except Exception: # noqa: BLE001 LOG.debug("no settings search strip for %s", screen.app_key, exc_info=True) try: from ..recipes import install as install_recipes install_recipes(screen) except Exception: # noqa: BLE001 LOG.debug("no recipes button for %s", screen.app_key, exc_info=True) try: from ..preview_registry import install as install_preview install_preview(screen) except Exception: # noqa: BLE001 LOG.debug("no preview for %s", screen.app_key, exc_info=True) class _VirtualStainApply(QWidget): """A button that applies a saved virtual-staining model to a folder. Asks for a model written by the Cellpose workbench's Virtual staining run and predicts its stain for every field of a folder on the CPU, writing ``<folder>/virtual_stain/<field>_virtual_c<target>.npy``. """ def __init__(self, folder=None, parent=None): """Build the button and its one-line note. The mounting screen names :attr:`button` and applies the alpha gate. :param folder: callable returning the screen's current folder. :param parent: owning widget. """ super().__init__(parent) from PySide6.QtWidgets import QHBoxLayout, QPushButton from ..job_runner import JobRunner self._folder = folder or (lambda: "") self.jobs = JobRunner(self, app_key="virtual stain") self.jobs.job_failed.connect(self._on_failed) row = QHBoxLayout(self) row.setContentsMargins(0, 0, 0, 0) self.button = QPushButton(tr("Apply virtual stain…"), self) self.button.setToolTip(tr( "Choose a virtual-staining model (.pt) saved by the Cellpose " "workbench and predict its stain for every .npy or .tif field " "of the folder, frame by frame for time stacks, on the CPU. " "Predictions are written to <folder>/virtual_stain. Default " "the screen's source folder.")) self.button.clicked.connect(lambda: self.apply()) row.addWidget(self.button) self.note = QLabel("", self) self.note.setWordWrap(True) self.note.setVisible(False) row.addWidget(self.note, 1) self.result = None def apply(self, model: str = "", folder: str = "") -> str: """Predict the model's stain for every field of a folder. :param model: saved model path; asks for one when empty. :param folder: field folder; the screen's folder, else asks. :returns: the folder used, or ``''`` when a dialog was dismissed. """ from PySide6.QtWidgets import QFileDialog if not model: model = QFileDialog.getOpenFileName( self, tr("Choose a virtual-staining model"), "", tr("Virtual-staining models (*.pt)"))[0] if not model: return "" folder = folder or str(self._folder() or "") if not folder or not os.path.isdir(folder): folder = QFileDialog.getExistingDirectory( self, tr("Choose a folder of fields")) if not folder: return "" model, folder = str(model), str(folder) self._say(tr("Applying the virtual stain…")) def work(): """Predict off the GUI thread.""" from ...deep_spacr import _apply_virtual_stain return _apply_virtual_stain(model, folder) self.jobs.submit(work, self._on_done) return folder def _say(self, text: str) -> None: """Show one line beside the button.""" self.note.setText(text) self.note.setVisible(True) def _on_done(self, written) -> None: """Say how many predictions were written.""" self.result = list(written) self._say(tr("Virtual stain written for {n} fields.").format( n=len(self.result))) def _on_failed(self, message: str) -> None: """Show why the job stopped.""" self._say(tr("Virtual staining failed: {error}").format( error=message))
[docs] class CellposeWorkbenchScreen(QWidget): """Train and Apply as two tabs of one module page. :ivar error_explain_requested: re-emitted from whichever tab raised, with ``(traceback, app_key)`` — the app key is the TAB's, so the AI console is asked about the module that actually failed. :ivar remote_submit_requested: re-emitted the same way, so submitting a run to Distributed Jobs from either tab carries that tab's key and that tab's settings. :param parent: parent widget. """ error_explain_requested = Signal(str, str) remote_submit_requested = Signal(str, dict) def __init__(self, parent: Optional[QWidget] = None): """Build the Cellpose workbench. The registry key is fixed rather than following the open tab: the window keys its screen table and its navigation by it, and a key that moved would make the page answer to one it is not filed under. :param parent: parent widget, or ``None``. """ super().__init__(parent) #: The registry key this page was opened under. Fixed, unlike #: :meth:`active_app_key` — the window keyed its screen table and #: its navigation by this, and a value that moved with the tab #: would make the page answer to a key it is not filed under. self.app_key = TRAIN_KEY outer = QVBoxLayout(self) outer.setContentsMargins(SPACING["lg"], SPACING["lg"], SPACING["lg"], SPACING["lg"]) outer.setSpacing(SPACING["md"]) self._header = ModuleHeader( tr(WORKBENCH_TITLE), description=tr(WORKBENCH_INTRO), instruction=tr(TABS[0][2]), app_key=TRAIN_KEY, ) outer.addWidget(self._header) self._tabs = QTabWidget(self) self._tabs.setObjectName("CellposeWorkbenchTabs") self._screens: List[AppScreen] = [] for app_key, label, _instruction in TABS: screen = AppScreen(app_key=app_key) header = getattr(screen, "_header", None) if header is not None: header.setVisible(False) _install_screen_seams(screen) screen.error_explain_requested.connect(self.error_explain_requested) screen.remote_submit_requested.connect(self.remote_submit_requested) self._screens.append(screen) self._tabs.addTab(screen, tr(label)) outer.addWidget(self._tabs, 1) #: Says what crossed when a checkpoint did. Hidden until one does — #: an empty reserved line reads as a thing that failed to load. self._carry_note = QLabel("", self) self._carry_note.setObjectName("Muted") self._carry_note.setWordWrap(True) self._carry_note.setTextInteractionFlags(Qt.TextSelectableByMouse) self._carry_note.setVisible(False) outer.addWidget(self._carry_note) self._add_virtual_stain(outer) self._current = self._tabs.currentIndex() self._tabs.currentChanged.connect(self._on_tab_changed) def _add_virtual_stain(self, outer) -> None: """The alpha button that learns to predict one channel from others.""" from PySide6.QtWidgets import QHBoxLayout, QPushButton from ..job_runner import JobRunner from ..preferences import _apply_alpha_widgets self._vs_jobs = JobRunner(self, app_key=TRAIN_KEY) self._vs_jobs.job_failed.connect(self._on_virtual_stain_failed) row = QHBoxLayout() self._vs_button = QPushButton(tr("Virtual staining…"), self) self._vs_button.setObjectName("CellposeWorkbenchVirtualStain") self._vs_button.setToolTip(tr( "Choose a folder of paired multichannel fields (.npy or .tif) and " "the channels to learn from and to predict, for example 1 > 0 to " "predict the nucleus stain from channel 1. A small U-Net is " "trained on the CPU, the last quarter of the fields is held out, " "and the real and predicted stains are segmented the same way and " "matched at IoU 0.5. The model, predictions and a score table are " "written to <folder>/virtual_stain. Default 20 epochs.")) self._vs_button.clicked.connect(lambda: self._virtual_stain()) row.addWidget(self._vs_button) row.addStretch(1) outer.addLayout(row) self._vs_note = QLabel("", self) self._vs_note.setWordWrap(True) self._vs_note.setTextInteractionFlags(Qt.TextSelectableByMouse) self._vs_note.setVisible(False) outer.addWidget(self._vs_note) _apply_alpha_widgets(self._vs_button) def _virtual_stain(self, folder: str = "", channels: str = "") -> str: """Train a virtual-staining model on a folder and score it. :param folder: folder of paired fields; asks for one when empty. :param channels: ``"<inputs> > <target>"``, for example ``"1,3 > 0"``, optionally followed by ``pix2pix`` to train the U-Net as a conditional GAN generator; asks when empty. :returns: the folder used, or ``''`` when a dialog was dismissed. """ from PySide6.QtWidgets import QFileDialog, QInputDialog if not folder: folder = QFileDialog.getExistingDirectory( self, tr("Choose a folder of paired multichannel fields")) if not folder: return "" if not channels: channels, ok = QInputDialog.getText( self, tr("Virtual staining"), tr("Input channels > channel to predict (add pix2pix for " "the adversarial model):"), text="1 > 0") if not ok: return "" inputs, _sep, goal = str(channels).partition(">") sources = [int(c) for c in inputs.replace(" ", "").split(",") if c] words = goal.split() target = int(words[0]) model_type = "pix2pix" if "pix2pix" in (w.lower() for w in words[1:]) \ else "unet" folder = str(folder) self._show_virtual_stain_note(tr("Training the virtual stain…")) def work(): """Train, predict and score off the GUI thread.""" from ...deep_spacr import _virtual_stain_from_folder return _virtual_stain_from_folder(folder, sources, target, epochs=20, model_type=model_type)[1] self._vs_jobs.submit(work, self._on_virtual_stain_done) return folder def _show_virtual_stain_note(self, text: str) -> None: """Show one line under the virtual-staining button.""" self._vs_note.setText(text) self._vs_note.setVisible(True) def _on_virtual_stain_done(self, summary) -> None: """Say how the predicted stain segments against the real one.""" self._vs_summary = dict(summary) self._show_virtual_stain_note(tr( "Virtual stain on {fields} held-out fields: F1 {f1:.2f} at IoU " "0.5 against the real stain's objects (input channel alone " "{base:.2f}), Pearson r {r:.2f}.").format( fields=int(summary["test_fields"]), f1=summary["predicted_f1_50"], base=summary["input_baseline_f1_50"], r=summary["predicted_pearson"])) def _on_virtual_stain_failed(self, message: str) -> None: """Show why the virtual-staining job stopped.""" self._show_virtual_stain_note( tr("Virtual staining failed: {error}").format(error=message))
[docs] def closeEvent(self, event): """Close both module pages before their owning workbench is destroyed. :param event: Qt close event; ignored when a page is finishing a write. """ for screen in self._screens: if not screen.close(): event.ignore() return super().closeEvent(event)
@property
[docs] def train_screen(self) -> AppScreen: """The Train tab's module page.""" return self._screens[0]
@property
[docs] def apply_screen(self) -> AppScreen: """The Apply tab's module page.""" return self._screens[1]
[docs] def screen_for(self, app_key: str) -> Optional[AppScreen]: """The tab that runs ``app_key``, or ``None``. :param app_key: the app key to find, compared as a string with each tab's ``app_key``. """ for screen in self._screens: if screen.app_key == str(app_key): return screen return None
[docs] def active_screen(self) -> AppScreen: """The module page the user is looking at.""" index = self._tabs.currentIndex() return self._screens[index if 0 <= index < len(self._screens) else 0]
[docs] def active_app_key(self) -> str: """The key of the module the user is looking at.""" return self.active_screen().app_key
@property def _settings_model(self): """The settings model of the visible tab. The name the shared tooling reaches for (the command palette's jump-to-setting, the recipes dialog, the walkthrough), so this page answers it with the form that is actually on screen rather than looking like a screen with no settings at all. """ return self.active_screen()._settings_model
[docs] def current_settings(self) -> Tuple[str, Dict]: """``(app_key, settings)`` for the visible tab.""" screen = self.active_screen() return screen.app_key, dict(screen._settings_model.collect())
[docs] def apply_settings_dict(self, settings: Dict) -> int: """Push ``settings`` into whichever tab the dict is for. A settings CSV, a restored session or a recipe belongs to one of the two modules, and applying it to both would write a training folder into the Apply tab's path field — the one thing the two-field split exists to prevent. The tab is chosen by which one owns more of the keys the OTHER one does not have, so ``n_epochs`` picks Train and ``flow_threshold`` picks Apply; a dict that distinguishes neither goes to the tab already on screen. The chosen tab is raised, so the settings are visible where they landed. :param settings: key/value pairs to apply. :returns: how many keys the chosen tab actually took. """ settings = dict(settings or {}) if not settings: return 0 target = self._tab_for_settings(settings) index = self._screens.index(target) if index != self._tabs.currentIndex(): self._tabs.setCurrentIndex(index) return target.apply_settings_dict(settings)
def _tab_for_settings(self, settings: Dict) -> AppScreen: """The tab a settings dict belongs to. Ties go to the visible one. Scored on the keys a tab does NOT share with the other, because the shared ones say nothing: ``diameter`` is both modules', ``n_epochs`` is only training's and ``flow_threshold`` is only applying's. """ given = set(settings) owned = [set(screen._settings_model._widgets) for screen in self._screens] best, best_score = self.active_screen(), 0 for index, screen in enumerate(self._screens): others = set().union(*(keys for position, keys in enumerate(owned) if position != index)) score = len(given & (owned[index] - others)) if score > best_score: best, best_score = screen, score return best
[docs] def apply_seed(self, seed: Dict) -> int: """Take a seed handed over by another screen. See :meth:`apply_settings_dict`, which decides where it lands. :param seed: settings name to value, passed unchanged to :meth:`apply_settings_dict`. """ return self.apply_settings_dict(seed)
def _on_tab_changed(self, index: int) -> None: """Carry the knobs into the tab being opened, and say what src means.""" previous, self._current = self._current, index if 0 <= previous < len(self._screens) and previous != index: try: self.carry(self._screens[previous], self.active_screen()) except Exception: # noqa: BLE001 LOG.exception("could not carry the Cellpose settings across") self._sync_instruction() def _sync_instruction(self) -> None: """Put the visible tab's reading of ``src`` under the title.""" index = self._tabs.currentIndex() _key, _label, instruction = TABS[index if 0 <= index < len(TABS) else 0] label = getattr(self._header, "instruction_label", None) if label is None: return label.setProperty("_spacr_i18n_text", instruction) label.setText(tr(instruction)) label.setVisible(True) help_label = getattr(self._header, "api_help", None) if help_label is not None and hasattr(help_label, "set_api_app_key"): help_label.set_api_app_key(self.active_app_key())
[docs] def carry(self, source: AppScreen, target: AppScreen) -> Dict: """Copy the shared knobs from ``source`` into ``target``. Only the keys :func:`carried_setting_keys` names, and only those ``target`` actually has a widget for — the two forms overlap partially, and a key the target does not offer would otherwise be written into its hidden values where nobody can see or change it. Entering the Apply tab additionally picks up the trained checkpoint. :param source: the tab being left. :param target: the tab being opened. :returns: what was written into ``target``. """ try: values = dict(source._settings_model.collect()) except Exception: # noqa: BLE001 LOG.debug("could not read %s's settings", source.app_key, exc_info=True) values = {} widgets = target._settings_model._widgets carried = {key: values[key] for key in carried_setting_keys() if key in values and widgets.get(key) is not None} if carried: target.apply_settings_dict(carried) if target is self.apply_screen: checkpoint = self.carry_trained_model() if checkpoint: carried["custom_model"] = checkpoint return carried
[docs] def carry_trained_model(self) -> str: """Point the Apply tab at the checkpoint the Train tab produced. Does nothing until a training run has actually written one: a ``custom_model`` naming a file that is not there stops :func:`spacr.spacr_cellpose.identify_masks_finetune` before it segments anything, which would break "run cpsam over this folder" for everyone who has never trained a model. :returns: the checkpoint path that was set, or ``""``. """ checkpoint = self.trained_checkpoint() if not checkpoint: return "" self.apply_screen.apply_settings_dict({"custom_model": checkpoint}) note = "Apply is set to the model you trained: {name}" self._carry_note.setProperty("_spacr_i18n_text", note) self._carry_note.setText( tr(note, name=os.path.basename(checkpoint))) self._carry_note.setVisible(True) return checkpoint
[docs] def trained_checkpoint(self) -> str: """The newest checkpoint the Train tab's settings would have written. Found by looking, not by rebuilding the file name: the training module stamps the architecture and epoch count into what it saves, and a second copy of that formula here would go stale the first time it changed. Everything under the training output folder whose name starts with the model name counts; a finished run's save is preferred over the periodic ones it made on the way, and among equals the most recently written wins. :returns: an absolute path, or ``""`` when there is nothing to find. """ try: values = self.train_screen._settings_model.collect() except Exception: # noqa: BLE001 LOG.debug("could not read the training settings", exc_info=True) return "" src = str(values.get("src") or "").strip() name = str(values.get("model_name") or "").strip() if not src or not name: return "" folder = (os.path.join(os.path.expanduser(values["save_path"]), "models") if values.get("save_path") else os.path.join(os.path.expanduser(src), *CHECKPOINT_DIR)) try: entries = sorted(os.listdir(folder)) except OSError: return "" prefix = f"{name}_" found = [os.path.join(folder, entry) for entry in entries if entry.startswith(prefix) and os.path.isfile(os.path.join(folder, entry))] finished = [path for path in found if EPOCH_INFIX not in os.path.basename(path)] candidates = finished or found if not candidates: return "" try: return max(candidates, key=os.path.getmtime) except OSError: return candidates[-1]
[docs] def build_screen(app_key: str = TRAIN_KEY, host=None) -> CellposeWorkbenchScreen: """Screen factory for the registry. Takes ``app_key`` and ``host`` because that is the contract ``spacr.qt.app._call_screen_factory`` offers; the key is fixed (this screen serves one row) and the host is used only to make the same two connections the window makes on a generic module page. """ screen = CellposeWorkbenchScreen() if host is not None: for signal_name, slot_name in ( ("error_explain_requested", "_on_explain_error"), ("remote_submit_requested", "_on_remote_submit_requested")): signal = getattr(screen, signal_name, None) slot = getattr(host, slot_name, None) if signal is not None and callable(slot): signal.connect(slot) return screen