Source code for spacr.segmentation_uncertainty

"""Headless segmentation-uncertainty scoring and lossless map export.

Four test-time transforms of the primary model define reference objects.
An optional second model contributes four equally weighted label sets;
near-miss probabilities and flow diagnostics remain those of the primary
model, whose threshold and object identities the user selected. Scores rank
review effort; they are not calibrated probabilities of biological error.
"""
from __future__ import annotations

import inspect
import json
import os
import tempfile
from pathlib import Path

import numpy as np

__all__ = ["compute_uncertainty", "save_uncertainty_map", "compute_queue_uncertainty"]


[docs] def compute_uncertainty(image, segment, *, second_segment=None, probability_threshold=0.0): """Score aligned TTA passes, optionally adding a second-model ensemble. :param image: one source field, unchanged by this function. :param segment: primary callable returning labels or labels/probability/flows. :param second_segment: optional second callable with the same contract. :param probability_threshold: the primary model's cell-probability threshold. :returns: the existing uncertainty score dictionary, including a float32 map. """ from .active_learning import _segmentation_uncertainty, _tta_passes passes = _tta_passes(image, segment) labels = list(passes["labels"]) if second_segment is not None: labels.extend(_tta_passes(image, second_segment)["labels"]) return _segmentation_uncertainty( labels, probabilities=passes["probabilities"], vectors=passes["vectors"], probability_threshold=probability_threshold)
[docs] def save_uncertainty_map(path, result, *, provenance=None, protected_paths=()): """Atomically save a float32 TIFF map with embedded JSON provenance. :param path: destination TIFF; parent directories are created if necessary. :param result: uncertainty score dictionary containing a finite 2-D map. :param provenance: source/model/settings metadata to embed in the TIFF. :param protected_paths: scientific source images or masks never overwritten. :returns: the written Path; failed writes retain the previous destination. :raises ValueError: for invalid maps or a protected destination. Metadata is validated before any file is opened; integer-keyed object scores are JSON-safe. """ import tifffile from .tiff_io import write_tiff destination = Path(path).expanduser().resolve() if destination.suffix.lower() not in {".tif", ".tiff"}: raise ValueError("An uncertainty map must be saved as .tif or .tiff.") for source in protected_paths: if source is not None and destination == Path(source).expanduser().resolve(): raise ValueError("An uncertainty map cannot overwrite a source image or mask.") if destination.exists(): try: with tifffile.TiffFile(destination) as existing: previous = json.loads(existing.pages[0].description) except (OSError, ValueError, KeyError, TypeError): previous = {} if not isinstance(previous, dict) or "spacr_uncertainty" not in previous: raise ValueError("Choose a new filename; this file is not an uncertainty map.") array = np.asarray(result["map"], dtype=np.float32) if (array.ndim != 2 or not array.size or not np.isfinite(array).all() or np.any(array < 0) or np.any(array > 1)): raise ValueError("An uncertainty map must be a finite 2-D array in [0, 1].") metadata = dict(provenance or {}) metadata["scores"] = {key: value for key, value in result.items() if key != "map"} metadata = json.loads(json.dumps(metadata, allow_nan=False, default=str)) destination.parent.mkdir(parents=True, exist_ok=True) fd, temporary = tempfile.mkstemp(prefix=".uncertainty-", suffix=".tif", dir=destination.parent) os.close(fd) try: write_tiff(temporary, array, photometric="minisblack", compression="zlib", metadata={"axes": "YX", "spacr_uncertainty": metadata}) os.replace(temporary, destination) finally: if os.path.exists(temporary): os.unlink(temporary) return destination
def _make_segmenter(model_name, device, parameters): """Load a repository-supported model without importing the Qt interface. :param model_name: stock model, local checkpoint or supported backend prefix. :param device: explicit torch device, normally cpu. :param parameters: Cellpose inference keyword arguments. :returns: a callable returning labels, cell-probability logits and flow vectors. """ from . import _segmentation_backends as backends from .spacr_cellpose import cellpose_channel_axis, parse_cellpose4_output prefixed = (backends._cellpose3_choice(model_name) is not None or backends._cellpose_dino_choice(model_name) is not None or backends._prefixed_backend(model_name) is not None) if prefixed: backend = str(model_name).split(":", 1)[0] model = backends._load_backend(backend, model_name=model_name, device=device) else: import torch from cellpose.models import CellposeModel from .utils import _resolve_cellpose_pretrained target = torch.device(device) kwargs = dict(pretrained_model=_resolve_cellpose_pretrained(model_name), device=target, gpu=target.type != "cpu") if "use_bfloat16" in inspect.signature(CellposeModel).parameters: kwargs["use_bfloat16"] = False model = CellposeModel(**kwargs) def segment(image): """Infer one transformed field through the shared output parser. :param image: 2-D field supplied by the TTA scorer. :returns: labels, logits and vector flows in the input orientation. """ kwargs = dict(parameters, channel_axis=cellpose_channel_axis(image), batch_size=1) signature = inspect.signature(model.eval).parameters if not any(p.kind == inspect.Parameter.VAR_KEYWORD for p in signature.values()): kwargs = {key: value for key, value in kwargs.items() if key in signature} masks, _rgb, vectors, probabilities, _other = parse_cellpose4_output(model.eval([image], **kwargs)) labels = np.asarray(masks[0] if np.asarray(masks).ndim > 2 else masks, dtype=np.int32) return (labels, probabilities[0] if probabilities else None, vectors[0] if vectors else None) return segment
[docs] def compute_queue_uncertainty(queue, *, model="cpsam", second_model=None, device="cpu", map_folder=None, parameters=None, progress=None, segmenter_factory=None): """Score a bounded curation queue headlessly and persist its ranking. :param queue: an existing CurationQueue; only its pending selected fields run. :param model: primary model name or checkpoint, default cpsam. :param second_model: optional distinct second model; absent means four passes. :param device: Device used for model predictions. The default is ``cpu``. :param map_folder: optional folder for lossless maps and embedded provenance. :param parameters: inference overrides for diameter, normalize and thresholds. :param progress: optional callback receiving one completed field's stem. :param segmenter_factory: injectable model/device/parameters loader for tests. :returns: Dictionary of score summaries keyed by image filename without its extension. Each summary excludes the pixel map. :raises ValueError: for duplicate ensemble models or a field that cannot run. This is pure image I/O: mask_engine imports no PySide6 and creates no Qt objects. Each completed field is saved as it finishes, so it is recoverable if a later field or model fails. """ from .curation_queue import _write_uncertainty from .qt.mask_engine import load_image_and_mask if second_model and str(second_model) == str(model): raise ValueError("The ensemble model must differ from the primary model.") if not queue.items: return {} protected = [path for item in queue.items for path in (item.image, item.mask, item.bundle) if path is not None] options = dict(diameter=None, normalize=True, flow_threshold=0.4, cellprob_threshold=0.0, min_size=0) options.update(parameters or {}) factory = segmenter_factory or _make_segmenter primary = factory(model, device, options) secondary = factory(second_model, device, options) if second_model else None scores = {} for item in queue.items: source = item.bundle if item.bundle is not None else item.image layout = {"masks_dir": str(item.mask.parent)} if item.mask is not None and item.bundle is None else {} image, _mask = load_image_and_mask(str(source.parent), source.name, **layout) result = compute_uncertainty(image, primary, second_segment=secondary, probability_threshold=options["cellprob_threshold"]) provenance = dict(source=str(source.resolve()), primary_model=str(model), second_model=second_model, device=device, parameters=options, transforms=["identity", "flip_lr", "flip_ud", "rot90"]) if map_folder is not None: save_uncertainty_map(Path(map_folder) / f"{item.stem}_uncertainty.tif", result, provenance=provenance, protected_paths=protected) summary = {key: value for key, value in result.items() if key != "map"} scores[item.stem] = summary _write_uncertainty(queue.folder, {item.stem: dict( uncertainty=result["field"], n_objects=result["n_objects"], passes=result["n_passes"], model=" + ".join(filter(None, (model, second_model))))}) if progress is not None: progress(item.stem) return scores