Source code for spacr.plot

"""Scientific plotting and statistical-annotation helpers."""

from __future__ import annotations
import contextlib
import os, random, cv2, glob, math, torch, itertools

from typing import Optional, Tuple, Union

import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from matplotlib.colors import ListedColormap
import matplotlib as mpl
import scipy.ndimage as ndi
import seaborn as sns
import scipy.stats as stats
import statsmodels.api as sm
import imageio.v2 as imageio
from IPython.display import display
from skimage import measure
from skimage.measure import find_contours, label, regionprops
from skimage.transform import resize as sk_resize
import scikit_posthocs as sp
from scipy.stats import chi2_contingency
import tifffile as tiff
from statsmodels.stats.multicomp import pairwise_tukeyhsd

from ipywidgets import IntSlider, interact
from IPython.display import Image as ipyimage
from matplotlib_venn import venn2

from .errors import RunLedger, raise_if_strict
from .image_colors import read_image_rgb, write_image_rgb
from .tiff_io import write_tiff

from .figures.style import (ROLES, TYPE_SCALE, WEIGHTS, Palette, descriptor,
                            figure_style, _group_colours, hide_unused,
                            _preference_deltas,
                            panel_letter,
                            reference_line, resolve_ink, rotate_ticks,
                            text_legend, theme_target)

#: Shipped answers, used when the preference store cannot be reached — a
#: pipeline run from the CLI, a notebook, a machine with no Qt installed.
#: They are the same values the Preferences dialog defaults to, so a headless
#: run and a GUI run that never touched the setting produce the same files.
DEFAULT_FIGURE_FORMAT = "pdf"
DEFAULT_FIGURE_DPI = 300

#: Formats a kept figure can be written in.
FIGURE_FORMATS = ("png", "pdf", "svg", "tiff")

#: matplotlib's Agg backend refuses a canvas above 2**16 px on either axis and
#: is unusable well before that. `spacrGraph._standerdize_figure_format` forces
#: a **10 inch minimum** square canvas and grows it with the group count, so
#: the DPI the preference offers is not always deliverable: a 50-inch grouped
#: graph at 1200 DPI is 60 000 px square, four gigapixels. `deliverable_dpi`
#: says so instead of letting matplotlib raise -- or worse, letting the file
#: appear at a resolution nobody asked for.
MAX_FIGURE_PX = 2 ** 16 - 1

#: Above this the raster is not a figure any more. 200 megapixels is a
#: 14 000 px square, already past any display or printer.
MAX_FIGURE_MEGAPIXELS = 200


_HOUSE_PANEL_INCHES = 5.6


def _montage_type_size(figuresize, tier='label'):
    """The house type tier for ``tier``, scaled to a montage-sized canvas.

    :param figuresize: the panel edge in inches the caller asked for.
    :param tier: a key of :data:`spacr.figures.style.TYPE_SCALE`.
    :returns: a point size, never smaller than the tier itself.
    """
    try:
        scale = max(1.0, float(figuresize) / _HOUSE_PANEL_INCHES)
    except (TypeError, ValueError):
        scale = 1.0
    return TYPE_SCALE[tier] * scale


[docs] def figure_output_preferences(): """Return ``(format, dpi)`` from the user's preferences. A format or DPI changed in the Preferences figure settings wins over the "Figure format" and "Resolution" preferences. Degrades to :data:`DEFAULT_FIGURE_FORMAT` / :data:`DEFAULT_FIGURE_DPI` rather than raising: the preference store is Qt's, and the pipelines that call this run headless from the CLI and from notebooks, where importing PySide6 to decide a file extension would be absurd. """ try: from .qt.preferences import get_figure_format, get_figure_png_dpi fmt = str(get_figure_format()).strip().lower() dpi = int(get_figure_png_dpi()) except Exception: return DEFAULT_FIGURE_FORMAT, DEFAULT_FIGURE_DPI from .figures.style import _preference_deltas chosen = _preference_deltas() if chosen.get("format"): fmt = str(chosen["format"]).strip().lower() try: dpi = int(chosen.get("dpi", dpi)) except (TypeError, ValueError): pass if fmt not in FIGURE_FORMATS: fmt = DEFAULT_FIGURE_FORMAT if dpi <= 0: dpi = DEFAULT_FIGURE_DPI return fmt, dpi
[docs] def deliverable_dpi(fig, dpi, path=None): """The DPI this figure can actually be written at, and a word if it is not the one that was asked for. A resolution preference is a request, not a guarantee. ``spacrGraph`` pins its canvas to at least 10 inches square and grows it with the number of groups, so 600 and 1200 DPI are not available for a large grouped figure -- the raster would be tens of thousands of pixels on a side. The old behaviour was to hand the number to matplotlib and find out. This returns the DPI that will be used and says, by name, when that is not the DPI that was requested. Appearing to accept a setting and then quietly delivering another one is the failure this avoids. :param fig: the figure about to be written. :param dpi: the requested dots per inch. :param path: destination, named in the message when there is one. :returns: the DPI to pass to ``savefig``. """ try: width_in, height_in = (float(v) for v in fig.get_size_inches()) except Exception: return int(dpi) longest = max(width_in, height_in, 0.01) area = max(width_in * height_in, 0.0001) by_edge = MAX_FIGURE_PX / longest by_area = ((MAX_FIGURE_MEGAPIXELS * 1_000_000) / area) ** 0.5 ceiling = int(min(by_edge, by_area)) if ceiling >= int(dpi): return int(dpi) ceiling = max(72, ceiling) where = f" for {path}" if path else "" print( f"Figure DPI: {int(dpi)} was requested but this figure is " f"{width_in:.0f}x{height_in:.0f} inches, so {int(dpi)} DPI would be " f"{int(width_in * dpi)}x{int(height_in * dpi)} pixels. Writing at " f"{ceiling} DPI instead{where}. Grouped graphs (spacrGraph) pin the " "canvas to at least 10 inches and grow it with the group count, so " "the highest DPI settings cannot be delivered for them.") return ceiling
def _with_extension(path, fmt): """``path`` with its extension replaced by ``fmt``. Only a *known figure extension* is replaced. ``os.path.splitext`` on its own would turn ``plate_2.5_umap`` into ``plate_2.pdf`` -- and ``plate_2.6_umap`` into the same name, which is two figures overwriting one file. """ text = str(path) stem, extension = os.path.splitext(text) known = {".png", ".pdf", ".svg", ".jpg", ".jpeg", ".tif", ".tiff", ".eps"} if extension.lower() in known: return f"{stem}.{fmt}" return f"{text}.{fmt}"
[docs] def figure_path(path, fmt=None): """``path`` with the extension the figure-format preference will write. THE ONE PLACE A NON-MATPLOTLIB RENDERER CAN ASK. ``save_figure`` rewrites the extension itself, which is why every matplotlib save has honoured the preference for months -- but pyqtgraph decides what it writes FROM THE NAME (``FastPlot.export`` branches on ``.pdf`` / ``.svg`` / else), so a renderer that is handed ``volcano.pdf`` writes a PDF whatever the user chose. The name has to be settled before the export sees it, and settling it twice in two modules is how the two drift apart. :param path: a destination, with or without an extension. :param fmt: force a format, bypassing the preference. Unknown formats fall back to the preference rather than raising: a run must not lose a figure over a typo. :returns: the path as a ``str``. """ preferred, _ = figure_output_preferences() chosen = str(fmt or preferred).strip().lower().lstrip(".") if chosen not in FIGURE_FORMATS: chosen = preferred return _with_extension(path, chosen)
def _chrome(fig, ax=None): """Yield colour accessors for a figure's non-data artists. Return ``(kind, artist, getter, setter)`` tuples for ground, text, spines, ticks, legends, and reference lines. Collections, images, bars, and data lines are excluded so print styling cannot change encoded data colours. Reference lines are distinguished from plotted series by their blended axes/data transform; ordinary series use ``ax.transData``. """ from matplotlib.axes import Axes if ax is None: yield "ground", fig.patch, fig.patch.get_facecolor, fig.patch.set_facecolor for text in fig.texts: yield "text", text, text.get_color, text.set_color for legend in getattr(fig, "legends", []): yield from _legend_chrome(legend) for axes in fig.axes: if isinstance(axes, Axes): yield from _chrome(fig, axes) return yield "ground", ax.patch, ax.patch.get_facecolor, ax.patch.set_facecolor for spine in ax.spines.values(): yield "chrome", spine, spine.get_edgecolor, spine.set_edgecolor yield "text", ax.title, ax.title.get_color, ax.title.set_color for axis in (ax.xaxis, ax.yaxis): label = axis.label yield "text", label, label.get_color, label.set_color for tick in list(axis.get_major_ticks()) + list(axis.get_minor_ticks()): for line in (tick.tick1line, tick.tick2line): yield "chrome", line, line.get_color, line.set_color for text in (tick.label1, tick.label2): yield "text", text, text.get_color, text.set_color grid = tick.gridline yield "grid", grid, grid.get_color, grid.set_color legend = ax.get_legend() if legend is not None: yield from _legend_chrome(legend) for text in ax.texts: yield "text", text, text.get_color, text.set_color arrow = getattr(text, "arrow_patch", None) if arrow is not None: yield "chrome", arrow, arrow.get_edgecolor, arrow.set_color for line in ax.lines: if line.get_transform() is not ax.transData: yield "chrome", line, line.get_color, line.set_color def _legend_chrome(legend): """Yield colour accessors for a legend's frame, ground, and text.""" frame = legend.get_frame() yield "chrome", frame, frame.get_edgecolor, frame.set_edgecolor yield "ground", frame, frame.get_facecolor, frame.set_facecolor for text in legend.get_texts(): yield "text", text, text.get_color, text.set_color
[docs] def data_colours(fig): """Every colour in ``fig`` that carries the CLAIM rather than the frame. :param fig: Matplotlib figure whose data artists are inspected. Used only to say when one of them stops working on paper (150 D), never to change one. A data line is identified the same way `_chrome` identifies a reference line, from the opposite side of the same test. """ from matplotlib.axes import Axes found = [] for axes in fig.axes: if not isinstance(axes, Axes): continue for line in axes.lines: if line.get_transform() is axes.transData: found.append(line.get_color()) for collection in axes.collections: for getter in ("get_facecolor", "get_edgecolor"): try: found.extend(list(getattr(collection, getter)())) except Exception: # noqa: BLE001 pass for patch in axes.patches: try: found.append(patch.get_facecolor()) except Exception: # noqa: BLE001 pass return found
[docs] def illegible_data_colours(fig, ground, floor=None): """The data colours a reader will not find on ``ground``, as hex. :param fig: Matplotlib figure whose data colours are checked. :param ground: background colour against which contrast is measured. The data deliberately does NOT flip, so a palette chosen against near-black can be illegible on paper -- and the honest answer is to NAME the colour, because a substitution the user did not ask for changes what the picture says. Deduplicated and sorted so the same sentence comes out of the same figure twice. """ from .figure_style import illegible_colours return illegible_colours(data_colours(fig), ground, floor)
@contextlib.contextmanager
[docs] def save_figure(fig, path, *, fmt=None, dpi=None, close=False, save_mode=None, announce_colours=True, integrity=None, **kwargs): """Write ``fig`` to ``path``, honouring the figure preferences. The single place a spaCR figure the user keeps gets written. Before this existed there were sixty-odd ``savefig`` calls, each with its own hard-coded format and DPI, and the "Figure format" and "Resolution" preferences reached exactly two of them -- both writing to a temp directory. Everything a pipeline saved into its results folder ignored both settings entirely. Three things are decided here rather than left to the caller or to matplotlib. **The format follows the preference**, and the file NAME follows the format: a PNG written to ``figure.pdf`` is a file no viewer opens. An explicit ``fmt=`` still wins, for the few callers that genuinely need one particular format. **Fonts are embedded as TrueType** (``pdf.fonttype = 42``) for the length of the save. matplotlib's default is Type 3, which draws every glyph as its own content stream: the file is still vector, but Illustrator and Inkscape open the text as unselectable outlines, and the preference that selects this path is labelled "PDF (vector, editable)". Scoped with ``rc_context`` so a caller that has deliberately chosen otherwise is not changed underneath it. **The DPI is passed**, always. A PDF page is resolution-independent, but spaCR figures are full of ``imshow`` panels -- cell montages, mask overlays, plate heatmaps -- and those are rasterised at the figure's own 100 DPI unless told otherwise. Without it, 100, 300 and 600 produced byte-identical files. What is passed is :func:`deliverable_dpi`, which says so out loud when the requested number is not achievable here. **The chrome is repainted for paper**, and only for the length of the write. A figure saved from a dark-themed session was white ink on a white page -- `spacr.qt.preferences.get_figure_colors` hands both renderers a white foreground on a dark theme and nothing inverted it at export time, so the file was a blank rectangle with some coloured dots in it, and a transparent PNG even looked right in a dark file manager and disappeared when it was pasted into a manuscript. :func:`print_ready` moves the furniture and puts it back; the DATA is never touched, because a white point turned black is, on a volcano, the colour of "not a hit". The figure first takes the settings changed in the Preferences figure settings (:func:`spacr.figures.style._apply_user_style`), and a second copy is written when "Also save" names another format. :param fig: a matplotlib ``Figure``. :param path: destination; its extension is corrected to the format. :param fmt: force a format, bypassing the preference. :param dpi: force a DPI, bypassing the preference. :param close: close the figure once written. :param save_mode: force one of :data:`spacr.figure_style.SAVE_MODES` -- ``'print'`` (light page, dark chrome), ``'screen'`` (exactly what is on screen, the old behaviour) or ``'transparent'``. None asks the preference, which defaults to ``'print'``. :param announce_colours: say when a data colour has stopped working on the light page. It is NAMED, never substituted. :param integrity: check the figure's image panels before writing -- display ranges that differ between panels meant for comparison, saturated or clipped pixels, repeated panels, a lossy format, and panels written with fewer pixels than they hold -- print any warning, stamp the provenance (source files, display settings, processing steps, spaCR version) into the PNG or PDF metadata and write it to a ``<file>.provenance.json`` sidecar beside the figure. None follows ``SPACR_FIGURE_INTEGRITY`` and then the Preferences toggle, which is off by default. A figure without image panels is written unchanged either way. :param kwargs: passed through to ``savefig`` (``bbox_inches`` etc.). An explicit ``facecolor`` still wins over the print ground. :returns: the path actually written, as a ``str``. """ preferred_fmt, preferred_dpi = figure_output_preferences() chosen_fmt = str(fmt or preferred_fmt).strip().lower().lstrip(".") if chosen_fmt not in FIGURE_FORMATS: chosen_fmt = preferred_fmt chosen_dpi = int(dpi) if dpi else preferred_dpi destination = _with_extension(path, chosen_fmt) directory = os.path.dirname(str(destination)) if directory: os.makedirs(directory, exist_ok=True) kwargs.pop("format", None) kwargs.pop("dpi", None) from matplotlib import rc_context from .figure_style import saved_figure_appearance write_dpi = deliverable_dpi(fig, chosen_dpi, destination) report = None if _figure_integrity_enabled(integrity): requested = str(fmt or os.path.splitext(str(path))[1] or chosen_fmt) try: report = _integrity_report( fig, fmt=chosen_fmt, dpi=write_dpi, requested_fmt=requested.strip().lower().lstrip("."), destination=destination) if report is not None: kwargs[_SAVEFIG_METADATA] = _integrity_metadata( report, chosen_fmt, kwargs.get(_SAVEFIG_METADATA)) except Exception as exc: print(f"Figure integrity: the check could not run ({exc}); " f"writing {destination} without it.") report = None from .figures.style import _apply_user_style, _preference_deltas _apply_user_style(fig) chosen_style = _preference_deltas() vector_text = bool(chosen_style.get("vector_text", True)) extra_fmt = str(chosen_style.get("also_save", "none") or "none").lower() look = saved_figure_appearance(save_mode) rc = {"pdf.fonttype": 42 if vector_text else 3, "svg.fonttype": "none" if vector_text else "path"} if look.ground is not None and "facecolor" not in kwargs: rc["savefig.facecolor"] = look.ground rc["savefig.edgecolor"] = look.ground rc["savefig.transparent"] = bool(look.transparent) with rc_context(rc), print_ready(fig, mode=look.mode, announce=announce_colours): targets = [(destination, chosen_fmt)] if (extra_fmt in FIGURE_FORMATS and extra_fmt != chosen_fmt and not fmt and report is None): targets.append((_with_extension(destination, extra_fmt), extra_fmt)) for target, target_fmt in targets: fig.savefig(target, format=target_fmt, dpi=write_dpi, **kwargs) if report is not None: try: _finish_integrity(report, destination) except Exception as exc: print(f"Figure integrity: the provenance sidecar for " f"{destination} could not be written ({exc}).") if close: plt.close(fig) return destination
def _checked_savefig(fig, path, *, dpi=None, integrity=None, **kwargs): """Write ``fig`` with ``savefig`` under the figure-integrity check. For writers that keep their own ``savefig`` call rather than going through :func:`save_figure`: when the check is on, the figure's image panels are checked, provenance is stamped into PNG and PDF files and a sidecar is written, exactly as :func:`save_figure` does. When it is off this is a plain ``savefig``. :param fig: the figure to write. :param path: destination; its extension names the format unless ``format`` is passed. :param dpi: resolution, passed on to ``savefig``. :param integrity: True/False forces the check; None follows the environment and the Preferences toggle. :returns: ``path``. """ fmt = str(kwargs.get("format") or os.path.splitext(str(path))[1] or "png").lower().lstrip(".") report = None if _figure_integrity_enabled(integrity): try: report = _integrity_report(fig, fmt=fmt, dpi=dpi or fig.dpi, destination=path) if report is not None: metadata = _integrity_metadata(report, fmt, kwargs.get(_SAVEFIG_METADATA)) if metadata is not None: kwargs[_SAVEFIG_METADATA] = metadata except Exception as exc: print(f"Figure integrity: the check could not run ({exc}); " f"writing {path} without it.") report = None if dpi is not None: kwargs["dpi"] = dpi fig.savefig(path, **kwargs) if report is not None: try: _finish_integrity(report, path) except Exception as exc: print(f"Figure integrity: the provenance sidecar for " f"{path} could not be written ({exc}).") return path _INTEGRITY_ENV = "SPACR_FIGURE_INTEGRITY" _SAVEFIG_METADATA = "metadata" _PANEL_TAG = "_spacr_panel_provenance" _PROVENANCE_SCHEMA = "spacr.figure_provenance/1" _PROVENANCE_SUFFIX = ".provenance.json" _PNG_PROVENANCE_KEY = "spaCR provenance" _RANGE_TOLERANCE = 0.05 _CLIP_WARN_FRACTION = 0.05 _SENSOR_WARN_FRACTION = 0.001 _DUPLICATE_CORRELATION = 0.97 _DUPLICATE_THUMB = 32 _DUPLICATE_DETAIL_SIZE = 96 _DUPLICATE_DETAIL = 0.2 _MIN_PANEL_SIDE = 16 _MIN_COMPARE_PIXELS = 64 * 64 _RESAMPLE_NOTE_FACTOR = 2.0 _SOURCE_HASH_LIMIT = 256 * 1024 * 1024 _MICROMANAGER_TAG = 51123 _SPATIAL_SIDE = 128 _SPATIAL_COMPRESSED_LIMIT = 24 * 1024 _SPATIAL_REPORT_LIMIT = 256 * 1024 _SPATIAL_SIDECAR_LIMIT = 1024 * 1024 _SPATIAL_SIDECAR_MARGIN = 4096 _SPATIAL_ATTEMPTS = 32 _SPATIAL_PAIRS = 12 _SPATIAL_SECONDS = 0.5 _SOURCE_REGION_PAIRS = 8 _SOURCE_REGION_SECONDS = 0.5 _SOURCE_REGION_BYTES = 8 * 1024 * 1024 _SOURCE_REGION_PIXELS = 1024 * 1024 _REGION_WORK = 128 _REGION_STEPS = 28 _REGION_REFINE = 256 _REGION_COARSE = 0.55 _REGION_SCORE = 0.7 _REGION_DETAIL = 0.25 _REGION_MAX_AREA = 0.85 _REGION_MIN_SIDE = 24 _REGION_MIN_BOX = 48 _REGION_MIN_FRACTION = 0.3 _REGION_PANELS = 16 _REGION_SECONDS = 2.0 _SPLICE_WORK = 512 _SPLICE_PANELS = 32 _SPLICE_STEP = 1.0 _SPLICE_SIGN = 0.55 _SPLICE_NOISE_RATIO = 1.6 _CLONE_GUARD = 8 _CLONE_PEAK = 0.006 _CLONE_AREA = 0.01 _CLONE_PEAKS = 16 _INDEX_NAME = ".spacr_figure_index.jsonl" _INDEX_ENV = "SPACR_FIGURE_INDEX" _INDEX_LIMIT = 4 * 1024 * 1024 _INDEX_HAMMING = 10 _INDEX_ENTRIES = 512 _INDEX_CONFIRM = 24 _LOSSY_FORMATS = frozenset({"jpg", "jpeg", "jpe", "jfif", "webp", "gif", "heic", "heif", "avif"}) _REPLAYABLE_OPS = frozenset({"select_channel", "crop", "max_project", "rescale", "to_uint8", "read_crop_png", "merged_crop", "overlay_composite", "combined_masks"}) def _figure_integrity_enabled(explicit=None): """Whether figure exports are checked and stamped with provenance. :param explicit: ``True`` or ``False`` decides outright; ``None`` asks the ``SPACR_FIGURE_INTEGRITY`` environment variable (``1``/``0``) and then the Preferences toggle, which is off on a fresh install. :returns: a bool. Any failure to read the preference reads as off, so a headless run without Qt exports exactly as before. """ if explicit is not None: return bool(explicit) env = os.environ.get(_INTEGRITY_ENV, "").strip().lower() if env in ("1", "true", "yes", "on"): return True if env in ("0", "false", "no", "off"): return False try: from .qt.preferences import _get_figure_integrity return bool(_get_figure_integrity()) except Exception: return False def _display_ranges(value): """``value`` as a list of ``[low, high]`` pairs, one per channel. :param value: a pair, a list of pairs, or ``None``. :returns: the list, or ``None`` when ``value`` is empty or malformed. """ if value is None: return None try: array = np.asarray(value, dtype=float) except (TypeError, ValueError): return None if array.ndim == 1 and array.size == 2: array = array.reshape(1, 2) if array.ndim != 2 or array.shape[1] != 2 or array.shape[0] == 0: return None return [[float(lo), float(hi)] for lo, hi in array] def _percentile_display(image, percentiles): """Percentile-stretch ``image`` for display and say where it was cut. A two-dimensional image is stretched as a whole; a stack is stretched channel by channel along its last axis into a float32 array. :param image: array to stretch. :param percentiles: ``(low, high)`` percentiles. :returns: ``(stretched, ranges)``, the stretched array in ``[0, 1]`` and the per-channel ``[low, high]`` it was cut at, in source units. """ image = np.asarray(image) if image.ndim == 2: lo, hi = np.percentile(image, percentiles) ranges = [[float(lo), float(hi)]] else: ranges = [] for c in range(image.shape[-1]): lo, hi = np.percentile(image[..., c], percentiles) ranges.append([float(lo), float(hi)]) return _apply_display_ranges(image, ranges), ranges def _apply_display_ranges(image, ranges): """Map ``image`` onto ``[0, 1]`` through fixed per-channel ranges. :param image: two-dimensional image, or a stack with channels last. :param ranges: ``[[low, high], ...]``, one pair per channel; a two-dimensional image uses the first. :returns: the clipped, rescaled array (float64 for a plane, float32 for a stack). """ image = np.asarray(image) if image.ndim == 2: lo, hi = ranges[0] return np.clip((image - lo) / (hi - lo), 0, 1) out = np.zeros_like(image, dtype=np.float32) for c in range(image.shape[-1]): lo, hi = ranges[min(c, len(ranges) - 1)] out[..., c] = np.clip((image[..., c] - lo) / (hi - lo), 0, 1) return out def _ome_pixels_for_ifd(root, ifd): """Resolve one TIFF IFD to one OME Pixels element without reading planes. :param root: parsed, size-bounded OME XML root. :param ifd: the zero-based IFD currently displayed by the image reader. :returns: the uniquely mapped Pixels element. :raises ValueError: the mapping is absent, ambiguous or unsupported. """ namespace = root.tag.rsplit('}', 1)[0] + '}' if (not root.tag.endswith('}OME') or not namespace.startswith('{http://www.openmicroscopy.org/Schemas/OME/')): raise ValueError('not an OME metadata document') matches = [] for pixels in root.findall(f'{namespace}Image/{namespace}Pixels'): sizes = {axis: int(pixels.attrib['Size' + axis]) for axis in 'ZTC'} if any(value <= 0 for value in sizes.values()): raise ValueError('invalid OME dimensions') channels = pixels.findall(namespace + 'Channel') samples = [int(channel.get('SamplesPerPixel', '1')) for channel in channels] if samples and (any(value <= 0 for value in samples) or sum(samples) != sizes['C']): raise ValueError('inconsistent OME channel dimensions') if samples and len(set(samples)) != 1: raise ValueError('mixed OME channel sample counts are unsupported') order = pixels.get('DimensionOrder', '') if order not in ('XYZCT', 'XYZTC', 'XYCTZ', 'XYCZT', 'XYTCZ', 'XYTZC'): raise ValueError('invalid OME dimension order') logical_sizes = dict(sizes, C=len(channels) or sizes['C']) plane_count = logical_sizes['Z'] * logical_sizes['T'] * logical_sizes['C'] for mapping in pixels.findall(namespace + 'TiffData'): uuid = mapping.find(namespace + 'UUID') if uuid is not None: file_uuid = (uuid.text or '').strip() if not file_uuid or not root.get('UUID'): raise ValueError('unresolved external TIFF identity') if file_uuid != root.get('UUID'): continue if 'PlaneCount' not in mapping.attrib and 'IFD' not in mapping.attrib: raise ValueError('OME TIFF mapping has no explicit IFD or PlaneCount') first = int(mapping.get('IFD', '0')) count = int(mapping.get('PlaneCount', '1')) if first < 0 or count < 0 or count > plane_count: raise ValueError('invalid OME TIFF plane mapping') coordinates = {axis: int(mapping.get('First' + axis, '0')) for axis in 'ZTC'} if any(not 0 <= coordinates[axis] < sizes[axis] for axis in 'ZTC'): raise ValueError('invalid OME plane coordinate') if samples and samples[0] > 1: if coordinates['C'] != 0: raise ValueError('nonzero packed-channel OME origin is unsupported') offset, stride = 0, 1 for axis in order[2:]: offset += coordinates[axis] * stride stride *= logical_sizes[axis] if offset + count > plane_count: raise ValueError('OME TIFF mapping exceeds the series dimensions') if first <= ifd < first + count: matches.append(pixels) if len(matches) != 1: raise ValueError('TIFF plane has no unique OME Pixels mapping') return matches[0] def _declared_sensor_range(raw, bits, metadata, record=None): """Validate a declared camera bit depth against the displayed source. :param raw: the unchanged source array. :param bits: the declared significant bits per pixel. :param metadata: where the declaration came from, in words. :param record: an existing record to update; a new one otherwise. :returns: the JSON-ready record. A declaration the pixels exceed is kept as a contradiction, with the storage ceiling still in force. """ raw = np.asarray(raw) if record is None: integer = np.issubdtype(raw.dtype, np.integer) record = {'ceiling': int(np.iinfo(raw.dtype).max) if integer else None, 'source': 'storage dtype', 'reason': None} record['declared_significant_bits'] = bits try: if raw.dtype.kind != 'u': raise ValueError(f'{metadata} needs unsigned integer pixels') bits = int(bits) if not 1 <= bits <= np.iinfo(raw.dtype).bits: raise ValueError(f'{metadata} is outside the storage type') ceiling = (1 << bits) - 1 observed = int(raw.max()) if raw.size else 0 if observed > ceiling: record.update(contradiction=True, observed_max=observed, metadata=metadata, declared_significant_bits=bits) raise ValueError('source pixels exceed the declared significant-bit ceiling') record.update(ceiling=ceiling, source=metadata, significant_bits=bits, reason=None) except (ValueError, TypeError) as error: record['reason'] = str(error) return record def _micromanager_bits(tags): """The camera bit depth in a Micro-Manager TIFF plane's metadata, or None. :raises ValueError: the metadata is oversized or malformed. """ import json text = tags.get(_MICROMANAGER_TAG) if text is None: return None if isinstance(text, (tuple, list)) and len(text) == 1: text = text[0] if isinstance(text, bytes): if len(text) > 1024 * 1024: raise ValueError('Micro-Manager metadata exceeds 1 MiB') text = text.decode('utf-8') if not isinstance(text, str) or len(text) > 1024 * 1024: raise ValueError('Micro-Manager metadata exceeds 1 MiB') data = json.loads(text) if not isinstance(data, dict): raise ValueError('Micro-Manager metadata is not an object') bits = data.get('BitDepth') if bits is None and isinstance(data.get('Summary'), dict): bits = data['Summary'].get('BitDepth') if bits is None: return None if isinstance(bits, bool) or not str(bits).strip().isdigit(): raise ValueError('Micro-Manager BitDepth is not a whole number') return int(str(bits).strip()) def _source_sensor_range(opened, raw): """Read bounded existing TIFF metadata for the displayed source plane. :param opened: already-open Pillow image at the displayed IFD. :param raw: already-decoded, unchanged source array. :returns: JSON-ready ceiling, evidence and fallback reason. Validated unsigned OME SignificantBits, or else a Micro-Manager BitDepth, overrides the storage dtype ceiling. A declaration the pixels exceed is recorded as a contradiction. """ import xml.etree.ElementTree as ET raw = np.asarray(raw) integer = np.issubdtype(raw.dtype, np.integer) record = {'ceiling': int(np.iinfo(raw.dtype).max) if integer else None, 'source': 'storage dtype', 'reason': 'no OME TIFF metadata'} tags = getattr(opened, 'tag_v2', None) if tags is None: return record try: record['ifd'] = int(opened.tell()) description = tags.get(270, '') if isinstance(description, bytes): if len(description) > 1024 * 1024: raise ValueError('OME metadata exceeds 1 MiB') description = description.decode('utf-8') if isinstance(description, str) and description: if len(description) > 1024 * 1024 or len(description.encode('utf-8')) > 1024 * 1024: raise ValueError('OME metadata exceeds 1 MiB') if '<!DOCTYPE' in description.upper() or '<!ENTITY' in description.upper(): raise ValueError('XML declarations with entities are unsupported') root = ET.fromstring(description) pixels = _ome_pixels_for_ifd(root, record['ifd']) record['pixels_id'] = pixels.get('ID') bits = pixels.get('SignificantBits') if bits is None: raise ValueError('OME SignificantBits is absent') record['declared_significant_bits'] = bits if (raw.dtype.kind != 'u' or pixels.get('Type') != raw.dtype.name or int(pixels.attrib['SizeX']) != raw.shape[1] or int(pixels.attrib['SizeY']) != raw.shape[0]): raise ValueError('OME pixel type or dimensions contradict the displayed array') bits = int(bits) if not 1 <= bits <= np.iinfo(raw.dtype).bits: raise ValueError('OME SignificantBits is outside the storage type') return _declared_sensor_range(raw, bits, 'OME Pixels SignificantBits', record) except (ValueError, TypeError, KeyError, IndexError, OSError, ET.ParseError) as error: record['reason'] = str(error) try: bits = _micromanager_bits(tags) except (ValueError, TypeError, UnicodeDecodeError) as error: record['micromanager_reason'] = str(error) return record if bits is not None: record['ome_reason'] = record.get('reason') return _declared_sensor_range(raw, bits, 'Micro-Manager BitDepth', record) return record def _raw_clip_stats(raw, ranges, *, sensor_ceiling=None): """How much of the source image a display range throws away. :param raw: the source pixels, before any display mapping. :param ranges: per-channel ``[low, high]`` in source units. :param sensor_ceiling: validated acquisition ceiling; None uses the dtype. :returns: ``{'clipped_high', 'clipped_low', 'sensor_saturated'}``, each the largest per-channel fraction of pixels above ``high``, below ``low``, or at the validated sensor ceiling (the integer storage limit when no explicit ceiling is supplied). """ raw = np.asarray(raw) planes = ([raw] if raw.ndim == 2 or not ranges or len(ranges) == 1 else [raw[..., c] for c in range(min(raw.shape[-1], len(ranges)))]) high = low = sensor = 0.0 ceiling = (sensor_ceiling if sensor_ceiling is not None else np.iinfo(raw.dtype).max if np.issubdtype(raw.dtype, np.integer) else None) for c, plane in enumerate(planes): if plane.size == 0: continue if ranges: lo, hi = ranges[min(c, len(ranges) - 1)] high = max(high, float(np.mean(plane > hi))) low = max(low, float(np.mean(plane < lo))) if ceiling is not None: sensor = max(sensor, float(np.mean(plane == ceiling))) return {"clipped_high": high, "clipped_low": low, "sensor_saturated": sensor} def _tag_panel(artist, source=None, steps=(), display_range=None, channel=None, compare=None, raw=None, sensor_range=None, significant_bits=None): """Attach provenance to an image artist so an export can trace it. Nothing is hashed or read here; the file hashes are taken only when a checked export writes the figure. :param artist: the ``AxesImage`` returned by ``imshow``. :param source: the source image path, or a list of paths. :param steps: the processing steps from source to the displayed array, as dicts with an ``op`` key (``select_channel``, ``crop``, ``max_project``, ``rescale``, ``to_uint8``). A step with any other op is recorded but makes the panel non-replayable. :param display_range: the ``[low, high]`` shown, in source units, or one pair per channel. :param channel: a channel label; panels with different labels are never compared for display range. :param compare: a comparison group name that overrides ``channel``. :param raw: the source pixels, used once to measure clipping and detector saturation. :param sensor_range: validated acquisition ceiling and metadata evidence. :param significant_bits: the camera bit depth, when the caller read it from acquisition metadata that :func:`_source_sensor_range` does not parse; it is validated against ``raw`` like any other declaration. :returns: the artist. """ if sensor_range is None and significant_bits is not None and raw is not None: sensor_range = _declared_sensor_range(raw, significant_bits, 'caller-declared bit depth') sources = ([] if source is None else [source] if isinstance(source, (str, os.PathLike)) else list(source)) ranges = _display_ranges(display_range) record = { "source": [os.path.abspath(str(path)) for path in sources], "steps": [dict(step) for step in (steps or ())], "display_range": ranges, "channel": None if channel is None else str(channel), "compare": None if compare is None else str(compare), } if sensor_range is not None: record["sensor_range"] = dict(sensor_range) if raw is not None: try: record["raw_stats"] = _raw_clip_stats( raw, ranges, sensor_ceiling=(sensor_range or {}).get("ceiling")) record["raw_dtype"] = str(np.asarray(raw).dtype) except Exception: record["raw_stats"] = None try: setattr(artist, _PANEL_TAG, record) except Exception: pass return artist def _array_digest(array): """SHA-256 of an array's dtype, shape and bytes, as 64 hex characters.""" import hashlib array = np.ascontiguousarray(array) digest = hashlib.sha256() digest.update(f"{array.dtype.str}|{array.shape}|".encode("ascii")) digest.update(array.tobytes()) return digest.hexdigest() def _source_record(path): """Path, size and full SHA-256 of one source file. Files above 256 MiB are listed with their size and modification time only, so a checked export never stalls on a whole-plate stack. """ record = {"path": str(path), "exists": os.path.isfile(path)} if not record["exists"]: return record try: size = os.path.getsize(path) record["bytes"] = int(size) record["mtime"] = float(os.path.getmtime(path)) if size <= _SOURCE_HASH_LIMIT: from .run_journal import hash_file record["sha256"] = hash_file(path, full=True) except Exception: pass return record def _panel_thumbnail(array): """Two z-scored grey signatures for repeat detection, or None. :returns: ``(layout, detail)``: a 32x32 area-filtered thumbnail that carries the arrangement of the image, and a 96x96 high-pass residual that carries its fine texture. Two different cells can share a layout; only a repeat of the same pixels shares the detail. """ from PIL import Image data = np.asarray(array, dtype=np.float32) if data.ndim == 3: data = data[..., :3].mean(axis=-1) if data.ndim != 2 or min(data.shape) < _MIN_PANEL_SIDE: return None if not np.isfinite(data).all(): data = np.nan_to_num(data, nan=float(np.nanmedian(data)), posinf=0.0, neginf=0.0) picture = Image.fromarray(np.ascontiguousarray(data), mode="F") def _zscored(values): """Standardize a signature, or return None when its spread is unusable.""" spread = values.std() if not np.isfinite(spread) or spread < 1e-9: return None return (values - values.mean()) / spread layout = _zscored(np.asarray(picture.resize( (_DUPLICATE_THUMB, _DUPLICATE_THUMB), Image.BILINEAR), dtype=np.float64)) if layout is None: return None fine = np.asarray(picture.resize( (_DUPLICATE_DETAIL_SIZE, _DUPLICATE_DETAIL_SIZE), Image.BILINEAR), dtype=np.float64) detail = _zscored(fine - ndi.uniform_filter(fine, 3, mode="reflect")) if detail is None: detail = np.zeros_like(fine) return layout, detail def _spatial_signature(array): """A bounded displayed-pixel thumbnail for later visual-reuse review. The signature makes no matching claim. Blank, nonfinite and small panels have no useful spatial evidence and are left out. """ import base64 import hashlib import zlib data = np.asarray(array) if data.dtype.kind not in "biuf": return None if data.ndim not in (2, 3) or min(data.shape[:2]) < 64: return None if data.ndim == 3 and data.shape[2] not in (3, 4): return None stride = max(1, int(math.ceil(max(data.shape[:2]) / 512))) sample = data[::stride, ::stride] if min(sample.shape[:2]) < 16: return None sample = sample.astype(np.float32) if not np.isfinite(sample).all(): return None sample = cv2.resize(sample, (_SPATIAL_SIDE, _SPATIAL_SIDE), interpolation=cv2.INTER_AREA) if sample.ndim == 3: sample = sample[..., :3].mean(axis=-1) low, high = np.percentile(sample, (1, 99)) if high - low < 1e-6: return None grey = np.clip((sample - low) * 255 / (high - low), 0, 255).astype(np.uint8) fine = grey.astype(np.float32) fine -= cv2.GaussianBlur(fine, (0, 0), 2) if float(fine.std()) < 1.0: return None raw = grey.tobytes() compressed = zlib.compress(raw, 6) return {"version": 1, "side": _SPATIAL_SIDE, "codec": "zlib+base64 gray-u8", "sha256": hashlib.sha256(raw).hexdigest(), "data": base64.b64encode(compressed).decode("ascii")} def _decode_spatial_signature(record): """Decode one bounded versioned signature or refuse malformed data.""" import base64 import binascii import hashlib import zlib if not isinstance(record, dict) or ( record.get("version"), record.get("side"), record.get("codec")) != ( 1, _SPATIAL_SIDE, "zlib+base64 gray-u8"): raise ValueError("unknown spatial signature version") data = record.get("data") if not isinstance(data, str) or len(data) > 4 * ( (_SPATIAL_COMPRESSED_LIMIT + 2) // 3): raise ValueError("spatial signature encoded size") try: compressed = base64.b64decode(data, validate=True) except (ValueError, TypeError, binascii.Error) as error: raise ValueError("malformed spatial signature") from error try: stream = zlib.decompressobj() raw = stream.decompress(compressed, _SPATIAL_SIDE ** 2 + 1) except zlib.error as error: raise ValueError("malformed spatial signature") from error if (len(raw) != _SPATIAL_SIDE ** 2 or not stream.eof or stream.unconsumed_tail or stream.unused_data): raise ValueError("spatial signature raw size") if hashlib.sha256(raw).hexdigest() != record.get("sha256"): raise ValueError("spatial signature digest") return np.frombuffer(raw, dtype=np.uint8).reshape( _SPATIAL_SIDE, _SPATIAL_SIDE) def _similar_displayed_region(first, second): """Find a strong localized displayed-pixel correlation, or abstain.""" first = _decode_spatial_signature(first) second = _decode_spatial_signature(second) for source, candidate in ((first, second), (second, first)): source_fine = source.astype(np.float32) source_fine -= cv2.GaussianBlur(source_fine, (0, 0), 2) if source_fine.std() < 1.0: continue for quarter in range(4): turned = np.rot90(candidate, quarter) for reflected in (turned, turned[:, ::-1]): for fraction in (0.625, 0.75, 0.875): side = int(round(_SPATIAL_SIDE * fraction)) small = cv2.resize(reflected, (side, side), interpolation=cv2.INTER_AREA) detail = small.astype(np.float32) detail -= cv2.GaussianBlur(detail, (0, 0), 2) if detail.std() < 1.0: continue _minimum, score, _low_at, (x, y) = cv2.minMaxLoc( cv2.matchTemplate(source_fine, detail, cv2.TM_CCOEFF_NORMED)) if score < 0.7: continue region = source[y:y + side, x:x + side].astype(np.float32) structure = float(cv2.matchTemplate( region, small.astype(np.float32), cv2.TM_CCOEFF_NORMED)[0, 0]) if structure >= 0.7: return {"detail": round(float(score), 4), "structure": round(structure, 4), "fraction": fraction} return None def _attach_spatial_signatures(report, arrays): """Attach opt-in evidence without pushing a sidecar past its read cap.""" import json used = 0 for panel, array in list(zip(report["panels"], arrays))[:_SPATIAL_ATTEMPTS]: try: signature = _spatial_signature(array) except (TypeError, ValueError, cv2.error): continue if signature is None: continue amount = len(signature["data"]) if used + amount > _SPATIAL_REPORT_LIMIT: break panel["spatial_v1"] = signature used += amount while len(json.dumps(report, indent=2, default=str).encode()) > ( _SPATIAL_SIDECAR_LIMIT - _SPATIAL_SIDECAR_MARGIN): for panel in reversed(report["panels"]): if panel.pop("spatial_v1", None) is not None: break else: break def _best_dihedral_correlation(first, second): """Highest correlation of ``first`` with any flip or rotation of ``second`` (both z-scored and square).""" best = -1.0 for k in range(4): turned = np.rot90(second, k) for variant in (turned, turned[:, ::-1]): best = max(best, float(np.mean(first * variant))) return best def _figure_panels(fig): """Every image panel in ``fig``: ``(axes_index, axes, artist)``.""" from matplotlib.image import AxesImage found = [] for axes_index, axes in enumerate(fig.get_axes()): for artist in axes.get_images(): if isinstance(artist, AxesImage): found.append((axes_index, axes, artist)) return found def _panel_record(index, axes_index, axes, artist, fig, dpi, source_cache=None): """The provenance and integrity measurements for one image panel. :returns: ``(record, displayed_array)``; the record is JSON-ready. """ data = np.asarray(np.ma.getdata(artist.get_array())) tag = getattr(artist, _PANEL_TAG, None) or {} kind = ("rgb" if data.ndim == 3 and data.shape[-1] in (3, 4) else "scalar" if data.ndim == 2 else "other") title = "" try: title = axes.get_title() or axes.get_ylabel() or "" except Exception: pass record = { "panel": index, "axes": axes_index, "title": str(title), "kind": kind, "shape": [int(n) for n in data.shape], "dtype": str(data.dtype), "interpolation": str(artist.get_interpolation()), "displayed_sha256": _array_digest(data), "source": [], "steps": list(tag.get("steps") or []), "channel": tag.get("channel"), "compare": tag.get("compare"), "tagged": bool(tag), } clipped_high = clipped_low = sensor = None if kind == "scalar": vmin, vmax = artist.get_clim() record["cmap"] = str(getattr(artist.get_cmap(), "name", "")) record["clim"] = [None if vmin is None else float(vmin), None if vmax is None else float(vmax)] finite = data[np.isfinite(data)] if data.dtype.kind == "f" else data if finite.size and vmin is not None and vmax is not None: clipped_high = float(np.mean(finite > vmax)) clipped_low = float(np.mean(finite < vmin)) if (np.issubdtype(data.dtype, np.integer) and np.iinfo(data.dtype).bits > 8 and data.size): sensor = float(np.mean(data == np.iinfo(data.dtype).max)) elif kind == "rgb" and data.size: top = (np.iinfo(data.dtype).max if np.issubdtype(data.dtype, np.integer) else 1.0) colour = data[..., :3] clipped_high = float(max(np.mean(colour[..., c] >= top) for c in range(3))) clipped_low = float(max(np.mean(colour[..., c] <= 0) for c in range(3))) ranges = _display_ranges(tag.get("display_range")) if ranges is not None: record["display_range"] = ranges record["display_range_units"] = "source" elif kind == "scalar" and None not in record["clim"]: record["display_range"] = [list(record["clim"])] record["display_range_units"] = "displayed array" else: record["display_range"] = None record["display_range_units"] = None raw_stats = tag.get("raw_stats") if raw_stats: clipped_high = raw_stats.get("clipped_high", clipped_high) clipped_low = raw_stats.get("clipped_low", clipped_low) sensor = raw_stats.get("sensor_saturated", sensor) record["raw_dtype"] = tag.get("raw_dtype") record["clipped_high"] = clipped_high record["clipped_low"] = clipped_low record["sensor_saturated"] = sensor if tag.get("sensor_range") is not None: record["sensor_range"] = dict(tag["sensor_range"]) record["clip_measured_on"] = ("source" if raw_stats else "displayed array") try: box = axes.get_window_extent() scale = float(dpi) / float(fig.dpi) record["exported_pixels"] = [int(round(box.width * scale)), int(round(box.height * scale))] except Exception: record["exported_pixels"] = None record["source"] = [] for path in tag.get("source") or (): if source_cache is None: record["source"].append(_source_record(path)) else: if path not in source_cache: source_cache[path] = _source_record(path) record["source"].append(source_cache[path]) ops = [str(step.get("op", "")) for step in record["steps"]] record["reproducible"] = bool(record["source"]) and all( op in _REPLAYABLE_OPS for op in ops) if any(op in ("overlay_composite", "combined_masks") for op in ops): record["reproducible"] = (record["reproducible"] and len(record["source"]) == 1 and bool(record["source"][0].get("sha256"))) return record, data def _range_findings(panels): """Panels meant for comparison that are shown through different ranges. Panels are compared within a group: an explicit comparison name, else a channel label, else (for untagged single-channel panels of at least 64x64 pixels) the colour map. Untagged colour panels carry their display range baked in and are not compared. A group is flagged when, for any channel, the lows or the highs spread by more than 5 % of the combined range. """ groups = {} for panel in panels: ranges = panel.get("display_range") if not ranges: continue if panel.get("compare"): key = ("compare", panel["compare"]) elif panel.get("tagged"): key = ("channel", panel.get("channel")) elif (panel["kind"] == "scalar" and int(np.prod(panel["shape"])) >= _MIN_COMPARE_PIXELS): key = ("cmap", panel.get("cmap")) else: continue groups.setdefault(key, []).append(panel) findings = [] for key, members in groups.items(): if len(members) < 2: continue width = min(len(p["display_range"]) for p in members) worst = 0.0 for c in range(width): lows = [p["display_range"][c][0] for p in members] highs = [p["display_range"][c][1] for p in members] span = max(highs) - min(lows) if not np.isfinite(span) or span <= 0: continue spread = max(max(lows) - min(lows), max(highs) - min(highs)) worst = max(worst, spread / span) if worst > _RANGE_TOLERANCE: units = members[0].get("display_range_units") or "" shown = ", ".join( f"{p['panel']}: " + "/".join( f"{lo:.4g}-{hi:.4g}" for lo, hi in p["display_range"]) for p in members[:8]) group = {"compare": f"group '{key[1]}'", "channel": (f"channel {key[1]}" if key[1] is not None else "panels traced to source files"), "cmap": f"colour map {key[1]}"}[key[0]] findings.append({ "check": "display_range", "severity": "warning", "panels": [p["panel"] for p in members], "spread": round(worst, 4), "message": ( f"{len(members)} panels in one comparison group " f"({group}) are shown through display ranges " f"that differ by {worst:.0%} of their combined range " f"({units} units; {shown}). Brightness is not comparable " f"between them; use one range, or say in the legend that " f"each panel was scaled on its own."), }) return findings def _saturation_findings(panels): """Panels with detector saturation or clipped highlights. Warns when more than 0.1 % of source pixels sit at the validated acquisition ceiling (the storage type limit without metadata), or more than 5 % are pushed above the top of the display range. On an untagged colour panel the measurement cannot tell a saturated pixel from an annotation colour, so it is a note there. """ findings = [] for panel in panels: sensor = panel.get("sensor_saturated") if sensor is not None and sensor > _SENSOR_WARN_FRACTION: sensor_range = panel.get("sensor_range") or {} limit = (f"the acquisition ceiling {sensor_range['ceiling']} declared " + ("by OME SignificantBits" if sensor_range.get("source") == "OME Pixels SignificantBits" else f"by {sensor_range.get('source')}") if sensor_range.get("significant_bits") is not None and sensor_range.get("reason") is None else "the largest value the image type can hold") findings.append({ "check": "saturation", "severity": "warning", "panels": [panel["panel"]], "fraction": round(sensor, 5), "message": ( f"Panel {panel['panel']}: {sensor:.2%} of the source " f"pixels are at {limit}. They are saturated at acquisition " f"and no display " f"setting recovers them."), }) high = panel.get("clipped_high") if high is not None and high > _CLIP_WARN_FRACTION: blind = panel["kind"] == "rgb" and not panel.get("tagged") findings.append({ "check": "saturation", "severity": "note" if blind else "warning", "panels": [panel["panel"]], "fraction": round(high, 5), "message": ( f"Panel {panel['panel']}: {high:.1%} of the pixels are " f"at or above the top of the display range and show as " f"one flat maximum." + (" This is measured on the finished colour image, so " "annotation colours count too." if blind else " Differences between them are hidden; raise the " "upper display limit.")), }) return findings def _duplicate_findings(panels, arrays): """Panels whose image content is the same, allowing flips, rotations and contrast changes. Identical displayed arrays are always flagged. Otherwise a pair is flagged when its 32x32 grey thumbnails correlate at 0.97 or more and its 96x96 high-pass texture at 0.2 or more, each at the best of the eight flips and rotations. The texture test is what keeps a montage of similar-looking but different cells from reading as repeats. Panels traced to the same source file are a declared reuse and are reported as a note. """ findings = [] signatures, owners = [], [] for panel, array in zip(panels, arrays): signature = _panel_thumbnail(array) if signature is not None: signatures.append(signature) owners.append(panel) if len(signatures) < 2: return findings count = len(signatures) layouts = np.stack([sig[0] for sig in signatures]) flat = layouts.reshape(count, -1) best = np.full((count, count), -1.0) for k in range(4): turned = np.rot90(layouts, k, axes=(1, 2)) for variant in (turned, turned[:, :, ::-1]): corr = flat @ variant.reshape(count, -1).T / flat.shape[1] best = np.maximum(best, corr) for i in range(count): for j in range(i + 1, count): a, b = owners[i], owners[j] same = a["displayed_sha256"] == b["displayed_sha256"] score = max(best[i, j], best[j, i]) if not same: if score < _DUPLICATE_CORRELATION: continue texture = _best_dihedral_correlation(signatures[i][1], signatures[j][1]) if texture < _DUPLICATE_DETAIL: continue declared = (a.get("source") and b.get("source") and [s["path"] for s in a["source"]] == [s["path"] for s in b["source"]]) findings.append({ "check": "duplicate", "severity": "note" if declared else "warning", "panels": [a["panel"], b["panel"]], "correlation": round(float(score), 4), "texture": None if same else round(float(texture), 4), "identical": bool(same), "message": ( f"Panels {a['panel']} and {b['panel']} show " + ("identical pixels" if same else f"the same image content (correlation " f"{score:.3f}, allowing flips, rotations and " f"contrast changes)") + (". They are traced to the same source file; say so " "in the legend." if declared else ". If one image is shown twice on purpose, say so in " "the legend.")), }) return findings def _panel_grey(array, side): """A finite grey float32 copy whose longer side is at most ``side``. :returns: ``(grey, scale)`` with ``scale`` the work pixels per displayed pixel, or None for an array that is not an image of usable size. """ data = np.asarray(array) if data.dtype.kind not in "biuf" or data.ndim not in (2, 3): return None if data.ndim == 3: if data.shape[2] not in (3, 4): return None data = data[..., :3].astype(np.float32).mean(axis=-1) else: data = data.astype(np.float32) if min(data.shape) < _MIN_PANEL_SIDE: return None finite = np.isfinite(data) if not finite.all(): fill = float(np.median(data[finite])) if finite.any() else 0.0 data = np.where(finite, data, fill).astype(np.float32) height, width = data.shape scale = min(1.0, side / max(height, width)) if scale < 1.0: data = cv2.resize(data, (max(1, int(round(width * scale))), max(1, int(round(height * scale)))), interpolation=cv2.INTER_AREA) return np.ascontiguousarray(data), scale def _band_pass(image, fine=False): """Texture band of a grey image; ``fine`` keeps only pixel-scale detail.""" if fine: return image - cv2.GaussianBlur(image, (0, 0), 1.0) return (cv2.GaussianBlur(image, (0, 0), 0.8) - cv2.GaussianBlur(image, (0, 0), 3.0)) def _region_scales(big, small, low=None, high=None, count=_REGION_STEPS, minimum=_REGION_MIN_SIDE): """Template scales that place ``small`` as a strict sub-region of ``big``. :param minimum: the shortest template side, in work pixels. """ big_h, big_w = big.shape small_h, small_w = small.shape top = min(big_w / small_w, big_h / small_h, math.sqrt(_REGION_MAX_AREA * big_w * big_h / (small_w * small_h))) bottom = max(max(_REGION_MIN_SIDE, minimum) / min(small_w, small_h), _REGION_MIN_FRACTION * min(big_w / small_w, big_h / small_h)) if low is not None: bottom, top = max(bottom, low), min(top, high) if bottom > top: return [] return list(np.geomspace(bottom, top, count)) if top > bottom else [bottom] def _region_variants(small): """The shape-preserving flips of a template: as is, mirrored, upturned.""" return {"none": small, "mirror": small[:, ::-1], "flip": small[::-1], "rot180": small[::-1, ::-1]} def _best_region(big, small, scales, variants): """Best band-pass template match of rescaled ``small`` within ``big``.""" target = _band_pass(big) if float(target.std()) < 1e-6: return None best = None for name in variants: turned = np.ascontiguousarray(_region_variants(small)[name]) for scale in scales: width = int(round(turned.shape[1] * scale)) height = int(round(turned.shape[0] * scale)) if (width > big.shape[1] or height > big.shape[0] or min(width, height) < _REGION_MIN_SIDE // 2): continue resized = cv2.resize(turned, (width, height), interpolation=(cv2.INTER_AREA if scale < 1 else cv2.INTER_LINEAR)) template = _band_pass(resized) if float(template.std()) < 1e-6: continue _low, score, _where, (x, y) = cv2.minMaxLoc(cv2.matchTemplate( target, template, cv2.TM_CCOEFF_NORMED)) if np.isfinite(score) and (best is None or score > best[0]): best = (float(score), float(scale), name, (x, y), resized) return best def _aligned_residual_correlation(window, template): """Correlation of the pixel-scale residuals once two patches are aligned. The template is first aligned to the window with sub-pixel accuracy (an affine ECC fit); the residual left after a 1.5-pixel blur is then the image's own noise and fine texture, which two different cells of the same shape do not share but a reused region does. :returns: the correlation, or None when the patches cannot be aligned or carry no residual. """ window = np.ascontiguousarray(window, dtype=np.float32) template = np.ascontiguousarray(template, dtype=np.float32) warp = np.eye(2, 3, dtype=np.float32) try: _found, warp = cv2.findTransformECC( window, template, warp, cv2.MOTION_AFFINE, (cv2.TERM_CRITERIA_EPS | cv2.TERM_CRITERIA_COUNT, 60, 1e-6), None, 1) except cv2.error: return None height, width = window.shape aligned = cv2.warpAffine(template, warp, (width, height), flags=cv2.INTER_LINEAR + cv2.WARP_INVERSE_MAP, borderMode=cv2.BORDER_REFLECT) margin = max(2, min(height, width) // 16) inner = (slice(margin, -margin), slice(margin, -margin)) smooth_first = cv2.GaussianBlur(window, (0, 0), 1.5) smooth_second = cv2.GaussianBlur(aligned, (0, 0), 1.5) slope = np.maximum( np.hypot(*np.gradient(cv2.GaussianBlur(window, (0, 0), 2.0))), np.hypot(*np.gradient(cv2.GaussianBlur(aligned, (0, 0), 2.0))))[inner] flat = slope <= np.percentile(slope, 50) first = (window - smooth_first)[inner][flat] second = (aligned - smooth_second)[inner][flat] if first.size < 64 or float(first.std()) <= 1e-3 * float(window.std()): return None gain = float(np.polyfit(smooth_second[inner].ravel(), smooth_first[inner].ravel(), 1)[0]) return 1.0 - float(np.var(first - gain * second)) / float(np.var(first)) def _region_in_panel(big, small): """Whether ``small`` shows a rescaled sub-region of ``big``, or None. Both are ``(coarse, fine)`` pairs of :func:`_panel_grey` results. The coarse search covers every scale and shape-preserving flip; the fine search confirms one candidate at a higher resolution, on the texture band and again on pixel-scale detail, which two different cells of similar shape do not share. """ (big_c, big_cs), (big_f, big_fs) = big (small_c, small_cs), (small_f, small_fs) = small scales = _region_scales(big_c, small_c, count=_REGION_STEPS, minimum=_REGION_MIN_BOX * big_cs) if not scales: return None candidates = [] for variant in _region_variants(small_c): found = _best_region(big_c, small_c, scales, [variant]) if found is not None and found[0] >= _REGION_COARSE: candidates.append(found) fine = None for coarse in sorted(candidates, key=lambda item: -item[0])[:2]: relative = coarse[1] * small_cs / big_cs centre = relative * big_fs / small_fs scales = _region_scales(big_f, small_f, centre / 1.05, centre * 1.05, 11, minimum=_REGION_MIN_BOX * big_fs) found = _best_region(big_f, small_f, scales, [coarse[2]]) if scales else None if found is not None and (fine is None or found[0] > fine[0]): fine = found if fine is None or fine[0] < _REGION_SCORE: return None score, scale, variant, (x, y), resized = fine window = big_f[y:y + resized.shape[0], x:x + resized.shape[1]] detail = _aligned_residual_correlation(window, resized) if detail is None or detail < _REGION_DETAIL: return None relative = scale * small_fs / big_fs return {"score": round(score, 4), "detail": round(detail, 4), "zoom": round(1.0 / relative, 3), "variant": variant, "box": [int(round(y / big_fs)), int(round((y + resized.shape[0]) / big_fs)), int(round(x / big_fs)), int(round((x + resized.shape[1]) / big_fs))]} def _region_reuse_findings(panels, arrays, known_pairs=()): """Panels that show a rescaled or cropped region of another panel. Each pair is searched both ways for the smaller-content panel inside the larger one, within a time budget; the report records how many pairs were checked and how many the budget left out. Pairs already reported as whole-panel repeats are not searched again. :returns: ``(findings, statistics)``. """ import time from .qt.i18n import tr started = time.monotonic() work = [] for panel, array in zip(panels, arrays): if len(work) >= _REGION_PANELS: break if int(np.prod(panel["shape"][:2])) < _MIN_COMPARE_PIXELS: continue coarse = _panel_grey(array, _REGION_WORK) fine = _panel_grey(array, _REGION_REFINE) if coarse is not None and fine is not None: work.append((panel, (coarse, fine))) known = {frozenset(pair) for pair in known_pairs} findings, checked, skipped = [], 0, 0 for big_panel, big in work: for small_panel, small in work: pair = frozenset((big_panel["panel"], small_panel["panel"])) if big_panel is small_panel or pair in known: continue if time.monotonic() - started > _REGION_SECONDS: skipped += 1 continue checked += 1 match = _region_in_panel(big, small) if match is None: continue known.add(pair) declared = bool( big_panel.get("source") and small_panel.get("source") and {s.get("path") for s in big_panel["source"]} & {s.get("path") for s in small_panel["source"]}) findings.append({ "check": "region_reuse", "severity": "note" if declared else "warning", "panels": [small_panel["panel"], big_panel["panel"]], "correlation": match["score"], "detail": match["detail"], "zoom": match["zoom"], "variant": match["variant"], "box": match["box"], "message": tr( "Panel {panel} shows a region of panel {other} " "(rows {top}-{bottom}, columns {left}-{right}) at " "{zoom}x magnification. If it is an inset or a zoom, " "mark the region and say so in the legend.", panel=small_panel["panel"], other=big_panel["panel"], top=match["box"][0], bottom=match["box"][1], left=match["box"][2], right=match["box"][3], zoom=f"{match['zoom']:.2f}"), }) return findings, {"pairs_checked": checked, "pairs_skipped": skipped} def _clone_region(grey): """A region repeated at another place inside one image, or None. A copied patch shows as a secondary peak in the autocorrelation of the texture band; the peak is confirmed only where both copies agree in 9x9 windows over a connected area of at least 1 % of the image. """ texture = _band_pass(grey, fine=True).astype(np.float64) texture -= texture.mean(axis=0, keepdims=True) texture -= texture.mean(axis=1, keepdims=True) energy = float((texture ** 2).sum()) height, width = texture.shape if energy <= 1e-9 or min(height, width) < 32: return None spectrum = np.fft.rfft2(texture, s=(2 * height, 2 * width)) corr = np.fft.irfft2(spectrum * np.conj(spectrum), s=(2 * height, 2 * width)) / energy corr = np.fft.fftshift(corr) cy, cx = height, width ys, xs = np.mgrid[-cy:cy, -cx:cx] masked = corr.copy() masked[((np.abs(ys) <= _CLONE_GUARD) & (np.abs(xs) <= _CLONE_GUARD)) | (ys < 0) | ((ys == 0) & (xs < 0)) | (np.abs(ys) > height - 16) | (np.abs(xs) > width - 16)] = 0 for _attempt in range(_CLONE_PEAKS): py, px = np.unravel_index(int(np.argmax(masked)), masked.shape) value = float(masked[py, px]) if value < _CLONE_PEAK: return None masked[py - 2:py + 3, px - 2:px + 3] = 0 found = _confirm_clone(texture, corr, (py, px), value) if found is not None: return found return None def _confirm_clone(texture, corr, at, value): """Confirm one autocorrelation peak as a copied patch, or None. A copy is a sharp peak with no echo at half or twice its shift (which a periodic pattern or pixel-replicated upscaling would have), and both copies agree in 9x9 windows over a connected area. """ height, width = texture.shape py, px = at dy, dx = int(py - height), int(px - width) around = corr[py - 2:py + 3, px - 2:px + 3].copy() around[2, 2] = -np.inf if float(around.max()) > 0.5 * value: return None echoes = [(2 * dy, 2 * dx), (dy / 2, dx / 2)] if dy and dx: echoes += [(dy, 0), (0, dx)] for fy, fx in echoes: if fy != int(fy) or fx != int(fx): continue ry, rx = int(height + fy), int(width + fx) if (0 <= ry < corr.shape[0] and 0 <= rx < corr.shape[1] and float(corr[ry, rx]) > 0.5 * value): return None first = texture[max(0, -dy):height - max(0, dy), max(0, -dx):width - max(0, dx)] second = texture[max(0, dy):height + min(0, dy), max(0, dx):width + min(0, dx)] window = (9, 9) product = cv2.blur((first * second).astype(np.float32), window) power = np.sqrt(cv2.blur((first ** 2).astype(np.float32), window) * cv2.blur((second ** 2).astype(np.float32), window)) noise = float(np.median(np.abs(texture))) + 1e-9 agree = (product / np.maximum(power, 1e-12) > 0.9) & (power > (0.5 * noise) ** 2) count, _labels, stats_, _centroids = cv2.connectedComponentsWithStats( agree.astype(np.uint8), connectivity=8) if count < 2: return None largest = 1 + int(np.argmax(stats_[1:, cv2.CC_STAT_AREA])) area = int(stats_[largest, cv2.CC_STAT_AREA]) if area < _CLONE_AREA * height * width or area < 16 * 16: return None x0, y0 = stats_[largest, cv2.CC_STAT_LEFT], stats_[largest, cv2.CC_STAT_TOP] w0, h0 = stats_[largest, cv2.CC_STAT_WIDTH], stats_[largest, cv2.CC_STAT_HEIGHT] oy, ox = max(0, -dy), max(0, -dx) return {"shift": [dy, dx], "peak": round(value, 4), "area_fraction": round(area / (height * width), 4), "box": [int(y0 + oy), int(y0 + oy + h0), int(x0 + ox), int(x0 + ox + w0)]} def _noise_like(image): """Whether an image's finest detail is acquisition noise. Rendered graphics, label masks and denoised pictures have pixel-scale detail that is smooth from one pixel to the next; camera noise is not. Noise-level comparisons and copied-patch searches need the latter. """ detail = image - cv2.GaussianBlur(image, (0, 0), 1.0) first, second = detail[:, :-1].ravel(), detail[:, 1:].ravel() if float(first.std()) < 1e-9 or float(second.std()) < 1e-9: return False lag = float(np.corrcoef(first, second)[0, 1]) return bool(np.isfinite(lag) and lag < 0.5) def _noise_by_intensity(left, right): """Median log ratio of matched-intensity noise across two strips, or None. Noise is compared only between pixels of similar local brightness, so shot noise that rises with signal does not read as a change of camera. """ ratios = [] pairs = [] for strip in (left, right): local = cv2.GaussianBlur(strip, (0, 0), 2.0) pairs.append((local.ravel(), (strip - cv2.GaussianBlur(strip, (0, 0), 1.0)).ravel())) both = np.concatenate([pairs[0][0], pairs[1][0]]) edges = np.quantile(both, np.linspace(0, 1, 6)) for low, high in zip(edges[:-1], edges[1:]): spreads = [] for local, residual in pairs: chosen = residual[(local >= low) & (local <= high)] if chosen.size < 100: break spreads.append(float(np.median(np.abs(chosen - np.median(chosen))))) if len(spreads) == 2 and min(spreads) > 1e-9: ratios.append(math.log(spreads[0] / spreads[1])) if len(ratios) < 2 or (min(ratios) < 0 < max(ratios)): return None return float(np.median(ratios)) def _seam(grey): """A straight seam where background level or noise changes, or None.""" best = None for axis, image in (("vertical", grey), ("horizontal", grey.T)): height, width = image.shape if width < 64 or height < 32: continue steps = np.diff(image, axis=1) centred = steps - np.median(steps) scale = 1.4826 * float(np.median(np.abs(centred))) residual = np.abs(image - cv2.GaussianBlur(image, (0, 0), 1.0)) profile = np.median(residual, axis=0) floor = float(np.median(profile)) if scale <= 1e-9 or floor <= 1e-9: continue band = max(8, width // 8) noisy = _noise_like(image) middle = np.median(steps, axis=0) agree = np.mean(np.sign(steps) == np.sign(middle)[None, :], axis=0) level = np.abs(middle) / scale for column in range(band, width - band): edge = column - 1 if (profile[column - band:column].min() < 0.2 * floor or profile[column:column + band].min() < 0.2 * floor): continue nearby = np.r_[level[max(0, edge - 10):edge - 1], level[edge + 2:edge + 11]] offset = (level[edge] >= _SPLICE_STEP and agree[edge] >= _SPLICE_SIGN and nearby.size and float(nearby.max()) < 0.3 * level[edge]) left = profile[column - band:column] right = profile[column:column + band] jump = abs(math.log((float(np.median(left)) + 1e-9) / (float(np.median(right)) + 1e-9))) if not offset and jump < math.log(_SPLICE_NOISE_RATIO) * 0.8: continue noise = None if not offset: if not noisy: continue noise = _noise_by_intensity(image[:, column - band:column], image[:, column:column + band]) if noise is None or abs(noise) < math.log(_SPLICE_NOISE_RATIO): continue strength = (float(level[edge]) if offset else 0.0) + abs(noise or 0.0) if best is None or strength > best["strength"]: best = {"axis": axis, "position": int(column), "offset": round(float(level[edge]), 3) if offset else None, "noise_ratio": (round(math.exp(abs(noise)), 3) if noise is not None else None), "strength": strength} return best def _splice_work(array): """A grey work image for splice checks, with pixel replication undone. An image enlarged by repeating pixels carries a regular lattice that reads as copied texture; rows and columns that only repeat their neighbour are dropped first when they make up 40 % or more of an axis. :returns: ``(grey, rows, columns, scale)``: the work image, the displayed-array row and column each kept line came from, and the work pixels per kept line; or None for a panel too small to check. """ whole = _panel_grey(array, float("inf")) if whole is None: return None grey = whole[0] rows, columns = np.arange(grey.shape[0]), np.arange(grey.shape[1]) for axis in (0, 1): same = np.all(np.diff(grey, axis=axis) == 0, axis=1 - axis) if same.size and float(same.mean()) >= 0.4: keep = np.r_[True, ~same] if axis == 0: grey, rows = grey[keep], rows[keep] else: grey, columns = grey[:, keep], columns[keep] height, width = grey.shape if min(height, width) < 64: return None scale = min(1.0, _SPLICE_WORK / max(height, width)) if scale < 1.0: grey = cv2.resize(grey, (max(1, int(round(width * scale))), max(1, int(round(height * scale)))), interpolation=cv2.INTER_AREA) return np.ascontiguousarray(grey), rows, columns, scale def _work_to_displayed(mapping, scale, value): """The displayed-array line a splice work-image coordinate came from. :param mapping: the displayed row or column of each kept line. :param scale: work pixels per kept line. :param value: the work-image coordinate; one past the last line maps to one past the last displayed line. """ position = int(round(value / scale)) if position >= len(mapping): return int(mapping[-1]) + 1 return int(mapping[max(0, position)]) def _splice_findings(panels, arrays): """Panels that look assembled from more than one image. Two signs are checked on each panel of at least 64x64 pixels: a region repeated elsewhere in the same panel (a cloned patch), and a straight full-length seam where the background level steps or the noise of equally bright pixels changes (two images butted together). Seams against flat padding are ignored. """ from .qt.i18n import tr findings = [] for panel, array in list(zip(panels, arrays))[:_SPLICE_PANELS]: if int(np.prod(panel["shape"][:2])) < _MIN_COMPARE_PIXELS: continue work = _splice_work(array) if work is None: continue grey, rows, columns, scale = work clone = _clone_region(grey) if _noise_like(grey) else None if clone is not None: box = [_work_to_displayed(rows, scale, clone["box"][0]), _work_to_displayed(rows, scale, clone["box"][1]), _work_to_displayed(columns, scale, clone["box"][2]), _work_to_displayed(columns, scale, clone["box"][3])] shift = [int(round(clone["shift"][0] * (rows[-1] + 1) / grey.shape[0])), int(round(clone["shift"][1] * (columns[-1] + 1) / grey.shape[1]))] findings.append({ "check": "splice", "kind": "clone", "severity": "warning", "panels": [panel["panel"]], "shift": shift, "box": box, "peak": clone["peak"], "area_fraction": clone["area_fraction"], "message": tr( "Panel {panel}: the region at rows {top}-{bottom}, columns " "{left}-{right} repeats {down} rows down and {across} " "columns across in the same image. Check the source image " "for a copied patch.", panel=panel["panel"], top=box[0], bottom=box[1], left=box[2], right=box[3], down=shift[0], across=shift[1]), }) seam = _seam(grey) if seam is not None: position = _work_to_displayed( columns if seam["axis"] == "vertical" else rows, scale, seam["position"]) findings.append({ "check": "splice", "kind": "seam", "severity": "warning", "panels": [panel["panel"]], "axis": seam["axis"], "position": position, "offset": seam["offset"], "noise_ratio": seam["noise_ratio"], "message": tr( "Panel {panel}: a straight {axis} seam at pixel {position} " "separates areas with a different background level or " "noise, as when two images are butted together. If the " "panel combines images, separate them with a visible line " "and say so in the legend.", panel=panel["panel"], axis=tr(seam["axis"]), position=position), }) return findings def _bit_depth_findings(panels): """Panels whose camera metadata contradicts their pixel values.""" from .qt.i18n import tr findings = [] for panel in panels: sensor = panel.get("sensor_range") or {} if not sensor.get("contradiction"): continue findings.append({ "check": "bit_depth", "severity": "warning", "panels": [panel["panel"]], "declared_bits": sensor.get("declared_significant_bits"), "metadata": sensor.get("metadata"), "message": tr( "Panel {panel}: the source metadata ({metadata}) declares " "{bits}-bit camera data, but pixel values reach {peak}. The " "image was rescaled after acquisition or the metadata is wrong; " "saturation was checked against the storage type instead.", panel=panel["panel"], metadata=sensor.get("metadata"), bits=sensor.get("declared_significant_bits"), peak=sensor.get("observed_max")), }) return findings def _lossy_findings(requested, written, panels): """A lossy file format asked for, or used, for image panels.""" if not panels: return [] findings = [] if written in _LOSSY_FORMATS: findings.append({ "check": "lossy_format", "severity": "warning", "panels": [], "message": ( f"{written.upper()} is a lossy format: it changes pixel " f"values in image panels. Use PNG, TIFF or PDF for figures " f"that carry intensities."), }) elif requested in _LOSSY_FORMATS: findings.append({ "check": "lossy_format", "severity": "note", "panels": [], "message": ( f"{requested.upper()} was asked for; the figure was written " f"as {written.upper()}, which keeps pixel values."), }) return findings def _resampling_findings(panels): """Panels written with fewer pixels than their image holds.""" findings = [] for panel in panels: exported = panel.get("exported_pixels") shape = panel.get("shape") or [] if not exported or len(shape) < 2 or min(exported) <= 0: continue factor = max(shape[1] / exported[0], shape[0] / exported[1]) if factor > _RESAMPLE_NOTE_FACTOR: findings.append({ "check": "resampling", "severity": "note", "panels": [panel["panel"]], "factor": round(factor, 2), "message": ( f"Panel {panel['panel']} holds {shape[1]}x{shape[0]} " f"pixels and is written at about {exported[0]}x" f"{exported[1]}; each written pixel stands for ~" f"{factor:.1f} image pixels per side. Raise the " f"resolution to keep single-pixel detail."), }) return findings def _spacr_version(): """The installed spaCR version string, or ``'unknown'``.""" try: from ._version import __version__ return str(__version__) except Exception: return "unknown" def _source_crop_window(panel): """Return one bounded crop on one hashed source, or abstain.""" sources, steps = panel.get("source") or [], panel.get("steps") or [] if len(sources) != 1 or not isinstance(sources[0], dict): return None source = sources[0] path, digest = source.get("path"), source.get("sha256") if (not isinstance(path, str) or not path.lower().endswith((".png", ".npy")) or not isinstance(digest, str) or len(digest) != 64 or not isinstance(source.get("bytes"), int) or source["bytes"] > _SOURCE_REGION_BYTES): return None try: int(digest, 16) except ValueError: return None if (path.lower().endswith(".npy") and isinstance(steps, list) and len(steps) == 1 and isinstance(steps[0], dict) and steps[0].get("op") == "merged_crop"): spec_data = steps[0].get("spec") if not isinstance(spec_data, dict) or spec_data.get("merged_path") != path: return None bbox = spec_data.get("bbox") size = spec_data.get("size") if (not isinstance(bbox, (list, tuple)) or len(bbox) != 4 or any(type(value) is not int for value in bbox) or not isinstance(size, (list, tuple)) or len(size) != 2 or any(type(value) is not int for value in size) or not (0 < (bbox[1] - bbox[0]) * (bbox[3] - bbox[2]) <= _SOURCE_REGION_PIXELS) or not (64 <= min(size) and size[0] * size[1] <= _SOURCE_REGION_PIXELS) or spec_data.get("object_type") == "cytoplasm" or spec_data.get("dilate") or spec_data.get("normalize_by", "png") == "fov"): return None try: from .crops import CropError, CropSpec, MergedField, _region_for spec = CropSpec(**spec_data) array = np.load(path, mmap_mode="r", allow_pickle=False) try: if not isinstance(array, np.memmap): return None field = MergedField(path, array=array, mask_dims=spec.mask_dims) height, width, _channels = field.shape centroid, _bounds, _mask = _region_for(field, spec) y0 = max(0, int(centroid[0]) - size[1] // 2) x0 = max(0, int(centroid[1]) - size[0] // 2) box = (y0, min(height, int(centroid[0]) - size[1] // 2 + size[1]), x0, min(width, int(centroid[1]) - size[0] // 2 + size[0])) if min(box[1] - box[0], box[3] - box[2]) < 64: return None return source, box finally: if isinstance(array, np.memmap): array._mmap.close() except (OSError, ValueError, TypeError, KeyError, IndexError, AttributeError, MemoryError, EOFError, OverflowError, CropError): return None if not path.lower().endswith(".png"): return None if not isinstance(steps, list) or sum( isinstance(step, dict) and step.get("op") == "crop" for step in steps) != 1: return None allowed = {"read_crop_png", "select_channel", "crop", "rescale", "to_uint8"} if any(not isinstance(step, dict) or step.get("op") not in allowed for step in steps): return None for index, step in enumerate(steps): if step["op"] == "read_crop_png" and ( index != 0 or step.get("path") != path): return None box = next(step.get("box") for step in steps if step["op"] == "crop") if (not isinstance(box, (list, tuple)) or len(box) != 4 or any(type(value) is not int for value in box)): return None y0, y1, x0, x1 = box if not (0 <= y0 < y1 and 0 <= x0 < x1 and min(y1 - y0, x1 - x0) >= 64): return None return source, tuple(box) def _source_regions_overlap(first, second): """Require most of both explicit source windows along both axes.""" for lo, hi in ((0, 1), (2, 3)): size_a, size_b = first[hi] - first[lo], second[hi] - second[lo] shared = min(first[hi], second[hi]) - max(first[lo], second[lo]) if (shared < 0.75 * max(size_a, size_b) or shared < 0.90 * min(size_a, size_b)): return False return True def _verified_source_region(panel, previous, first, second, deadline): """Check bounded raw-source texture and exact replay for both exports.""" import time from PIL import Image merged = first[0]["path"].lower().endswith(".npy") for source, box in (first, second): if time.monotonic() >= deadline: return False path = source["path"] if not merged: try: with Image.open(path) as image: width, height = image.size if (width * height > _SOURCE_REGION_PIXELS or not (box[1] <= height and box[3] <= width)): return False except Image.DecompressionBombError: return False current = _source_record(path) if (current.get("sha256") != source["sha256"] or current.get("bytes") != source["bytes"]): return False image = (np.load(first[0]["path"], mmap_mode="r", allow_pickle=False) if merged else _read_panel_source(first[0]["path"])) y0 = max(first[1][0], second[1][0]) y1 = min(first[1][1], second[1][1]) x0 = max(first[1][2], second[1][2]) x1 = min(first[1][3], second[1][3]) try: if merged: spec = panel["steps"][0]["spec"] channel = int(spec["channels"][0]) region = np.asarray(image[y0:y1, x0:x1, channel], dtype=np.float32) else: region = np.asarray(image[y0:y1, x0:x1], dtype=np.float32) if region.ndim == 3: region = region[..., :3].mean(axis=-1) finally: if isinstance(image, np.memmap): image._mmap.close() if (region.ndim != 2 or not np.isfinite(region).all() or float(region.std()) < 1.0): return False for record in (panel, previous): if time.monotonic() >= deadline: return False _rebuilt, matches = _reproduce_panel({"panels": [record]}, record["panel"]) if not matches: return False for source in (first[0], second[0]): if _source_record(source["path"]).get("sha256") != source["sha256"]: return False return time.monotonic() < deadline def _prior_figure_findings(panels, destination): """Find exact reuse and bounded displayed-region similarity nearby. Only closed, small sidecars for existing figures are considered. This is an advisory comparison of two validated displayed-pixel signatures, not an acquisition-source or manipulation verdict. A stale figure cannot report a repeat after replacement, even with coarse modification times. """ import heapq import json import stat import time destination = os.path.abspath(os.fspath(destination)) folder = os.path.dirname(destination) suffix = _PROVENANCE_SUFFIX recent = [] try: with os.scandir(folder) as entries: for entry in entries: if not entry.name.endswith(suffix) or entry.path == destination + suffix: continue try: info = entry.stat(follow_symlinks=False) except OSError: continue if not stat.S_ISREG(info.st_mode) or info.st_size > 1024 * 1024: continue previous_figure = entry.path[:-len(suffix)] try: figure_info = os.stat(previous_figure, follow_symlinks=False) except OSError: continue if (not stat.S_ISREG(figure_info.st_mode) or figure_info.st_mtime_ns > info.st_mtime_ns): continue item = (info.st_mtime_ns, entry.path, previous_figure, figure_info.st_size) if len(recent) < 64: heapq.heappush(recent, item) elif item > recent[0]: heapq.heapreplace(recent, item) except OSError: return [] current = {panel["displayed_sha256"]: panel for panel in panels} current_crops = {} for panel in panels: sources, steps = panel.get("source") or [], panel.get("steps") or [] if len(sources) != 1 or not steps: continue source, step = sources[0], steps[0] if not isinstance(source, dict) or not isinstance(step, dict): continue path, digest = source.get("path"), source.get("sha256") if not path or not digest or step.get("op") not in ( "read_crop_png", "merged_crop"): continue step_path = (step.get("path") if step["op"] == "read_crop_png" else (step.get("spec") or {}).get("merged_path")) if step_path != path: continue key = (path, digest, json.dumps(step, sort_keys=True)) current_crops.setdefault(key, []).append(panel) source_started = time.monotonic() current_windows = [(panel, _source_crop_window(panel)) for panel in panels[:_SPATIAL_ATTEMPTS] if time.monotonic() - source_started < _SOURCE_REGION_SECONDS] findings, seen, seen_crops, seen_regions = [], set(), set(), set() spatial_pairs = 0 spatial_started = time.monotonic() source_pairs = 0 for _mtime, sidecar, old_figure, old_bytes in sorted(recent, reverse=True): try: with open(sidecar, "r", encoding="utf-8") as handle: old = json.load(handle) if old.get("schema") != _PROVENANCE_SCHEMA: continue old_panels = old.get("panels", []) crop_matches = [] for previous in old_panels: sources, steps = (previous.get("source") or [], previous.get("steps") or []) if len(sources) != 1 or not steps: continue source, step = sources[0], steps[0] if not isinstance(source, dict) or not isinstance(step, dict): continue path, digest = source.get("path"), source.get("sha256") if not path or not digest or step.get("op") not in ( "read_crop_png", "merged_crop"): continue step_path = (step.get("path") if step["op"] == "read_crop_png" else (step.get("spec") or {}).get("merged_path")) if step_path != path: continue key = (path, digest, json.dumps(step, sort_keys=True)) if key in current_crops and key not in seen_crops: crop_matches.append((key, previous)) source_regions = [] if (old_bytes <= 64 * 1024 ** 2 and time.monotonic() - source_started < _SOURCE_REGION_SECONDS): for panel, first in current_windows: if first is None or panel["panel"] in seen_regions: continue for previous in old_panels[:128]: if time.monotonic() - source_started >= _SOURCE_REGION_SECONDS: break prior_sources = previous.get("source") or [] if (len(prior_sources) != 1 or not isinstance(prior_sources[0], dict) or prior_sources[0].get("sha256") != first[0]["sha256"] or os.path.splitext(str(prior_sources[0].get("path")))[1].lower() != os.path.splitext(first[0]["path"])[1].lower()): continue second = _source_crop_window(previous) if (second is None or first[0]["sha256"] != second[0]["sha256"] or first[1] == second[1] or panel["displayed_sha256"] == previous.get("displayed_sha256") or not _source_regions_overlap(first[1], second[1])): continue source_regions.append((panel, previous, first, second)) if len(source_regions) + source_pairs >= _SOURCE_REGION_PAIRS: break if len(source_regions) + source_pairs >= _SOURCE_REGION_PAIRS: break similar = [] for panel in panels: if (not panel.get("spatial_v1") or panel["panel"] in seen_regions or panel["displayed_sha256"] in seen or old_bytes > 64 * 1024 ** 2): continue for previous in old_panels[:128]: if (spatial_pairs >= _SPATIAL_PAIRS or time.monotonic() - spatial_started >= _SPATIAL_SECONDS): break if (not isinstance(previous, dict) or not previous.get("spatial_v1") or previous.get("displayed_sha256") == panel["displayed_sha256"]): continue spatial_pairs += 1 try: matched = _similar_displayed_region( panel["spatial_v1"], previous["spatial_v1"]) except (ValueError, cv2.error): continue if matched: similar.append((panel, previous, matched)) if (not crop_matches and not source_regions and not similar and not any( p.get("displayed_sha256") in current for p in old_panels)): continue from .run_journal import hash_file if old.get("figure_sha256") != hash_file(old_figure, full=True): continue for previous in old_panels: digest = previous.get("displayed_sha256") if digest not in current or digest in seen: continue panel = current[digest] seen.add(digest) prior_sources = [s.get("path") for s in previous.get("source", [])] sources = [s.get("path") for s in panel.get("source", [])] declared = bool(sources and sources == prior_sources) findings.append({ "check": "cross_figure_duplicate", "severity": "note" if declared else "warning", "panels": [panel["panel"]], "prior_figure": os.path.basename(old_figure), "prior_panel": previous.get("panel"), "identical": True, "message": ( f"Panel {panel['panel']} has identical pixels to panel " f"{previous.get('panel')} in an earlier export " f"({os.path.basename(old_figure)}). " + ("They share a recorded source; explain the reuse in " "the legend." if declared else "If this reuse is intentional, explain it in the legend.")), }) for key, previous in crop_matches: for panel in current_crops[key]: identity = (key, panel["panel"]) if (identity in seen_crops or panel["displayed_sha256"] in seen): continue seen_crops.add(identity) findings.append({ "check": "cross_figure_source_crop", "severity": "note", "panels": [panel["panel"]], "prior_figure": os.path.basename(old_figure), "prior_panel": previous.get("panel"), "identical": False, "message": ( f"Panel {panel['panel']} records the same source " f"bytes and crop recipe as panel {previous.get('panel')} " f"in an earlier export ({os.path.basename(old_figure)}), " "but its displayed pixels differ. Explain the reused " "source crop and the display change in the legend."), }) for panel, previous, first, second in source_regions: if (source_pairs >= _SOURCE_REGION_PAIRS or time.monotonic() - source_started >= _SOURCE_REGION_SECONDS): break if (panel["panel"] in seen_regions or panel["displayed_sha256"] in seen): continue source_pairs += 1 try: verified = _verified_source_region( panel, previous, first, second, source_started + _SOURCE_REGION_SECONDS) except (OSError, ValueError, TypeError, AttributeError, KeyError, IndexError, MemoryError): continue if not verified: continue from .qt.i18n import tr seen_regions.add(panel["panel"]) findings.append({ "check": "cross_figure_source_region", "severity": "note", "panels": [panel["panel"]], "prior_figure": os.path.basename(old_figure), "prior_panel": previous.get("panel"), "source_sha256": first[0]["sha256"], "source_boxes": [list(first[1]), list(second[1])], "message": tr( "Panel {panel} and panel {prior_panel} in {prior_figure} " "show overlapping regions of the same verified source " "pixels. This is a source reuse note, not an image " "manipulation verdict.", panel=panel["panel"], prior_panel=previous.get("panel"), prior_figure=os.path.basename(old_figure)), }) if similar: from .qt.i18n import tr for panel, previous, matched in similar: if panel["displayed_sha256"] in seen or panel["panel"] in seen_regions: continue seen_regions.add(panel["panel"]) findings.append({ "check": "similar_displayed_region", "severity": "note", "panels": [panel["panel"]], "prior_figure": os.path.basename(old_figure), "prior_panel": previous.get("panel"), "detail": matched["detail"], "structure": matched["structure"], "fraction": matched["fraction"], "message": tr( "Panel {panel} has a similar displayed region to panel " "{prior_panel} in {prior_figure}. This is a bounded " "pixel similarity note, not an acquisition or image " "manipulation verdict.", panel=panel["panel"], prior_panel=previous.get("panel"), prior_figure=os.path.basename(old_figure)), }) except (OSError, ValueError, TypeError, AttributeError, KeyError): continue if len(seen) == len(current): break return findings def _index_hash(grey): """A 64-bit block-mean hash of a 128x128 grey signature, as 16 hex digits.""" blocks = np.asarray(grey, dtype=np.float32).reshape(8, 16, 8, 16).mean(axis=(1, 3)) bits = (blocks > np.median(blocks)).ravel() return f"{int(''.join('1' if bit else '0' for bit in bits), 2):016x}" def _index_hashes(grey): """The hashes of every flip and rotation of a grey signature.""" found = set() for quarter in range(4): turned = np.rot90(grey, quarter) found.add(_index_hash(turned)) found.add(_index_hash(turned[:, ::-1])) return found def _figure_index_paths(destination): """Every index file a checked export reads and appends to. One beside the figure, one in the open run's folder, and the file named by ``SPACR_FIGURE_INDEX`` when that is set, so a project can share one index across output folders. """ folder = os.path.dirname(os.path.abspath(os.fspath(destination))) paths = [os.path.join(folder, _INDEX_NAME)] try: from .run_journal import current_run active = current_run() if active is not None: paths.append(os.path.join(str(active.dir), _INDEX_NAME)) except Exception: pass shared = os.environ.get(_INDEX_ENV, "").strip() if shared: paths.append(os.path.abspath(shared)) unique = [] for path in paths: if path not in unique: unique.append(path) return unique def _read_figure_index(path): """The newest valid entry per figure in one index file, newest first.""" import json try: with open(path, "rb") as handle: handle.seek(0, os.SEEK_END) size = handle.tell() handle.seek(max(0, size - _INDEX_LIMIT)) data = handle.read(_INDEX_LIMIT) except OSError: return [] lines = data.split(b"\n") if size > _INDEX_LIMIT: lines = lines[1:] latest = {} for line in lines: try: entry = json.loads(line) except ValueError: continue if (not isinstance(entry, dict) or entry.get("v") != 1 or not isinstance(entry.get("figure"), str) or not isinstance(entry.get("panels"), list)): continue latest.pop(entry["figure"], None) latest[entry["figure"]] = entry return list(reversed(list(latest.values()))) def _record_in_index(report, figure_path, sidecar): """Append one checked export's panel signatures to its index files. An index that grows past its size limit keeps its newest half. """ import json entry = { "v": 1, "figure": os.path.abspath(str(figure_path)), "figure_sha256": report.get("figure_sha256"), "sidecar": os.path.abspath(str(sidecar)), "created": report.get("created"), "panels": [{"panel": panel["panel"], "displayed_sha256": panel.get("displayed_sha256"), "hash": panel.get("index_hash"), "sources": [source.get("path") for source in panel.get("source") or [] if isinstance(source, dict)]} for panel in report.get("panels", [])], } line = (json.dumps(entry, separators=(",", ":"), default=str) + "\n").encode() if len(line) > _INDEX_LIMIT // 8: return for path in _figure_index_paths(figure_path): try: os.makedirs(os.path.dirname(path), exist_ok=True) with open(path, "ab") as handle: handle.write(line) size = handle.tell() if size > _INDEX_LIMIT: with open(path, "rb") as handle: handle.seek(size - _INDEX_LIMIT // 2) kept = handle.read().split(b"\n", 1)[-1] temporary = f"{path}.tmp" with open(temporary, "wb") as handle: handle.write(kept) os.replace(temporary, path) except OSError: continue def _indexed_prior_panel(entry, previous, figure_hashes): """The prior panel record behind an index entry, after verification. The prior figure must still hash to what the index and its sidecar recorded, so a replaced or edited file is never used as evidence. """ import json import stat from .run_journal import hash_file figure, sidecar = entry.get("figure"), entry.get("sidecar") if not isinstance(sidecar, str) or not isinstance(figure, str): return None try: info = os.stat(sidecar) if not stat.S_ISREG(info.st_mode) or info.st_size > _SPATIAL_SIDECAR_LIMIT: return None figure_info = os.stat(figure) if not stat.S_ISREG(figure_info.st_mode) or figure_info.st_size > 256 * 1024 ** 2: return None with open(sidecar, encoding="utf-8") as handle: old = json.load(handle) except (OSError, ValueError): return None if figure not in figure_hashes: figure_hashes[figure] = hash_file(figure, full=True) digest = figure_hashes[figure] if (not isinstance(old, dict) or digest != entry.get("figure_sha256") or digest != old.get("figure_sha256")): return None for record in old.get("panels") or []: if (isinstance(record, dict) and record.get("panel") == previous.get("panel") and record.get("displayed_sha256") == previous.get("displayed_sha256")): return record return None def _index_findings(panels, destination, findings): """Repeats of this figure's panels in any figure the index has seen. Candidates come from identical pixel digests and from block-mean hashes within a few bits under any flip or rotation; an approximate candidate is a repeat only when the stored signatures pass the same layout and texture test as panels within one figure. :param findings: findings so far; a prior panel already reported by the folder scan is not reported twice. """ from .qt.i18n import tr destination = os.path.abspath(os.fspath(destination)) reported = {(finding.get("prior_figure"), finding.get("prior_panel"), finding["panels"][0]) for finding in findings if finding.get("check", "").startswith("cross_figure")} current, greys = [], {} for panel in panels: variants = set() if panel.get("spatial_v1"): try: grey = _decode_spatial_signature(panel["spatial_v1"]) except ValueError: grey = None if grey is not None: greys[panel["panel"]] = grey variants = _index_hashes(grey) current.append((panel, {int(value, 16) for value in variants})) entries, seen_figures = [], set() for path in _figure_index_paths(destination): for entry in _read_figure_index(path): if entry["figure"] in seen_figures or entry["figure"] == destination: continue seen_figures.add(entry["figure"]) entries.append(entry) candidates = [] for entry in entries[:_INDEX_ENTRIES]: for previous in entry["panels"]: if not isinstance(previous, dict): continue try: prior_hash = int(previous.get("hash") or "", 16) except ValueError: prior_hash = None for panel, hashes in current: exact = previous.get("displayed_sha256") == panel["displayed_sha256"] near = prior_hash is not None and any( bin(prior_hash ^ value).count("1") <= _INDEX_HAMMING for value in hashes) if exact or near: candidates.append((entry, previous, panel, exact)) found, figure_hashes, confirmed = [], {}, set() for entry, previous, panel, exact in candidates[:_INDEX_CONFIRM]: name = os.path.basename(entry["figure"]) identity = (name, previous.get("panel"), panel["panel"]) if identity in reported or panel["panel"] in confirmed: continue record = _indexed_prior_panel(entry, previous, figure_hashes) if record is None: continue score = None if not exact: try: prior = _decode_spatial_signature(record.get("spatial_v1")) except ValueError: continue first = _panel_thumbnail(greys[panel["panel"]]) second = _panel_thumbnail(prior) if first is None or second is None: continue score = _best_dihedral_correlation(first[0], second[0]) if (score < _DUPLICATE_CORRELATION or _best_dihedral_correlation(first[1], second[1]) < _DUPLICATE_DETAIL): continue prior_sources = [source.get("path") for source in record.get("source") or [] if isinstance(source, dict)] sources = [source.get("path") for source in panel.get("source") or []] declared = bool(sources and sources == prior_sources) reported.add(identity) confirmed.add(panel["panel"]) found.append({ "check": "cross_figure_index", "severity": "note" if declared else "warning", "panels": [panel["panel"]], "prior_figure": name, "prior_figure_path": entry["figure"], "prior_panel": previous.get("panel"), "identical": bool(exact), "correlation": None if score is None else round(float(score), 4), "message": tr( "Panel {panel} shows the same image as panel {prior_panel} in " "an earlier export ({prior_figure}){how}. If the reuse is " "intentional, explain it in the legend.", panel=panel["panel"], prior_panel=previous.get("panel"), prior_figure=name, how="" if exact else tr( ", allowing flips, rotations and contrast changes")), }) return found def _integrity_report(fig, *, fmt, requested_fmt=None, dpi=None, destination=None): """Check ``fig``'s image panels and assemble its provenance. :param fig: the figure about to be written. :param fmt: the format that will be written. :param requested_fmt: the format the caller asked for, if different. :param dpi: the resolution it will be written at. :param destination: optional output path for checking earlier exports in the same folder for an identical image panel. :returns: a JSON-ready report, or ``None`` when the figure holds no image panel of at least 16x16 pixels. """ import datetime import platform dpi = float(dpi or fig.dpi) panels, arrays, source_cache = [], [], {} for axes_index, axes, artist in _figure_panels(fig): record, data = _panel_record(len(panels), axes_index, axes, artist, fig, dpi, source_cache) if data.ndim < 2 or min(data.shape[:2]) < _MIN_PANEL_SIDE: continue panels.append(record) arrays.append(data) if not panels: return None written = str(fmt or "").lower().lstrip(".") requested = str(requested_fmt or written).lower().lstrip(".") duplicates = _duplicate_findings(panels, arrays) regions, region_search = _region_reuse_findings( panels, arrays, [finding["panels"] for finding in duplicates]) findings = (_range_findings(panels) + _saturation_findings(panels) + _bit_depth_findings(panels) + duplicates + regions + _splice_findings(panels, arrays) + _lossy_findings(requested, written, panels) + _resampling_findings(panels)) run = None try: from .run_journal import current_run active = current_run() if active is not None: run = {"dir": str(active.dir), "app": str(active.app_key), "manifest": str(os.path.join(str(active.dir), "manifest.json"))} except Exception: run = None import matplotlib report = { "schema": _PROVENANCE_SCHEMA, "created": datetime.datetime.now(datetime.timezone.utc).isoformat( timespec="seconds"), "software": {"spacr": _spacr_version(), "matplotlib": matplotlib.__version__, "numpy": np.__version__, "python": platform.python_version()}, "format": written, "requested_format": requested, "dpi": dpi, "figure_inches": [float(v) for v in fig.get_size_inches()], "run": run, "panels": panels, "integrity": { "checks": ["display_range", "saturation", "duplicate", "lossy_format", "resampling", "cross_figure_duplicate", "cross_figure_source_crop", "cross_figure_source_region", "similar_displayed_region", "region_reuse", "splice", "bit_depth", "cross_figure_index"], "region_search": region_search, "warnings": sum(f["severity"] == "warning" for f in findings), "notes": sum(f["severity"] == "note" for f in findings), "findings": findings, }, } _attach_spatial_signatures(report, arrays) for panel in panels: if panel.get("spatial_v1"): try: panel["index_hash"] = _index_hash( _decode_spatial_signature(panel["spatial_v1"])) except ValueError: pass if destination is not None: findings += _prior_figure_findings(panels, destination) try: findings += _index_findings(panels, destination, findings) except (OSError, ValueError, TypeError, KeyError, AttributeError): pass report["integrity"]["warnings"] = sum( f["severity"] == "warning" for f in findings) report["integrity"]["notes"] = sum( f["severity"] == "note" for f in findings) _attach_spatial_signatures(report, ()) retained = {panel["panel"] for panel in panels if "spatial_v1" in panel} findings[:] = [finding for finding in findings if finding["check"] != "similar_displayed_region" or finding["panels"][0] in retained] report["integrity"]["notes"] = sum( f["severity"] == "note" for f in findings) return report def _integrity_metadata(report, fmt, existing=None): """Metadata that stamps ``report`` into the written file. PNG files carry the whole provenance as a ``spaCR provenance`` text chunk; PDF files carry it in the document ``Subject``. Keys the caller already set are kept. Other formats get no stamp and rely on the sidecar. :returns: the metadata dict to pass to ``savefig``, or ``existing``. """ import json merged = dict(existing or {}) text = json.dumps(report, separators=(",", ":"), default=str) if fmt == "png": merged.setdefault(_PNG_PROVENANCE_KEY, text) elif fmt == "pdf": merged.setdefault("Subject", f"{_PNG_PROVENANCE_KEY}: {text}") else: return existing return merged def _provenance_sidecar_path(figure_path): """Where the sidecar for ``figure_path`` is written.""" return f"{figure_path}{_PROVENANCE_SUFFIX}" def _finish_integrity(report, figure_path): """Write the sidecar for a figure just written and announce findings. Warnings are printed, and recorded on the open run journal when there is one; the sidecar holds every finding and the provenance of every panel, plus the written file's SHA-256. :returns: the sidecar path, or ``None`` if it could not be written. """ import json from .run_journal import hash_file figure_path = str(figure_path) report = dict(report) report["figure"] = os.path.basename(figure_path) report["figure_sha256"] = hash_file(figure_path, full=True) sidecar = _provenance_sidecar_path(figure_path) try: temporary = f"{sidecar}.tmp" with open(temporary, "w", encoding="utf-8") as handle: json.dump(report, handle, indent=2, default=str) os.replace(temporary, sidecar) except OSError: sidecar = None findings = report["integrity"]["findings"] run = None try: from .run_journal import current_run run = current_run() except Exception: run = None for finding in findings: if finding["severity"] != "warning": continue line = (f"Figure integrity ({os.path.basename(figure_path)}): " f"{finding['message']}") print(line) if run is not None: try: run.record_warning(line) except Exception: pass if findings or sidecar: count = report["integrity"]["warnings"] print(f"Figure integrity: {count} warning(s), " f"{report['integrity']['notes']} note(s) for {figure_path}" + (f"; provenance in {sidecar}" if sidecar else "")) if run is not None and sidecar: try: run.record_output(sidecar, setting_key="figure_provenance") except Exception: pass if sidecar: _record_in_index(report, figure_path, sidecar) return sidecar def _read_panel_source(path): """Read one source image as an array, by extension.""" extension = os.path.splitext(str(path))[1].lower() if extension == ".npy": return np.load(path, allow_pickle=False) if extension in (".tif", ".tiff"): return tiff.imread(path) from PIL import Image with Image.open(path) as image: return np.array(image) def _overlay_replay_colored(mask, seed): """Apply the original deterministic label palette to one replay mask.""" count = int(mask.max() + 1) if count > 65536: raise ValueError("overlay colour table exceeds replay budget") if count <= 0: cmap = ListedColormap(np.array([[0, 0, 0]])) else: rng = np.random.default_rng(seed) hues = np.linspace(0, 1, count, endpoint=False) rng.shuffle(hues) sats = rng.uniform(0.70, 1.00, size=count) vals = rng.uniform(0.85, 1.00, size=count) colors = mpl.colors.hsv_to_rgb(np.column_stack([hues, sats, vals])) cmap = ListedColormap(np.vstack([[0, 0, 0], colors])) result = cmap(mask / (mask.max() + 1e-5)) result[..., 3] = np.where(mask > 0, 1, 0) return result def _replay_overlay(image, step): """Rebuild a merged-stack overlay from its recorded plane and filter recipe.""" if np.asarray(image).ndim != 3 or np.asarray(image).dtype.kind not in "uif": raise ValueError("overlay source is not a numeric merged stack") height, width, depth = image.shape names = step.get("mask_order") planes = step.get("mask_planes") if (not isinstance(names, list) or len(names) > 16 or any(not isinstance(name, str) for name in names) or len(names) != len(set(names)) or not isinstance(planes, dict) or set(names) != set(planes)): raise ValueError("invalid overlay mask recipe") if (not height or not width or height * width * image.dtype.itemsize > 256 * 1024 ** 2 or height * width * (max(image.dtype.itemsize, 4) * ( len(names) + 5) + 32) > 1536 * 1024 ** 2): raise ValueError("overlay plane exceeds the replay budget") from .object_roles import ORGANELLE_ROLES if set(names) - ({"cell", "nucleus", "pathogen"} | set(ORGANELLE_ROLES)): raise ValueError("unknown overlay mask role") indices = {} for name in names: index = planes[name] if not isinstance(name, str) or type(index) is not int or not 0 <= index < depth: raise ValueError("invalid overlay mask plane") indices[name] = index filters = step.get("filters") or {} channels = step.get("mask_channels") or {} if not isinstance(filters, dict) or not isinstance(channels, dict): raise ValueError("invalid overlay filter recipe") if set(filters) - set(names) or set(channels) - set(names): raise ValueError("unknown overlay mask role") outlines = [] for name in names: mask = np.take(image, indices[name], axis=2) if image.dtype in (np.uint8, np.uint16): mask = mask.astype(np.float32) if name in filters: channel = channels.get(name) bounds = filters[name] if (type(channel) is not int or not 0 <= channel < depth or len(bounds) != 2 or any(len(pair) != 2 for pair in bounds)): raise ValueError("invalid overlay filter or intensity channel") intensity = np.take(image, channel, axis=2) if image.dtype in (np.uint8, np.uint16): intensity = intensity.astype(np.float32) original_dtype = mask.dtype mask_int = mask.astype(np.int64) intensity = intensity.astype(np.float64) kept = np.zeros_like(mask_int) for label in np.unique(mask_int): if label == 0: continue selected = mask_int == label area = np.sum(selected) mean = np.mean(intensity[selected]) if (bounds[0][0] <= area <= bounds[0][1] and bounds[1][0] <= mean <= bounds[1][1]): kept[selected] = label mask = kept.astype(original_dtype) outlines.append(mask) if step["op"] == "combined_masks": if not outlines: raise ValueError("combined panel has no masks") combined = np.zeros_like(outlines[0], dtype=np.int64) offset = 0 for outline in outlines: labels = outline.astype(np.int64) selected = labels > 0 if np.any(selected): combined[selected] = labels[selected] + offset offset += int(labels.max()) rgba = _overlay_replay_colored(combined, 9999) blank = np.zeros((*combined.shape, 3)) return np.clip(blank * (1 - rgba[..., 3:]) + rgba[..., :3] * rgba[..., 3:], 0, 1) channel = step.get("channel") if type(channel) is not int or not 0 <= channel < depth: raise ValueError("invalid overlay image channel") plane = np.take(image, channel, axis=2) if image.dtype in (np.uint8, np.uint16): plane = plane.astype(np.float32) percentiles = step.get("percentiles") if (not isinstance(percentiles, list) or len(percentiles) != 2 or any(not isinstance(value, (int, float)) or not np.isfinite(value) for value in percentiles) or not 0 <= percentiles[0] < percentiles[1] <= 100): raise ValueError("invalid overlay percentiles") low, high = np.percentile(plane, percentiles) grey = np.clip((plane - low) / (high - low + 1e-5), 0, 1) rendered = np.dstack([grey] * 3) colors = outline_palette_colours(step.get("outline_palette")) roles = {"cell": colors["cell"], "nucleus": colors["nucleus"], "pathogen": colors["pathogen"], "organelle": colors["organelle"]} for index, name in enumerate(name for name in names if name.startswith("organelle") and name != "organelle"): roles[name] = _organelle_slot_colour(step.get("outline_palette"), index) mapped = {channels[name]: index for index, name in enumerate(names) if name in channels} if step.get("all_on_all"): selected = list(range(len(names))) elif channel in mapped: selected = [mapped[channel]] elif step.get("all_outlines"): selected = list(range(len(names))) else: selected = [] for index in selected: mask = outlines[index] if step.get("mode") == "outlines": for label in np.unique(mask): if label == 0: continue contours, _ = cv2.findContours( (mask == label).astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) cv2.drawContours(rendered, contours, -1, mpl.colors.to_rgb(roles[names[index]]), int(step["thickness"])) else: seed = (1000 + index if step.get("all_on_all") or channel not in mapped else channels[names[index]] + 1000) if channel in mapped and not step.get("all_on_all"): seed = 1000 + index rgba = _overlay_replay_colored(mask, seed) rendered = np.clip(rendered * (1 - rgba[..., 3:]) + rgba[..., :3] * rgba[..., 3:], 0, 1) return rendered def _replay_steps(image, steps): """Apply recorded display steps to a source array. :raises ValueError: for an op that cannot be replayed. """ for step in steps: op = step.get("op") if op == "select_channel": image = np.asarray(image)[..., int(step["index"])] elif op == "crop": y0, y1, x0, x1 = (int(v) for v in step["box"]) image = np.asarray(image)[y0:y1, x0:x1] elif op == "max_project": image = np.asarray(image).max(axis=int(step.get("axis", 0))) elif op == "rescale": image = _apply_display_ranges(image, _display_ranges(step["ranges"])) elif op == "to_uint8": image = (np.asarray(image) * 255).astype(np.uint8) elif op == "read_crop_png": from .crops import read_crop_png image = read_crop_png(step["path"], fmt=int(step["format"])) elif op == "merged_crop": from .crops import CropSpec, extract_crop, png_view spec = CropSpec(**step["spec"]) image = png_view(extract_crop(spec.merged_path, spec=spec)) elif op in ("overlay_composite", "combined_masks"): image = _replay_overlay(image, step) else: raise ValueError(f"step {op!r} cannot be replayed") return np.asarray(image) def _reproduce_panel(sidecar, panel): """Rebuild one exported panel from its source files and sidecar. :param sidecar: path to a ``.provenance.json`` sidecar, or its loaded dict. :param panel: the panel number in the sidecar. :returns: ``(array, matches)`` -- the rebuilt displayed array, and whether it is bit-identical to what was exported. :raises ValueError: when the panel names no source or holds a step that cannot be replayed. """ import json if not isinstance(sidecar, dict): with open(sidecar, encoding="utf-8") as handle: sidecar = json.load(handle) record = next(p for p in sidecar["panels"] if p["panel"] == int(panel)) if not record.get("source"): raise ValueError(f"panel {panel} names no source file") path = os.path.abspath(record["source"][0]["path"]) for step in record.get("steps") or []: if step.get("op") == "read_crop_png": recorded = os.path.abspath(step["path"]) elif step.get("op") == "merged_crop": recorded = os.path.abspath(step["spec"]["merged_path"]) else: continue if recorded != path: raise ValueError("crop recipe names a different source file") overlay = any(step.get("op") in ("overlay_composite", "combined_masks") for step in record.get("steps") or []) if overlay: if len(record["source"]) != 1 or not path.lower().endswith(".npy"): raise ValueError("overlay recipe needs one merged NPY source") source = record["source"][0] if not source.get("sha256"): raise ValueError("overlay source has no full export digest") current = _source_record(path) for key in ("exists", "bytes", "mtime", "sha256"): if key in source and current.get(key) != source[key]: raise ValueError("overlay source changed since export") image = np.load(path, mmap_mode="r", allow_pickle=False) if not isinstance(image, np.memmap): if hasattr(image, "close"): image.close() raise ValueError("overlay source is not a merged NPY array") try: rebuilt = _replay_steps(image, record.get("steps") or []) finally: image._mmap.close() if _source_record(path).get("sha256") != source["sha256"]: raise ValueError("overlay source changed during replay") else: image = _read_panel_source(path) rebuilt = _replay_steps(image, record.get("steps") or []) return rebuilt, _array_digest(rebuilt) == record["displayed_sha256"] #: Outline colours for the published overlay figure, per palette. #: #: ``default`` is what spaCR has always drawn and stays the default: changing #: it silently would change every figure a user has already made and every one #: their paper's methods section describes. #: #: IT IS ALSO NOT LEGIBLE TO EVERYONE, and the numbers say how badly. Worst #: confusable pair across the four outlines, 0-255, under Brettel-style #: simulation -- the pair a reader would have to tell apart and could not: #: #: palette normal deuteranope protanope tritanope #: default 255 27 77 39 #: colourblind 148 142 125 134 #: #: 27 is not a small number, it is invisible. cell is drawn RED and pathogen #: GREEN, which is the one pair red-green deficiency removes, in the figure #: that goes into the paper. #: #: ``colourblind`` is the four of the Okabe-Ito set that scored best on that #: worst-pair measure. It gives up separation under normal vision (148 against #: 255) to buy five times as much under every deficiency, which is the right #: trade for a figure with more than one reader. OUTLINE_PALETTES = { 'default': {'cell': 'red', 'nucleus': 'blue', 'pathogen': 'green', 'organelle': 'yellow'}, 'colourblind': {'cell': '#D55E00', 'nucleus': '#56B4E9', 'pathogen': '#009E73', 'organelle': '#F0E442'}, }
[docs] def outline_palette_colours(palette): """The four outline colours for ``palette``. :param palette: a key of :data:`OUTLINE_PALETTES`. Anything unknown -- including ``None`` -- falls back to ``default`` rather than raising: a figure drawn in the historic colours is a far smaller problem than a pipeline that stops at the plotting step. :returns: ``{object_name: colour}``. """ name = str(palette or 'default').strip().lower() return dict(OUTLINE_PALETTES.get(name, OUTLINE_PALETTES['default']))
#: Outline colours for organelle slots 2 onward, cycled (item 76, #: 2026-09-30). The colourblind list is the rest of the Okabe-Ito set the #: four fixed colours are taken from. _ORGANELLE_SLOT_COLOURS = { 'default': ('magenta', 'cyan', 'orange', 'white', 'purple', 'lime'), 'colourblind': ('#CC79A7', '#0072B2', '#E69F00', '#FFFFFF'), } def _organelle_slot_colour(palette, index): """Outline colour of the ``index``-th organelle slot after the first. :param palette: a key of :data:`OUTLINE_PALETTES`; unknown means default. :param index: 0 for the second slot, 1 for the third, and so on. :returns: a matplotlib colour. """ name = str(palette or 'default').strip().lower() cycle = _ORGANELLE_SLOT_COLOURS.get( name, _ORGANELLE_SLOT_COLOURS['default']) return cycle[index % len(cycle)] def _extra_organelle_slots(organelle_channels): """``(role, channel)`` for organelle slots 2 onward, in slot order. :param organelle_channels: ``{role: channel}`` or ``None``. The first slot and roles that are not organelle slots are ignored, since the first slot has its own ``organelle_channel`` argument. :returns: the slots whose channel is set. """ from .object_roles import ORGANELLE_ROLES given = dict(organelle_channels or {}) return [(role, given[role]) for role in ORGANELLE_ROLES[1:] if given.get(role) is not None] def _overlay_mask_dims(file, names, n_planes): """Which plane of a merged stack holds each object's mask. The merged folder's plane layout sidecar is the record, and is used when it names every object drawn: counting planes back from the end put every mask one plane off whenever the stack held an object the caller did not ask for, such as a second organelle slot. Without a usable sidecar the masks are taken to be the last ``len(names)`` planes, in order, as before. :param file: path of the merged ``.npy`` stack. :param names: object roles to draw, in mask-plane order. :param n_planes: number of planes in the stack. :returns: ``{role: plane index}``. """ import json from .crops import MERGED_LAYOUT_SIDECAR sidecar = os.path.join(os.path.dirname(str(file)), MERGED_LAYOUT_SIDECAR) try: with open(sidecar, 'r', encoding='utf-8') as handle: dims = dict(json.load(handle).get('mask_dims') or {}) except (OSError, ValueError, AttributeError): dims = {} if names and all(name in dims and 0 <= int(dims[name]) < n_planes for name in names): return {name: int(dims[name]) for name in names} base = n_planes - len(names) return {name: base + offset for offset, name in enumerate(names)}
[docs] def plot_image_mask_overlay( file, channels, cell_channel, nucleus_channel, pathogen_channel, organelle_channel=None, figuresize=10, percentiles=(2, 98), thickness=3, save_pdf=True, mode='outlines', export_tiffs=False, all_on_all=False, all_outlines=False, filter_dict=None, outline_palette='default', organelle_channels=None ): """Plot image and mask overlays. Loads the merged ``.npy`` stack, draws one panel per requested channel with the object masks applied as contours or filled labels, and closes with a panel showing every object combined. :param file: Path to the merged ``.npy`` stack for one field of view. :param channels: Indices of the image channels to draw, one panel each. :param cell_channel: Intensity channel the cell mask belongs to, or ``None`` when there is no cell mask. :param nucleus_channel: Intensity channel the nucleus mask belongs to, or ``None``. :param pathogen_channel: Intensity channel the pathogen mask belongs to, or ``None``. :param organelle_channel: Intensity channel the organelle mask belongs to, or ``None``. Default ``None``. :param figuresize: Figure height in inches; the figure is drawn four times as wide. Default ``10``. :param percentiles: Two-element percentile pair used to normalise each channel. Default ``(2, 98)``. :param thickness: Contour line width in pixels. Default ``3``. :param save_pdf: If True, save the figure into ``results/overlay/`` two directories above ``file``, in the configured figure format rather than always as PDF. Default ``True``. :param mode: ``'outlines'`` draws mask contours; any other value overlays filled, randomly coloured labels. Default ``'outlines'``. :param export_tiffs: If True, also write every stack plane as a grayscale TIFF into ``results/<stem>/tiff/`` alongside it. Default ``False``. :param all_on_all: If True, draw every mask on every channel. Default ``False``. :param all_outlines: If True, draw every mask on the channels that own no mask themselves. Default ``False``. :param filter_dict: Optional per-object limits keyed by ``'cell'``, ``'nucleus'``, ``'pathogen'`` or ``'organelle'``, each holding ``((min_area, max_area), (min_intensity, max_intensity))``; objects outside the limits are dropped before plotting. :param outline_palette: which outline colours to draw, a key of :data:`OUTLINE_PALETTES`. ``'default'`` is what spaCR has always drawn; ``'colourblind'`` is legible under red-green and blue-yellow deficiency, where the default's worst pair scores 27 out of 255 -- cell is drawn red and pathogen green, the one pair the commonest deficiency removes. Default ``'default'``, because changing every figure a user has already made would be worse than the defect. :param organelle_channels: Optional ``{slot role: channel}`` for the organelle slots after the first (``{'organelleb': 3}``), each drawn in its own colour. Default ``None`` draws the first slot only, as before. Mask planes are located through the merged folder's plane layout sidecar when one is present, so a slot left out here no longer shifts the planes of the objects that are drawn. :returns: The generated matplotlib ``Figure``. """ def random_color_cmap(n_labels, seed=None): """Generate a random-looking but deterministic colormap with a unique seed. :param n_labels: How many object colours to draw. Index 0 of the returned colormap is forced to black for background, so the map holds ``n_labels + 1`` entries; callers here pass ``int(outline.max() + 1)`` per object type, or ``int(combined_mask.max() + 1)`` for the merged panel. A value ``<= 0`` short-circuits to a black-only colormap rather than raising. :param seed: Seed for a local ``default_rng``; the same seed always produces the same hue assignment, which is why each object type is given its own fixed seed and so keeps its colours across panels. ``None`` draws fresh entropy and colours change per call. :returns: A ``ListedColormap`` of vivid, well-separated hues. """ if n_labels <= 0: return ListedColormap(np.array([[0, 0, 0]])) rng = np.random.default_rng(seed) hues = np.linspace(0, 1, n_labels, endpoint=False) rng.shuffle(hues) sats = rng.uniform(0.70, 1.00, size=n_labels) vals = rng.uniform(0.85, 1.00, size=n_labels) rand_colors = mpl.colors.hsv_to_rgb(np.column_stack([hues, sats, vals])) rand_colors = np.vstack([[0, 0, 0], rand_colors]) return ListedColormap(rand_colors) def _plot_merged_plot( image, outlines, outline_colors, figuresize, thickness, percentiles, mode='outlines', all_on_all=False, all_outlines=False, channels=None, channel_to_outline=None, channel_to_label=None, save_pdf=True ): """Plot the merged plot with overlay, image channels, and masks.""" def _generate_colored_mask(mask, cmap): """Generate a colored mask using the given colormap.""" mask_norm = mask / (mask.max() + 1e-5) colored_mask = cmap(mask_norm) colored_mask[..., 3] = np.where(mask > 0, 1, 0) return colored_mask def _overlay_mask(image, mask): """Overlay the colored mask onto the original image.""" combined = np.clip(image * (1 - mask[..., 3:]) + mask[..., :3] * mask[..., 3:], 0, 1) return combined def _normalize_image(image, percentiles): """Normalize the image based on given percentiles.""" v_min, v_max = np.percentile(image, percentiles) image_normalized = np.clip((image - v_min) / (v_max - v_min + 1e-5), 0, 1) return image_normalized, (float(v_min), float(v_max)) def _generate_contours(mask): """Generate contours from the mask using OpenCV.""" contours, _ = cv2.findContours( mask.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE ) return contours def _apply_contours(image, mask, color, thickness): """Apply contours to the image.""" unique_labels = np.unique(mask) for label in unique_labels: if label == 0: continue label_mask = (mask == label).astype(np.uint8) contours = _generate_contours(label_mask) cv2.drawContours( image, contours, -1, mpl.colors.to_rgb(color), thickness ) return image with figure_style(theme_target()): num_channels = image.shape[-1] fig, ax = plt.subplots(1, num_channels + 1, figsize=(4 * figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, image, kind="overlay") ax = np.atleast_1d(ax).ravel() channels_with_outlines = set(channel_to_outline.keys()) if channel_to_outline is not None else set() for v in range(num_channels): channel_image = image[..., v] channel_image_normalized, display_range = _normalize_image( channel_image, percentiles) channel_image_rgb = np.dstack([channel_image_normalized] * 3) current_channel = channels[v] if all_on_all: for idx, (outline, color) in enumerate(zip(outlines, outline_colors)): if mode == 'outlines': channel_image_rgb = _apply_contours( channel_image_rgb, outline, color, thickness ) else: cmap = random_color_cmap( int(outline.max() + 1), seed=1000 + idx ) mask = _generate_colored_mask(outline, cmap) channel_image_rgb = _overlay_mask(channel_image_rgb, mask) elif current_channel in channels_with_outlines: outline_info = channel_to_outline[current_channel] outline = outline_info['mask'] color = outline_info['color'] cmap_seed = outline_info.get('cmap_seed', current_channel + 1000) if mode == 'outlines': channel_image_rgb = _apply_contours( channel_image_rgb, outline, color, thickness ) else: cmap = random_color_cmap( int(outline.max() + 1), seed=cmap_seed ) mask = _generate_colored_mask(outline, cmap) channel_image_rgb = _overlay_mask(channel_image_rgb, mask) else: if all_outlines: for idx, (outline, color) in enumerate(zip(outlines, outline_colors)): if mode == 'outlines': channel_image_rgb = _apply_contours( channel_image_rgb, outline, color, thickness ) else: cmap = random_color_cmap( int(outline.max() + 1), seed=1000 + idx ) mask = _generate_colored_mask(outline, cmap) channel_image_rgb = _overlay_mask(channel_image_rgb, mask) title = channel_to_label.get(current_channel, f'channel {current_channel}') artist = ax[v].imshow(channel_image_rgb) _tag_panel( artist, source=file, steps=({"op": "overlay_composite", "channel": int(current_channel), "percentiles": [float(value) for value in percentiles], "mode": str(mode), "thickness": int(thickness), "all_on_all": bool(all_on_all), "all_outlines": bool(all_outlines), "outline_palette": str(outline_palette), "mask_order": [name for name, _channel, _color in present_objects], "mask_channels": {name: int(channel) for name, channel, _color in present_objects}, "mask_planes": {name: int(index) for name, index in mask_dims.items()}, "filters": filter_recipe},), display_range=display_range, channel=current_channel, raw=channel_image) ax[v].set_title(title) ax[v].axis('off') if len(outlines) > 0: combined_mask = np.zeros_like(outlines[0], dtype=np.int64) label_offset = 0 for outline in outlines: outline_int = outline.astype(np.int64) object_pixels = outline_int > 0 if np.any(object_pixels): combined_mask[object_pixels] = outline_int[object_pixels] + label_offset label_offset += int(outline_int.max()) cmap = random_color_cmap(int(combined_mask.max() + 1), seed=9999) mask = _generate_colored_mask(combined_mask, cmap) blank_image = np.zeros((*combined_mask.shape, 3)) filled_image = _overlay_mask(blank_image, mask) artist = ax[-1].imshow(filled_image) _tag_panel(artist, source=file, steps=({"op": "combined_masks", "mask_planes": {name: int(index) for name, index in mask_dims.items()}, "mask_order": [name for name, _channel, _color in present_objects], "mask_channels": {name: int(channel) for name, channel, _color in present_objects}, "filters": filter_recipe, "outline_palette": str(outline_palette)},)) ax[-1].set_title('combined objects') ax[-1].axis('off') else: ax[-1].imshow(np.zeros((*image.shape[:2], 3))) ax[-1].set_title('no objects') ax[-1].axis('off') plt.tight_layout() if save_pdf: pdf_dir = os.path.join( os.path.dirname(os.path.dirname(file)), 'results', 'overlay' ) os.makedirs(pdf_dir, exist_ok=True) pdf_path = os.path.join( pdf_dir, os.path.basename(file).replace('.npy', '.pdf') ) pdf_path = save_figure(fig, pdf_path) plt.show() return fig def _save_channels_as_tiff(stack, save_dir, filename): """Save each channel in the stack as a grayscale TIFF.""" os.makedirs(save_dir, exist_ok=True) for i in range(stack.shape[-1]): channel = stack[..., i] tiff_path = os.path.join(save_dir, f"{filename}_channel_{i}.tiff") write_tiff(tiff_path, channel.astype(np.uint16)) print(f"Saved {tiff_path}") def _filter_object(mask, intensity_image, min_max_area=(0, 10000000), min_max_intensity=(0, 65000), type_='object'): """ Filter objects in a mask based on their area (size) and mean intensity. Args: mask (ndarray): The input mask. intensity_image (ndarray): The corresponding intensity image. min_max_area (tuple): A tuple (min_area, max_area) specifying the minimum and maximum area thresholds. min_max_intensity (tuple): A tuple (min_intensity, max_intensity) specifying the minimum and maximum intensity thresholds. Returns: ndarray: The filtered mask. """ original_dtype = mask.dtype mask_int = mask.astype(np.int64) intensity_image = intensity_image.astype(np.float64) unique_labels = np.unique(mask_int) unique_labels = unique_labels[unique_labels != 0] num_objects_before = len(unique_labels) areas = [] mean_intensities = [] labels_to_keep = [] for label in unique_labels: label_mask = (mask_int == label) area = np.sum(label_mask) mean_intensity = np.mean(intensity_image[label_mask]) areas.append(area) mean_intensities.append(mean_intensity) if (min_max_area[0] <= area <= min_max_area[1]) and (min_max_intensity[0] <= mean_intensity <= min_max_intensity[1]): labels_to_keep.append(label) areas = np.array(areas) mean_intensities = np.array(mean_intensities) num_objects_after = len(labels_to_keep) avg_area_before = areas.mean() if num_objects_before > 0 else 0 avg_intensity_before = mean_intensities.mean() if num_objects_before > 0 else 0 areas_after = areas[np.isin(unique_labels, labels_to_keep)] mean_intensities_after = mean_intensities[np.isin(unique_labels, labels_to_keep)] avg_area_after = areas_after.mean() if num_objects_after > 0 else 0 avg_intensity_after = mean_intensities_after.mean() if num_objects_after > 0 else 0 print(f"Before filtering {type_}: {num_objects_before} objects") print(f"Average area {type_}: {avg_area_before:.2f} pixels, Average intensity: {avg_intensity_before:.2f}") print(f"After filtering {type_}: {num_objects_after} objects") print(f"Average area {type_}: {avg_area_after:.2f} pixels, Average intensity: {avg_intensity_after:.2f}") mask_filtered = np.zeros_like(mask_int) for label in labels_to_keep: mask_filtered[mask_int == label] = label mask_filtered = mask_filtered.astype(original_dtype) return mask_filtered stack = np.load(file) if export_tiffs: save_dir = os.path.join( os.path.dirname(os.path.dirname(file)), 'results', os.path.splitext(os.path.basename(file))[0], 'tiff' ) filename = os.path.splitext(os.path.basename(file))[0] _save_channels_as_tiff(stack, save_dir, filename) if stack.dtype in (np.uint16, np.uint8): stack = stack.astype(np.float32) image = stack[..., channels] outlines = [] outline_colors = [] colours = outline_palette_colours(outline_palette) object_specs = [ ('cell', cell_channel, colours['cell']), ('nucleus', nucleus_channel, colours['nucleus']), ('pathogen', pathogen_channel, colours['pathogen']), ('organelle', organelle_channel, colours['organelle']), ] for index, (role, channel) in enumerate( _extra_organelle_slots(organelle_channels)): object_specs.append( (role, channel, _organelle_slot_colour(outline_palette, index))) present_objects = [(name, channel, color) for name, channel, color in object_specs if channel is not None] mask_dims = _overlay_mask_dims( file, [name for name, _channel, _color in present_objects], stack.shape[2]) channel_to_outline = {} channel_to_label = {} for mask_offset, (name, channel, color) in enumerate(present_objects): mask_dim = mask_dims[name] outline = np.take(stack, mask_dim, axis=2) if filter_dict is not None and name in filter_dict: intensity = np.take(stack, channel, axis=2) outline = _filter_object( outline, intensity, filter_dict[name][0], filter_dict[name][1], type_=name ) outlines.append(outline) outline_colors.append(color) channel_to_outline[channel] = { 'mask': outline, 'color': color, 'cmap_seed': 1000 + mask_offset } channel_to_label[channel] = f'{name} (channel {channel})' for ch in channels: if ch not in channel_to_label: channel_to_label[ch] = f'channel {ch}' filter_recipe = (None if filter_dict is None else { name: [[float(value) for value in bounds] for bounds in filter_dict[name]] for name in mask_dims if name in filter_dict}) fig = _plot_merged_plot( image=image, outlines=outlines, outline_colors=outline_colors, figuresize=figuresize, thickness=thickness, percentiles=percentiles, mode=mode, all_on_all=all_on_all, all_outlines=all_outlines, channels=channels, channel_to_outline=channel_to_outline, channel_to_label=channel_to_label, save_pdf=save_pdf ) return fig
[docs] def plot_image_mask_overlay_magenta_outlines( file, channels, cell_channel, nucleus_channel, pathogen_channel, figuresize=10, percentiles=(2, 98), thickness=3, save_pdf=True, mode='outlines', export_tiffs=False, all_on_all=False, all_outlines=False, filter_dict=None ): """Plot image and mask overlays, outlining each channel's own mask in magenta. Variant of :func:`plot_image_mask_overlay` with no ``organelle_channel``: when ``mode`` is ``'outlines'`` and ``all_on_all`` is False, the mask belonging to a channel is outlined in magenta rather than in that object's colour. In every other mode it falls back to filled, randomly coloured labels as that function does, but seeded per call rather than per object, so the colours differ between runs and between panels. :param file: Path to the merged ``.npy`` stack for one field of view. :param channels: Indices of the image channels to draw, one panel each. :param cell_channel: Intensity channel the cell mask belongs to, or ``None`` when there is no cell mask. :param nucleus_channel: Intensity channel the nucleus mask belongs to, or ``None``. :param pathogen_channel: Intensity channel the pathogen mask belongs to, or ``None``. :param figuresize: Figure height in inches; the figure is drawn four times as wide. Default ``10``. :param percentiles: Two-element percentile pair used to normalise each channel. Default ``(2, 98)``. :param thickness: Contour line width in pixels. Default ``3``. :param save_pdf: If True, save the figure into ``results/overlay/`` two directories above ``file``, in the configured figure format rather than always as PDF. Default ``True``. :param mode: ``'outlines'`` draws mask contours; any other value overlays filled, randomly coloured labels. Default ``'outlines'``. :param export_tiffs: If True, also write every stack plane as a grayscale TIFF into ``results/<stem>/tiff/`` alongside it. Default ``False``. :param all_on_all: If True, draw every mask on every channel in its own colour. Default ``False``. :param all_outlines: If True, draw every mask on the channels that own no mask themselves. Default ``False``. :param filter_dict: Optional per-object limits with a ``'cell'``, ``'nucleus'`` and ``'pathogen'`` entry, each holding ``((min_area, max_area), (min_intensity, max_intensity))``; objects outside the limits are dropped before plotting. :returns: The generated matplotlib ``Figure``. """ def random_color_cmap(n_labels, seed): """Generates a random color map for a given number of labels. :param n_labels: How many object colours to draw. Index 0 is prepended as black for background, so the map holds ``n_labels + 1`` entries; callers here pass ``int(outline.max() + 1)`` per object type, or ``int(combined_mask.max() + 1)`` for the merged panel. Colours are drawn as uniform RGB, so unlike the HSV variant in :func:`plot_image_mask_overlay` some come out dark and low-contrast against the image. :param seed: Seeds the *global* ``numpy.random`` state, not a local generator, so passing it also shifts every later ``np.random`` draw in the process. Callers here pass a fresh ``random.randint(0, 100)`` per panel, which is why the same object gets a different colour in each panel and each run. ``None`` leaves the global state alone. :returns: A ``ListedColormap``. """ np.random.seed(seed) rand_colors = np.random.rand(n_labels, 3) rand_colors = np.vstack([[0, 0, 0], rand_colors]) cmap = ListedColormap(rand_colors) return cmap def _plot_merged_plot( image, outlines, outline_colors, figuresize, thickness, percentiles, mode='outlines', all_on_all=False, all_outlines=False, channels=None, cell_channel=None, nucleus_channel=None, pathogen_channel=None, cell_outlines=None, nucleus_outlines=None, pathogen_outlines=None, save_pdf=True ): """Plot the merged plot with overlay, image channels, and masks.""" def _generate_colored_mask(mask, cmap): """Generate a colored mask using the given colormap.""" mask_norm = mask / (mask.max() + 1e-5) colored_mask = cmap(mask_norm) colored_mask[..., 3] = np.where(mask > 0, 1, 0) return colored_mask def _overlay_mask(image, mask): """Overlay the colored mask onto the original image.""" combined = np.clip(image * (1 - mask[..., 3:]) + mask[..., :3] * mask[..., 3:], 0, 1) return combined def _normalize_image(image, percentiles): """Normalize the image based on given percentiles.""" v_min, v_max = np.percentile(image, percentiles) image_normalized = np.clip((image - v_min) / (v_max - v_min + 1e-5), 0, 1) return image_normalized def _generate_contours(mask): """Generate contours from the mask using OpenCV.""" contours, _ = cv2.findContours( mask.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE ) return contours def _apply_contours(image, mask, color, thickness): """Apply contours to the image.""" unique_labels = np.unique(mask) for label in unique_labels: if label == 0: continue label_mask = (mask == label).astype(np.uint8) contours = _generate_contours(label_mask) cv2.drawContours( image, contours, -1, mpl.colors.to_rgb(color), thickness ) return image with figure_style(theme_target()): num_channels = image.shape[-1] fig, ax = plt.subplots(1, num_channels + 1, figsize=(4 * figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, image, kind="overlay") outlines_by_channel = {} if pathogen_channel is not None: outlines_by_channel[pathogen_channel] = pathogen_outlines if nucleus_channel is not None: outlines_by_channel[nucleus_channel] = nucleus_outlines if cell_channel is not None: outlines_by_channel[cell_channel] = cell_outlines for v in range(num_channels): channel_image = image[..., v] channel_image_normalized = _normalize_image(channel_image, percentiles) channel_image_rgb = np.dstack([channel_image_normalized] * 3) current_channel = channels[v] if all_on_all: for outline, color in zip(outlines, outline_colors): if mode == 'outlines': channel_image_rgb = _apply_contours( channel_image_rgb, outline, color, thickness ) else: cmap = random_color_cmap(int(outline.max() + 1), random.randint(0, 100)) mask = _generate_colored_mask(outline, cmap) channel_image_rgb = _overlay_mask(channel_image_rgb, mask) elif current_channel in outlines_by_channel: outline = outlines_by_channel[current_channel] if mode == 'outlines': channel_image_rgb = _apply_contours( channel_image_rgb, outline, '#FF00FF', thickness ) else: cmap = random_color_cmap(int(outline.max() + 1), random.randint(0, 100)) mask = _generate_colored_mask(outline, cmap) channel_image_rgb = _overlay_mask(channel_image_rgb, mask) else: if all_outlines: for outline, color in zip(outlines, outline_colors): if mode == 'outlines': channel_image_rgb = _apply_contours( channel_image_rgb, outline, color, thickness ) else: cmap = random_color_cmap(int(outline.max() + 1), random.randint(0, 100)) mask = _generate_colored_mask(outline, cmap) channel_image_rgb = _overlay_mask(channel_image_rgb, mask) ax[v].imshow(channel_image_rgb) ax[v].set_title(f'Image - Channel {current_channel}') if len(outlines) > 0: combined_mask = np.zeros_like(outlines[0]) for outline in outlines: combined_mask = np.maximum(combined_mask, outline) cmap = random_color_cmap(int(combined_mask.max() + 1), random.randint(0, 100)) mask = _generate_colored_mask(combined_mask, cmap) blank_image = np.zeros((*combined_mask.shape, 3)) filled_image = _overlay_mask(blank_image, mask) ax[-1].imshow(filled_image) ax[-1].set_title('Combined Objects Image') else: ax[-1].imshow(np.zeros((*image.shape[:2], 3))) ax[-1].set_title('no objects') plt.tight_layout() if save_pdf: pdf_dir = os.path.join( os.path.dirname(os.path.dirname(file)), 'results', 'overlay' ) os.makedirs(pdf_dir, exist_ok=True) pdf_path = os.path.join( pdf_dir, os.path.basename(file).replace('.npy', '.pdf') ) pdf_path = save_figure(fig, pdf_path) plt.show() return fig def _save_channels_as_tiff(stack, save_dir, filename): """Save each channel in the stack as a grayscale TIFF.""" os.makedirs(save_dir, exist_ok=True) for i in range(stack.shape[-1]): channel = stack[..., i] tiff_path = os.path.join(save_dir, f"{filename}_channel_{i}.tiff") write_tiff(tiff_path, channel.astype(np.uint16)) print(f"Saved {tiff_path}") def _filter_object(mask, intensity_image, min_max_area=(0, 10000000), min_max_intensity=(0, 65000), type_='object'): """ Filter objects in a mask based on their area (size) and mean intensity. Args: mask (ndarray): The input mask. intensity_image (ndarray): The corresponding intensity image. min_max_area (tuple): A tuple (min_area, max_area) specifying the minimum and maximum area thresholds. min_max_intensity (tuple): A tuple (min_intensity, max_intensity) specifying the minimum and maximum intensity thresholds. Returns: ndarray: The filtered mask. """ original_dtype = mask.dtype mask_int = mask.astype(np.int64) intensity_image = intensity_image.astype(np.float64) unique_labels = np.unique(mask_int) unique_labels = unique_labels[unique_labels != 0] num_objects_before = len(unique_labels) areas = [] mean_intensities = [] labels_to_keep = [] for label in unique_labels: label_mask = (mask_int == label) area = np.sum(label_mask) mean_intensity = np.mean(intensity_image[label_mask]) areas.append(area) mean_intensities.append(mean_intensity) if (min_max_area[0] <= area <= min_max_area[1]) and (min_max_intensity[0] <= mean_intensity <= min_max_intensity[1]): labels_to_keep.append(label) areas = np.array(areas) mean_intensities = np.array(mean_intensities) num_objects_after = len(labels_to_keep) avg_area_before = areas.mean() if num_objects_before > 0 else 0 avg_intensity_before = mean_intensities.mean() if num_objects_before > 0 else 0 areas_after = areas[np.isin(unique_labels, labels_to_keep)] mean_intensities_after = mean_intensities[np.isin(unique_labels, labels_to_keep)] avg_area_after = areas_after.mean() if num_objects_after > 0 else 0 avg_intensity_after = mean_intensities_after.mean() if num_objects_after > 0 else 0 print(f"Before filtering {type_}: {num_objects_before} objects") print(f"Average area {type_}: {avg_area_before:.2f} pixels, Average intensity: {avg_intensity_before:.2f}") print(f"After filtering {type_}: {num_objects_after} objects") print(f"Average area {type_}: {avg_area_after:.2f} pixels, Average intensity: {avg_intensity_after:.2f}") mask_filtered = np.zeros_like(mask_int) for label in labels_to_keep: mask_filtered[mask_int == label] = label mask_filtered = mask_filtered.astype(original_dtype) return mask_filtered stack = np.load(file) if export_tiffs: save_dir = os.path.join( os.path.dirname(os.path.dirname(file)), 'results', os.path.splitext(os.path.basename(file))[0], 'tiff' ) filename = os.path.splitext(os.path.basename(file))[0] _save_channels_as_tiff(stack, save_dir, filename) if stack.dtype in (np.uint16, np.uint8): stack = stack.astype(np.float32) image = stack[..., channels] outlines = [] outline_colors = [] cell_outlines = None nucleus_outlines = None pathogen_outlines = None if pathogen_channel is not None: pathogen_mask_dim = -1 pathogen_outlines = np.take(stack, pathogen_mask_dim, axis=2) if not filter_dict is None: pathogen_intensity = np.take(stack, pathogen_channel, axis=2) pathogen_outlines = _filter_object(pathogen_outlines, pathogen_intensity, filter_dict['pathogen'][0], filter_dict['pathogen'][1], type_='pathogen') outlines.append(pathogen_outlines) outline_colors.append('green') if nucleus_channel is not None: nucleus_mask_dim = -2 if pathogen_channel is not None else -1 nucleus_outlines = np.take(stack, nucleus_mask_dim, axis=2) if not filter_dict is None: nucleus_intensity = np.take(stack, nucleus_channel, axis=2) nucleus_outlines = _filter_object(nucleus_outlines, nucleus_intensity, filter_dict['nucleus'][0], filter_dict['nucleus'][1], type_='nucleus') outlines.append(nucleus_outlines) outline_colors.append('blue') if cell_channel is not None: if nucleus_channel is not None and pathogen_channel is not None: cell_mask_dim = -3 elif nucleus_channel is not None or pathogen_channel is not None: cell_mask_dim = -2 else: cell_mask_dim = -1 cell_outlines = np.take(stack, cell_mask_dim, axis=2) if not filter_dict is None: cell_intensity = np.take(stack, cell_channel, axis=2) cell_outlines = _filter_object(cell_outlines, cell_intensity, filter_dict['cell'][0], filter_dict['cell'][1], type_='cell') outlines.append(cell_outlines) outline_colors.append('red') fig = _plot_merged_plot( image=image, outlines=outlines, outline_colors=outline_colors, figuresize=figuresize, thickness=thickness, percentiles=percentiles, mode=mode, all_on_all=all_on_all, all_outlines=all_outlines, channels=channels, cell_channel=cell_channel, nucleus_channel=nucleus_channel, pathogen_channel=pathogen_channel, cell_outlines=cell_outlines, nucleus_outlines=nucleus_outlines, pathogen_outlines=pathogen_outlines, save_pdf=save_pdf ) return fig
[docs] def plot_cellpose4_output(batch, masks, flows, cmap='inferno', figuresize=10, nr=1, print_object_number=True): """Display per-channel images, label mask and flow field for Cellpose v4 outputs. :param batch: Image batch of shape ``(N, H, W, C)``. :param masks: Label masks, one per image. :param flows: Flow arrays, one per image. :param cmap: Colormap for image channels. Default ``'inferno'``. :param figuresize: Base figure size. Default ``10``. :param nr: Maximum number of images to plot. Default ``1``. :param print_object_number: If True, annotate each object with its label ID. Default ``True``. :returns: None """ from .utils import _generate_mask_random_cmap font = figuresize/2 index = 0 for image, mask, flow in zip(batch, masks, flows): random_cmap = _generate_mask_random_cmap(mask) if index < nr: index += 1 chans = image.shape[-1] with figure_style(theme_target()): fig, ax = plt.subplots(1, image.shape[-1] + 2, figsize=(4 * figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, [image, mask], kind="mask") for v in range(0, image.shape[-1]): ax[v].imshow(image[..., v], cmap=cmap, interpolation='nearest') ax[v].set_title('Image - Channel'+str(v)) ax[chans].imshow(mask, cmap=random_cmap, interpolation='nearest') ax[chans].set_title('Mask') if print_object_number: unique_objects = np.unique(mask) unique_objects = unique_objects[unique_objects != 0] for obj in unique_objects: cy, cx = ndi.center_of_mass(mask == obj) ax[chans].text(cx, cy, str(obj), color='white', fontsize=font, ha='center', va='center') ax[chans+1].imshow(flow, cmap='viridis', interpolation='nearest') ax[chans+1].set_title('Flow') plt.show() return
[docs] def plot_organelle_output(img_batch, masks, settings, cmap='inferno', figuresize=10, nr=1, print_object_number=True): """Plot organelle segmentation results: raw channel, label mask, morphology-specific diagnostic. :param img_batch: Single-channel image batch of shape ``(N, H, W)``. :param masks: Label masks, one per image. :param settings: Organelle settings dict; ``organelle_morphology`` and ``organelle_method`` drive the diagnostic panel. :param cmap: Colormap for the raw channel. Default ``'inferno'``. :param figuresize: Base figure size. Default ``10``. :param nr: Maximum number of images to plot. Default ``1``. :param print_object_number: If True, annotate each object with its label ID. Default ``True``. :returns: None """ from .utils import _generate_mask_random_cmap, _organelle_diagnostic morphology = settings.get('organelle_morphology', 'spots') method = settings.get('organelle_method', 'otsu') font = figuresize / 2 for idx in range(min(len(masks), nr, img_batch.shape[0])): img = img_batch[idx] mask = masks[idx] random_cmap = _generate_mask_random_cmap(mask) num_objects = len(np.unique(mask)) - (1 if 0 in mask else 0) diag_img, diag_title = _organelle_diagnostic(img, morphology, method, settings) with figure_style(theme_target()): n_panels = 3 fig, ax = plt.subplots(1, n_panels, figsize=(n_panels * figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, [img, mask], kind="mask") ax[0].imshow(img, cmap=cmap, interpolation='nearest') ax[0].set_title(f'Organelle channel ({morphology}/{method})') ax[1].imshow(mask, cmap=random_cmap, interpolation='nearest') ax[1].set_title(f'Mask ({num_objects} objects)') if print_object_number: unique_objects = np.unique(mask) unique_objects = unique_objects[unique_objects != 0] for obj in unique_objects: cy, cx = ndi.center_of_mass(mask == obj) ax[1].text(cx, cy, str(obj), color='white', fontsize=font, ha='center', va='center') ax[2].imshow(diag_img, cmap='viridis', interpolation='nearest') ax[2].set_title(diag_title) for a in ax: a.axis('off') plt.tight_layout() plt.show() return
[docs] def plot_masks(batch, masks, flows, cmap='inferno', figuresize=10, nr=1, file_type='.npz', print_object_number=True): """Display per-channel images, label masks and flow fields for a batch. :param batch: Image batch — shape ``(N, H, W, C)`` or a single image of shape ``(H, W, C)``. :param masks: Label masks, one per image (list or ndarray). :param flows: Flow arrays, one per image. :param cmap: Colormap for image channels. Default ``'inferno'``. :param figuresize: Base figure size. Default ``10``. :param nr: Maximum number of images to plot. Default ``1``. :param file_type: Source file type of ``flows`` — ``'png'`` selects the first element of each flow entry. Default ``'.npz'``. :param print_object_number: If True, annotate each object with its label ID. Default ``True``. :returns: None """ if len(batch.shape) == 3: batch = np.expand_dims(batch, axis=0) if not isinstance(masks, list): masks = np.asarray(masks) masks = [masks] if masks.ndim == 2 else list(masks) if not isinstance(flows, list): flows = [flows] else: flows = flows[0] if file_type == 'png': flows = [f[0] for f in flows] font = figuresize/2 index = 0 for image, mask, flow in zip(batch, masks, flows): unique_labels = np.unique(mask) num_objects = len(unique_labels[unique_labels != 0]) random_colors = np.random.rand(num_objects+1, 4) random_colors[:, 3] = 1 random_colors[0, :] = [0, 0, 0, 1] random_cmap = mpl.colors.ListedColormap(random_colors) if index < nr: index += 1 chans = image.shape[-1] with figure_style(theme_target()): fig, ax = plt.subplots(1, image.shape[-1] + 2, figsize=(4 * figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, [image, mask], kind="mask") for v in range(0, image.shape[-1]): ax[v].imshow(image[..., v], cmap=cmap) ax[v].set_title('Image - Channel'+str(v)) ax[chans].imshow(mask, cmap=random_cmap) ax[chans].set_title('Mask') if print_object_number: unique_objects = np.unique(mask) unique_objects = unique_objects[unique_objects != 0] for obj in unique_objects: cy, cx = ndi.center_of_mass(mask == obj) ax[chans].text(cx, cy, str(obj), color='white', fontsize=font, ha='center', va='center') ax[chans+1].imshow(flow, cmap='viridis') ax[chans+1].set_title('Flow') plt.show() return
def _plot_4D_arrays(src, figuresize=10, cmap='inferno', nr_npz=1, nr=1): """ Plot 4D arrays from .npz files. Args: src (str): The directory path where the .npz files are located. figuresize (int, optional): The size of the figure. Defaults to 10. cmap (str, optional): The colormap to use for image visualization. Defaults to 'inferno'. nr_npz (int, optional): The number of .npz files to plot. Defaults to 1. nr (int, optional): The number of images to plot from each .npz file. Defaults to 1. """ paths = [os.path.join(src, file) for file in os.listdir(src) if file.endswith('.npz')] paths = random.sample(paths, min(nr_npz, len(paths))) for path in paths: with np.load(path) as data: stack = data['data'] num_images = stack.shape[0] num_channels = stack.shape[3] for i in range(min(nr, num_images)): img = stack[i] with figure_style(theme_target()): if num_channels == 1: fig, axs = plt.subplots(1, 1, figsize=(figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, img, kind="image") axs = [axs] else: fig, axs = plt.subplots(1, num_channels, figsize=(num_channels * figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, img, kind="image") for c in range(num_channels): axs[c].imshow(img[:, :, c], cmap=cmap) axs[c].set_title(f'Channel {c}', size=_montage_type_size(figuresize)) axs[c].axis('off') fig.tight_layout() plt.show() return
[docs] def generate_mask_random_cmap(mask): """Return a random ``ListedColormap`` sized to the labels in ``mask``. :param mask: Label mask array (0 = background). :returns: Random colormap where index 0 is black and remaining entries are random opaque RGBA colours. """ unique_labels = np.unique(mask) num_objects = len(unique_labels[unique_labels != 0]) random_colors = np.random.rand(num_objects+1, 4) random_colors[:, 3] = 1 random_colors[0, :] = [0, 0, 0, 1] random_cmap = mpl.colors.ListedColormap(random_colors) return random_cmap
[docs] def random_cmap(num_objects=100): """Return a random ``ListedColormap`` with ``num_objects + 1`` colours. :param num_objects: Number of foreground colours to generate. Default ``100``. :returns: Colormap with index 0 = black and remaining indices random opaque RGBA colours. """ random_colors = np.random.rand(num_objects+1, 4) random_colors[:, 3] = 1 random_colors[0, :] = [0, 0, 0, 1] random_cmap = mpl.colors.ListedColormap(random_colors) return random_cmap
def _generate_mask_random_cmap(mask): """ Generate a random colormap based on the unique labels in the given mask. Parameters: mask (ndarray): The mask array containing unique labels. Returns: ListedColormap: A random colormap generated based on the unique labels in the mask. """ unique_labels = np.unique(mask) num_objects = len(unique_labels[unique_labels != 0]) random_colors = np.random.rand(num_objects+1, 4) random_colors[:, 3] = 1 random_colors[0, :] = [0, 0, 0, 1] random_cmap = mpl.colors.ListedColormap(random_colors) return random_cmap def _get_colours_merged(outline_color): """ Get the merged outline colors based on the specified outline color format. Parameters: outline_color (str): The outline color format. Can be one of 'rgb', 'bgr', 'gbr', or 'rbg'. Returns: list: A list of merged outline colors based on the specified format. """ if outline_color == 'rgb': outline_colors = [[1, 0, 0], [0, 1, 0], [0, 0, 1]] elif outline_color == 'bgr': outline_colors = [[0, 0, 1], [0, 1, 0], [1, 0, 0]] elif outline_color == 'gbr': outline_colors = [[0, 1, 0], [0, 0, 1], [1, 0, 0]] elif outline_color == 'rbg': outline_colors = [[1, 0, 0], [0, 0, 1], [0, 1, 0]] else: outline_colors = [[1, 0, 0], [0, 0, 1], [0, 1, 0]] return outline_colors
[docs] def plot_images_and_arrays(folders, lower_percentile=1, upper_percentile=99, threshold=1000, extensions=None, overlay=False, max_nr=None, randomize=True): """Show side-by-side images and arrays found across multiple folders. Each image is either percentile-normalised (values below ``threshold``) or shown as a label mask. Optionally overlays object outlines from a matching mask file. :param folders: Folders to scan for image/array files. :param lower_percentile: Lower percentile clip. Default ``1``. :param upper_percentile: Upper percentile clip. Default ``99``. :param threshold: Values <= threshold are treated as label data instead of intensity. Default ``1000``. :param extensions: File extensions to include. Default ``['.npy', '.tif', '.tiff', '.png']``. :param overlay: If True, overlay object outlines. Default ``False``. :param max_nr: Maximum number of key groups to plot. :param randomize: If True, shuffle key order before plotting. Default ``True``. :returns: None """ if extensions is None: extensions = ['.npy', '.tif', '.tiff', '.png'] def normalize_image(image, lower=1, upper=99): """Percentile-clip and rescale ``image`` to ``[0, 1]``. :param image: Any numeric array; normalisation is over the whole array at once, so a multi-channel stack is scaled by a single pair of percentiles rather than per channel. :param lower: Lower percentile, in 0-100. Default ``1``. :param upper: Upper percentile, in 0-100, and must be strictly greater than ``lower``: when the two percentiles evaluate equal (a flat image) the rescale divides by zero and returns ``nan`` rather than a blank frame, and swapping the two inverts the image instead of raising. Default ``99``. :returns: A float array clipped to ``[0, 1]``. """ p2, p98 = np.percentile(image, (lower, upper)) return np.clip((image - p2) / (p98 - p2), 0, 1) def find_files(folders, extensions): """Return a dict keyed by base filename mapping to files with the requested extensions. :param folders: Folder paths, each walked recursively. Grouping is by basename without extension, and only names found in *every* folder survive the final filter — one missing file drops that name from the result entirely, and two files with the same basename under one folder keep only the last one walked. :param extensions: Extensions to accept, matched with ``str.endswith`` so they must include the dot and match case. :returns: ``{basename: {folder: path}}`` for complete groups only. """ file_dict = {} for folder in folders: for root, _, files in os.walk(folder): for file in files: if any(file.endswith(ext) for ext in extensions): file_name_wo_ext = os.path.splitext(file)[0] file_path = os.path.join(root, file) if file_name_wo_ext not in file_dict: file_dict[file_name_wo_ext] = {} file_dict[file_name_wo_ext][folder] = file_path filtered_dict = {k: v for k, v in file_dict.items() if len(v) == len(folders)} return filtered_dict def plot_from_file_dict(file_dict, threshold=1000, lower_percentile=1, upper_percentile=99, overlay=False): """Show image/mask pairs collected in ``file_dict`` side-by-side. :param file_dict: ``{filename: {folder: path}}`` produced by ``find_files``. :param threshold: Values above this unique-count are treated as intensity images; otherwise as label masks. Default ``1000``. :param lower_percentile: Lower percentile clip. Default ``1``. :param upper_percentile: Upper percentile clip. Default ``99``. :param overlay: If True, overlay mask outlines on the image. Default ``False``. :returns: None """ for filename, folder_paths in file_dict.items(): image_data = None mask_data = None for folder, path in folder_paths.items(): if path.endswith('.npy'): data = np.load(path) elif path.endswith('.tif') or path.endswith('.tiff'): data = imageio.imread(path) else: continue unique_values = np.unique(data) if len(unique_values) > threshold: image_data = normalize_image(data, lower_percentile, upper_percentile) else: mask_data = data if image_data is not None and mask_data is not None: with figure_style(theme_target()): fig, axes = plt.subplots(1, 2, figsize=(15, 7)) from .figures.bundle import _register_figure_data _register_figure_data(fig, [image_data, mask_data], kind="mask", title=str(filename)) cmap = random_cmap(num_objects=len(np.unique(mask_data))) axes[0].imshow(mask_data, cmap=cmap) axes[0].set_title(f"{filename} - Mask") axes[0].axis('off') axes[1].imshow(image_data, cmap='gray') if overlay: labeled_mask = label(mask_data) for region in regionprops(labeled_mask): if region.image.shape[0] >= 2 and region.image.shape[1] >= 2: contours = find_contours(region.image, 0.75) for contour in contours: contour[:, 0] += region.bbox[0] contour[:, 1] += region.bbox[1] axes[1].plot(contour[:, 1], contour[:, 0], linewidth=2, color='magenta') axes[1].set_title(f"{filename} - Normalized Image") axes[1].axis('off') plt.tight_layout() plt.show() if overlay: print(f'Overlay will only work on the first two folders in the list') file_dict = find_files(folders, extensions) items = list(file_dict.items()) if randomize: random.shuffle(items) if isinstance(max_nr, (int, float)): items = items[:int(max_nr)] file_dict = dict(items) plot_from_file_dict(file_dict, threshold, lower_percentile, upper_percentile, overlay) return
def _filter_objects_in_plot(stack, cell_mask_dim, nucleus_mask_dim, pathogen_mask_dim, mask_dims, filter_min_max, nuclei_limit, pathogen_limit): """ Filters objects in a plot based on various criteria. Args: stack (numpy.ndarray): The input stack of masks. cell_mask_dim (int): The dimension index of the cell mask. nucleus_mask_dim (int): The dimension index of the nucleus mask. pathogen_mask_dim (int): The dimension index of the pathogen mask. mask_dims (list): A list of dimension indices for additional masks. filter_min_max (list): A list of minimum and maximum area values for each mask. nuclei_limit (bool): Whether to include multinucleated cells. pathogen_limit (bool): Whether to include multiinfected cells. Returns: numpy.ndarray: The filtered stack of masks. """ from .utils import _remove_outside_objects, _remove_multiobject_cells stack = _remove_outside_objects(stack, cell_mask_dim, nucleus_mask_dim, pathogen_mask_dim) _role_index = {} for _position, _dim in enumerate((cell_mask_dim, nucleus_mask_dim, pathogen_mask_dim)): if _dim is not None and _dim not in _role_index: _role_index[_dim] = _position for mask_dim in mask_dims: if filter_min_max is None: min_max = [0, 100000000] else: _position = _role_index.get(mask_dim) if _position is None or _position >= len(filter_min_max): min_max = [0, 100000000] else: min_max = filter_min_max[_position] mask = np.take(stack, mask_dim, axis=2) props = measure.regionprops_table(mask, properties=['label', 'area']) avg_size_before = np.mean(props['area']) total_count_before = len(props['label']) if not filter_min_max is None: valid_labels = props['label'][np.logical_and(props['area'] > min_max[0], props['area'] < min_max[1])] stack[:, :, mask_dim] = np.isin(mask, valid_labels) * mask props_after = measure.regionprops_table(stack[:, :, mask_dim], properties=['label', 'area']) avg_size_after = np.mean(props_after['area']) total_count_after = len(props_after['label']) if mask_dim == cell_mask_dim: if nuclei_limit is False and nucleus_mask_dim is not None: stack = _remove_multiobject_cells(stack, mask_dim, cell_mask_dim, nucleus_mask_dim, pathogen_mask_dim, object_dim=nucleus_mask_dim) if pathogen_limit is False and cell_mask_dim is not None and pathogen_mask_dim is not None: stack = _remove_multiobject_cells(stack, mask_dim, cell_mask_dim, nucleus_mask_dim, pathogen_mask_dim, object_dim=pathogen_mask_dim) cell_area_before = avg_size_before cell_count_before = total_count_before cell_area_after = avg_size_after cell_count_after = total_count_after if mask_dim == nucleus_mask_dim: nucleus_area_before = avg_size_before nucleus_count_before = total_count_before nucleus_area_after = avg_size_after nucleus_count_after = total_count_after if mask_dim == pathogen_mask_dim: pathogen_area_before = avg_size_before pathogen_count_before = total_count_before pathogen_area_after = avg_size_after pathogen_count_after = total_count_after if cell_mask_dim is not None: print(f'removed {cell_count_before-cell_count_after} cells, cell size from {cell_area_before} to {cell_area_after}') if nucleus_mask_dim is not None: print(f'removed {nucleus_count_before-nucleus_count_after} nucleus, nucleus size from {nucleus_area_before} to {nucleus_area_after}') if pathogen_mask_dim is not None: print(f'removed {pathogen_count_before-pathogen_count_after} pathogens, pathogen size from {pathogen_area_before} to {pathogen_area_after}') return stack
[docs] def plot_arrays(src, figuresize=10, cmap='inferno', nr=1, normalize=True, q1=1, q2=99): """Plot random ``.npy`` / ``.npz`` arrays from ``src``, one channel per subplot. :param src: Directory or single ``.npy``/``.npz`` path. :param figuresize: Base figure size. Default ``10``. :param cmap: Matplotlib colormap. Default ``'inferno'``. :param nr: Maximum number of arrays to plot. Default ``1``. :param normalize: If True, percentile-normalise before display. Default ``True``. :param q1: Lower percentile for normalisation. Default ``1``. :param q2: Upper percentile for normalisation. Default ``99``. :returns: None """ from .utils import normalize_to_dtype from .io import _listdir_visible paths = [] if src.endswith('.npz') or src.endswith('.npy'): paths = [src] else: paths = [os.path.join(src, f) for f in _listdir_visible(src) if f.endswith(('.npy', '.npz'))] paths = random.sample(paths, min(nr, len(paths))) for path in paths: print(f'Image path: {path}') if path.endswith('.npz'): with np.load(path) as data: key = list(data.keys())[0] img = data[key][0] else: img = np.load(path) if normalize: if img.ndim == 2: img = normalize_to_dtype(array=img[:, :, np.newaxis], p1=q1, p2=q2)[:, :, 0] else: img = normalize_to_dtype(array=img, p1=q1, p2=q2) with figure_style(theme_target()): if img.ndim == 3: array_nr = img.shape[2] fig, axs = plt.subplots(1, array_nr, figsize=(figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, img, kind="image") if array_nr == 1: axs = [axs] for channel in range(array_nr): i = img[:, :, channel] axs[channel].imshow(i, cmap=plt.get_cmap(cmap)) axs[channel].set_title(f'Channel {channel}', size=_montage_type_size(figuresize)) axs[channel].axis('off') else: fig, ax = plt.subplots(1, 1, figsize=(figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, img, kind="image") ax.imshow(img, cmap=plt.get_cmap(cmap)) ax.set_title('Channel 0', size=_montage_type_size(figuresize)) ax.axis('off') fig.tight_layout() plt.show()
def _normalize_and_outline(image, remove_background, normalize, normalization_percentiles, overlay, overlay_chans, mask_dims, outline_colors, outline_thickness): """ Normalize and outline an image. Args: image (ndarray): The input image. remove_background (bool): Flag indicating whether to remove the background. backgrounds (list): List of background values for each channel. normalize (bool): Flag indicating whether to normalize the image. normalization_percentiles (list): List of percentiles for normalization. overlay (bool): Flag indicating whether to overlay outlines onto the image. overlay_chans (list): List of channel indices to overlay. mask_dims (list): List of dimensions to use for masking. outline_colors (list): List of colors for the outlines. outline_thickness (int): Thickness of the outlines. Returns: tuple: A tuple containing the overlayed image, the original image, and a list of outlines. """ from .utils import normalize_to_dtype, _outline_and_overlay, _gen_rgb_image raw_masks = {d: image[:, :, d].copy() for d in mask_dims} if remove_background: backgrounds = np.percentile(image, 1, axis=(0, 1)) backgrounds = backgrounds[:, np.newaxis, np.newaxis] mask = np.zeros_like(image, dtype=bool) for chan_index in range(image.shape[-1]): if chan_index not in mask_dims: mask[:, :, chan_index] = image[:, :, chan_index] < backgrounds[chan_index] image[mask] = 0 if normalize: image = normalize_to_dtype(array=image, p1=normalization_percentiles[0], p2=normalization_percentiles[1]) else: image = normalize_to_dtype(array=image, p1=0, p2=100) rgb_image = _gen_rgb_image(image, channels=overlay_chans) for d, raw in raw_masks.items(): image[:, :, d] = raw if overlay: overlayed_image, outlines, image = _outline_and_overlay(image, rgb_image, mask_dims, outline_colors, outline_thickness) return overlayed_image, image, outlines else: channels_to_keep = [i for i in range(image.shape[-1]) if i not in mask_dims] image = np.take(image, channels_to_keep, axis=-1) return [], image, [] def _plot_merged_plot(overlay, image, stack, mask_dims, figuresize, overlayed_image, outlines, cmap, outline_colors, print_object_number, mask_names=None): """ Plot the merged plot with overlay, image channels, and masks. Args: overlay (bool): Flag indicating whether to overlay the image with outlines. image (ndarray): Input image array. stack (ndarray): Stack of masks. mask_dims (list): List of mask dimensions. figuresize (float): Size of the figure. overlayed_image (ndarray): Overlayed image array. outlines (list): List of outlines. cmap (str): Colormap for the masks. outline_colors (list): List of outline colors. print_object_number (bool): Flag indicating whether to print object numbers on the masks. Returns: fig (Figure): The generated matplotlib figure. """ with figure_style(theme_target()): if overlay: fig, ax = plt.subplots(1, image.shape[-1] + len(mask_dims) + 1, figsize=(4 * figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, image, kind="overlay") ax[0].imshow(overlayed_image) ax[0].set_title('Overlayed Image') ax_index = 1 else: fig, ax = plt.subplots(1, image.shape[-1] + len(mask_dims), figsize=(4 * figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, image, kind="overlay") ax_index = 0 for v in range(0, image.shape[-1]): channel_image = image[..., v] channel_image_normalized = channel_image.astype(float) channel_image_normalized -= channel_image_normalized.min() channel_image_normalized /= channel_image_normalized.max() channel_image_rgb = np.dstack((channel_image_normalized, channel_image_normalized, channel_image_normalized)) for outline, color in zip(outlines, outline_colors): for j in np.unique(outline)[1:]: channel_image_rgb[outline == j] = mpl.colors.to_rgb(color) ax[v + ax_index].imshow(channel_image_rgb) ax[v + ax_index].set_title(f'Channel {v + 1}') for i, mask_dim in enumerate(mask_dims): mask = np.take(stack, mask_dim, axis=2) random_cmap = _generate_mask_random_cmap(mask) ax[i + image.shape[-1] + ax_index].imshow(mask, cmap=random_cmap) n_obj = int(len(np.unique(mask)) - 1) cls = (mask_names[i] if mask_names and i < len(mask_names) else f'Mask {i + 1}') ax[i + image.shape[-1] + ax_index].set_title( f'{cls} - {n_obj} object' + ('' if n_obj == 1 else 's')) if print_object_number: unique_objects = np.unique(mask)[1:] for obj in unique_objects: cy, cx = ndi.center_of_mass(mask == obj) ax[i + image.shape[-1] + ax_index].text(cx, cy, str(obj), color='white', fontsize=8, ha='center', va='center') plt.tight_layout() plt.show() return fig
[docs] def plot_merged(src, settings): """Show multi-channel image stacks with per-object outlines overlaid. :param src: Folder containing ``.npy`` merged stacks. :param settings: Plot settings dict — includes channel/mask dims, overlay colours, normalisation, filter and object-count keys. :returns: The last generated ``Figure`` when ``settings['nr']`` is exceeded; otherwise ``None``. """ from .utils import _remove_noninfected outline_colors = _get_colours_merged(settings['outline_color']) index = 0 _mask_dim_pairs = [('Cell Mask', settings['cell_mask_dim']), ('Nucleus Mask', settings['nucleus_mask_dim']), ('Pathogen Mask', settings['pathogen_mask_dim'])] mask_dims = [dim for _name, dim in _mask_dim_pairs if dim is not None] mask_names = [name for name, dim in _mask_dim_pairs if dim is not None] if settings['verbose']: display(settings) if settings['pathogen_mask_dim'] is None: settings['pathogen_limit'] = True fig = None for file in os.listdir(src): path = os.path.join(src, file) stack = np.load(path) print(f'Loaded: {path}') if settings['pathogen_limit'] > 0: if settings['pathogen_mask_dim'] is not None and settings['cell_mask_dim'] is not None: stack = _remove_noninfected(stack, settings['cell_mask_dim'], settings['nucleus_mask_dim'], settings['pathogen_mask_dim']) if settings['pathogen_limit'] is not True or settings['nuclei_limit'] is not True or settings['filter_min_max'] is not None: stack = _filter_objects_in_plot(stack, settings['cell_mask_dim'], settings['nucleus_mask_dim'], settings['pathogen_mask_dim'], mask_dims, settings['filter_min_max'], settings['nuclei_limit'], settings['pathogen_limit']) overlayed_image, image, outlines = _normalize_and_outline(image=stack, remove_background=settings['remove_background'], normalize=settings['normalize'], normalization_percentiles=settings['normalization_percentiles'], overlay=settings['overlay'], overlay_chans=settings['overlay_chans'], mask_dims=mask_dims, outline_colors=outline_colors, outline_thickness=settings['outline_thickness']) if index < settings['nr']: index += 1 fig = _plot_merged_plot(overlay=settings['overlay'], image=image, stack=stack, mask_dims=mask_dims, figuresize=settings['figuresize'], overlayed_image=overlayed_image, outlines=outlines, cmap=settings['cmap'], outline_colors=outline_colors, print_object_number=settings['print_object_number'], mask_names=mask_names) else: return fig
def _plot_images_on_grid(image_files, channel_indices, um_per_pixel, scale_bar_length_um=5, fontsize=8, show_filename=True, channel_names=None, plot=False): """ Plots a grid of images with optional scale bar and channel names. Args: image_files (list): List of image file paths. channel_indices (list): List of channel indices to select from the images. um_per_pixel (float): Micrometers per pixel. scale_bar_length_um (float, optional): Length of the scale bar in micrometers. Defaults to 5. fontsize (int, optional): Font size for the image titles. Defaults to 8. show_filename (bool, optional): Whether to show the image file names as titles. Defaults to True. channel_names (list, optional): Names for the legend, **in FILE order** — entry 0 is the file's red plane, 1 green, 2 blue. Defaults to None. plot (bool, optional): Whether to display the plot. Defaults to False. Returns: matplotlib.figure.Figure: The generated figure object. .. note:: ``channel_names`` is in FILE order, not source-channel order, and the distinction is the one that caused the crop-colour episode (INVARIANTS 13). The legend colours entry *i* with ``['red', 'green', 'blue'][i]``, which is right because the image is read as RGB — but a caller passing names in SOURCE order (``['DAPI', 'GFP', 'RFP']`` for channels 0, 1, 2) would get DAPI labelled red while it is rendered blue, and the figure would look entirely reasonable. To go from the source order a user thinks in to the file order this wants, ask ``spacr.crops.resolve_png_channel_mapping(settings)`` and read off ``r``, ``g``, ``b``. No live caller passes this argument today (the one call site passes ``channel_names=None``), so this is a contract being written down before it is relied on rather than a bug being fixed. """ print(f'scale bar represents {scale_bar_length_um} um') nr_of_images = len(image_files) cols = int(np.ceil(np.sqrt(nr_of_images))) rows = np.ceil(nr_of_images / cols) with figure_style(theme_target()): fig, axes = plt.subplots(int(rows), int(cols), figsize=(20, 20), squeeze=False) from .figures.bundle import _register_figure_data _register_figure_data(fig, None, kind="montage", files=[str(v) for v in image_files]) axes = axes.flatten() scale_bar_length_px = int(scale_bar_length_um / um_per_pixel) channel_colors = ['red','green','blue'] for i, image_file in enumerate(image_files): img_array = read_image_rgb(image_file, cv2.IMREAD_UNCHANGED) if channel_indices is not None: if len(channel_indices) == 1: img_array = img_array[:, :, channel_indices[0]] cmap = 'gray' elif len(channel_indices) == 2: img_array = np.mean(img_array[:, :, channel_indices], axis=2) cmap = 'gray' else: img_array = img_array[:, :, channel_indices] cmap = None else: cmap = None if img_array.ndim == 3 else 'gray' if img_array.dtype == np.uint16: img_array = img_array.astype(np.float32) / 65535.0 elif img_array.dtype == np.uint8: img_array = img_array.astype(np.float32) / 255.0 ax = axes[i] ax.imshow(img_array, cmap=cmap) ax.axis('off') if show_filename: ax.set_title(os.path.basename(image_file), color=resolve_ink(theme_target()), fontsize=fontsize, pad=20) ax.plot([10, 10 + scale_bar_length_px], [img_array.shape[0] - 10] * 2, lw=2, color='white') initial_offset = 0.02 increment = 0.05 if channel_names: current_offset = initial_offset for ci, channel_name in enumerate(channel_names): color = (channel_colors[ci] if ci < len(channel_colors) else resolve_ink(theme_target())) fig.text(current_offset, 0.99, channel_name, color=color, fontsize=fontsize, verticalalignment='top', horizontalalignment='left') current_offset += increment for j in range(nr_of_images, len(axes)): axes[j].axis('off') plt.tight_layout(pad=3) if plot: plt.show() return fig def _save_scimg_plot(src, nr_imgs=16, channel_indices=None, um_per_pixel=0.1, scale_bar_length_um=10, standardize=True, fontsize=8, show_filename=True, channel_names=None, dpi=300, plot=False, i=1, all_folders=1): """ Save and visualize single-cell images. Args: src (str): The source directory path. nr_imgs (int, optional): The number of images to visualize. Defaults to 16. channel_indices (list, optional): List of channel indices to visualize. Defaults to [0,1,2]. um_per_pixel (float, optional): Micrometers per pixel. Defaults to 0.1. scale_bar_length_um (float, optional): Length of the scale bar in micrometers. Defaults to 10. standardize (bool, optional): Whether to standardize the image sizes. Defaults to True. fontsize (int, optional): Font size for the filename. Defaults to 8. show_filename (bool, optional): Whether to show the filename on the image. Defaults to True. channel_names (list, optional): List of channel names. Defaults to None. dpi (int, optional): Dots per inch for the saved image. Defaults to 300. plot (bool, optional): Whether to plot the images. Defaults to False. Returns: None """ if channel_indices is None: channel_indices = [0,1,2] from .io import _save_figure def _visualize_scimgs(src, channel_indices=None, um_per_pixel=0.1, scale_bar_length_um=10, show_filename=True, standardize=True, nr_imgs=None, fontsize=8, channel_names=None, plot=False): """ Visualize single-cell images. Args: src (str): The source directory path. channel_indices (list, optional): List of channel indices to visualize. Defaults to None. um_per_pixel (float, optional): Micrometers per pixel. Defaults to 0.1. scale_bar_length_um (float, optional): Length of the scale bar in micrometers. Defaults to 10. show_filename (bool, optional): Whether to show the filename on the image. Defaults to True. standardize (bool, optional): Whether to standardize the image sizes. Defaults to True. nr_imgs (int, optional): The number of images to visualize. Defaults to None. fontsize (int, optional): Font size for the filename. Defaults to 8. channel_names (list, optional): List of channel names. Defaults to None. plot (bool, optional): Whether to plot the images. Defaults to False. Returns: matplotlib.figure.Figure: The figure object containing the plotted images. """ from .utils import _find_similar_sized_images def _generate_filelist(src): """ Generate a list of image files in the specified directory. Args: src (str): The source directory path. Returns: list: A list of image file paths. """ files = glob.glob(os.path.join(src, '*')) image_files = [f for f in files if f.lower().endswith(('.png', '.jpg', '.jpeg', '.bmp', '.tif', '.tiff', '.gif'))] return image_files def _random_sample(file_list, nr_imgs=None): """ Randomly selects a subset of files from the given file list. Args: file_list (list): A list of file names. nr_imgs (int, optional): The number of files to select. If None, all files are selected. Defaults to None. Returns: list: A list of randomly selected file names. """ if nr_imgs is not None and nr_imgs < len(file_list): random.seed(42) file_list = random.sample(file_list, nr_imgs) return file_list image_files = _generate_filelist(src) if standardize: image_files = _find_similar_sized_images(image_files) if nr_imgs is not None: image_files = _random_sample(image_files, nr_imgs) fig = _plot_images_on_grid(image_files, channel_indices, um_per_pixel, scale_bar_length_um, fontsize, show_filename, channel_names, plot) return fig fig = _visualize_scimgs(src, channel_indices, um_per_pixel, scale_bar_length_um, show_filename, standardize, nr_imgs, fontsize, channel_names, plot) _save_figure(fig, src, text='all_channels') for channel in channel_indices: channel_indices=[channel] fig = _visualize_scimgs(src, channel_indices, um_per_pixel, scale_bar_length_um, show_filename, standardize, nr_imgs, fontsize, channel_names=None, plot=plot) _save_figure(fig, src, text=f'channel_{channel}') return def _plot_cropped_arrays(stack, filename, figuresize=10, cmap='inferno', threshold=500): """ Plot cropped arrays. Args: stack (ndarray): The array to be plotted, 2D (one panel) or 3D with the channels last (one panel per ``stack.shape[2]``). A 1D array matches neither branch and raises ``UnboundLocalError`` on the return rather than being rejected up front. filename (str): Accepted and ignored -- the only reference to it is a commented-out print, so it never reaches the figure or a file. figuresize (int, optional): Both width and height of the figure, in inches; the multi-channel case does not widen it per channel, so panels get thinner as channels are added. Defaults to 10. cmap (str, optional): Name resolved with ``plt.get_cmap`` and used only for the planes treated as intensity images. Defaults to 'inferno'. threshold (int, optional): A plane with this many distinct values or fewer is drawn as a label mask with a random colormap and an object count in its title. Defaults to 500. Returns: Figure: The figure that was drawn. The 2D case also calls ``plt.show()`` before returning; the multi-channel case does not. """ dim = stack.shape def plot_single_array(array, ax, title, chosen_cmap): """Render one channel from ``stack`` onto ``ax``. No colorbar is drawn. Args: array (ndarray): One 2D plane of ``stack``. Its count of distinct values, not its dtype, is what decides whether it is treated as an intensity image or as a label mask -- so a uint8 plane, which can hold at most 256 distinct values, is always taken for a mask under the default ``threshold`` of 500. ax (matplotlib.axes.Axes): Axes drawn on in place; its frame and ticks are switched off, the title is fixed at size 18, and nothing is returned. title (str): Panel title. When the plane is treated as a mask, the object count is appended as ``", N (obj.)"``. That count is the number of distinct non-zero values, so the background value 0 is never counted, but a negative value is counted as an object. chosen_cmap (Colormap): Colormap for the intensity case only. It is discarded when the plane has no more than ``threshold`` unique values (the enclosing function's argument, default ``500``), because a random colormap -- black at index 0, one random opaque colour per non-zero label -- is generated instead so neighbouring objects stay distinguishable. """ unique_values = np.unique(array) num_unique_values = len(unique_values) if num_unique_values <= threshold: num_objects = int(np.count_nonzero(unique_values)) chosen_cmap = _generate_mask_random_cmap(array) title = f'{title}, {num_objects} (obj.)' ax.imshow(array, cmap=chosen_cmap) ax.set_title(title, size=_montage_type_size(figuresize)) ax.axis('off') with figure_style(theme_target()): if len(dim) == 2: fig, ax = plt.subplots(1, 1, figsize=(figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, stack, kind="image") plot_single_array(stack, ax, 'Channel one', plt.get_cmap(cmap)) fig.tight_layout() plt.show() elif len(dim) > 2: num_channels = dim[2] fig, axs = plt.subplots(1, num_channels, figsize=(figuresize, figuresize)) from .figures.bundle import _register_figure_data _register_figure_data(fig, stack, kind="image") axs = np.atleast_1d(axs) for channel in range(num_channels): plot_single_array(stack[:, :, channel], axs[channel], f'C. {channel}', plt.get_cmap(cmap)) fig.tight_layout() return fig def _visualize_and_save_timelapse_stack_with_tracks(masks, tracks_df, save, src, name, plot, filenames, object_type, mode='btrack', interactive=False): """ Visualizes and saves a timelapse stack with tracks. Args: masks (list): List of binary masks representing each frame of the timelapse stack. tracks_df (pandas.DataFrame): DataFrame containing track information. save (bool): Flag indicating whether to save the timelapse stack. src (str): Source file path. name (str): Name of the timelapse stack. plot (bool): Flag indicating whether to plot the timelapse stack. filenames (list): List of filenames corresponding to each frame of the timelapse stack. object_type (str): Type of object being tracked. mode (str, optional): Tracking mode. Defaults to 'btrack'. interactive (bool, optional): Flag indicating whether to display the timelapse stack interactively. Defaults to False. """ from .io import _save_mask_timelapse_as_gif, _mask_movie_frame_geometry highest_label = max(np.max(mask) for mask in masks) random_colors = np.random.rand(highest_label + 1, 4) random_colors[:, 3] = 1 random_colors[0] = [0, 0, 0, 1] cmap = plt.cm.colors.ListedColormap(random_colors) norm = plt.cm.colors.Normalize(vmin=0, vmax=highest_label) geometry = _mask_movie_frame_geometry(masks) def _view_frame_with_tracks(frame=0): """ Display the frame with tracks overlaid. Parameters: frame (int): The frame number to display. Returns: None """ with figure_style(theme_target()): fig, ax = plt.subplots(figsize=geometry['figsize'], dpi=geometry['dpi']) from .figures.bundle import _register_figure_data _register_figure_data(fig, masks[frame], kind="mask", frame=int(frame)) current_mask = masks[frame] ax.imshow(current_mask, cmap=cmap, norm=norm) ax.set_title(f'Frame: {frame}', fontsize=geometry['title_pt']) for label_value in np.unique(current_mask): if label_value == 0: continue y, x = np.mean(np.where(current_mask == label_value), axis=1) ax.text(x, y, str(label_value), color='white', fontsize=geometry['label_pt'], ha='center', va='center') for track in tracks_df['track_id'].unique(): _track = tracks_df[tracks_df['track_id'] == track] ax.plot(_track['x'], _track['y'], '-k', linewidth=1) ax.axis('off') plt.show() if plot: if interactive: interact(_view_frame_with_tracks, frame=IntSlider(min=0, max=len(masks)-1, step=1, value=0)) if save: gif_path = os.path.join(os.path.dirname(src), 'movies', 'gif') os.makedirs(gif_path, exist_ok=True) save_path_gif = os.path.join(gif_path, f'timelapse_masks_{object_type}_{name}.gif') _save_mask_timelapse_as_gif(masks, tracks_df, save_path_gif, cmap, norm, filenames) if plot: if not interactive: _display_gif(save_path_gif) def _display_gif(path): """ Display a GIF image from the given path. Parameters: path (str): The path to the GIF image file. Returns: None """ with open(path, 'rb') as file: display(ipyimage(file.read(), format='gif')) def _plot_recruitment(df, df_type, channel_of_interest, columns=None, figuresize=10): """ Plot recruitment data for different conditions and pathogens. Args: df (DataFrame): The input DataFrame containing the recruitment data. df_type (str): The type of DataFrame (e.g., 'train', 'test'). channel_of_interest (str): The channel of interest for plotting. target (str): The target variable for plotting. columns (list, optional): Additional columns to plot. Defaults to an empty list. figuresize (int, optional): The size of the figure. Defaults to 50. Returns: None """ if columns is None: columns = [] color_list = [(55/255, 155/255, 155/255), (155/255, 55/255, 155/255), (55/255, 155/255, 255/255), (255/255, 55/255, 155/255)] with mpl.rc_context({'axes.prop_cycle': mpl.cycler(color=color_list)}), \ figure_style(theme_target()): font = figuresize/2 width=figuresize height=figuresize/4 fig, axes = plt.subplots(nrows=1, ncols=4, figsize=(width, height)) from .figures.bundle import _register_figure_data _register_figure_data( fig, df, x="condition", y=f"cell_channel_{channel_of_interest}_mean_intensity", hue="pathogen", kind="bar", grid=(1, 4), panels=[{"x": "condition", "hue": "pathogen", "y": f"{name}_channel_{channel_of_interest}_mean_intensity"} for name in ("cell", "nucleus", "cytoplasm", "pathogen")]) sns.barplot(ax=axes[0], data=df, x='condition', y=f'cell_channel_{channel_of_interest}_mean_intensity', hue='pathogen', capsize=.1, errorbar='sd', dodge=False) axes[0].set_xlabel(f'pathogen {df_type}', fontsize=font) axes[0].set_ylabel(f'cell_channel_{channel_of_interest}_mean_intensity', fontsize=font) sns.barplot(ax=axes[1], data=df, x='condition', y=f'nucleus_channel_{channel_of_interest}_mean_intensity', hue='pathogen', capsize=.1, errorbar='sd', dodge=False) axes[1].set_xlabel(f'pathogen {df_type}', fontsize=font) axes[1].set_ylabel(f'nucleus_channel_{channel_of_interest}_mean_intensity', fontsize=font) sns.barplot(ax=axes[2], data=df, x='condition', y=f'cytoplasm_channel_{channel_of_interest}_mean_intensity', hue='pathogen', capsize=.1, errorbar='sd', dodge=False) axes[2].set_xlabel(f'pathogen {df_type}', fontsize=font) axes[2].set_ylabel(f'cytoplasm_channel_{channel_of_interest}_mean_intensity', fontsize=font) sns.barplot(ax=axes[3], data=df, x='condition', y=f'pathogen_channel_{channel_of_interest}_mean_intensity', hue='pathogen', capsize=.1, errorbar='sd', dodge=False) axes[3].set_xlabel(f'pathogen {df_type}', fontsize=font) axes[3].set_ylabel(f'pathogen_channel_{channel_of_interest}_mean_intensity', fontsize=font) handles, labels = axes[3].get_legend_handles_labels() axes[3].legend(handles, labels, bbox_to_anchor=(1.05, 0.5), loc='center left') for i in [0,1,2,3]: axes[i].tick_params(axis='both', which='major', labelsize=font) rotate_ticks(axes[i]) fig.tight_layout() plt.show() columns = columns + [f'pathogen_channel_{channel_of_interest}_cytoplasm_{stat}_ratio' for stat in ('mean', 'q75', 'periphery_mean', 'outside_mean', 'outside_q75')] width = figuresize*2 columns_per_row = math.ceil(len(columns) / 2) height = (figuresize*2)/columns_per_row fig, axes = plt.subplots(nrows=2, ncols=columns_per_row, figsize=(width, height * 2)) from .figures.bundle import _register_figure_data _register_figure_data( fig, df, x="condition", y=str(columns[0]) if len(columns) else "", hue="pathogen", kind="bar", grid=(2, columns_per_row), panels=[{"x": "condition", "y": str(column), "hue": "pathogen"} for column in columns]) axes = axes.flatten() print(f'{columns}') for i, col in enumerate(columns): ax = axes[i] sns.barplot(ax=ax, data=df, x='condition', y=f'{col}', hue='pathogen', capsize=.1, errorbar='sd', dodge=False) ax.set_xlabel(f'pathogen {df_type}', fontsize=font) ax.set_ylabel(f'{col}', fontsize=int(font*2)) if ax.get_legend() is not None: ax.legend_.remove() ax.tick_params(axis='both', which='major', labelsize=font) rotate_ticks(ax) if i <= 5: ax.set_ylim(1, None) hide_unused(axes[len(columns):]) fig.tight_layout() plt.show() def _plot_controls(df, mask_chans, channel_of_interest, figuresize=5): """ Plot controls for different channels and conditions. Args: df (pandas.DataFrame): The DataFrame containing the data. mask_chans (list): The list of channels to include in the plot. channel_of_interest (int): The channel of interest. figuresize (int, optional): The size of the figure. Defaults to 5. Returns: None """ mask_chans.append(channel_of_interest) if len(mask_chans) == 4: mask_chans = [0,1,2,3] if len(mask_chans) == 3: mask_chans = [0,1,2] if len(mask_chans) == 2: mask_chans = [0,1] if len(mask_chans) == 1: mask_chans = [0] controls_cols = [] for chan in mask_chans: controls_cols_c = [] controls_cols_c.append(f'cell_channel_{chan}_mean_intensity') controls_cols_c.append(f'nucleus_channel_{chan}_mean_intensity') controls_cols_c.append(f'pathogen_channel_{chan}_mean_intensity') controls_cols_c.append(f'cytoplasm_channel_{chan}_mean_intensity') controls_cols.append(controls_cols_c) unique_conditions = df['condition'].unique().tolist() if len(unique_conditions) ==1: unique_conditions=unique_conditions+unique_conditions color_list = [ROLES['data']] * 4 with figure_style(theme_target()): fig, axes = plt.subplots(len(unique_conditions), len(mask_chans)+1, figsize=(figuresize*len(mask_chans), figuresize*len(unique_conditions))) from .figures.bundle import _register_figure_data _register_figure_data( fig, df, kind="bar", grid=(len(unique_conditions), len(mask_chans) + 1), panels=[{ "slot": condition_index * (len(mask_chans) + 1) + channel_index, "x": "component", "y": "mean_intensity", "title": f"Condition: {condition} - Channel {channel_index}", "measurement": f"Condition: {condition} - Channel {channel_index}", "xlabel": "Component", "ylabel": "Mean Intensity", "where": {"condition": [condition]}, "melt": { "id_vars": ["condition"], "columns": [column for column in channel_columns if column in df.columns], "var_name": "component", "value_name": "mean_intensity", "labels": {column: column.split("_channel_")[0] for column in channel_columns}, }, } for condition_index, condition in enumerate(unique_conditions) for channel_index, channel_columns in enumerate(controls_cols)]) for idx_condition, condition in enumerate(unique_conditions): df_temp = df[df['condition'] == condition] for idx_channel, control_cols_c in enumerate(controls_cols): names = [] data = [] std_dev = [] colors = [] for color, control_col in zip(color_list, control_cols_c): if control_col in df_temp.columns: mean_intensity = df_temp[control_col].mean() mean_intensity = 0 if np.isnan(mean_intensity) else mean_intensity names.append(control_col.split('_channel_')[0]) data.append(mean_intensity) std_dev.append(df_temp[control_col].std()) colors.append(color) current_axis = axes[idx_condition][idx_channel] current_axis.bar(names, data, yerr=std_dev, capsize=4, color=colors, ecolor=resolve_ink(theme_target()), error_kw={'lw': WEIGHTS['reference']}) current_axis.set_xlabel('Component') current_axis.set_ylabel('Mean Intensity') descriptor(current_axis, f'Condition: {condition} - Channel {idx_channel}') fig.tight_layout() plt.show() def _imshow(img, labels, nrow=20, color='white', fontsize=12): """ Display multiple images in a grid with corresponding labels. Args: img (list): List of images to display. labels (list): List of labels corresponding to each image. nrow (int, optional): Number of images per row in the grid. Defaults to 20. color (str, optional): Color of the label text. Defaults to 'white'. fontsize (int, optional): Font size of the label text. Defaults to 12. """ n_images = len(labels) n_col = nrow n_row = int(np.ceil(n_images / n_col)) img_height = img[0].shape[1] img_width = img[0].shape[2] canvas = np.zeros((img_height * n_row, img_width * n_col, 3)) for i in range(n_row): for j in range(n_col): idx = i * n_col + j if idx < n_images: canvas[i * img_height:(i + 1) * img_height, j * img_width:(j + 1) * img_width] = np.transpose(img[idx], (1, 2, 0)) with figure_style(theme_target()): fig = plt.figure(figsize=(50, 50)) from .figures.bundle import _register_figure_data _register_figure_data(fig, canvas, kind="montage", labels=[str(v) for v in labels]) plt.imshow(canvas) plt.axis("off") for i, label in enumerate(labels): row = i // n_col col = i % n_col x = col * img_width + 2 y = row * img_height + 15 plt.text(x, y, label, color=color, fontsize=fontsize, fontweight='bold') return fig def _imshow_gpu(img, labels, nrow=20, color='white', fontsize=12): """ Display multiple images in a grid with corresponding labels. Args: img (torch.Tensor): A batch of images as a tensor. labels (list): List of labels corresponding to each image. nrow (int, optional): Number of images per row in the grid. Defaults to 20. color (str, optional): Color of the label text. Defaults to 'white'. fontsize (int, optional): Font size of the label text. Defaults to 12. """ if img.is_cuda: img = img.cpu() n_images = len(labels) n_col = nrow n_row = int(np.ceil(n_images / n_col)) img_height = img.shape[2] img_width = img.shape[3] canvas = torch.zeros((img_height * n_row, img_width * n_col, 3)) for i in range(n_row): for j in range(n_col): idx = i * n_col + j if idx < n_images: canvas[i * img_height:(i + 1) * img_height, j * img_width:(j + 1) * img_width] = img[idx].permute(1, 2, 0) canvas = canvas.numpy() with figure_style(theme_target()): fig = plt.figure(figsize=(50, 50)) from .figures.bundle import _register_figure_data _register_figure_data(fig, canvas, kind="montage", labels=[str(v) for v in labels]) plt.imshow(canvas) plt.axis("off") for i, label in enumerate(labels): row = i // n_col col = i % n_col x = col * img_width + 2 y = row * img_height + 15 plt.text(x, y, label, color=color, fontsize=fontsize, fontweight='bold') return fig def _plot_histograms_and_stats(df): """Print prediction statistics and show a histogram for each condition.""" conditions = df['condition'].unique() for condition in conditions: subset = df[df['condition'] == condition] mean_pred = subset['pred'].mean() over_0_5 = sum(subset['pred'] > 0.5) under_0_5 = sum(subset['pred'] <= 0.5) print(f"Condition: {condition}") print(f"Number of rows: {len(subset)}") print(f"Mean of pred: {mean_pred}") print(f"Count of pred values over 0.5: {over_0_5}") print(f"Count of pred values under 0.5: {under_0_5}") print(f"Percent positive: {(over_0_5/(over_0_5+under_0_5))*100}") print(f"Percent negative: {(under_0_5/(over_0_5+under_0_5))*100}") print('-'*40) with figure_style(theme_target()): fig, ax = plt.subplots(figsize=(10, 10)) from .figures.bundle import _register_figure_data _register_figure_data(fig, subset, y="pred", kind="hist", title=f"Condition: {condition}") ax.hist(subset['pred'], bins=30, color=ROLES['fill'], edgecolor='none') mean_line = reference_line(ax, x=mean_pred) mean_line.set_label(f"Mean = {mean_pred:.2f}") descriptor(ax, f'Histogram for pred - Condition: {condition}') ax.set_xlabel('Pred Value') ax.set_ylabel('Count') ax.legend() plt.show() def _show_residules(model): """Draw the three residual diagnostics and print the Shapiro-Wilk test. :param model: anything exposing ``resid`` and ``fittedvalues`` -- a fitted statsmodels result, or a stand-in carrying the same two arrays. :returns: None. Three figures are shown: the residual histogram, the QQ plot, and residuals against fitted values. """ residuals = model.resid with figure_style(theme_target()): fig, ax = plt.subplots() from .figures.bundle import _register_figure_data _register_figure_data(fig, lambda: pd.DataFrame({"residual": np.asarray(residuals)}), y="residual", kind="hist", histogram=dict(bins=30, color=ROLES['fill'], edgecolor='none')) ax.hist(residuals, bins=30, color=ROLES['fill'], edgecolor='none') descriptor(ax, 'Histogram of Residuals') ax.set_xlabel('Residual Value') ax.set_ylabel('Frequency') plt.show() qq_fig, qq_ax = plt.subplots() from .figures.bundle import _register_figure_data _register_figure_data(qq_fig, lambda: pd.DataFrame({"residual": np.asarray(residuals)}), y="residual", kind="qq", qq=dict(fit=True, line='45', data_color=ROLES['data'], reference_color=ROLES['reference'], reference_width=WEIGHTS['reference'])) sm.qqplot(residuals, fit=True, line='45', ax=qq_ax) for line in qq_ax.lines: if line.get_linestyle() == 'None': line.set_color(ROLES['data']) line.set_markerfacecolor(ROLES['data']) line.set_markeredgecolor('none') else: line.set_color(ROLES['reference']) line.set_linewidth(WEIGHTS['reference']) line.set_linestyle((0, (4, 3))) descriptor(qq_ax, 'QQ Plot') plt.show() resid_fig, resid_ax = plt.subplots() from .figures.bundle import _register_figure_data _register_figure_data(resid_fig, lambda: pd.DataFrame({"fitted": np.asarray(model.fittedvalues), "residual": np.asarray(residuals)}), x="fitted", y="residual", kind="scatter", scatter=dict(s=8, color=ROLES['data'], edgecolors='none'), references=[dict(axis='y', value=0, color=ROLES['reference'], linewidth=WEIGHTS['reference'], dashes=[4, 3])]) resid_ax.scatter(model.fittedvalues, residuals, s=8, color=ROLES['data'], edgecolors='none') resid_ax.set_xlabel('Fitted values') resid_ax.set_ylabel('Residuals') descriptor(resid_ax, 'Residuals vs. Fitted Values') reference_line(resid_ax, y=0) plt.show() W, p_value = stats.shapiro(residuals) print(f'Shapiro-Wilk Test W-statistic: {W}, p-value: {p_value}') def _reg_v_plot(df, grouping=None, variable=None, plate_number=None): """Show the legacy regression volcano and label significant rows. ``-log10(p)`` is added to ``df`` in place. The other arguments are retained only for compatibility with historical call sites. """ df['-log10(p)'] = -np.log10(df['p']) called = np.asarray(df['p'] < 0.05) effect = np.asarray(df['effect'], dtype=float) colours = np.where(called & (effect >= 0), ROLES['up'], np.where(called, ROLES['down'], ROLES['data'])) with figure_style(theme_target()): fig, ax = plt.subplots(figsize=(5.6, 4.4)) from .figures.bundle import _register_figure_data _register_figure_data(fig, lambda: df.assign(effect=np.asarray(effect)), x="effect", y="-log10(p)", kind="scatter") ax.scatter(effect, df['-log10(p)'], c=colours, s=12, edgecolors='none') descriptor(ax, 'Volcano Plot') ax.set_xlabel('Coefficient') ax.set_ylabel('-log10(P-value)') for idx, row in df.iterrows(): if row['p'] < 0.05: ax.text(row['effect'], -np.log10(row['p']), idx, fontsize=TYPE_SCALE['annotation'], ha='center', va='bottom', color=resolve_ink(theme_target())) reference_line(ax, y=-np.log10(0.05)) plt.show() def _well_axis_labels(tokens, parse, render): """Map raw row (or column) tokens onto ``(index, canonical label)``. Two rules, and both of them matter: * **The letter walk is borrowed, never rewritten.** ``parse`` is :func:`spacr.plate_qc.parse_row_label` / ``parse_column_label`` and ``render`` is :func:`spacr.schema.row_id` / ``column_id``, so ``'AA'``, ``'r27'`` and ``'row27'`` are one row here, in QC and in the database. A hand-rolled ``chr(ord('A') + n)`` produces ``'['`` for row 27, which is how this class of bug started. * **Each distinct token is parsed once.** A measurement frame is a million object rows over a few hundred wells; parsing per row would run the regexes a million times for a few hundred answers. :param tokens: sequence of raw tokens as they appear in ``prc``. :param parse: label reader returning a 1-based index or ``None``. :param render: id builder turning that index into ``'r<N>'``/``'c<N>'``. :returns: ``(indices, labels)`` object arrays. An unreadable token has index ``None`` and keeps its raw text as its label, so a caller can name it in a report instead of dropping it namelessly. """ cache = {} for token in set(tokens): index = parse(token) cache[token] = (index, token if index is None else render(int(index))) indices = np.array([cache[t][0] for t in tokens], dtype=object) labels = np.array([cache[t][1] for t in tokens], dtype=object) return indices, labels
[docs] def generate_plate_heatmap(df, plate_number, variable, grouping, min_max, min_count): """Aggregate a well-level DataFrame into a plate-shaped heatmap. The grid is **read off the data**. It used to be pinned to ``r1..r16`` by ``c1..c27``, so every well of a 1536 plate past row P or past column 27 fell outside the ``Categorical``, became NaN, and was dropped by the groupby — measured, in the database, and absent from the figure with nothing said. Rows and columns now go through :func:`spacr.plate_qc.parse_row_label` / ``parse_column_label`` (which is :mod:`spacr.schema`'s letter walk, so ``AA``…``AF`` and beyond are real rows), and the axes span exactly the wells present: a 96 plate is still 8x12 and a 384 still 16x24, because nothing is padded out to the largest format that exists. A well that genuinely cannot be placed — a ``prc`` with too few parts, or a row/column token holding no position — is reported through :func:`spacr.errors.raise_if_strict` (an ``ERROR`` on ``spacr.errors``, or a raise under ``SPACR_STRICT_ERRORS``) naming the identifiers concerned. Replacing a silent drop with a quieter silent drop would fix nothing. :param df: Long-format DataFrame with a ``prc`` (plate_row_column) identifier and the requested ``variable`` column. :param plate_number: Plate ID selecting the subset to display. :param variable: Column to aggregate. Ignored when ``grouping='count'``. :param grouping: Aggregation — ``'count'``, ``'mean'`` or ``'sum'``. :param min_max: Colour scale spec — ``'all'``, ``'allq'``, or a two-element list ``[vmin, vmax]`` (floats treated as quantiles). :param min_count: Drop wells with fewer than this many rows. :returns: ``(plate_map, (vmin, vmax))`` — the pivoted matrix, indexed ``'r<N>'`` by ``'c<N>'``, and the colour-limit tuple. :raises ValueError: if ``grouping`` is not one of the accepted values. :raises KeyError: if ``variable`` is missing and required. """ from . import plate_qc as _plate_qc from . import schema as _schema if not isinstance(min_count, (int, float)): min_count = 0 prc_text = df['prc'].astype(str) parts = [text.split(_schema.KEY_SEPARATOR) for text in prc_text] plate_token = np.array( [p[0] if len(p) == 3 else str(plate_number) for p in parts], dtype=object) row_token = np.array([p[-2] if len(p) >= 3 else '' for p in parts], dtype=object) col_token = np.array([p[-1] if len(p) >= 3 else '' for p in parts], dtype=object) if not all(len(p) == 3 for p in parts): df = df.copy() row_index, row_label = _well_axis_labels( row_token, _plate_qc.parse_row_label, _schema.row_id) col_index, col_label = _well_axis_labels( col_token, _plate_qc.parse_column_label, _schema.column_id) df['plateID'], df['rowID'], df['columnID'] = plate_token, row_label, col_label on_plate = np.asarray(plate_token == str(plate_number), dtype=bool) placeable = np.array([r is not None and c is not None for r, c in zip(row_index, col_index)], dtype=bool) lost = on_plate & ~placeable if lost.any(): names = sorted(set(prc_text.to_numpy()[lost])) shown = ', '.join(names[:12]) + (' …' if len(names) > 12 else '') raise_if_strict( f"plate {plate_number!r}: {int(lost.sum())} row(s) covering " f"{len(names)} identifier(s) hold no well position and are " f"missing from the heatmap: {shown}. A prc must be " f"<plate>_<row>_<column> with a readable row ('r3', 'C', 'AA') " f"and column ('c7', '7'); a well drawn nowhere is " f"indistinguishable from a well that was never measured.") keep = on_plate & placeable df = df[keep].copy() df['_row_index'] = row_index[keep].astype(int) df['_col_index'] = col_index[keep].astype(int) keys = ['_row_index', '_col_index'] df['_well_count'] = df.groupby( keys, observed=False)['_row_index'].transform('count') if min_count > 0: df = df[df['_well_count'] >= min_count] grouped = df.groupby(keys, observed=False) if grouping == 'count': plate = grouped.size().reset_index(name='value') elif grouping in ('mean', 'sum'): if variable not in df.columns: raise KeyError(f"variable '{variable}' not in df") vals = pd.to_numeric(df[variable], errors='coerce') tmp = df.assign(__val__=vals) if grouping == 'mean': plate = tmp.groupby( keys, observed=False)['__val__'].mean().reset_index(name='value') else: plate = tmp.groupby( keys, observed=False)['__val__'].sum().reset_index(name='value') else: raise ValueError("grouping must be 'count', 'sum', or 'mean'") plate_map = pd.pivot_table(plate, values='value', index='_row_index', columns='_col_index').fillna(0) plate_map.index = pd.Index([_schema.row_id(int(i)) for i in plate_map.index], name='rowID') plate_map.columns = pd.Index([_schema.column_id(int(i)) for i in plate_map.columns], name='columnID') if plate_map.values.size == 0: return plate_map, (0.0, 1.0) if min_max == 'all': vmin, vmax = float(np.nanmin(plate_map.values)), float(np.nanmax(plate_map.values)) elif min_max == 'allq': vmin, vmax = np.quantile(plate_map.values, [0.02, 0.98]) elif isinstance(min_max, (list, tuple)) and len(min_max) == 2: if all(isinstance(x, float) for x in min_max): vmin, vmax = np.quantile(plate_map.values, [min_max[0], min_max[1]]) else: vmin, vmax = float(min_max[0]), float(min_max[1]) else: vmin, vmax = float(np.nanmin(plate_map.values)), float(np.nanmax(plate_map.values)) if vmin == vmax: vmax = vmin + 1e-6 return plate_map, (vmin, vmax)
#: Legacy default colormap. It is treated as unset so internal calls adopt the #: current house palette; any other colormap is treated as an explicit choice #: and honoured as supplied. LEGACY_PLATE_CMAP = 'viridis'
[docs] def plot_plates(df, variable, grouping, min_max, cmap, min_count=0, verbose=True, dst=None): """Render every plate of a screen as ONE panel, wells square, on one colour scale. The layout, the colour scale and the treatment of unmeasured wells live in :mod:`spacr.figures.plates`; this function is the call the pipeline already makes, kept at its own signature. WHAT CHANGED, AND WHY (the design -- "the lpates look super small on the collected figure"): the plates were laid out four-per-row on a 40 x 5 inch figure, an 8:1 strip that uses about an eighth of a square tile in the figure grid, with wells 1.14:1 rather than square. They are now a SMALL MULTIPLE -- 2 x 2 for a four-plate screen -- which is a 1.3:1 composite, and the figure is sized from the well grid so the wells come out exactly square. Two things that were wrong with the picture and not only with its size: each plate carried its OWN colour scale, so the same blue meant a different number on the plate beside it; and a well that was never measured was drawn as a measurement of zero, which on a screen with 155 of 384 wells used is more than half the panel — and set the bottom of the scale. One scale is now shared across the plates, and an unmeasured well is drawn as a neutral wash and left out of the scale. :param df: Long-format DataFrame with a ``prc`` column of the form ``plateID_rowID_columnID`` and the column named by ``variable``. :param variable: Column to aggregate (see :func:`generate_plate_heatmap`). :param grouping: Aggregation mode — ``'count'``, ``'mean'`` or ``'sum'``. :param min_max: Color-scale spec (``'all'``, ``'allq'``, ``[vmin, vmax]``), applied ONCE over every plate rather than once per plate. :param cmap: Matplotlib colormap name or object. ``None`` — or the legacy ``'viridis'`` literal — uses the house single-hue ramp. :param min_count: Drop wells with fewer than this many rows before plotting. Default ``0``. :param verbose: If True, call ``plt.show()`` after building the figure. Default ``True``. :param dst: If given, save the figure as ``<dst>/plate_heatmap_<variable>.pdf``. :returns: The generated matplotlib ``Figure``. Example: .. code-block:: python from spacr.plot import plot_plates fig = plot_plates( df, variable='recruitment', grouping='mean', min_max='allq', cmap=None, min_count=20, ) See Also: :func:`spacr.figures.plates.build_plates` — the panel itself, which also returns the legend sentence for it. :func:`spacr.ml.generate_ml_scores` — produces score dataframes typically fed to this plotter. """ from .figures.plates import build_plates if isinstance(cmap, str) and cmap.strip().lower() == LEGACY_PLATE_CMAP: cmap = None fig, panel = build_plates(df, variable, grouping=grouping, min_max=min_max, min_count=min_count, cmap=cmap) if not panel.drawn and verbose: print(f'No plate heatmap drawn: {panel.reason}') if dst is not None: from .figures.plates import plate_figure_name filename = os.path.join(dst, plate_figure_name(variable)) filename = save_figure(fig, filename) print(f'Saved heatmap to {filename}') if verbose: plt.show() return fig
[docs] def plot_resize(images, resized_images, labels, resized_labels): """Show original vs. resized image/label pairs in a 2x2 grid. :param images: Sequence of original images (first element shown). :param resized_images: Sequence of resized images. :param labels: Sequence of original label arrays. :param resized_labels: Sequence of resized label arrays. :returns: None """ def prepare_image(img): """Return ``(display_array, cmap)`` handling 2D/3D input shapes. :param img: A 2D array, or a 3D array in channels-last order. One channel is squeezed to 2D and three or four are passed through as RGB/RGBA with a ``None`` colormap; any other channel count (a 5-channel spaCR stack, or a channels-first array read straight off disk) falls back to the mean across the last axis, which is a legal but usually misleading picture. :returns: ``(array, cmap)`` to hand straight to ``imshow``, where ``cmap`` is ``None`` for true-colour data. :raises ValueError: if ``img`` is neither 2D nor 3D. """ if img.ndim == 2: return img, 'gray' elif img.ndim == 3: if img.shape[-1] == 1: return np.squeeze(img, axis=-1), 'gray' elif img.shape[-1] == 3: return img, None elif img.shape[-1] == 4: return img, None else: return np.mean(img, axis=-1), 'gray' else: raise ValueError(f"Unsupported image shape: {img.shape}") with figure_style(theme_target()): fig, ax = plt.subplots(2, 2, figsize=(20, 20)) from .figures.bundle import _register_figure_data _register_figure_data( fig, [images[0], resized_images[0], labels[0], resized_labels[0]], kind="image") img, cmap = prepare_image(images[0]) ax[0, 0].imshow(img, cmap=cmap) ax[0, 0].set_title('Original Image') img, cmap = prepare_image(resized_images[0]) ax[0, 1].imshow(img, cmap=cmap) ax[0, 1].set_title('Resized Image') lbl, cmap = prepare_image(labels[0]) ax[1, 0].imshow(lbl, cmap=cmap) ax[1, 0].set_title('Original Label') lbl, cmap = prepare_image(resized_labels[0]) ax[1, 1].imshow(lbl, cmap=cmap) ax[1, 1].set_title('Resized Label') plt.tight_layout() plt.show()
[docs] def normalize_and_visualize(image, normalized_image, title=""): """Show the original and the normalised image side by side in grayscale. Multi-channel inputs are averaged over their channels for display. :param image: Original image, 2D or ``(H, W, C)``. :param normalized_image: Normalised counterpart to compare against. :param title: Suffix appended to both panel titles. Default ``""``. :returns: None """ with figure_style(theme_target()): fig, ax = plt.subplots(1, 2, figsize=(12, 6)) from .figures.bundle import _register_figure_data _register_figure_data(fig, [image, normalized_image], kind="image", title=str(title)) if image.ndim == 3: ax[0].imshow(np.mean(image, axis=-1), cmap='gray') else: ax[0].imshow(image, cmap='gray') ax[0].set_title("Original " + title) ax[0].axis('off') if normalized_image.ndim == 3: ax[1].imshow(np.mean(normalized_image, axis=-1), cmap='gray') else: ax[1].imshow(normalized_image, cmap='gray') ax[1].set_title("Normalized " + title) ax[1].axis('off') plt.show()
[docs] def visualize_masks(mask1, mask2, mask3, title="Masks Comparison"): """Show three masks side by side with random colormaps. :param mask1: First label mask. :param mask2: Second label mask. :param mask3: Third label mask. :param title: Figure suptitle. Default ``"Masks Comparison"``. :returns: None """ with figure_style(theme_target()): fig, axs = plt.subplots(1, 3, figsize=(30, 10)) from .figures.bundle import _register_figure_data _register_figure_data(fig, [mask1, mask2, mask3], kind="mask", title=str(title)) for ax, mask, panel_title in zip(axs, [mask1, mask2, mask3], ['Mask 1', 'Mask 2', 'Mask 3']): cmap = generate_mask_random_cmap(mask) if np.isin(mask, [0, 1]).all(): ax.imshow(mask, cmap=cmap) else: norm = plt.Normalize(vmin=0, vmax=mask.max()) ax.imshow(mask, cmap=cmap, norm=norm) ax.set_title(panel_title) ax.axis('off') plt.suptitle(title) plt.show()
[docs] def visualize_cellpose_masks(masks, titles=None, filename=None, save=False, src=None): """Display several Cellpose-style label masks side by side for a quick visual QC. Handy for sanity-checking the masks produced by :func:`spacr.core.preprocess_generate_masks` (e.g. compare the cell, nucleus and pathogen masks of the same field, or two runs against each other). Each mask is rendered with a random-color palette so neighbouring objects stay distinguishable. :param masks: Sequence of 2D label mask arrays. :param titles: Titles paired positionally with ``masks``. Falls back to ``'Mask 1'``, ``'Mask 2'``, ... :param filename: Used in the figure suptitle and, when ``save``, as the output PDF filename. :param save: If True, save the figure under ``<src>/results/<filename>.pdf``. Default ``False``. :param src: Root folder for saving. Defaults to the current working directory. :returns: None. Displays (and optionally writes) the figure. :raises AssertionError: if ``titles`` and ``masks`` have different lengths. Example: .. code-block:: python from spacr.plot import visualize_cellpose_masks visualize_cellpose_masks( [cell_mask, nucleus_mask, pathogen_mask], titles=['cell','nucleus','pathogen'], filename='field_001', save=True, src='/data/plate01', ) See Also: :func:`spacr.core.preprocess_generate_masks` — produces the masks visualized here. """ comparison_title=f"Masks Comparison for {filename}" if titles is None: titles = [f'Mask {i+1}' for i in range(len(masks))] assert len(titles) == len(masks), "Number of titles and masks must match" with figure_style(theme_target()): num_masks = len(masks) fig, axs = plt.subplots(1, num_masks, figsize=(10 * num_masks, 10)) from .figures.bundle import _register_figure_data _register_figure_data(fig, list(masks), kind="mask", title=str(comparison_title)) axs = np.atleast_1d(axs) for ax, mask, title in zip(axs, masks, titles): cmap = generate_mask_random_cmap(mask) norm = plt.Normalize(vmin=0, vmax=mask.max()) ax.imshow(mask, cmap=cmap, norm=norm) ax.set_title(title) ax.axis('off') plt.suptitle(comparison_title) plt.show() if save: if src is None: src = os.getcwd() results_dir = os.path.join(src, 'results') os.makedirs(results_dir, exist_ok=True) fig_path = os.path.join(results_dir, f'{filename}.pdf') fig_path = save_figure(fig, fig_path) print(f'Saved figure to {fig_path}') return
[docs] def plot_comparison_results(comparison_results): """Plot Jaccard, Dice, boundary-F1 and average-precision distributions per comparison. :param comparison_results: Iterable of dicts with per-file metrics (each key ending in ``jaccard``/``dice``/``boundary_f1``/ ``average_precision``). :returns: The generated ``Figure``. """ df = pd.DataFrame(comparison_results) df_melted = pd.melt(df, id_vars=['filename'], var_name='metric', value_name='value') df_jaccard = df_melted[df_melted['metric'].str.contains('jaccard')] df_dice = df_melted[df_melted['metric'].str.contains('dice')] df_boundary_f1 = df_melted[df_melted['metric'].str.contains('boundary_f1')] df_ap = df_melted[df_melted['metric'].str.contains('average_precision')] panels = ( (df_jaccard, 'Jaccard Index by Comparison', 'Jaccard Index'), (df_dice, 'Dice Coefficient by Comparison', 'Dice Coefficient'), (df_boundary_f1, 'Boundary F1 Score by Comparison', 'Boundary F1 Score'), (df_ap, 'Average Precision by Comparison', 'Average Precision'), ) with figure_style(theme_target()): fig, axs = plt.subplots(1, 4, figsize=(40, 10)) from .figures.bundle import _register_figure_data _register_figure_data( fig, df_melted, x="metric", y="value", kind="box_strip", grid=(1, 4), panels=[{ "title": title, "xlabel": "Comparison", "ylabel": ylabel, "measurement": ylabel, "where": {"metric": frame["metric"].unique().tolist()}, } for frame, title, ylabel in panels]) for index, (frame, title, ylabel) in enumerate(panels): ax = axs[index] sns.boxplot(data=frame, x='metric', y='value', ax=ax, color=ROLES['data'], linecolor=resolve_ink(theme_target()), linewidth=WEIGHTS['spine'], fliersize=2.0) sns.stripplot(data=frame, x='metric', y='value', ax=ax, jitter=True, color=Palette.GREY_DARK, size=3.0, linewidth=0) descriptor(ax, title) rotate_ticks(ax) ax.set_xlabel('Comparison') ax.set_ylabel(ylabel) panel_letter(ax, 'ABCD'[index]) fig.tight_layout() plt.show() return fig
[docs] def plot_object_outlines(src, objects=None, channels=None, max_nr=10): """Overlay mask outlines on the matching channel image for each object type. :param src: Experiment root; ``masks/<object>_mask_stack`` and channel folders live directly under it. :param objects: Object types to plot. Default ``['nucleus', 'cell', 'pathogen']``. :param channels: Channel indices paired with ``objects`` (channel folders are named ``<channel + 1>``). Default ``[0, 1, 2]``. :param max_nr: Maximum number of images to plot per object. Default ``10``. :returns: None """ if objects is None: objects = ['nucleus','cell','pathogen'] if channels is None: channels = [0,1,2] for object_, channel in zip(objects, channels): folders = [os.path.join(src, 'masks', f'{object_}_mask_stack'), os.path.join(src,f'{channel+1}')] print(folders) plot_images_and_arrays(folders, lower_percentile=2, upper_percentile=99.5, threshold=1000, extensions=['.npy', '.tif', '.tiff', '.png'], overlay=True, max_nr=max_nr, randomize=True)
[docs] def plot_histogram(df, column, dst=None): """Plot a histogram of ``df[column]`` and optionally save it as PDF. :param df: DataFrame containing ``column``. :param column: Column to plot. :param dst: If set, save under ``<dst>/<column>_histogram.pdf``. :returns: None """ with figure_style(theme_target()): fig, ax = plt.subplots(figsize=(10, 10)) from .figures.bundle import _register_figure_data _register_figure_data(fig, df, y=column, kind="hist") sns.histplot(df[column], kde=False, color=ROLES['fill'], edgecolor=None, ax=ax) descriptor(ax, f'Histogram of {column}') ax.set_xlabel(column) ax.set_ylabel('Frequency') if not dst is None: filename = os.path.join(dst, f'{column}_histogram.pdf') filename = save_figure(fig, filename) print(f'Saved histogram to {filename}') plt.show()
[docs] def plot_lorenz_curves(csv_files, name_column='grna_name', value_column='count', remove_keys=None, x_lim=None, y_lim=None, remove_outliers=False, save=True): """Overlay Lorenz curves from multiple gRNA count CSVs with per-plate Gini coefficients. :param csv_files: Paths to per-plate CSVs, each with columns ``name_column`` and ``value_column``. :param name_column: Identifier column used for outlier filtering. Default ``'grna_name'``. :param value_column: Column whose distribution is analysed. Default ``'count'``. :param remove_keys: Names to exclude before analysis. Default ``[]`` (exclude nothing). :param x_lim: X-axis limits ``[lo, hi]``. Default ``[0.0, 1]``. :param y_lim: Y-axis limits ``[lo, hi]``. Default ``[0, 1]``. :param remove_outliers: If True, drop names whose number of observations falls outside a fence extending 1.5 times the 5th-to-95th-percentile spread of group sizes. Count values do not enter this filter. Default ``False``. :param save: If True, save the figure alongside the first CSV under ``results/lorenz_curve_with_gini.pdf``. Default ``True``. :returns: None """ if remove_keys is None: remove_keys = [] if x_lim is None: x_lim = [0.0, 1] if y_lim is None: y_lim = [0, 1] def lorenz_curve(data): """Calculate Lorenz curve. :param data: 1D array of non-negative counts; it is sorted here, so the caller's order does not matter. The curve is normalised by the running total's last element, so an all-zero input divides by zero and an input mixing signs is not a Lorenz curve at all. Must be non-empty — an empty array indexes past the end. :returns: ``len(data) + 1`` cumulative shares rising from 0 to 1, one longer than the input because the origin is prepended. """ sorted_data = np.sort(data) cumulative_data = np.cumsum(sorted_data) lorenz_curve = cumulative_data / cumulative_data[-1] lorenz_curve = np.insert(lorenz_curve, 0, 0) return lorenz_curve def gini_coefficient(data): """Calculate Gini coefficient from data. :param data: 1D array of non-negative counts, sorted internally. It is normalised by ``np.sum(data)``, so an all-zero input yields ``nan``. Unlike ``lorenz_curve``, an empty array does not raise here — it silently returns ``1.0``, the value for maximum inequality — so filter empty plates out upstream. :returns: The Gini coefficient as a float, from ``0.0`` for a perfectly even distribution up towards ``1.0`` as the counts concentrate on a few gRNAs. The area is taken with the trapezoid rule, so an even distribution reports exactly ``0.0`` rather than ``1/n``. """ sorted_data = np.sort(data) n = len(data) cumulative_data = np.cumsum(sorted_data) / np.sum(sorted_data) cumulative_data = np.insert(cumulative_data, 0, 0) gini = 1 - np.sum((cumulative_data[:-1] + cumulative_data[1:]) * np.diff(np.linspace(0, 1, n + 1))) return gini def remove_outliers_by_wells(data, name_col, wells_col): """Remove outliers based on 95% confidence interval for well counts. :param data: DataFrame with one row per well-and-name observation. Whole names are kept or dropped together, never individual rows, so the surviving frame still has every well of every name it keeps. :param name_col: Column identifying the gRNA (or other name). Rows are grouped on it and the group *size* — the number of wells a name appears in — is what the fence is applied to, so the count column's values play no part in this filter. :param wells_col: Accepted so the call reads symmetrically with the enclosing function's ``value_column``, but never read: the well count is derived from the group sizes above. Passing a wrong or missing column name changes nothing. :returns: ``data`` restricted to the names inside the fence. The fence is ``1.5 *`` the 5th-to-95th-percentile spread, not the interquartile range, so it is far wider than a textbook IQR rule and its lower edge is usually negative — in practice only unusually widespread names are removed. """ well_counts = data.groupby(name_col, observed=False).size() q1 = well_counts.quantile(0.05) q3 = well_counts.quantile(0.95) iqr_range = q3 - q1 lower_bound = q1 - 1.5 * iqr_range upper_bound = q3 + 1.5 * iqr_range valid_names = well_counts[(well_counts >= lower_bound) & (well_counts <= upper_bound)].index return data[data[name_col].isin(valid_names)] combined_data = [] gini_values = {} source_frames = [] curves = [] entries = [] with figure_style(theme_target()): fig, ax = plt.subplots(figsize=(10, 10)) for idx, csv_file in enumerate(csv_files): df = pd.read_csv(csv_file) source_frames.append(df.copy()) for remove in remove_keys: df = df[df[name_column] != remove] if remove_outliers: df = remove_outliers_by_wells(df, name_column, value_column) values = df[value_column].values combined_data.extend(values) curves.append(dict(input=idx, path=str(csv_file), rows=df.index.tolist(), label=f"plate {idx+1}", color=ROLES['data'], linestyle='-')) lorenz = lorenz_curve(values) gini = gini_coefficient(values) gini_values[f"plate {idx+1}"] = gini name = f"plate {idx+1} (Gini: {gini:.4f})" ax.plot(np.linspace(0, 1, len(lorenz)), lorenz, label=name, color=ROLES['data']) entries.append((name, ROLES['data'])) combined_lorenz = lorenz_curve(np.array(combined_data)) combined_gini = gini_coefficient(np.array(combined_data)) gini_values["Combined"] = combined_gini combined_label = f"Combined (Gini: {combined_gini:.4f})" ax.plot(np.linspace(0, 1, len(combined_lorenz)), combined_lorenz, label=combined_label, linestyle='--', color=ROLES['highlight']) entries.append((combined_label, ROLES['highlight'])) from .figures.bundle import _register_figure_data source = pd.concat(source_frames, ignore_index=True) identity_columns = [] for name in ("_spacr_lorenz_input", "_spacr_lorenz_row"): while name in source.columns: name += "_" identity_columns.append(name) source[identity_columns[0]] = np.concatenate([ np.full(len(frame), idx) for idx, frame in enumerate(source_frames)]) source[identity_columns[1]] = np.concatenate([ np.arange(len(frame)) for frame in source_frames]) _register_figure_data(fig, source, y=value_column, kind="lorenz", keep_limits=True, gini=dict(gini_values), lorenz=dict(input_column=identity_columns[0], row_column=identity_columns[1], name_column=name_column, remove_keys=list(remove_keys), remove_outliers=bool(remove_outliers), curves=curves, combined_color=ROLES['highlight'])) ax.set_xlim(x_lim) ax.set_ylim(y_lim) descriptor(ax, 'Lorenz Curves') ax.set_xlabel('Cumulative Share of Individuals') ax.set_ylabel('Cumulative Share of Value') text_legend(ax, entries) if save: save_path = os.path.join(os.path.dirname(csv_files[0]), 'results') os.makedirs(save_path, exist_ok=True) save_file_path = os.path.join(save_path, 'lorenz_curve_with_gini.pdf') save_file_path = save_figure(fig, save_file_path, bbox_inches='tight') print(f"Saved Lorenz Curve: {save_file_path}") plt.show() for plate, gini in gini_values.items(): print(f"{plate}: Gini Coefficient = {gini:.4f}")
[docs] def plot_permutation(permutation_df): """Plot a horizontal bar chart of permutation feature importances with error bars. :param permutation_df: DataFrame with columns ``feature``, ``importance_mean`` and ``importance_std``. :returns: The generated ``Figure``. """ num_features = len(permutation_df) fig_height = max(8, num_features * 0.3) fig_width = 10 font_size = max(10, 12 - num_features * 0.2) with figure_style(theme_target()): fig, ax = plt.subplots(figsize=(fig_width, fig_height)) from .figures.bundle import _register_figure_data _register_figure_data(fig, permutation_df, x="feature", y="importance_mean", kind="bar") ax.barh(permutation_df['feature'], permutation_df['importance_mean'], xerr=permutation_df['importance_std'], color=ROLES['data'], align="center", ecolor=resolve_ink(theme_target()), error_kw={'lw': WEIGHTS['reference']}) if float(np.nanmin(np.asarray( permutation_df['importance_mean'], dtype=float))) < 0: reference_line(ax, x=0) ax.set_xlabel('Permutation Importance', fontsize=font_size) ax.tick_params(axis='both', which='major', labelsize=font_size) fig.tight_layout() return fig
[docs] def plot_feature_importance(feature_importance_df, title=""): """Plot a horizontal bar chart of raw feature importances. :param feature_importance_df: DataFrame with columns ``feature`` and ``importance``. :param title: what the bars MEAN, when it is not the model's own importances. Four of the classifiers spaCR offers expose no ``feature_importances_`` and are drawn from permutation importance instead -- a different quantity, measuring what the fitted model loses when a column is shuffled -- and a panel that did not say so would be passing one off as the other. :returns: The generated ``Figure``. """ num_features = len(feature_importance_df) fig_height = max(8, num_features * 0.3) fig_width = 10 font_size = max(10, 12 - num_features * 0.2) with figure_style(theme_target()): fig, ax = plt.subplots(figsize=(fig_width, fig_height)) from .figures.bundle import _register_figure_data _register_figure_data(fig, feature_importance_df, x="feature", y="importance", kind="bar") ax.barh(feature_importance_df['feature'], feature_importance_df['importance'], color=ROLES['data'], align="center") if float(np.nanmin(np.asarray( feature_importance_df['importance'], dtype=float))) < 0: reference_line(ax, x=0) ax.set_xlabel('Feature Importance', fontsize=font_size) if title: ax.set_title(str(title), fontsize=font_size + 1) ax.tick_params(axis='both', which='major', labelsize=font_size) fig.tight_layout() return fig
[docs] def read_and_plot__vision_results(base_dir, y_axis='accuracy', name_split='_time', y_lim=None): """Aggregate vision-model test CSVs under ``base_dir`` and plot mean score per model. :param base_dir: Root directory containing ``*_test_result.csv`` files nested per epoch. :param y_axis: Metric column to average. Default ``'accuracy'``. :param name_split: Substring that splits filename into model name and epoch info. Default ``'_time'``. :param y_lim: Y-axis limits ``[lo, hi]``. Default ``[0.8, 0.9]``. :returns: None """ if y_lim is None: y_lim = [0.8, 0.9] data_frames = [] dst = os.path.join(base_dir, 'result') os.makedirs(dst, exist_ok=True) for root, dirs, files in os.walk(base_dir): for file in files: if file.endswith("_test_result.csv"): file_path = os.path.join(root, file) file_name = os.path.basename(file_path) model = file_name.split(f'{name_split}')[0] base_folder = os.path.dirname(file_path) epoch = os.path.basename(base_folder) df = pd.read_csv(file_path) df['model'] = model df['epoch'] = epoch data_frames.append(df) if data_frames: result_df = pd.concat(data_frames, ignore_index=True) avg_metric = result_df.groupby( 'model', observed=False)[y_axis].mean().reset_index() avg_metric = avg_metric.sort_values(by=y_axis) print(avg_metric) colours = [ROLES['data']] * len(avg_metric) if len(colours) > 1: colours[-1] = ROLES['highlight'] with figure_style(theme_target()): fig, ax = plt.subplots(figsize=(10, 6)) from .figures.bundle import _register_figure_data _register_figure_data(fig, result_df, x="model", y=y_axis, kind="bar") ax.bar(avg_metric['model'], avg_metric[y_axis], color=colours) ax.set_xlabel('Model') ax.set_ylabel(f'{y_axis}') descriptor(ax, f'Average {y_axis.capitalize()} per Model') rotate_ticks(ax) fig.tight_layout() ax.set_ylim(y_lim) plt.show() else: print("No CSV files found in the specified directory.")
[docs] def jitterplot_by_annotation(src, x_column, y_column, plot_title='Jitter Plot', output_path=None, filter_column=None, filter_values=None): """Read measurements + annotation from a spacr DB and plot a class-balanced jitter plot. :param src: Path to a spacr experiment directory containing ``measurements/measurements.db``. :param x_column: Column used as grouping variable (x-axis). :param y_column: Numeric column plotted on the y-axis. :param plot_title: Title for the plot. Default ``'Jitter Plot'``. :param output_path: If set, save the figure to this path; otherwise show it. :param filter_column: Optional column (or list of columns) to filter rows on before plotting. :param filter_values: Values (or list of value lists) accepted per ``filter_column``. :returns: Balanced ``DataFrame`` used for the plot. :raises KeyError: if required plate/row/col columns are missing. """ def join_measurments_and_annotation(src, tables): """Join per-object measurement tables with the ``png_list`` annotation table. :param src: spaCR experiment directory; the database is read from ``<src>/measurements/measurements.db`` and no other layout is supported — pass the experiment folder, not the ``.db`` file. :param tables: Object tables to merge, joined on ``prcfo``. Every name listed must exist in the database. :returns: One row per object, with the ``png_list`` crop path attached by a left join. :raises pandas.errors.MergeError: if ``png_list`` holds more than one crop per ``prcfo``; the join is validated ``one_to_one`` precisely so duplicated crops cannot silently multiply the measurement rows and inflate the jitter plot. """ from .io import _read_and_merge_data, _read_db db_loc = [src+'/measurements/measurements.db'] loc = src+'/measurements/measurements.db' df, _ = _read_and_merge_data(db_loc, tables, verbose=True, nuclei_limit=True, pathogen_limit=True) paths_df = _read_db(loc, tables=['png_list']) merged_df = pd.merge(df, paths_df[0], on='prcfo', how='left', validate='one_to_one') return merged_df df = join_measurments_and_annotation(src, tables=['cell', 'nucleus', 'pathogen', 'cytoplasm']) print(f"Generated dataframe with: {df.shape[1]} columns and {df.shape[0]} rows") df[x_column] = df[x_column].fillna('NaN') if not filter_column is None: if isinstance(filter_column, str): df = df[df[filter_column].isin(filter_values)] if isinstance(filter_column, list): for i,val in enumerate(filter_column): print(f'hello {len(df)}') df = df[df[val].isin(filter_values[i])] def _resolve_well_column(frame, *bases): """Return the first ``_x``, bare, or ``_y`` well-column spelling.""" for base in bases: for candidate in (f'{base}_x', base, f'{base}_y'): if candidate in frame.columns: return candidate return None required_columns = [_resolve_well_column(df, 'plateID', 'plate'), _resolve_well_column(df, 'rowID', 'row'), _resolve_well_column(df, 'columnID', 'col')] if any(column is None for column in required_columns): raise KeyError("DataFrame does not contain the necessary columns: ['plateID', 'rowID', 'columnID']") non_nan_df = df[df[x_column] != 'NaN'] retained_rows = df[df[required_columns].apply(tuple, axis=1).isin(non_nan_df[required_columns].apply(tuple, axis=1))] min_count = retained_rows[x_column].value_counts().min() print(f'Found {min_count} annotated images') balanced_df = retained_rows.groupby( x_column, observed=False, group_keys=False ).sample(n=min_count, random_state=42).reset_index(drop=True) groups = list(pd.unique(balanced_df[x_column])) with figure_style(theme_target()): fig, ax = plt.subplots(figsize=(10, 6)) from .figures.bundle import _register_figure_data _register_figure_data(fig, balanced_df, x=x_column, y=y_column, kind="strip") sns.stripplot(data=balanced_df, x=x_column, y=y_column, hue=x_column, jitter=True, dodge=False, legend=False, size=3.0, linewidth=0, palette={group: ROLES['data'] for group in groups}) for index, group in enumerate(groups): mean = float(balanced_df.loc[balanced_df[x_column] == group, y_column].mean()) ax.plot([index - 0.28, index + 0.28], [mean, mean], color=Palette.GREY_DARK, lw=WEIGHTS['data'], solid_capstyle='butt', zorder=3) descriptor(ax, plot_title) ax.set_xlabel(x_column) ax.set_ylabel(y_column) rotate_ticks(ax) if output_path: output_path = save_figure(fig, output_path, bbox_inches='tight') print(f"Jitter plot saved to {output_path}") else: plt.show() return balanced_df
[docs] def create_grouped_plot(df, grouping_column, data_column, graph_type='jitter_box', summary_func='mean', order=None, colors=None, output_dir='./output', save=False, y_lim=None, error_bar_type='std'): """Plot grouped observations and run assumption-aware comparisons. Pairwise tests are chosen independently by :func:`spacr.figures.stats.compare`. Student's t, Welch's t, or Mann-Whitney U is used according to the normality and equal-variance checks; an underpowered assumption check selects the rank test. When at least three groups jointly pass normality, Tukey HSD rows are added. The ``'jitter_box'`` default is a STATISTICAL CORRECTION, not a presentation preference: it shows the observations and their distribution instead of reducing each group to a mean bar. Two groups can have the same mean while having different spreads. The box summarizes the distribution; the jitter stays because the points are the evidence. Parameters ---------- df : pandas.DataFrame Source observations. grouping_column : str Categorical column defining groups. data_column : str Numeric column to plot and compare. graph_type : {'bar', 'violin', 'jitter', 'box', 'jitter_box'}, optional Plot representation. The default shows every observation together with median, quartiles, and whiskers. summary_func : str or callable, optional Aggregation used by the bar representation. order : sequence of str, optional Group order. By default, sort observed group values. colors : palette-like, optional Colours passed to seaborn. By default, use the house data colour. output_dir : path-like, optional Directory for saved output. save : bool, optional Save the figure and ``test_results.csv`` when true. y_lim : sequence of float, optional Two-element vertical-axis limits. error_bar_type : {'std', 'sem'}, optional Error statistic for bar plots. Returns ------- figure : matplotlib.figure.Figure Displayed figure. It carries the source recipe used by the interactive representation menu. results_df : pandas.DataFrame Normality, pairwise, and optional Tukey HSD results. Raises ------ ValueError If a bar plot receives an unsupported ``error_bar_type``. """ from .figures.stats import _clean, check_normality, compare from .sp_stats import _ENGINE_TEST_NAMES df = df.dropna(subset=[grouping_column]) if save: os.makedirs(output_dir, exist_ok=True) if order: df[grouping_column] = pd.Categorical(df[grouping_column], categories=order, ordered=True) else: df[grouping_column] = pd.Categorical(df[grouping_column], categories=sorted(df[grouping_column].unique()), ordered=True) unique_groups = df[grouping_column].unique() test_results = [] grouped_data = {group: _clean(df.loc[df[grouping_column] == group, data_column]) for group in unique_groups} is_normal = check_normality(list(grouped_data.values())).passed for group, values in grouped_data.items(): check = check_normality([values]) test_results.append({ 'Comparison': f'Normality test for {group}', 'Test Statistic': check.statistic, 'p-value': check.p_value, 'Test Name': 'Normality test' }) from .figures.stats import MIN_N_FOR_TEST singletons = sum(1 for values in grouped_data.values() if len(values) < MIN_N_FOR_TEST) if len(unique_groups) > 20 and singletons > 0.9 * len(unique_groups): raise ValueError( f"{grouping_column!r} is not a grouping: it has " f"{len(unique_groups)} distinct values and {singletons} of them " f"appear on a single row, so there is nothing to compare. That " f"is what a column of measurements looks like -- group by a " f"category (a well, a plate, a condition, a gene) and put the " f"measurement on the y axis.") comparisons = list(itertools.combinations(unique_groups, 2)) chosen_names = [] untestable = [] for (group1, group2) in comparisons: pair = {group1: grouped_data[group1], group2: grouped_data[group2]} try: result = compare(pair) except ValueError as refusal: untestable.append((group1, group2, str(refusal))) test_results.append({ 'Comparison': f'{group1} vs {group2}', 'Test Statistic': float('nan'), 'p-value': float('nan'), 'Test Name': 'not testable'}) continue name = _ENGINE_TEST_NAMES.get(result.test, result.test) chosen_names.append(name) test_results.append({'Comparison': f'{group1} vs {group2}', 'Test Statistic': result.statistic, 'p-value': result.p_value, 'Test Name': name}) if untestable: thin = sorted(group for group, values in grouped_data.items() if len(values) < MIN_N_FOR_TEST) print(f"{len(untestable)} of {len(comparisons)} comparison(s) could " f"not be tested, because {len(thin)} group(s) have fewer than " f"{MIN_N_FOR_TEST} usable observations: " f"{', '.join(str(g) for g in thin[:5])}" f"{'...' if len(thin) > 5 else ''}. " f"They are in the results table marked 'not testable'.") test_name = ', '.join(dict.fromkeys(chosen_names)) or 'not testable' if is_normal and len(unique_groups) > 2: tukey_result = pairwise_tukeyhsd(df[data_column], df[grouping_column], alpha=0.05) for comparison, p_value in zip(tukey_result._results_table.data[1:], tukey_result.pvalues): test_results.append({ 'Comparison': f'{comparison[0]} vs {comparison[1]}', 'Test Statistic': None, 'p-value': p_value, 'Test Name': 'Tukey HSD Post-hoc' }) with figure_style(theme_target()): fig = plt.figure(figsize=(10, 6)) from .figures.bundle import _register_figure_data _register_figure_data(fig, df, x=grouping_column, y=data_column, kind=graph_type) if colors: color_palette = colors else: color_palette = (_group_colours(len(unique_groups)) or [ROLES['data']] * len(unique_groups)) if graph_type == 'bar': summary_df = df.groupby( grouping_column, observed=False)[data_column].agg( [summary_func, 'std', 'sem']) if error_bar_type == 'std': error_bars = summary_df['std'] elif error_bar_type == 'sem': error_bars = summary_df['sem'] else: raise ValueError(f"Invalid error_bar_type: {error_bar_type}. Choose either 'std' or 'sem'.") sns.barplot( x=grouping_column, y=summary_func, hue=grouping_column, data=summary_df.reset_index(), errorbar=None, order=order, palette=color_palette, legend=False) plt.errorbar(x=np.arange(len(summary_df)), y=summary_df[summary_func], yerr=error_bars, fmt='none', c=resolve_ink(theme_target()), capsize=5, lw=WEIGHTS['reference']) elif graph_type == 'violin': sns.violinplot( x=grouping_column, y=data_column, hue=grouping_column, data=df, order=order, palette=color_palette, legend=False) elif graph_type == 'jitter': sns.stripplot( x=grouping_column, y=data_column, hue=grouping_column, data=df, jitter=True, order=order, palette=color_palette, legend=False) elif graph_type == 'box': sns.boxplot( x=grouping_column, y=data_column, hue=grouping_column, data=df, order=order, palette=color_palette, legend=False) elif graph_type == 'jitter_box': _ink = resolve_ink(theme_target()) sns.boxplot( x=grouping_column, y=data_column, hue=grouping_column, data=df, order=order, palette=color_palette, legend=False, showfliers=False, boxprops={'facecolor': 'none', 'edgecolor': _ink, 'linewidth': WEIGHTS['spine']}, whiskerprops={'color': _ink, 'linewidth': WEIGHTS['spine']}, capprops={'color': _ink, 'linewidth': WEIGHTS['spine']}, medianprops={'color': _ink, 'linewidth': WEIGHTS['data']}) sns.stripplot(x=grouping_column, y=data_column, data=df, jitter=True, color=ROLES['data'], size=3.0, linewidth=0, order=order) elif graph_type == 'jitter_bar': _ink = resolve_ink(theme_target()) summary_df = df.groupby( grouping_column, observed=False)[data_column].agg( [summary_func, 'std', 'sem']) spread = summary_df['sem' if error_bar_type == 'sem' else 'std'] sns.barplot( x=grouping_column, y=summary_func, hue=grouping_column, data=summary_df.reset_index(), errorbar=None, order=order, palette=color_palette, legend=False, fill=False, edgecolor=_ink, linewidth=WEIGHTS['spine']) plt.errorbar(x=np.arange(len(summary_df)), y=summary_df[summary_func], yerr=spread, fmt='none', c=_ink, capsize=5, lw=WEIGHTS['reference']) sns.stripplot(x=grouping_column, y=data_column, data=df, jitter=True, color=ROLES['data'], size=3.0, linewidth=0, order=order) elif graph_type in ('line', 'line_std'): _ink = resolve_ink(theme_target()) summary_df = df.groupby( grouping_column, observed=False)[data_column].agg( [summary_func, 'std', 'sem']) if order: summary_df = summary_df.reindex(order) spread = summary_df['sem' if error_bar_type == 'sem' else 'std'] positions = np.arange(len(summary_df)) plt.errorbar(x=positions, y=summary_df[summary_func], yerr=spread.fillna(0.0), marker='o', markersize=6, capsize=4, lw=WEIGHTS['data'], color=color_palette[0] if color_palette else _ink) plt.xticks(positions, [str(v) for v in summary_df.index]) plt.xlabel(str(grouping_column)) plt.ylabel(str(data_column)) else: raise ValueError( f"graph_type={graph_type!r} is not one of bar, violin, " f"jitter, box, jitter_box, jitter_bar, line") results_df = pd.DataFrame(test_results) if isinstance(y_lim, list) and len(y_lim) == 2: plt.ylim(y_lim) axis = plt.gca() descriptor(axis, f'{test_name} results for {graph_type} plot') rotate_ticks(axis) plt.tight_layout() if save: plot_path = os.path.join(output_dir, 'grouped_plot') plot_path = save_figure(plt.gcf(), plot_path) print(f"Plot saved to {plot_path}") results_path = os.path.join(output_dir, 'test_results.csv') results_df.to_csv(results_path, index=False) print(f"Test results saved to {results_path}") plt.show() figure = plt.gcf() try: figure._spacr_replot = { "df": df, "grouping_column": grouping_column, "data_column": data_column, "graph_type": graph_type, "summary_func": summary_func, "order": order, "colors": colors, "y_lim": y_lim, "error_bar_type": error_bar_type, } except Exception: # noqa: BLE001 pass return figure, results_df
def _finite_p_value(value): """``float(value)`` when it is a real p-value, else ``None``. A results row can carry ``None`` (a test that was skipped) or ``nan`` (a test that was refused), and both mean the comparison has no answer. """ try: number = float(value) except (TypeError, ValueError): return None return number if np.isfinite(number) else None def _significance_marker(p_value): """Return the conventional plot annotation for a statistical p-value.""" if p_value <= 0.001: return '***' if p_value <= 0.01: return '**' if p_value <= 0.05: return '*' return 'ns'
[docs] class spacrGraph: """Grouped plot + statistical-test helper for spacr experiment DataFrames. Wraps preprocessing (aggregation by object / well / plate), normality and variance testing, group-wise pairwise stats, and plot rendering (bar / jitter / box / violin / jitter_box / jitter_bar / line / line_std) in a single object whose output can optionally be persisted alongside a CSV of stats. :param df: Input DataFrame. :param grouping_column: Categorical grouping variable. :param data_column: Metric column (or list of columns) to summarise. :param graph_type: Plot type. Default ``'jitter_box'`` -- a box with the points over it. See :func:`create_grouped_plot` for why that default is a correction rather than a taste. :param summary_func: Aggregator for well/plate level. Default ``'mean'``. :param order: Explicit ordering of groups. :param colors: Optional colour palette. :param output_dir: Save location when ``save=True``. :param save: If True, persist plot and stats. :param y_lim: Two-element y-axis limits. :param log_y: Use log scale for y-axis. :param log_x: Use log scale for x-axis. :param error_bar_type: ``'std'`` or ``'sem'``. Default ``'std'``. :param remove_outliers: Drop 1.5*IQR outliers per group before plotting. :param theme: Seaborn palette name. Default ``'pastel'``. :param representation: Aggregation level — ``'object'``, ``'well'`` or ``'plate'``. Default ``'object'``. :param paired: Treat groups as paired samples where applicable. :param all_to_all: Run every pairwise comparison; ``False`` compares each group to ``compare_group``. :param compare_group: Reference group when ``all_to_all=False``. :param graph_name: Prefix for saved file names. :param annotate_stats: Draw a bracket over each pairwise comparison with its asterisks (or ``ns``) above it. Default ``False``: the tests are run and written to the results table on every plot, but with ``all_to_all=True`` an N-group plot has N(N-1)/2 comparisons and a stack of that many brackets buries the data it is about. Ask for them when the comparisons are few enough to read. Only drawn for a single ``data_column``; see :meth:`_draw_comparison_lines`. """ def __init__(self, df, grouping_column, data_column, graph_type='jitter_box', summary_func='mean', order=None, colors=None, output_dir='./output', save=False, y_lim=None, log_y=False, log_x=False, error_bar_type='std', remove_outliers=False, theme='pastel', representation='object', paired=False, all_to_all=True, compare_group=None, graph_name=None, annotate_stats=False): """Store configuration, set the theme, and preprocess the DataFrame.""" self.df = df self.grouping_column = grouping_column self.order = order or sorted(df[self.grouping_column].dropna().unique().tolist()) self.data_column = data_column if isinstance(data_column, list) else [data_column] self.graph_type = graph_type self.summary_func = summary_func self.colors = colors self.output_dir = output_dir self.save = save self.error_bar_type = error_bar_type self.remove_outliers = remove_outliers self.theme = theme self.representation = representation self.paired = paired self.all_to_all = all_to_all self.compare_group = compare_group self.y_lim = y_lim self.graph_name = graph_name self.log_x = log_x self.log_y = log_y self.annotate_stats = bool(annotate_stats) self.results_df = pd.DataFrame() self.sns_palette = None self.fig = None self.results_name = str(self.graph_name)+'_'+str(self.data_column[0])+'_'+str(self.grouping_column)+'_'+str(self.graph_type) self._set_theme() self.raw_df = self.df.copy() self.df = self.preprocess_data() def _set_theme(self): """Set the Seaborn theme and reorder colors if necessary.""" integer_list = list(range(1, 81)) color_order = [7,9,4,0,3,6,2] + integer_list self.sns_palette = self._set_reordered_theme(self.theme, color_order, 100) def _set_reordered_theme(self, theme='deep', order=None, n_colors=100, show_theme=False): """Set and reorder the Seaborn color palette.""" palette = sns.color_palette(theme, n_colors) if order: reordered_palette = [palette[i] for i in order] else: reordered_palette = palette if show_theme: sns.palplot(reordered_palette) plt.show() return reordered_palette def _user_style(self): """The grouped-graph settings the user changed in Preferences.""" return _preference_deltas("jitter_bar") def _jitter(self): """Horizontal spread of the overlaid points, in category widths.""" return float(self._user_style().get("jitter_width", self.bar_width)) def _point_look(self, size): """``(alpha, size)`` for overlaid points; ``size`` is the house size.""" chosen = self._user_style() alpha = float(chosen.get("point_alpha", 0.6)) if "marker_size" in chosen: size = float(chosen["marker_size"]) ** 0.5 return alpha, size def _points_overlaid(self): """Whether points are drawn over bars and boxes.""" return bool(self._user_style().get("point_overlay", True)) def _error_half_width(self, row): """The error bar's half width for one summary row, or ``None``. The Preferences error-bar style decides when the user changed it: ``sem``, ``sd``, ``ci`` at the chosen confidence level, ``ci95`` or ``none``. Otherwise ``error_bar_type`` decides. """ chosen = str(self._user_style().get("error_bars", "")).lower() if not chosen: return row[self.error_bar_type] if chosen == "none": return None if chosen == "sd": return row["std"] if chosen in ("ci", "ci95"): level = 95.0 if chosen == "ci95" else float( self._user_style().get("ci_level", 95.0)) count = int(row.get("count", 0) or 0) if count < 2: return None from scipy.stats import t as student_t return row["sem"] * student_t.ppf(0.5 + level / 200.0, count - 1) return row["sem"] def _plot_palette(self, count): """The colours for ``count`` drawn series, under the house rule. EVERYTHING IS GREY EXCEPT WHAT THE SENTENCE IS ABOUT. With a single data column the hue is the grouping column -- which is already the x axis -- so painting each group a different colour argues nothing and every series is the one grey. The reordered seaborn theme stays in charge of the case where a categorical palette is genuinely the data: several measurements overlaid on one axis, where the colour is the only thing saying which measurement a mark belongs to. ``colors`` wins over both. It was a documented constructor parameter that was stored and never read, so a caller who passed a palette got the theme's anyway. THE USER'S ``mark_colouring`` comes next: ``uniform`` and ``random`` replace the house rule (see :func:`spacr.figures.style._group_colours`), and ``group``, the default, is the house rule itself. :param count: how many series will be drawn. :returns: a list of ``count`` colour specs. """ count = max(1, int(count)) if self.colors: chosen = list(self.colors) return [chosen[index % len(chosen)] for index in range(count)] picked = self._user_style().get("palette") if picked: from .figure_style import palette_colours chosen = palette_colours(picked) if chosen: return [chosen[index % len(chosen)] for index in range(count)] ruled = _group_colours(count, list(self.sns_palette or ())) if ruled is not None: return ruled if len(self.data_column) == 1: return [ROLES['data']] * count return list(self.sns_palette[:count])
[docs] def preprocess_data(self): """Return a new DataFrame aggregated to the configured representation. Drops rows with NaN in the grouping or data columns, aggregates the data columns with ``summary_func`` per well (``'prc'``) or per plate (``'plateID'``, split out of ``prc`` when needed) — or leaves them per object — and makes the grouping column an ordered Categorical. :returns: The preprocessed DataFrame; ``__init__`` assigns it back to ``self.df`` rather than the frame being modified in place. :raises KeyError: if ``representation='plate'`` and neither a ``plateID`` nor a ``prc`` column is available. :raises ValueError: if ``representation`` is not ``'object'``, ``'well'`` or ``'plate'``. """ df = self.df.dropna(subset=[self.grouping_column] + self.data_column) if self.representation == 'object': group_cols = None elif self.representation == 'well': group_cols = ['prc', self.grouping_column] elif self.representation == 'plate': if 'plateID' not in df.columns: if 'prc' in df.columns: df[['plateID', 'rowID', 'columnID']] = df['prc'].str.split('_', expand=True) else: raise KeyError( "Representation is 'plateID', but no 'plateID' column found. " "Also cannot split from 'prc' because 'prc' column is missing." ) if self.grouping_column == 'plateID': group_cols = ['plateID'] else: group_cols = ['plateID', self.grouping_column] else: raise ValueError(f"Unknown representation: {self.representation}, use object, well, or plate") if group_cols is not None: df = df.groupby( group_cols, observed=False)[self.data_column].agg( self.summary_func).reset_index() if self.order and (self.grouping_column in df.columns): df[self.grouping_column] = pd.Categorical( df[self.grouping_column], categories=self.order, ordered=True ) else: df[self.grouping_column] = pd.Categorical( df[self.grouping_column], categories=sorted(df[self.grouping_column].unique()), ordered=True ) return df
[docs] def remove_outliers_from_plot(self): """Remove outliers from the plot but keep them in the data.""" filtered_df = self.df.copy() unique_groups = filtered_df[self.grouping_column].unique() drop_index = pd.Index([]) for group in unique_groups: group_mask = filtered_df[self.grouping_column] == group for col in self.data_column: group_data = filtered_df.loc[group_mask, col] q1 = group_data.quantile(0.25) q3 = group_data.quantile(0.75) iqr = q3 - q1 lower_bound = q1 - 1.5 * iqr upper_bound = q3 + 1.5 * iqr outliers = group_mask & ((filtered_df[col] < lower_bound) | (filtered_df[col] > upper_bound)) drop_index = drop_index.union(filtered_df.index[outliers]) return filtered_df.drop(drop_index)
def _grouped_values(self, column, unique_groups): """``{group: finite values}`` for one measurement column. Cleaning is delegated to the engine's own ``_clean`` rather than repeated here. Two spellings of "which values count" is the same class of defect as two spellings of "which test applies". """ from .figures.stats import _clean return {group: _clean(self.df.loc[ self.df[self.grouping_column] == group, column]) for group in unique_groups}
[docs] def perform_normality_tests(self): """Evaluate normality for each requested column and group. Shapiro-Wilk results and the overall verdict come from :func:`spacr.figures.stats.check_normality`, including its Bonferroni correction across groups. Groups with fewer than three observations are reported as ``"Skipped"``. Checks below the configured information threshold report ``Informative=False`` rather than treating a failure to reject as evidence of normality. Returns ------- is_normal : bool ``True`` only when every requested column passes. results : list of dict Per-group test statistics, sample sizes, and verdicts. """ from .figures.stats import check_normality unique_groups = self.df[self.grouping_column].unique() normality_results = [] column_verdicts = [] for column in self.data_column: groups = self._grouped_values(column, unique_groups) for group, data in groups.items(): n_samples = int(data.size) if n_samples < 3: print(f"Skipping normality test for group '{group}' on " f"column '{column}' - not enough data.") normality_results.append({ 'Comparison': f'Normality test for {group} on {column}', 'Test Statistic': None, 'p-value': None, 'Test Name': 'Skipped', 'Column': column, 'n': n_samples, 'Informative': False, 'Verdict': (f'{n_samples} observations, too few to run ' f'a normality test at all'), }) continue check = check_normality([data]) normality_results.append({ 'Comparison': f'Normality test for {group} on {column}', 'Test Statistic': check.statistic, 'p-value': check.p_value, 'Test Name': check.name, 'Column': column, 'n': n_samples, 'Informative': check.informative, 'Verdict': check.verdict, }) column_verdicts.append( check_normality(list(groups.values())).passed) is_normal = bool(column_verdicts) and all(column_verdicts) return is_normal, normality_results
[docs] def perform_levene_test(self, unique_groups): """Levene's test for equal variance on ``data_column[0]``, MEDIAN-centred. Delegates to :func:`spacr.figures.stats.check_equal_variance`. Two things moved when it did, and both change the number a caller writes into a CSV: * The centring is the median (Brown-Forsythe), not SciPy's default mean. Median centring is less sensitive to non-normal data, and this function is called before the normality verdict is known. * Below :data:`spacr.figures.stats.MIN_N_FOR_ASSUMPTIONS` observations in the smallest group the result is ``(nan, nan)``. On three replicates Levene has almost no power, so "p = 0.7, variances are equal" means "we could not tell", and printing 0.7 into a results table invites exactly the reading that publishes a difference that is not there. :param unique_groups: Groups to compare. :returns: ``(statistic, p_value)``, both NaN when the check had no power. """ from .figures.stats import check_equal_variance groups = self._grouped_values(self.data_column[0], unique_groups) check = check_equal_variance(list(groups.values())) return check.statistic, check.p_value
#: The parametric test each rank test replaces, so a caller who has #: already decided the data are not normal can only make the choice MORE #: conservative, never less. Read by :meth:`perform_statistical_tests`. RANK_EQUIVALENT = { "Student's t": 'Mann-Whitney U', "Welch's t": 'Mann-Whitney U', 'paired t': 'Wilcoxon signed-rank', 'one-way ANOVA': 'Kruskal-Wallis', "Welch's ANOVA": 'Kruskal-Wallis', }
[docs] def perform_statistical_tests(self, unique_groups, is_normal): """Run one supported group comparison per data column. Parameters ---------- unique_groups : sequence Groups to compare. Two groups produce a pairwise test; larger sets produce an omnibus test. is_normal : bool External normality verdict. ``False`` forces a rank test; ``True`` still requires informative engine-level assumption checks. Returns ------- list of dict Test name, statistic, p-value, sample counts, effect size, and selection rationale for each data column. Untestable comparisons use ``Test Name='not testable'`` and include the reason. Notes ----- Test selection is delegated to :func:`spacr.figures.stats.compare`. Two-group comparisons may use Student's t, Welch's t, or Mann-Whitney U; larger comparisons may use one-way ANOVA, Welch's ANOVA, or Kruskal-Wallis. Paired data use the paired t-test or Wilcoxon signed-rank test. """ from .figures.stats import compare from .sp_stats import _ENGINE_TEST_NAMES labels = [str(group) for group in unique_groups] comparison_label = ' vs '.join(labels[:2]) or 'no groups' test_results = [] for column in self.data_column: groups = self._grouped_values(column, unique_groups) arrays = list(groups.values()) refusal = None if self.paired and len(arrays) == 2: first, second = arrays if first.size != second.size: refusal = (f'paired groups of {first.size} and ' f'{second.size} cannot be matched up') elif first.size and float(np.ptp(first - second)) == 0.0: refusal = ('the matched differences have no spread at ' 'all, so a paired test has no standard error ' 'to divide by') result = None if refusal is None: try: result = compare(groups, paired=self.paired) except ValueError as engine_refusal: refusal = str(engine_refusal) if result is not None and not is_normal: rank_test = self.RANK_EQUIVALENT.get(result.test) if rank_test is not None: result = compare(groups, paired=self.paired, force=rank_test) if result is None: test_name = 'not testable' stat, p = np.nan, np.nan effect, effect_name, why = np.nan, '', refusal else: test_name = _ENGINE_TEST_NAMES.get(result.test, result.test) stat, p = result.statistic, result.p_value effect, effect_name = result.effect_size, result.effect_name why = result.reason test_results.append({ 'Comparison': f'{comparison_label} ({column})', 'Test Statistic': stat, 'p-value': p, 'Test Name': test_name, 'Effect Size': effect, 'Effect': effect_name, 'Why This Test': why, 'Column': column, 'n_object': sum( len(self.raw_df[self.raw_df[self.grouping_column] == group] [column].dropna()) for group in unique_groups), 'n_well': sum( len(self.df[self.df[self.grouping_column] == group]) for group in unique_groups)}) return test_results
[docs] def perform_posthoc_tests(self, is_normal, unique_groups): """Perform post-hoc tests for multiple groups based on all_to_all flag. :param is_normal: Outcome of the normality check, which selects the family of test: True runs Tukey HSD, False runs Dunn's test with an automatically chosen p-adjustment. It only matters when post-hoc testing runs at all — see ``unique_groups``. It must be the verdict :meth:`perform_normality_tests` returned, which is :func:`spacr.figures.stats.check_normality`'s. A hand-computed one puts the omnibus test and the pairwise tests on different footing — Kruskal-Wallis across the groups followed by Tukey between them is two different assumptions about one dataset — and it is how the power floor gets bypassed: three replicates buy Dunn's, not Tukey. :param unique_groups: The distinct group labels. Only its *length* is read; the comparisons themselves are rebuilt from ``self.df[self.grouping_column]``, so reordering or renaming entries has no effect. Fewer than three groups returns an empty list, as does ``self.all_to_all`` being False, because pairwise correction is meaningless for a single comparison. :returns: A list of per-comparison dicts with ``Comparison``, ``Test Statistic`` (always ``None`` — neither test reports one), ``p-value``, ``Test Name`` and the ``n_object`` / ``n_well`` counts; empty when no post-hoc test was warranted. Only ``self.data_column[0]`` is tested, so extra data columns are ignored here. """ from .sp_stats import choose_p_adjust_method posthoc_results = [] if is_normal and len(unique_groups) > 2 and self.all_to_all: tukey_result = pairwise_tukeyhsd(self.df[self.data_column[0]], self.df[self.grouping_column], alpha=0.05) posthoc_results = [] for comparison, p_value in zip(tukey_result._results_table.data[1:], tukey_result.pvalues): raw_data1 = self.raw_df[self.raw_df[self.grouping_column] == comparison[0]][self.data_column] raw_data2 = self.raw_df[self.raw_df[self.grouping_column] == comparison[1]][self.data_column] posthoc_results.append({ 'Comparison': f'{comparison[0]} vs {comparison[1]}', 'Test Statistic': None, 'p-value': p_value, 'Test Name': 'Tukey HSD Post-hoc', 'n_object': len(raw_data1) + len(raw_data2), 'n_well': len(self.df[self.df[self.grouping_column] == comparison[0]]) + len(self.df[self.df[self.grouping_column] == comparison[1]])}) return posthoc_results elif len(unique_groups) > 2 and self.all_to_all: print('performing_dunns') long_data = self.df[[self.data_column[0], self.grouping_column]].dropna() p_adjust_method = choose_p_adjust_method(num_groups=len(long_data[self.grouping_column].unique()),num_data_points=len(long_data) // len(long_data[self.grouping_column].unique())) dunn_result = sp.posthoc_dunn( long_data, val_col=self.data_column[0], group_col=self.grouping_column, p_adjust=p_adjust_method ) for group_a, group_b in zip(*np.triu_indices_from(dunn_result, k=1)): raw_data1 = self.raw_df[self.raw_df[self.grouping_column] == dunn_result.index[group_a]][self.data_column] raw_data2 = self.raw_df[self.raw_df[self.grouping_column] == dunn_result.columns[group_b]][self.data_column] posthoc_results.append({ 'Comparison': f"{dunn_result.index[group_a]} vs {dunn_result.columns[group_b]}", 'Test Statistic': None, 'p-value': dunn_result.iloc[group_a, group_b], 'Test Name': "Dunn's Post-hoc", 'p_adjust_method': p_adjust_method, 'n_object': len(raw_data1) + len(raw_data2), 'n_well': len(self.df[self.df[self.grouping_column] == dunn_result.index[group_a]]) + len(self.df[self.df[self.grouping_column] == dunn_result.columns[group_b]])}) return posthoc_results return posthoc_results
[docs] def create_plot(self, ax=None): """Build the plot for the chosen graph type onto ``self.fig``. Nothing is displayed: retrieve the figure with :meth:`get_figure` (and the statistics with :meth:`get_results`), or call ``plt.show()``. :param ax: Existing ``Axes`` to draw into, for placing this graph in a panel of a larger figure. ``self.fig`` is then set to that axes' parent figure, so a later ``save=True`` writes the whole enclosing figure, not this panel alone. ``None`` creates a fresh figure sized from the group count and ``bar_width`` — and note that with a single ``data_column`` the standardisation pass still calls ``ax.figure.set_size_inches``, which resizes a shared figure underneath its other panels. """ def _generate_tabels(unique_groups): """Generate row labels and a symbol table for multi-level grouping.""" row_labels = [self.grouping_column] + self.data_column table_data = [] grouping_row = [] for _ in self.data_column: for group in unique_groups: grouping_row.append(group) table_data.append(grouping_row) for column in self.data_column: column_row = [] for data_col in self.data_column: for group in unique_groups: if column == data_col: column_row.append('+') else: column_row.append('-') table_data.append(column_row) transposed_table = list(map(list, zip(*table_data))) return row_labels, transposed_table def _place_symbols(row_labels, transposed_table, x_positions, ax): """ Places symbols and row labels aligned under the bars or jitter points on the graph. Parameters: - row_labels: List of row titles to be displayed along the y-axis. - transposed_table: Data to be placed under each bar/jitter as symbols. - x_positions: X-axis positions for each group to align the symbols. - ax: The matplotlib Axes object where the plot is drawn. """ y_axis_min = ax.get_ylim()[0] symbol_start_y = y_axis_min - 0.05 * (ax.get_ylim()[1] - y_axis_min) y_spacing = 0.04 label_x_pos = ax.get_xlim()[0] - 0.3 for row_idx, title in enumerate(row_labels): y_pos = symbol_start_y - (row_idx * y_spacing) ax.text(label_x_pos, y_pos, title, ha='right', va='center', fontsize=TYPE_SCALE['annotation'], fontweight='regular') for idx, (x_pos, column_data) in enumerate(zip(x_positions, transposed_table)): for row_idx, text in enumerate(column_data): y_pos = symbol_start_y - (row_idx * y_spacing) ax.text(x_pos, y_pos, text, ha='center', va='center', fontsize=TYPE_SCALE['annotation'], fontweight='regular') ax.figure.canvas.draw() def _get_positions(self, ax): """Return plotted group centers in left-to-right table order.""" if self.graph_type in ['bar','jitter_bar']: x_positions = [np.mean(bar.get_paths()[0].vertices[:, 0]) for bar in ax.collections if hasattr(bar, 'get_paths')] elif self.graph_type == 'violin': x_positions = [np.mean(violin.get_paths()[0].vertices[:, 0]) for violin in ax.collections if hasattr(violin, 'get_paths')] elif self.graph_type in ['box', 'jitter_box']: x_positions = sorted({line.get_xdata().mean() for line in ax.lines if line.get_linestyle() == '-'}) elif self.graph_type == 'jitter': x_positions = [np.mean(collection.get_offsets()[:, 0]) for collection in ax.collections if collection.get_offsets().size > 0] else: x_positions = [] return x_positions stats_df = self.df self.df_melted = pd.melt(stats_df, id_vars=[self.grouping_column], value_vars=self.data_column,var_name='Data Column', value_name='Value') unique_groups = stats_df[self.grouping_column].unique() is_normal, normality_results = self.perform_normality_tests() test_results = self.perform_statistical_tests(unique_groups, is_normal) posthoc_results = self.perform_posthoc_tests(is_normal, unique_groups) self.results_df = pd.DataFrame(normality_results + test_results + posthoc_results) if self.remove_outliers: self.df = self.remove_outliers_from_plot() self.results_df['outliers_removed_from_plot_only'] = True trimmed = len(stats_df) - len(self.df) if trimmed > 0: print(f"remove_outliers: {trimmed} of {len(stats_df)} points " f"are hidden from the plot. THE STATISTICS ABOVE USED " f"ALL {len(stats_df)}.") self.df_melted = pd.melt( self.df, id_vars=[self.grouping_column], value_vars=self.data_column, var_name='Data Column', value_name='Value') num_groups = len(self.df[self.grouping_column].unique()) self.bar_width = 0.4 spacing_between_groups = self.bar_width/0.5 self.fig_width = (num_groups * self.bar_width) + (spacing_between_groups * num_groups) self.fig_height = self.fig_width/2 if self.graph_type in ['line','line_std']: self.fig_height, self.fig_width = 10, 10 with figure_style(theme_target()): if ax is None: self.fig, ax = plt.subplots(figsize=(self.fig_height, self.fig_width)) from .figures.bundle import _register_figure_data _register_figure_data(self.fig, self.df, x=self.grouping_column, y=self.data_column[0] if self.data_column else "", kind=str(getattr(self, "graph_type", "") or "")) else: self.fig = ax.figure if len(self.data_column) == 1: self.hue=self.grouping_column self.jitter_bar_dodge = False else: self.hue='Data Column' self.jitter_bar_dodge = True if self.graph_type == 'bar': self._create_bar_plot(ax) elif self.graph_type == 'jitter': self._create_jitter_plot(ax) elif self.graph_type == 'box': self._create_box_plot(ax) elif self.graph_type == 'violin': self._create_violin_plot(ax) elif self.graph_type == 'jitter_box': self._create_jitter_box_plot(ax) elif self.graph_type == 'jitter_bar': self._create_jitter_bar_plot(ax) elif self.graph_type == 'line': self._create_line_graph(ax) elif self.graph_type == 'line_std': self._create_line_with_std_area(ax) else: raise ValueError(f"Unknown graph type: {self.graph_type}") if len(self.data_column) == 1: num_groups = len(self.df[self.grouping_column].unique()) self._standerdize_figure_format(ax=ax, num_groups=num_groups, graph_type=self.graph_type) if isinstance(self.y_lim, list): if len(self.y_lim) == 2: ax.set_ylim(self.y_lim[0], self.y_lim[1]) elif len(self.y_lim) == 1: ax.set_ylim(self.y_lim[0], None) sns.despine(ax=ax, top=True, right=True) handles, labels = ax.get_legend_handles_labels() if handles: ax.legend( handles, labels, loc='center left', bbox_to_anchor=(1, 0.5), title='Data Column') if not self.graph_type in ['line','line_std']: ax.set_xlabel('') x_positions = _get_positions(self, ax) if len(self.data_column) == 1 and not self.graph_type in ['line','line_std']: legend = ax.get_legend() if legend is not None: legend.remove() rotate_ticks(ax) elif len(self.data_column) > 1 and not self.graph_type in ['line','line_std']: ax.set_xticks([]) ax.tick_params(bottom=False) ax.set_xticklabels([]) legend_ax = self.fig.add_axes([0.1, -0.2, 0.62, 0.2]) legend_ax.set_axis_off() row_labels, table_data = _generate_tabels(unique_groups) _place_symbols(row_labels, table_data, x_positions, ax) if self.annotate_stats: self._draw_comparison_lines(ax) if self.save: self._save_results() ax.margins(x=0.12)
def _comparison_pairs(self): """The pairwise rows of ``results_df``, as ``(group_a, group_b, p)``. ``results_df`` mixes three row shapes: a per-group normality row (``'Normality test for X on Y'``), one omnibus/two-group row per data column (``'A vs B (column)'``) and the post-hoc rows (``'A vs B'``). Only the last two name two groups, and a bracket can only be drawn between two groups -- reading a normality row as a comparison is how the old annotation pass died on its first row. A trailing ``(column)`` is stripped only when it names one of this plot's data columns, so a group genuinely called ``'x (y)'`` keeps its name. :returns: list of ``(group_a, group_b, p_value)``, in table order. """ pairs = [] if self.results_df.empty or 'Comparison' not in self.results_df: return pairs suffixes = tuple(f' ({column})' for column in self.data_column) for _index, row in self.results_df.iterrows(): label = str(row.get('Comparison', '')) parts = label.split(' vs ') if len(parts) != 2: continue first, second = parts[0].strip(), parts[1].strip() for suffix in suffixes: if second.endswith(suffix): second = second[:-len(suffix)].strip() break pairs.append((first, second, row.get('p-value'))) return pairs @staticmethod def _tick_positions(ax): """``{drawn group label: x}``, read off the axis's own ticks. THE AXIS IS ASKED, NOT THE FRAME. The groups are drawn in the order of the ordered Categorical `preprocess_data` builds, which is not the order ``DataFrame.unique()`` returns them in -- indexing a list of x positions by a group's position in ``unique()`` put brackets over the wrong pair whenever the two orders disagreed, which is silent and looks exactly like a real result. """ labels = [text.get_text() for text in ax.get_xticklabels()] return {label: float(x) for label, x in zip(labels, list(ax.get_xticks())) if label} def _draw_comparison_lines(self, ax): """Bracket each pairwise comparison and mark it with its asterisks. Drawn only for a single data column: with several columns the x axis is one position per (column, group) pair and the symbol table below it already says which is which, so a bracket has no unambiguous pair of ends to sit on. A comparison naming a group that is not on this axis is skipped rather than guessed at. The stack sits above the data, and the top of the view is raised to make room for it UNLESS the caller pinned ``y_lim``, which is an instruction about the window and wins. :param ax: the axes the plot was drawn on. :returns: how many brackets were drawn. """ pairs = self._comparison_pairs() if not pairs: print("No comparisons available to annotate.") return 0 if len(self.data_column) != 1 or self.graph_type in ('line', 'line_std'): return 0 positions = self._tick_positions(ax) drawable = [(a, b, p) for a, b, p in pairs if a in positions and b in positions and _finite_p_value(p) is not None] if not drawable: print("No comparisons available to annotate.") return 0 bottom, top = ax.get_ylim() data_top = ax.dataLim.y1 if not np.isfinite(data_top): data_top = top base = min(float(data_top), float(top)) step = 0.08 * (float(top) - float(bottom) or 1.0) ink = resolve_ink(theme_target()) for index, (first, second, p_value) in enumerate(drawable): x1, x2 = positions[first], positions[second] line_y = base + step * (index + 1) tick = step * 0.15 ax.plot([x1, x1, x2, x2], [line_y - tick, line_y, line_y, line_y - tick], lw=WEIGHTS['spine'], c=ink) ax.text((x1 + x2) / 2, line_y, _significance_marker(p_value), ha='center', va='bottom', fontsize=TYPE_SCALE['annotation']) if self.y_lim is None: ax.set_ylim(bottom, base + step * (len(drawable) + 1)) return len(drawable) def _standerdize_figure_format(self, ax, num_groups, graph_type): """ Adjusts the figure layout (size, bar width, jitter, and spacing) based on the number of groups. Parameters: - ax: The matplotlib Axes object. - num_groups: Number of unique groups. - graph_type: The type of graph (e.g., 'bar', 'jitter', 'box', etc.). Returns: - None. Modifies the figure and Axes in place. """ if graph_type in ['line', 'line_std']: print("Skipping layout adjustment for line graphs.") return correction_factor = 4 fig_size = max(6, num_groups * 2) / correction_factor if fig_size < 10: fig_size = 10 ax.figure.set_size_inches(fig_size, fig_size) bar_width = min(0.8, 1.5 / num_groups) / correction_factor jitter_amount = min(0.1, 0.2 / num_groups) / correction_factor jitter_size = max(50 / num_groups, 200) ax.set_xlim(-0.5, num_groups - 0.5) rotate_ticks(ax) if graph_type == 'bar': for bar in ax.patches: bar.set_width(bar_width) bar.set_x(bar.get_x() - bar_width / 2) elif graph_type in ['jitter', 'jitter_bar', 'jitter_box']: for coll in ax.collections: offsets = coll.get_offsets() offsets[:, 0] += jitter_amount coll.set_offsets(offsets) coll.set_sizes([jitter_size] * len(offsets)) elif graph_type in ['box', 'violin']: for artist in ax.artists: artist.set_width(bar_width) ax.tick_params(axis='x', labelsize=max(10, 15 - num_groups // 2)) ax.tick_params(axis='y', labelsize=max(10, 15 - num_groups // 2)) if ax.get_legend(): ax.get_legend().set_bbox_to_anchor((1.05, 1)) ax.get_legend().prop.set_size(max(8, 12 - num_groups // 3)) ax.figure.canvas.draw() def _create_bar_plot(self, ax): """Helper method to create a bar plot with consistent bar thickness and centered error bars.""" if len(self.data_column) > 1: self.df_melted['Combined Group'] = (self.df_melted[self.grouping_column].astype(str) + " - " + self.df_melted['Data Column'].astype(str)) x_axis_column = 'Combined Group' hue = None plot_order = [f"{g} - {c}" for g in self.order for c in self.data_column] ax.set_ylabel('Value') else: x_axis_column = self.grouping_column ax.set_ylabel(self.data_column[0]) hue = self.hue plot_order = self.order plot_hue = hue or x_axis_column plot_palette = self._plot_palette(len(plot_order)) show_legend = hue is not None summary_df = self.df_melted.groupby( [x_axis_column], observed=False ).agg(mean=('Value', 'mean'), std=('Value', 'std'), sem=('Value', 'sem'), count=('Value', 'count')).reset_index() self.summary_df = summary_df.copy() sns.barplot( data=self.df_melted, x=x_axis_column, y='Value', hue=plot_hue, palette=plot_palette, legend=show_legend, ax=ax, dodge=self.jitter_bar_dodge, errorbar=None, order=plot_order) if len(self.data_column) > 1: bars = [bar for bar in ax.patches if isinstance(bar, plt.Rectangle)] target_width = self.bar_width * 2 for bar in bars: bar.set_width(target_width) bar.set_x(bar.get_x() - target_width / 2) bars = [bar for bar in ax.patches if isinstance(bar, plt.Rectangle)] for bar, (_, row) in zip(bars, summary_df.iterrows()): x_bar = bar.get_x() + bar.get_width() / 2 err = self._error_half_width(row) if err is None: continue capsize = float(self._user_style().get("error_capsize", 5)) ax.errorbar(x=x_bar, y=bar.get_height(), yerr=err, fmt='none', c=resolve_ink(theme_target()), capsize=capsize, lw=WEIGHTS['data']) ax.set_xlabel(self.grouping_column) if self.log_y: ax.set_yscale('log') if self.log_x: ax.set_xscale('log') def _create_jitter_plot(self, ax): """Helper method to create a jitter plot (strip plot) with consistent spacing.""" if len(self.data_column) > 1: self.df_melted['Combined Group'] = (self.df_melted[self.grouping_column].astype(str) + " - " + self.df_melted['Data Column'].astype(str)) x_axis_column = 'Combined Group' hue = None plot_order = [f"{g} - {c}" for g in self.order for c in self.data_column] ax.set_ylabel('Value') else: x_axis_column = self.grouping_column ax.set_ylabel(self.data_column[0]) hue = self.hue plot_order = self.order plot_hue = hue or x_axis_column plot_palette = self._plot_palette(len(plot_order)) show_legend = hue is not None self.summary_df = self.df_melted.copy() alpha, size = self._point_look(16) sns.stripplot( data=self.df_melted, x=x_axis_column, y='Value', hue=plot_hue, palette=plot_palette, legend=show_legend, dodge=self.jitter_bar_dodge, jitter=self._jitter(), ax=ax, alpha=alpha, size=size, order=plot_order) ax.set_xlabel(self.grouping_column) handles, labels = ax.get_legend_handles_labels() unique_labels = dict(zip(labels, handles)) if unique_labels: ax.legend(unique_labels.values(), unique_labels.keys(), loc='best') if self.log_y: ax.set_yscale('log') if self.log_x: ax.set_xscale('log') def _create_line_graph(self, ax): """Helper method to create a line graph with one line per group. TWO SHAPES, because "line" means two different pictures here. With two data columns it is the one this was written for -- epochs against accuracy, one line per group. With ONE data column there is no second column to put on x, and the group itself is the x axis: a line across the groups, which is what every other graph type in the menu draws the groups as. The one-column case used to index ``data_column[1]`` regardless and raise IndexError, which `create_grouped_plot` swallowed -- so choosing "Line" on an ordinary grouped plot returned an EMPTY figure and said nothing. Same for a caller asking for graph_type='line' directly. """ if len(self.data_column) < 2: self._create_line_across_groups(ax) return x_axis_column = self.data_column[0] y_axis_column = self.data_column[1] if self.log_y: self.df[y_axis_column] = np.log10(self.df[y_axis_column]) if self.log_x: self.df[x_axis_column] = np.log10(self.df[x_axis_column]) hue = self.grouping_column required_columns = [x_axis_column, y_axis_column, self.grouping_column] for col in required_columns: if col not in self.df.columns: raise ValueError(f"Column '{col}' not found in DataFrame.") self.summary_df = self.df.copy() line_palette = self._plot_palette( self.df[hue].nunique(dropna=True)) sns.lineplot( data=self.df, x=x_axis_column, y=y_axis_column, hue=hue, palette=line_palette, ax=ax, marker='o', linewidth=1, markersize=6) ax.set_xlabel(f"{x_axis_column}") ax.set_ylabel(f"{y_axis_column}") def _create_line_across_groups(self, ax): """A line over the groups, for the single-data-column case. The point on each group is the summary the rest of the panel uses (:attr:`summary_func`), so the line agrees with what the bar chart of the same data would show, and the error bar is the same spread. """ value_column = self.data_column[0] frame = self.df_melted if getattr(self, "df_melted", None) is not None \ else self.df column = 'Value' if 'Value' in getattr(frame, "columns", []) \ else value_column order = list(self.order) if self.order else \ list(dict.fromkeys(frame[self.grouping_column])) summary = (frame.groupby([self.grouping_column], observed=False) .agg(centre=(column, self.summary_func or 'mean'), spread=(column, 'std')) .reindex(order).reset_index()) self.summary_df = summary colour = self._plot_palette(1) colour = colour[0] if colour is not None and len(colour) else None ax.errorbar(range(len(summary)), summary['centre'], yerr=summary['spread'].fillna(0.0), marker='o', linewidth=1, markersize=6, capsize=3, color=colour) ax.set_xticks(range(len(summary))) ax.set_xticklabels([str(v) for v in summary[self.grouping_column]]) ax.set_xlabel(str(self.grouping_column)) ax.set_ylabel(str(value_column)) if self.log_y: ax.set_yscale('log') def _create_line_with_std_area(self, ax): """Helper method to create a line graph with shaded area representing standard deviation.""" x_axis_column = self.data_column[0] y_axis_column = self.data_column[1] y_axis_column_mean = f"mean_{y_axis_column}" y_axis_column_std = f"std_{y_axis_column_mean}" if self.log_y: self.df[y_axis_column] = np.log10(self.df[y_axis_column]) if self.log_x: self.df[x_axis_column] = np.log10(self.df[x_axis_column]) summary_df = self.df.pivot_table(index=x_axis_column,values=y_axis_column,aggfunc=['mean', 'std']).reset_index() summary_df.columns = [x_axis_column, y_axis_column_mean, y_axis_column_std] self.summary_df = summary_df.copy() sns.lineplot(data=summary_df,x=x_axis_column,y=y_axis_column_mean,ax=ax,marker='o',linewidth=WEIGHTS['data'],markersize=0,color=ROLES['highlight'],label=y_axis_column_mean) ax.fill_between(summary_df[x_axis_column],summary_df[y_axis_column_mean] - summary_df[y_axis_column_std],summary_df[y_axis_column_mean] + summary_df[y_axis_column_std],color=ROLES['highlight'], alpha=0.25 ) ax.set_xlabel(f"{x_axis_column}") ax.set_ylabel(f"{y_axis_column}") def _create_box_plot(self, ax): """Helper method to create a box plot with consistent spacing.""" if len(self.data_column) > 1: self.df_melted['Combined Group'] = (self.df_melted[self.grouping_column].astype(str) + " - " + self.df_melted['Data Column'].astype(str)) x_axis_column = 'Combined Group' hue = None plot_order = [f"{g} - {c}" for g in self.order for c in self.data_column] ax.set_ylabel('Value') else: x_axis_column = self.grouping_column ax.set_ylabel(self.data_column[0]) hue = self.hue plot_order = self.order plot_hue = hue or x_axis_column plot_palette = self._plot_palette(len(plot_order)) show_legend = hue is not None self.summary_df = self.df_melted.copy() sns.boxplot( data=self.df_melted, x=x_axis_column, y='Value', hue=plot_hue, palette=plot_palette, legend=show_legend, ax=ax, order=plot_order) ax.set_xlabel(self.grouping_column) handles, labels = ax.get_legend_handles_labels() unique_labels = dict(zip(labels, handles)) if unique_labels: ax.legend(unique_labels.values(), unique_labels.keys(), loc='best') if self.log_y: ax.set_yscale('log') if self.log_x: ax.set_xscale('log') def _create_violin_plot(self, ax): """Helper method to create a violin plot with consistent spacing.""" if len(self.data_column) > 1: self.df_melted['Combined Group'] = (self.df_melted[self.grouping_column].astype(str) + " - " + self.df_melted['Data Column'].astype(str)) x_axis_column = 'Combined Group' hue = None plot_order = [f"{g} - {c}" for g in self.order for c in self.data_column] ax.set_ylabel('Value') else: x_axis_column = self.grouping_column ax.set_ylabel(self.data_column[0]) hue = self.hue plot_order = self.order plot_hue = hue or x_axis_column plot_palette = self._plot_palette(len(plot_order)) show_legend = hue is not None self.summary_df = self.df_melted.copy() sns.violinplot( data=self.df_melted, x=x_axis_column, y='Value', hue=plot_hue, palette=plot_palette, legend=show_legend, ax=ax, order=plot_order) ax.set_xlabel(self.grouping_column) ax.set_ylabel('Value') handles, labels = ax.get_legend_handles_labels() unique_labels = dict(zip(labels, handles)) if unique_labels: ax.legend(unique_labels.values(), unique_labels.keys(), loc='best') if self.log_y: ax.set_yscale('log') if self.log_x: ax.set_xscale('log') def _create_jitter_bar_plot(self, ax): """Helper method to create a bar plot with consistent bar thickness and centered error bars.""" if len(self.data_column) > 1: self.df_melted['Combined Group'] = (self.df_melted[self.grouping_column].astype(str) + " - " + self.df_melted['Data Column'].astype(str)) x_axis_column = 'Combined Group' hue = None plot_order = [f"{g} - {c}" for g in self.order for c in self.data_column] ax.set_ylabel('Value') else: x_axis_column = self.grouping_column ax.set_ylabel(self.data_column[0]) hue = self.hue plot_order = self.order plot_hue = hue or x_axis_column plot_palette = self._plot_palette(len(plot_order)) show_legend = hue is not None summary_df = self.df_melted.groupby( [x_axis_column], observed=False ).agg(mean=('Value', 'mean'), std=('Value', 'std'), sem=('Value', 'sem')).reset_index() self.summary_df = summary_df sns.barplot( data=self.df_melted, x=x_axis_column, y='Value', hue=plot_hue, palette=plot_palette, legend=show_legend, ax=ax, dodge=self.jitter_bar_dodge, errorbar=None, order=plot_order) alpha, size = self._point_look(16) if self._points_overlaid(): sns.stripplot( data=self.df_melted, x=x_axis_column, y='Value', hue=plot_hue, palette=plot_palette, legend=show_legend, dodge=self.jitter_bar_dodge, jitter=self._jitter(), ax=ax, alpha=alpha, edgecolor='none', linewidth=0, size=size, order=plot_order) if len(self.data_column) > 1: bars = [bar for bar in ax.patches if isinstance(bar, plt.Rectangle)] target_width = self.bar_width * 2 for bar in bars: bar.set_width(target_width) bar.set_x(bar.get_x() - target_width / 2) ax.set_xlabel(self.grouping_column) if self.log_y: ax.set_yscale('log') if self.log_x: ax.set_xscale('log') def _create_jitter_box_plot(self, ax): """Helper method to create a box plot with consistent spacing.""" if len(self.data_column) > 1: self.df_melted['Combined Group'] = (self.df_melted[self.grouping_column].astype(str) + " - " + self.df_melted['Data Column'].astype(str)) x_axis_column = 'Combined Group' hue = None plot_order = [f"{g} - {c}" for g in self.order for c in self.data_column] ax.set_ylabel('Value') else: x_axis_column = self.grouping_column ax.set_ylabel(self.data_column[0]) hue = self.hue plot_order = self.order plot_hue = hue or x_axis_column plot_palette = self._plot_palette(len(plot_order)) show_legend = hue is not None self.summary_df = self.df_melted.copy() sns.boxplot( data=self.df_melted, x=x_axis_column, y='Value', hue=plot_hue, palette=plot_palette, legend=show_legend, ax=ax, order=plot_order) alpha, size = self._point_look(12) if self._points_overlaid(): sns.stripplot( data=self.df_melted, x=x_axis_column, y='Value', hue=plot_hue, palette=plot_palette, legend=show_legend, dodge=self.jitter_bar_dodge, jitter=self._jitter(), ax=ax, alpha=alpha, edgecolor='none', linewidth=0, size=size, order=plot_order) ax.set_xlabel(self.grouping_column) handles, labels = ax.get_legend_handles_labels() unique_labels = dict(zip(labels, handles)) if unique_labels: ax.legend(unique_labels.values(), unique_labels.keys(), loc='best') if self.log_y: ax.set_yscale('log') if self.log_x: ax.set_xscale('log') def _save_results(self): """Save figure, stats, and all data used to generate the plot.""" os.makedirs(self.output_dir, exist_ok=True) plot_path = os.path.join(self.output_dir, f"{self.results_name}.pdf") plot_path = save_figure(self.fig, plot_path, bbox_inches='tight', transparent=True) stats_path = os.path.join(self.output_dir, f"{self.results_name}_stats.csv") self.results_df.to_csv(stats_path, index=False) data_path = os.path.join(self.output_dir, f"{self.results_name}_data.csv") self.df.to_csv(data_path, index=False) if hasattr(self, 'summary_df') and self.summary_df is not None: data_path = os.path.join(self.output_dir, f"{self.results_name}_summary.csv") self.summary_df.to_csv(data_path, index=False) print(f"Data -> {data_path}") print(f"Plot -> {plot_path}") print(f"Stats -> {stats_path}")
[docs] def get_results(self): """Return the results dataframe.""" return self.results_df
[docs] def get_figure(self): """Return the generated figure.""" return self.fig
[docs] def plot_data_from_db(settings): """Read one or more measurement DBs, annotate conditions and render a ``spacrGraph`` plot. Concatenates results across source directories, derives the ``recruitment`` column if requested, drops missing rows, then hands the data to :class:`spacrGraph` for statistics + plotting. :param settings: Settings dict. See ``settings.set_default_plot_data_from_db`` for accepted keys (notably ``src``, ``database``, ``table_names``, ``data_column``, ``grouping_column``, ``graph_type``, ``graph_name``). :returns: The plotted DataFrame, or ``None`` when the requested data or grouping column is missing. :raises ValueError: if ``src`` is neither a string nor a list. """ from .io import _read_db, _read_and_merge_data from .utils import annotate_conditions, save_settings from .settings import set_default_plot_data_from_db """ Extracts the specified table from the SQLite database and plots a specified column. Args: db_path (str): The path to the SQLite database. table_names (str): The name of the table to extract. data_column (str): The column to plot from the table. Returns: df (pd.DataFrame): The extracted table as a DataFrame. """ settings = set_default_plot_data_from_db(settings) if isinstance(settings['src'], str): srcs = [settings['src']] elif isinstance(settings['src'], list): srcs = settings['src'] else: raise ValueError("src must be a string or a list of strings.") if isinstance(settings['database'], str): settings['database'] = [settings['database'] for _ in range(len(srcs))] settings['dst'] = os.path.join(srcs[0], 'results') save_settings(settings, name=f"{settings['graph_name']}_plot_settings_db", show=True) dfs = [] for i, src in enumerate(srcs): db_loc = os.path.join(src, 'measurements', settings['database'][i]) print(f"Database: {db_loc}") if settings['table_names'] in ['saliency_image_correlations']: print(f"Database table: {settings['table_names']}") [df1] = _read_db(db_loc, tables=[settings['table_names']]) else: df1, _ = _read_and_merge_data(locs=[db_loc], tables = settings['table_names'], verbose=settings['verbose'], nuclei_limit=settings['nuclei_limit'], pathogen_limit=settings['pathogen_limit']) dft = annotate_conditions(df1, cells=settings['cell_types'], cell_loc=settings['cell_plate_metadata'], pathogens=settings['pathogen_types'], pathogen_loc=settings['pathogen_plate_metadata'], treatments=settings['treatments'], treatment_loc=settings['treatment_plate_metadata']) dfs.append(dft) df = pd.concat(dfs, axis=0) df['prc'] = df['plateID'].astype(str) + '_' + df['rowID'].astype(str) + '_' + df['columnID'].astype(str) annotation_ledger = RunLedger('plot_data_from_db:annotation') for meta_key, column, label in ( ('cell_plate_metadata', 'host_cells', 'host_cell'), ('pathogen_plate_metadata', 'pathogen', 'pathogen'), ('treatment_plate_metadata', 'treatment', 'treatment')): if settings[meta_key] != None: try: df = df.dropna(subset=column) except Exception as e: print(f"Could not drop NaN values from '{label}' column: {e}") annotation_ledger.record_failure(column, stage='annotate_conditions', exc=e) raise_if_strict( f"{meta_key} was set but the {column!r} column was never " f"created, so rows cannot be filtered to the requested " f"conditions; every group in this plot is suspect.", exc=e, settings=settings) else: annotation_ledger.record_success(column, stage='annotate_conditions') annotation_ledger.finalize() if settings['data_column'] == 'recruitment': pahtogen_measurement = df[f"pathogen_channel_{settings['channel_of_interest']}_mean_intensity"] cytoplasm_measurement = df[f"cytoplasm_channel_{settings['channel_of_interest']}_mean_intensity"] df['recruitment'] = pahtogen_measurement / cytoplasm_measurement if settings['data_column'] not in df.columns: print(f"Data column {settings['data_column']} not found in DataFrame.") print(f'Please use one of the following columns:') for col in df.columns: print(col) display(df) return None df = df.dropna(subset=settings['data_column']) if settings['grouping_column'] not in df.columns: print(f"Grouping column {settings['grouping_column']} not found in DataFrame.") print(f'Please use one of the following columns:') for col in df.columns: print(col) display(df) return None df = df.dropna(subset=settings['grouping_column']) src = srcs[0] dst = os.path.join(src, 'results', settings['graph_name']) os.makedirs(dst, exist_ok=True) spacr_graph = spacrGraph( df=df, grouping_column=settings['grouping_column'], data_column=settings['data_column'], graph_type=settings['graph_type'], graph_name=settings['graph_name'], summary_func='mean', colors=None, output_dir=dst, save=settings['save'], y_lim=settings['y_lim'], error_bar_type='std', representation=settings['representation'], theme=settings['theme'], ) spacr_graph.create_plot() fig = spacr_graph.get_figure() plt.show() results_df = spacr_graph.get_results() return fig, results_df, df
[docs] def plot_data_from_csv(settings): """Load per-plate CSVs, filter/outlier-clean and render a ``spacrGraph`` plot. :param settings: Settings dict — see ``settings.get_plot_data_from_csv_default_settings`` for keys (``src``, ``data_column``, ``grouping_column``, ``keep_groups``, ``remove_outliers``, ``graph_type``, ``graph_name``, ...). :returns: ``(fig, results_df, df)`` — the figure, stats DataFrame and plotted DataFrame. :raises ValueError: if ``src`` is not a string or list. """ from .utils import remove_outliers_by_group """ Extracts the specified table from the SQLite database and plots a specified column. Args: db_path (str): The path to the SQLite database. table_names (str): The name of the table to extract. data_column (str): The column to plot from the table. Returns: df (pd.DataFrame): The extracted table as a DataFrame. """ def filter_rows_by_column_values(df: pd.DataFrame, column: str, values: list) -> pd.DataFrame: """Return a filtered DataFrame where only rows with the column value in the list are kept. :param df: Frame to filter; it is not modified, and the result is a ``.copy()`` so later assignment to it raises no ``SettingWithCopyWarning``. :param column: Column to test. Must exist, or ``KeyError`` is raised — here it is the caller's ``grouping_column``. :param values: Values to keep, matched with ``isin`` so comparison is exact and type-sensitive: the string ``'1'`` will not match an integer ``1`` read from the CSV. An empty list keeps nothing and yields an empty frame rather than passing everything through. :returns: A new filtered DataFrame. """ return df[df[column].isin(values)].copy() if isinstance(settings['src'], str): srcs = [settings['src']] elif isinstance(settings['src'], list): srcs = settings['src'] else: raise ValueError("src must be a string or a list of strings.") dfs = [] for i, src in enumerate(srcs): dft = pd.read_csv(src) if 'plateID' not in dft.columns: dft['plateID'] = f"plate{i+1}" dft['common'] = 'spacr' dfs.append(dft) df = pd.concat(dfs, axis=0) if 'prc' in df.columns: if not all(col in df.columns for col in ['plate', 'rowID', 'columnID']): try: df[['plateID', 'rowID', 'columnID']] = df['prc'].str.split('_', expand=True) except Exception as e: print(f"Could not split the prc column: {e}") raise_if_strict( "The 'prc' column could not be split into " "plateID/rowID/columnID; any grouping in this plot is " "computed on the wrong keys.", exc=e, settings=settings) if 'keep_groups' in settings.keys(): if isinstance(settings['keep_groups'], str): settings['keep_groups'] = [settings['keep_groups']] elif isinstance(settings['keep_groups'], list): df = filter_rows_by_column_values(df, settings['grouping_column'], settings['keep_groups']) if settings['remove_outliers']: df = remove_outliers_by_group(df, settings['grouping_column'], settings['data_column'], method='iqr', threshold=1.5) if settings['verbose']: display(df) df = df.dropna(subset=settings['data_column']) df = df.dropna(subset=settings['grouping_column']) src = srcs[0] dst = os.path.join(os.path.dirname(src), 'results', settings['graph_name']) os.makedirs(dst, exist_ok=True) spacr_graph = spacrGraph( df=df, grouping_column=settings['grouping_column'], data_column=settings['data_column'], graph_type=settings['graph_type'], graph_name=settings['graph_name'], summary_func='mean', colors=None, output_dir=dst, save=settings['save'], y_lim=settings['y_lim'], log_y=settings['log_y'], log_x=settings['log_x'], error_bar_type='std', representation=settings['representation'], theme=settings['theme'], ) spacr_graph.create_plot() fig = spacr_graph.get_figure() plt.show() results_df = spacr_graph.get_results() return fig, results_df
[docs] def plot_region(settings): """Render mask overlay, cropped PNG grid and activation-map grid for one FOV. Reads the FOV's merged NPY, resolves its PNG crops and activation maps from the measurements and activation DBs, and writes the three figures under ``<src>/results/<name>/`` when possible — in the configured figure format, so PDF only while that is the preference. :param settings: Settings dict with ``src``, ``name``, ``channels``, ``cell_channel``, ``nucleus_channel``, ``pathogen_channel``, ``percentiles``, ``activation_mode``, ``activation_db``, ``mode``, ``export_tiffs``. :returns: Tuple ``(fig_mask_overlay, fig_png_grid, fig_activation_grid)`` — any element may be ``None`` when the corresponding assets were not found. """ def _sort_paths_by_basename(paths): """Return ``paths`` sorted by their basename.""" return sorted(paths, key=lambda path: os.path.basename(path)) def save_figure_as_pdf(fig, path): """Save ``fig`` in the user's chosen figure format. Named for the format it used to hard-code; it follows the preference now, like every other figure the user keeps, and `save_figure` creates the parent directory itself. :param fig: Figure to write. It is left open, so the caller can still return it to the notebook after saving. :param path: Destination path. Its extension is rewritten to whichever format the preference selected, so passing a ``.pdf`` name does not force PDF; missing parent directories are created. The path actually written is printed, not returned. """ path = save_figure(fig, path, bbox_inches='tight') print(f"Saved {path}") from .io import _read_db from .utils import correct_paths fov_path = os.path.join(settings['src'], 'merged', settings['name']) name = os.path.splitext(settings['name'])[0] db_path = os.path.join(settings['src'], 'measurements', 'measurements.db') paths_df = _read_db(db_path, tables=['png_list'])[0] paths_df, _ = correct_paths(df=paths_df, base_path=settings['src'], folder='data') paths_df = paths_df[paths_df['png_path'].str.contains(name, na=False)] activation_mode = f"{settings['activation_mode']}_list" activation_db_path = os.path.join(settings['src'], 'measurements', settings['activation_db']) activation_paths_df = _read_db(activation_db_path, tables=[activation_mode])[0] activation_db = os.path.splitext(settings['activation_db'])[0] base_path=os.path.join(settings['src'], 'datasets',activation_db) activation_paths_df, _ = correct_paths(df=activation_paths_df, base_path=base_path, folder=settings['activation_mode']) activation_paths_df = activation_paths_df[activation_paths_df['png_path'].str.contains(name, na=False)] png_paths = _sort_paths_by_basename(paths_df['png_path'].tolist()) activation_paths = _sort_paths_by_basename(activation_paths_df['png_path'].tolist()) if activation_paths: fig_3 = plot_image_grid(image_paths=activation_paths, percentiles=settings['percentiles']) else: fig_3 = None print(f"Could not find any cropped PNGs") if png_paths: fig_2 = plot_image_grid(image_paths=png_paths, percentiles=settings['percentiles']) else: fig_2 = None print(f"Could not find any activation maps") print('fov_path', fov_path) fig_1 = plot_image_mask_overlay(file=fov_path, channels=settings['channels'], cell_channel=settings['cell_channel'], nucleus_channel=settings['nucleus_channel'], pathogen_channel=settings['pathogen_channel'], figuresize=10, percentiles=settings['percentiles'], thickness=3, save_pdf=True, mode=settings['mode'], export_tiffs=settings['export_tiffs']) dst = os.path.join(settings['src'], 'results', name) if not fig_1 == None: save_figure_as_pdf(fig_1, os.path.join(dst, f"{name}_mask_overlay.pdf")) if not fig_2 == None: save_figure_as_pdf(fig_2, os.path.join(dst, f"{name}_png_grid.pdf")) if not fig_3 == None: save_figure_as_pdf(fig_3, os.path.join(dst, f"{name}_activation_grid.pdf")) return fig_1, fig_2, fig_3
[docs] def plot_image_grid(image_paths, percentiles): """Render a square grid of percentile-normalised images with a black background. Each tile carries its source file and the per-channel display range it was stretched to, so a checked export (see :func:`save_figure`) can say when tiles are scaled differently and write a provenance sidecar that rebuilds every tile from its file. :param image_paths: Image files to display; extra tiles are filled black. :param percentiles: Two-element percentile pair used to normalise each channel. :returns: The generated ``Figure``. """ from PIL import Image import matplotlib.pyplot as plt import math N = len(image_paths) grid_size = math.ceil(math.sqrt(N)) with figure_style(theme_target()): fig, axs = plt.subplots( grid_size, grid_size, figsize=(grid_size * 2, grid_size * 2), facecolor='black', squeeze=False ) from .figures.bundle import _register_figure_data _register_figure_data(fig, None, kind="montage", files=[str(v) for v in image_paths]) axs = axs.flatten() for i, img_path in enumerate(image_paths): ax = axs[i] with Image.open(img_path) as opened: raw = np.array(opened) sensor_range = _source_sensor_range(opened, raw) stretched, ranges = _percentile_display(raw, percentiles) shown = Image.fromarray((stretched * 255).astype(np.uint8)) artist = ax.imshow(shown) _tag_panel(artist, source=img_path, display_range=ranges, raw=raw, sensor_range=sensor_range, steps=[{"op": "rescale", "ranges": ranges, "percentiles": [float(p) for p in percentiles]}, {"op": "to_uint8"}]) ax.axis('off') for j in range(i + 1, len(axs)): axs[j].imshow([[0, 0, 0]], cmap='gray') axs[j].axis('off') plt.subplots_adjust(wspace=0, hspace=0, left=0, right=1, top=1, bottom=0) return fig
[docs] def overlay_masks_on_images(img_folder, normalize=True, resize=True, save=False, plot=False, thickness=2): """Overlay ``masks/*`` outlines onto matching images from ``img_folder``. :param img_folder: Folder containing images; masks live in ``img_folder/masks`` with matching filenames. :param normalize: If True, percentile-normalise images before blending. Default ``True``. :param resize: If True, resize the blended overlay to 1000x1000. Default ``True``. :param save: If True, write PNGs to ``img_folder/overlay/``. Default ``False``. :param plot: If True, show each overlay via matplotlib. Default ``False``. :param thickness: Contour line thickness in pixels. Default ``2``. :returns: ``{'written': int, 'failed': [(filename, reason)]}``. A field that cannot be read is named and skipped rather than ending the run, so a folder holding one truncated TIFF still produces every other overlay -- and the caller can tell which ones are missing. """ def normalize_image(image): """Normalize the image to the 1st and 99th percentiles. :param image: Image array of any numeric dtype, typically the raw 16-bit TIFF. The percentiles are taken over the whole array, so a multi-channel image is stretched by one shared window rather than per channel, and the brightest and darkest 1% saturate. A flat or near-constant image puts both percentiles on the same value; the rescale then divides by zero and the ``uint8`` cast turns the resulting ``nan`` into an undefined value, so guard empty fields upstream. :returns: A ``uint8`` array on 0-255, ready to blend with the mask overlay. """ lower, upper = np.percentile(image, [1, 99]) image = np.clip((image - lower) / (upper - lower), 0, 1) return (image * 255).astype(np.uint8) mask_folder = os.path.join(img_folder,'masks') overlay_folder = os.path.join(img_folder, "overlay") if save and not os.path.exists(overlay_folder): os.makedirs(overlay_folder) image_filenames = set(os.listdir(img_folder)) mask_filenames = set(os.listdir(mask_folder)) common_filenames = image_filenames.intersection(mask_filenames) if not common_filenames: print("No matching filenames found in both folders.") return failed = [] written = 0 for filename in sorted(common_filenames): try: img_path = os.path.join(img_folder, filename) mask_path = os.path.join(mask_folder, filename) image = tiff.imread(img_path) mask = tiff.imread(mask_path) if normalize: image = normalize_image(image) mask = (mask > 0).astype(np.uint8) if mask.shape != image.shape[:2]: mask = cv2.resize(mask, (image.shape[1], image.shape[0]), interpolation=cv2.INTER_NEAREST) contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE) if image.ndim == 2: image_rgb = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB) else: image_rgb = image.copy() overlay = image_rgb.copy() cv2.drawContours(overlay, contours, -1, (255, 0, 0), thickness) blended = cv2.addWeighted(overlay, 0.7, image_rgb, 0.3, 0) if resize: blended = cv2.resize(blended, (1000, 1000), interpolation=cv2.INTER_AREA) if save: save_path = os.path.join(overlay_folder, filename) write_image_rgb(save_path, blended) if plot: with figure_style(theme_target()): plt.figure(figsize=(10, 10)) from .figures.bundle import _register_figure_data _register_figure_data(plt.gcf(), blended, kind="overlay", title=str(filename)) plt.imshow(blended) plt.title(f"Overlay: {filename}") plt.axis('off') plt.show() written += 1 except Exception as error: # noqa: BLE001 failed.append((filename, f"{type(error).__name__}: {error}")) if failed: print(f"overlay_masks_on_images: {written} of " f"{len(common_filenames)} overlaid; {len(failed)} failed.") for name, why in failed[:10]: print(f" {name}: {why}") if len(failed) > 10: print(f" ...and {len(failed) - 10} more") return {"written": written, "failed": failed}
[docs] def graph_importance(settings): """Concatenate feature-importance CSVs and hand off to :class:`spacrGraph` for plotting. :param settings: Settings dict with ``csvs`` (single path or list), ``grouping_column``, ``data_column``, ``graph_type``, ``save``. :returns: None (side-effects: plot shown, artefacts saved). """ from .settings import set_graph_importance_defaults from .utils import save_settings if isinstance(settings['csvs'], (str, os.PathLike)): settings['csvs'] = [settings['csvs']] settings['src'] = os.path.dirname(settings['csvs'][0]) settings = set_graph_importance_defaults(settings) save_settings(settings, name='graph_importance') dfs = [] for path in settings['csvs']: dft = pd.read_csv(path) dfs.append(dft) df = pd.concat(dfs) if not all(col in df.columns for col in (settings['grouping_column'], settings['data_column'])): print(f"grouping {settings['grouping_column']} and data {settings['data_column']} columns must be in {df.columns.to_list()}") return output_dir = os.path.dirname(settings['csvs'][0]) spacr_graph = spacrGraph( df=df, grouping_column=settings['grouping_column'], data_column=settings['data_column'], graph_type=settings['graph_type'], graph_name=settings['grouping_column'], summary_func='mean', colors=None, output_dir=output_dir, save=settings['save'], y_lim=None, error_bar_type='std', representation='object', theme='muted', ) spacr_graph.create_plot() plt.show()
#: Which column carries the unit of replication for each declared level. #: `level` chose the FIGURE and nothing else; it now chooses the denominator #: of the test as well, which is the whole point of declaring it. REPLICATION_UNIT = {"well": None, "plate": "plateID", "plateid": "plateID"} def _unit_column(level, prc_column): """The column whose distinct values are the independent observations. ``well`` resolves to whatever the caller passed as ``prc_column`` -- the well identifier is not always spelled ``prc`` -- and ``plate`` to ``plateID``. """ key = str(level or "object").strip().lower() if key == "object": return None return REPLICATION_UNIT.get(key, prc_column) or prc_column
[docs] def proportions_per_unit(df, group_column, bin_column, unit_column): """Each unit's share of every bin, one row per unit. :param df: object-level observations carrying group, bin, and unit fields. :param group_column: column naming the conditions to compare. :param bin_column: categorical outcome column whose shares are computed. :param unit_column: column naming independent wells, plates, or other replication units. :returns: a frame with ``group_column``, ``unit_column`` and one column per bin holding a proportion in [0, 1]. Units contributing no objects do not appear. """ keys = list(dict.fromkeys([group_column, unit_column, bin_column])) counts = (df.groupby(keys, observed=True).size() .unstack(fill_value=0)) totals = counts.sum(axis=1) proportions = counts.div(totals.where(totals > 0), axis=0) proportions = proportions.dropna(how="all") clashing = [name for name in proportions.index.names if name in proportions.columns] if clashing: proportions = proportions.drop(columns=clashing) return proportions.reset_index()
def _compare_groups(samples): """The test the one engine picks for this shape, under its name. Delegates to :func:`spacr.figures.stats.compare` -- Student's t, Welch's t or Mann-Whitney U for two groups; one-way ANOVA, Welch's ANOVA or Kruskal-Wallis for more -- and translates the name through the same map :mod:`spacr.sp_stats` uses, so one package does not spell "Mann-Whitney U test" two ways. THIS IS THE CALL WHERE THE POWER FLOOR MATTERS MOST. Its samples are PER-UNIT proportions, so n is the number of wells or plates: three to twelve, never the object count. It used to run its own Shapiro and Levene and read "did not reject" as "the assumption holds", which on six wells is a decision the data cannot support -- and the whole purpose of :func:`proportion_test_by_unit` is to stop overstating exactly this kind of result. :param samples: one 1-D array-like per group. :returns: ``(name, statistic, p)``, or ``('too few units', nan, nan)`` when the engine refuses -- which is itself the answer worth printing. """ from .figures.stats import compare from .sp_stats import _ENGINE_TEST_NAMES groups = {index: values for index, values in enumerate(samples)} try: result = compare(groups) except ValueError: return "too few units", float("nan"), float("nan") return (_ENGINE_TEST_NAMES.get(result.test, result.test), float(result.statistic), float(result.p_value))
[docs] def proportion_test_by_unit(df, group_column, bin_column, unit_column): """Compare conditions on their PER-UNIT proportions, one row per bin. :param df: object-level observations carrying group, bin, and unit fields. :param group_column: column naming the conditions to compare. :param bin_column: categorical outcome column whose bins are tested. :param unit_column: column naming independent replication units. The object-level chi-squared asks whether 20,000 objects came from one distribution. Objects in a well share a treatment, a transfection, an imaging session and a monolayer, so that is not the question anyone asked, and its p-value is smaller than the experiment supports by orders of magnitude. This asks the question the design supports: do the WELLS differ, with n = the number of wells. """ if unit_column == group_column: return pd.DataFrame([{ "test": f"not applicable: the groups ARE the {unit_column}s", "bin": None, "unit": unit_column, "n": int(df[unit_column].nunique()) if unit_column in df else 0, "n_per_group": "1 each", "statistic": float("nan"), "p_value": float("nan"), }]) table = proportions_per_unit(df, group_column, bin_column, unit_column) bins = [c for c in table.columns if c not in (group_column, unit_column)] groups = list(dict.fromkeys(table[group_column].tolist())) rows = [] for bin_value in bins: samples = [table.loc[table[group_column] == g, bin_value].to_numpy() for g in groups] name, stat, p = _compare_groups(samples) rows.append({ "test": f"{name} on per-{unit_column} proportions", "bin": bin_value, "unit": unit_column, "n": int(table[unit_column].nunique()), "n_per_group": ", ".join(f"{g}={len(s)}" for g, s in zip(groups, samples)), "statistic": stat, "p_value": p, }) return pd.DataFrame(rows)
[docs] def proportion_mixed_model(df, group_column, bin_column, unit_column): """A binomial GLM on the per-object outcome, standard errors clustered by unit. :param df: object-level observations carrying group, bin, and unit fields. :param group_column: column naming the conditions to compare. :param bin_column: categorical outcome column modelled one bin at a time. :param unit_column: column naming clusters used for robust standard errors. The proportions test throws away how many objects each well contributed; this keeps them while still charging the degrees of freedom the DESIGN supports, by clustering on the unit. Reported beside the other two because when it disagrees with them, the disagreement is the finding. """ if unit_column == group_column: return pd.DataFrame([{ "test": f"not applicable: the groups ARE the {unit_column}s", "bin": None, "unit": unit_column, "n": int(df[unit_column].nunique()) if unit_column in df else 0, "n_per_group": f"objects={len(df)}", "statistic": float("nan"), "p_value": float("nan"), }]) bins = list(dict.fromkeys(df[bin_column].dropna().tolist())) groups = list(dict.fromkeys(df[group_column].dropna().tolist())) rows = [] for bin_value in bins: outcome = (df[bin_column] == bin_value).astype(float).to_numpy() design = pd.get_dummies(df[group_column].astype(str), drop_first=True, dtype=float) if design.empty or len(groups) < 2: rows.append({"test": "binomial GLM, clustered by " + unit_column, "bin": bin_value, "unit": unit_column, "n": int(df[unit_column].nunique()), "n_per_group": f"objects={len(df)}", "statistic": float("nan"), "p_value": float("nan")}) continue design = sm.add_constant(design, has_constant="add") try: fit = sm.GLM(outcome, design.to_numpy(), family=sm.families.Binomial()).fit( cov_type="cluster", cov_kwds={"groups": df[unit_column].astype(str).to_numpy()}) terms = [i for i, name in enumerate(design.columns) if name != "const"] wald = fit.wald_test(np.eye(len(design.columns))[terms], scalar=True) statistic, p = float(wald.statistic), float(wald.pvalue) except Exception as error: print(f"mixed model for bin {bin_value!r} did not fit: {error}") statistic = p = float("nan") rows.append({ "test": f"binomial GLM, standard errors clustered by {unit_column}", "bin": bin_value, "unit": unit_column, "n": int(df[unit_column].nunique()), "n_per_group": f"objects={len(df)}", "statistic": statistic, "p_value": p, }) return pd.DataFrame(rows)
[docs] def plot_proportion_stacked_bars(settings, df, group_column, bin_column, prc_column='prc', level='object', cmap='viridis'): """Plot stacked proportion bars per group with chi-squared and pairwise stats. :param settings: Settings dict — ``verbose`` toggles pairwise chi-squared verbosity. :param df: Long-format DataFrame with categorical ``group_column`` and ``bin_column``. :param group_column: Group axis of the stacked bars. :param bin_column: Categorical column stacked within each bar. :param prc_column: Per-well identifier used when aggregating at the well or plate level. Default ``'prc'``. :param level: Aggregation level — ``'object'`` for direct counts, or ``'well'`` / ``'plateID'`` for per-well means with SD bars. :param cmap: Matplotlib colormap. Default ``'viridis'``. :returns: ``(results_df, pairwise_results, fig)`` — chi-squared summary, pairwise comparison table and the plot figure. """ from .sp_stats import chi_pairwise if isinstance(cmap, str) and cmap.strip().lower() == LEGACY_PLATE_CMAP: cmap = Palette.SEQUENTIAL raw_counts = df.groupby([group_column, bin_column], observed=True).size().unstack(fill_value=0) chi2, p, dof, expected = chi2_contingency(raw_counts) print(f"Chi-squared test statistic (raw data): {chi2:.4f}") print(f"p-value (raw data): {p:.4e}") pairwise_results = chi_pairwise(raw_counts, verbose=settings.get('verbose', False)) _level = str(level or 'object').strip().lower() _AGGREGATED = {'well': prc_column, 'plate': 'plateID', 'plateid': 'plateID'} if _level not in _AGGREGATED and _level != 'object': raise ValueError( f"level={level!r} is not one of 'object', 'well' or 'plate'. " f"Pooling every object would have answered a different question " f"than the one asked.") if _level in _AGGREGATED: prc_column = _AGGREGATED[_level] if prc_column not in df.columns: raise ValueError( f"level={level!r} groups by {prc_column!r}, which this table " f"does not have. Available: {sorted(df.columns)[:12]}") well_proportions = ( df.groupby([group_column, prc_column, bin_column], observed=True) .size() .groupby(level=[0, 1], observed=False) .apply(lambda x: x / x.sum()) .unstack(fill_value=0) ) mean_proportions = well_proportions.groupby( group_column, observed=False).mean() std_proportions = well_proportions.groupby( group_column, observed=False).std() with figure_style(theme_target()): if _level in _AGGREGATED: axis = mean_proportions.plot( kind='bar', stacked=True, yerr=std_proportions, capsize=5, colormap=cmap, figsize=(12, 8) ) title = (f'Proportion of Volume Bins by Group (Mean ± SD across ' f'{"plates" if _level != "well" else "wells"})') else: group_counts = df.groupby([group_column, bin_column], observed=True).size() group_totals = group_counts.groupby( level=0, observed=False).sum() proportions = group_counts / group_totals proportion_df = proportions.unstack(fill_value=0) axis = proportion_df.plot(kind='bar', stacked=True, colormap=cmap, figsize=(12, 8)) title = 'Proportion of Volume Bins by Group' descriptor(axis, title) axis.set_xlabel('Group') axis.set_ylabel('Proportion') axis.legend(title='Classes', bbox_to_anchor=(1.05, 1), loc='upper left') rotate_ticks(axis) axis.set_ylim(0, 1) fig = axis.figure results_df = pd.DataFrame({ 'chi_squared_stat': [chi2], 'p_value': [p], 'degrees_of_freedom': [dof], 'test': ['chi-squared on object counts'], 'unit': ['object'], 'n': [int(len(df))], 'statistic': [float(chi2)], }) unit_column = _unit_column(level, prc_column) or prc_column if unit_column in df.columns: extra = [proportion_test_by_unit(df, group_column, bin_column, unit_column), proportion_mixed_model(df, group_column, bin_column, unit_column)] results_df = pd.concat([results_df] + extra, ignore_index=True) for _, row in pd.concat(extra, ignore_index=True).iterrows(): print(f"{row['test']} [bin {row['bin']}, n={row['n']} " f"{row['unit']}]: p = {row['p_value']:.4e}") else: print(f"no {unit_column!r} column, so only the object-level " f"chi-squared could be computed; objects in one well are not " f"independent and this p-value is smaller than the experiment " f"supports") return results_df, pairwise_results, fig
[docs] def create_venn_diagram(file1, file2, gene_column="gene", filter_coeff=0.1, save=True, save_path=None): """Compute a two-set gene overlap from CSVs and draw its Venn diagram. :param file1: First CSV file. :param file2: Second CSV file. :param gene_column: Column identifying genes. Default ``'gene'``. :param filter_coeff: Threshold on the ``coefficient`` column — positive filters ``> threshold``, negative filters ``< threshold``. :param save: If True, save as PDF; requires ``save_path``. :param save_path: Output PDF path when ``save`` is True. :returns: ``{'overlap', 'unique_to_file1', 'unique_to_file2'}`` lists. :raises ValueError: if ``save`` is True but ``save_path`` is missing. """ from .tabular import read_table df1 = read_table(file1) df2 = read_table(file2) original_frames = (df1.copy(), df2.copy()) input_column = "_spacr_venn_input" while any(input_column in frame.columns for frame in original_frames): input_column += "_" source_data = pd.concat([ frame.assign(**{input_column: index}) for index, frame in enumerate(original_frames)], ignore_index=True) if filter_coeff is not None: df1 = df1[df1['coefficient'] > filter_coeff] if filter_coeff >= 0 else df1[df1['coefficient'] < filter_coeff] df2 = df2[df2['coefficient'] > filter_coeff] if filter_coeff >= 0 else df2[df2['coefficient'] < filter_coeff] genes1 = set(df1[gene_column].dropna()) genes2 = set(df2[gene_column].dropna()) overlapping_genes = genes1.intersection(genes2) unique_to_file1 = genes1.difference(genes2) unique_to_file2 = genes2.difference(genes1) with figure_style(theme_target()): fig, ax = plt.subplots(figsize=(8, 6)) from .figures.bundle import _register_figure_data _register_figure_data(fig, source_data, kind="venn", venn={ "input_column": input_column, "gene_column": gene_column, "filter_coeff": filter_coeff, "inputs": [{"input": index, "path": str(path)} for index, path in enumerate((file1, file2))], "labels": ["File 1 Genes", "File 2 Genes"], "colors": {"10": ROLES["data"], "01": Palette.GREY_DARK, "11": ROLES["highlight"]}, "fontsize": TYPE_SCALE["annotation"], "ink": resolve_ink(theme_target())}) diagram = venn2([genes1, genes2], ('File 1 Genes', 'File 2 Genes'), ax=ax) for region, colour in (('10', ROLES['data']), ('01', Palette.GREY_DARK), ('11', ROLES['highlight'])): patch = diagram.get_patch_by_id(region) if patch is not None: patch.set_color(colour) patch.set_alpha(1.0) patch.set_edgecolor('none') for label in list(diagram.set_labels or []) + list( diagram.subset_labels or []): if label is not None: label.set_fontsize(TYPE_SCALE['annotation']) label.set_color(resolve_ink(theme_target())) descriptor(ax, "Venn Diagram of Overlapping Genes") if save: if save_path is None: raise ValueError("save_path must be provided when save=True.") save_path = save_figure(fig, save_path, bbox_inches="tight") print(f"Venn diagram saved to {save_path}") else: plt.show() return { "overlap": list(overlapping_genes), "unique_to_file1": list(unique_to_file1), "unique_to_file2": list(unique_to_file2) }
[docs] def volcano_plot( data: Union[str, pd.DataFrame], *, fold_change_col: str, p_value_col: str, name_col: Optional[str] = None, x_transform: str = "none", y_transform: str = "-log10", fold_change_threshold: Optional[float] = None, p_value_threshold: Optional[float] = None, annotate: bool = True, annotate_max: Optional[int] = None, point_size: float = 20.0, alpha: float = 0.7, figsize: Tuple[float, float] = (8.0, 6.0), title: Optional[str] = None, xlim: Optional[Tuple[float, float]] = None, ylim: Optional[Tuple[float, float]] = None, threshold_line_kwargs: Optional[dict] = None, scatter_kwargs: Optional[dict] = None, text_kwargs: Optional[dict] = None, save_path: Optional[str] = None, show: bool = True, ax: Optional[plt.Axes] = None, sheet_name: Union[int, str] = 0, ) -> Tuple[plt.Figure, plt.Axes, list]: """Read a table (CSV/TSV/XLS/XLSX or a DataFrame) and render a volcano plot. Auto-detects file type from extension (.csv, .tsv/.tab, .xls/.xlsx) and applies the requested x/y transforms before drawing. :param data: Path to table file or a pandas ``DataFrame``. :param fold_change_col: Column of raw fold change (or logFC when ``x_transform='none'``). :param p_value_col: Column of p-values. :param name_col: Optional column supplying point labels. :param x_transform: One of ``'none'``, ``'log2'``, ``'log10'``, ``'ln'``. Use ``'none'`` when the column already stores logFC (may be negative). :param y_transform: One of ``'none'``, ``'-log10'``, ``'-ln'``, ``'log10'``, ``'ln'``. Default ``'-log10'``. :param fold_change_threshold: Threshold on x — in plotted units when ``x_transform='none'``, otherwise in raw FC units. :param p_value_threshold: Threshold on raw p; drawn as a dashed horizontal line in plotted units. :param annotate: Annotate significant points when a name column is supplied. :param annotate_max: Cap on the number of annotated points (highest y first). :param point_size: Scatter marker size. :param alpha: Scatter marker alpha. :param figsize: Figure size in inches. :param title: Optional figure title. :param xlim: Optional x-axis limits. :param ylim: Optional y-axis limits. :param threshold_line_kwargs: Extra kwargs for threshold lines. :param scatter_kwargs: Extra kwargs for the scatter call. :param text_kwargs: Extra kwargs for label texts. :param save_path: If given, save the figure to this path. :param show: Call ``plt.show()`` at the end. Default ``True``. :param ax: Existing axes to draw on; a new figure is created if None. :param sheet_name: Excel sheet index/name for .xls/.xlsx inputs. :returns: ``(fig, ax, hits)`` where ``hits`` are the labels drawn. :raises ValueError: on unknown transforms, or numeric columns that cannot be coerced. """ def _read_table_auto(path: str) -> pd.DataFrame: """Read Excel or delimited text, sniffing comma versus tab as fallback.""" lower = path.lower() if lower.endswith((".xls", ".xlsx")): try: return pd.read_excel(path, sheet_name=sheet_name) except ImportError as e: raise ImportError( "Reading Excel requires an engine.\n" "For .xlsx: pip install openpyxl\n" "For .xls: pip install xlrd\n" ) from e if lower.endswith((".tsv", ".tab")): return pd.read_csv(path, sep="\t") if lower.endswith(".csv"): return pd.read_csv(path) with open(path, "r", encoding="utf-8", errors="ignore") as f: head = f.read(4096) comma = head.count(",") tab = head.count("\t") sep = "\t" if tab > comma else "," return pd.read_csv(path, sep=sep) def _as_numeric(s: pd.Series, colname: str) -> np.ndarray: """Coerce a column to floats, refusing an entirely nonnumeric result.""" arr = pd.to_numeric(s, errors="coerce").to_numpy(dtype=float) if np.all(np.isnan(arr)): raise ValueError(f"Column '{colname}' could not be converted to numeric.") return arr def _transform_x(x: np.ndarray, mode: str) -> np.ndarray: """Apply the selected x transform, requiring positive log inputs.""" mode = mode.lower() if mode == "none": return x if np.any(x <= 0): raise ValueError( f"x_transform='{mode}' requires all fold changes > 0. " f"If your column is already logFC (can be negative), use x_transform='none'." ) if mode == "log2": return np.log2(x) if mode == "log10": return np.log10(x) if mode in ("ln", "log"): return np.log(x) raise ValueError(f"Unknown x_transform: {mode}") def _transform_y(p: np.ndarray, mode: str) -> np.ndarray: """Apply the selected y transform after clipping logarithm inputs.""" mode = mode.lower() if mode == "none": return p tiny = np.finfo(float).tiny p2 = np.clip(p, tiny, 1.0) if mode == "-log10": return -np.log10(p2) if mode == "-ln": return -np.log(p2) if mode == "log10": return np.log10(p2) if mode in ("ln", "log"): return np.log(p2) raise ValueError(f"Unknown y_transform: {mode}") def _threshold_x_in_plot_units(thresh: float) -> float: """Convert a raw fold-change threshold to absolute plotted units.""" t = float(thresh) if x_transform.lower() == "none": return abs(t) if t <= 0: raise ValueError("fold_change_threshold must be > 0 when using a log x_transform.") if x_transform.lower() == "log2": return abs(np.log2(t)) if x_transform.lower() == "log10": return abs(np.log10(t)) return abs(np.log(t)) def _threshold_y_in_plot_units(pthresh: float) -> float: """Validate and transform a raw p-value threshold for the y axis.""" pt = float(pthresh) if pt <= 0: raise ValueError("p_value_threshold must be > 0.") return float(_transform_y(np.array([pt], dtype=float), y_transform)[0]) df = data.copy() if isinstance(data, pd.DataFrame) else _read_table_auto(str(data)) if fold_change_col not in df.columns: raise KeyError(f"fold_change_col '{fold_change_col}' not found in columns.") if p_value_col not in df.columns: raise KeyError(f"p_value_col '{p_value_col}' not found in columns.") if name_col is not None and name_col not in df.columns: raise KeyError(f"name_col '{name_col}' not found in columns.") x_raw = _as_numeric(df[fold_change_col], fold_change_col) p_raw = _as_numeric(df[p_value_col], p_value_col) keep = ~np.isnan(x_raw) & ~np.isnan(p_raw) df = df.loc[keep].copy() x_raw = x_raw[keep] p_raw = p_raw[keep] x = _transform_x(x_raw, x_transform) y = _transform_y(p_raw, y_transform) mask = np.ones(len(df), dtype=bool) x_thr_plot = None if fold_change_threshold is not None: x_thr_plot = _threshold_x_in_plot_units(fold_change_threshold) mask &= (np.abs(x) >= x_thr_plot) y_thr_plot = None if p_value_threshold is not None: y_thr_plot = _threshold_y_in_plot_units(p_value_threshold) if y_transform.lower() == "none": mask &= (p_raw <= float(p_value_threshold)) else: if y_transform.lower().startswith("-"): mask &= (y >= y_thr_plot) else: mask &= (y <= y_thr_plot) with figure_style(theme_target()): if ax is None: fig, ax = plt.subplots(figsize=figsize) from .figures.bundle import _register_figure_data _register_figure_data(fig, lambda: pd.DataFrame({"effect": np.asarray(x, dtype=float), "significance": np.asarray(y, dtype=float)}), x="effect", y="significance", kind="scatter") else: fig = ax.figure scatter_defaults = dict(s=point_size, alpha=alpha, edgecolors="none") if scatter_kwargs: scatter_defaults.update(scatter_kwargs) if (fold_change_threshold is not None) or (p_value_threshold is not None): colors = np.where(mask & (x >= 0), ROLES["up"], np.where(mask & (x < 0), ROLES["down"], ROLES["data"])) else: colors = ROLES["data"] ax.scatter(x, y, c=colors, **scatter_defaults) xlab = fold_change_col if x_transform.lower() == "none" else f"{x_transform}({fold_change_col})" ylab = p_value_col if y_transform.lower() == "none" else f"{y_transform}({p_value_col})" ax.set_xlabel(xlab) ax.set_ylabel(ylab) if title: ax.set_title(title) line_defaults = dict(color=ROLES["reference"], linestyle=(0, (4, 3)), linewidth=WEIGHTS["reference"], alpha=1.0) if threshold_line_kwargs: line_defaults.update(threshold_line_kwargs) if x_thr_plot is not None: ax.axvline(-x_thr_plot, **line_defaults) ax.axvline(+x_thr_plot, **line_defaults) if y_thr_plot is not None: ax.axhline(y_thr_plot, **line_defaults) if xlim is not None: ax.set_xlim(xlim) if ylim is not None: ax.set_ylim(ylim) ax.spines["right"].set_visible(False) ax.spines["top"].set_visible(False) reference_line(ax, x=0) hits: list = [] if annotate and (name_col is not None): eligible = mask.copy() if (fold_change_threshold is None) and (p_value_threshold is None) and (annotate_max is None): eligible[:] = False if np.any(eligible): idx = np.where(eligible)[0] if annotate_max is not None and len(idx) > int(annotate_max): idx = idx[np.argsort(y[idx])[::-1][: int(annotate_max)]] try: from adjustText import adjust_text except ImportError as e: raise ImportError( "Annotation requires the 'adjustText' package. Install with:\n" " pip install adjustText" ) from e tkw = dict(fontsize=TYPE_SCALE["annotation"], ha="center", va="bottom", color=resolve_ink(theme_target())) if text_kwargs: tkw.update(text_kwargs) texts = [] for i in idx: label = str(df.iloc[i][name_col]) hits.append(label) texts.append(ax.text(x[i], y[i], label, **tkw)) adjust_text( texts, ax=ax, arrowprops=dict(arrowstyle="-", color=ROLES["reference"], lw=WEIGHTS["reference"], alpha=1.0), ) if save_path: save_path = save_figure(fig, save_path, bbox_inches="tight") if show: plt.show() return fig, ax, hits