"""Shared image, model, database, statistics, and pipeline utilities."""
import os, re, sqlite3, torch, torchvision, random, shutil, cv2, tarfile, glob, psutil, platform, gzip, subprocess, time, requests, ast, traceback, logging
import numpy as np
_trapezoid = getattr(np, 'trapezoid', None) or np.trapz
import pandas as pd
from contextlib import contextmanager, nullcontext
from functools import partial
class _DeferredModule:
"""Small module proxy for dependencies used by one distant code path.
:param name: dotted module name, imported on FIRST ATTRIBUTE ACCESS and
cached from then on. Nothing checks it here, so a name with a typo
in it costs nothing until the distant path is finally taken -- which
is the trade this proxy exists to make.
"""
def __init__(self, name):
"""Record the module name without importing it."""
self.__dict__['_name'] = name
self.__dict__['_module'] = None
def _load(self):
"""Import on first use and cache the module."""
module = self.__dict__['_module']
if module is None:
from importlib import import_module
module = import_module(self.__dict__['_name'])
self.__dict__['_module'] = module
return module
def __getattr__(self, name):
"""Forward attribute reads, importing the module if needed."""
return getattr(self._load(), name)
def __setattr__(self, name, value):
"""Forward attribute writes ONTO THE REAL MODULE, importing it first.
The proxy keeps its own state in ``__dict__`` directly for this reason: an
ordinary assignment here would set an attribute on the imported module,
not on the proxy.
"""
setattr(self._load(), name, value)
def __repr__(self):
"""Name the module and say whether it has been imported yet.
Deliberately does NOT import it -- inspecting a proxy in a debugger must
not be the thing that triggers the import it exists to defer.
"""
state = (
'loaded' if self.__dict__['_module'] is not None
else 'not yet imported'
)
return f"<deferred module {self.__dict__['_name']!r} ({state})>"
cp_models = _DeferredModule('cellpose.models')
from skimage import morphology
from skimage.measure import label, regionprops_table, regionprops
import skimage.measure as measure
from skimage.transform import resize as resizescikit
from skimage.morphology import dilation
try:
from skimage.morphology import footprint_rectangle
except ImportError:
from skimage.morphology import square as _legacy_square
def _square_footprint(size):
"""A square structuring element, on scikit-image 0.22 to 0.24.
`square(n)` was removed in 0.25 in favour of
`footprint_rectangle((n, n))`. Both spellings are wrapped rather
than pinning a version, because spaCR is installed alongside
whatever cellpose and torch have already chosen.
"""
return _legacy_square(size)
else:
def _square_footprint(size):
"""A square structuring element, on scikit-image 0.25 and later.
The modern spelling. See the ImportError branch above for why both
exist.
"""
return footprint_rectangle((size, size))
from skimage.measure import find_contours
from skimage.segmentation import clear_border, find_boundaries
from scipy.stats import pearsonr
from skimage.filters import (gaussian, frangi, sato, meijering, difference_of_gaussians, apply_hysteresis_threshold)
from skimage.morphology import white_tophat, disk
from skimage.feature import blob_log, blob_dog
from collections import defaultdict, OrderedDict, Counter
from PIL import Image
from statsmodels.stats.outliers_influence import variance_inflation_factor
from statsmodels.stats.stattools import durbin_watson
import statsmodels.formula.api as smf
import statsmodels.api as sm
from statsmodels.stats.multitest import multipletests
from itertools import combinations
from functools import reduce
from .figures.style import figure_style, theme_target
try:
from IPython.display import display
except Exception:
[docs]
def display(*args, **kwargs):
"""Do nothing: IPython is unavailable, so there is nowhere to display to.
THE FALLBACK IS THE POINT. `IPython.display.display` is imported at
module scope, and IPython can be mid-init -- partially imported by
another thread -- which makes that import raise. Letting it propagate
would make importing this module fail for a reason that has nothing to
do with what the module does. spaCR only calls `display` from notebook
contexts; the Qt GUI ignores it.
:param args: whatever the caller would have displayed.
:param kwargs: likewise.
"""
pass
from typing import Optional, Any
from .image_colors import read_image_rgb, write_image_rgb
from .measurement_schema import MEASUREMENT_STAMP_COLUMNS
from multiprocessing import cpu_count, set_start_method, get_start_method
from .resource_log import _parallel_pool as Pool
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
from torch.utils.data import Subset
from torch.autograd import grad
from torchvision import models
from torchvision.models.resnet import ResNet18_Weights, ResNet34_Weights, ResNet50_Weights, ResNet101_Weights, ResNet152_Weights
import torchvision.transforms as transforms
from torchvision.models import resnet50
from torchvision import models as tv_models
import seaborn as sns
import matplotlib.pyplot as plt
from matplotlib.offsetbox import OffsetImage, AnnotationBbox
import matplotlib as mpl
from scipy import stats
import scipy.ndimage as ndi
from scipy.spatial import distance
from scipy.stats import fisher_exact, f_oneway, kruskal
from scipy.ndimage import gaussian_filter
from scipy.spatial import ConvexHull
from scipy.interpolate import splprep, splev
from scipy import ndimage
from scipy.ndimage import binary_dilation, binary_fill_holes
from skimage.exposure import rescale_intensity
LOG = logging.getLogger(__name__)
@contextmanager
def _preserve_batchnorm_running_stats(module: nn.Module):
"""Prevent checkpoint recomputation from updating BatchNorm twice.
Non-reentrant activation checkpointing replays the forward operation
during backward. BatchNorm must still run in training mode for identical
gradients, but its running buffers should reflect one batch, not two.
"""
snapshots = []
for child in module.modules():
if isinstance(child, nn.modules.batchnorm._BatchNorm):
for name in ("running_mean", "running_var", "num_batches_tracked"):
value = getattr(child, name, None)
if value is not None:
snapshots.append((value, value.detach().clone()))
try:
yield
finally:
with torch.no_grad():
for target, saved in snapshots:
target.copy_(saved)
def _checkpoint_module(module: nn.Module, function, *args):
"""Checkpoint ``function`` while preserving stateful normalization buffers."""
def contexts():
"""Return forward and recomputation contexts for non-reentrant checkpointing."""
return nullcontext(), _preserve_batchnorm_running_stats(module)
return checkpoint(function, *args, use_reentrant=False,
context_fn=contexts)
from sklearn.metrics import auc, precision_recall_curve
from sklearn.linear_model import Lasso, Ridge
from sklearn.preprocessing import OneHotEncoder, StandardScaler
from sklearn.cluster import KMeans, DBSCAN
from sklearn.manifold import TSNE, Isomap, SpectralEmbedding
from sklearn.decomposition import PCA
from sklearn.ensemble import RandomForestClassifier
from huggingface_hub import list_repo_files
def _run_random_state(default=None):
"""Return the active run's seed, for an estimator's ``random_state=``.
Imported inside the call rather than at module scope: :mod:`spacr.runctx`
reaches :mod:`spacr.settings`, which reaches back here, and a top-level
import would be a cycle. Outside a run this is whatever ``default`` was,
which is the literal these call sites used to hard-code.
:param default: the value to use when no run is open.
:returns: the run seed, or ``default``.
"""
from .runctx import random_state
return random_state(default)
spacr_path = os.path.join(os.path.dirname(__file__), '__init__.py')
#: Import roots spaCR refuses to let an optional dependency drag in.
#: TensorFlow is not a spaCR dependency -- setup.py has it commented out --
#: it is merely installed in some environments. It costs ~2.6 s of import
#: (TF plus the Keras it pulls), prints its cpu_feature_guard banner over the
#: run log, and is a known off-main-thread segfault vector in a GUI process.
_TF_BACKED_ROOTS = ('tensorflow', 'keras', 'tf_keras')
class _TensorFlowIsNotADependency(ImportError):
"""Raised instead of importing TensorFlow inside a spaCR import."""
[docs]
class OptionalDependencyCompatibilityError(ImportError):
"""An installed optional dependency is too old for spaCR's API contract."""
def _distribution_version(name):
"""Return an installed distribution version without importing its package."""
from importlib.metadata import version
return version(name)
def _release_version(value):
"""Return the numeric release segment of a PEP 440-style version."""
match = re.match(r"^\s*(\d+(?:\.\d+)*)", str(value))
if match is None:
return ()
return tuple(int(part) for part in match.group(1).split("."))
class _BlockTensorFlowFinder:
"""``sys.meta_path`` finder that refuses TF-backed imports.
Installed only for the duration of one wrapped import and removed
immediately afterwards, so it can never affect code that genuinely wants
TensorFlow. Optional-dependency probes already handle ``ImportError``
-- ``umap/__init__.py`` catches it and substitutes a stub
``ParametricUMAP`` -- so raising one simply gives them the behaviour they
have on a machine where TF was never installed.
"""
def find_spec(self, fullname, path=None, target=None):
"""Raise for a TF-backed root; defer to the next finder otherwise."""
if fullname.split('.')[0] in _TF_BACKED_ROOTS:
raise _TensorFlowIsNotADependency(
f"{fullname} is not a spaCR dependency and is never imported "
f"by spaCR; see spacr.utils._BlockTensorFlowFinder.")
return None
class _LazyModule:
"""Import a module the first time an attribute is read off it.
``import umap.umap_ as umap`` at module scope makes every importer of
``spacr.utils`` pay for umap, and umap pays for numba, pynndescent and --
through ``umap.parametric_umap`` -- TensorFlow when it happens to be
installed. Measured on a developer box that is **6.5 s and ~1.4 GB**, and
it lands on processes that will never embed anything: every field-measuring
worker of a ``spawn`` or ``forkserver`` pool re-imports the whole chain
from a cold interpreter, so the cost is paid once *per worker*.
Deferring it keeps the two real call sites
(:func:`reduction_and_clustering` and :func:`generate_image_umap`) written
exactly as they were -- ``umap.UMAP(...)`` still works -- while an
``ImportError`` now surfaces where UMAP is actually asked for rather than
at ``import spacr.utils``.
Deferring alone is not enough for umap, though: the TensorFlow import is
postponed, not prevented, and reappears the moment anything reads
``umap.UMAP``. ``block_roots`` closes that -- the wrapped import runs with
those roots refused, which is why spaCR can use umap without TensorFlow
ever entering the process.
:param name: dotted module name to import on first attribute access.
:param block_roots: import roots refused for the duration of that import.
:param minimum_distribution: ``(distribution, version, reason)``, or
``None`` for no check. Guards against an INSTALLED BUT TOO OLD
package, not a missing one -- a distribution that cannot be found at
all is left to the import below, so the caller gets Python's normal
error rather than a version complaint about something absent.
"""
def __init__(self, name, block_roots=(), minimum_distribution=None):
"""Record the module name and its import guards without importing."""
self.__dict__['_name'] = name
self.__dict__['_module'] = None
self.__dict__['_block_roots'] = tuple(block_roots)
self.__dict__['_minimum_distribution'] = minimum_distribution
def reset(self):
"""Forget the cached module so the next access performs a fresh import.
This is intentionally narrower than deleting entries from
:data:`sys.modules`: other code may legitimately hold the imported
package. The proxy itself returns to its pristine lazy state, which
gives dependency probes and tests an explicit, order-independent
reset point.
"""
self.__dict__['_module'] = None
def _load(self):
"""Import and cache the wrapped module, blocking ``block_roots``."""
module = self.__dict__['_module']
name = self.__dict__['_name']
root = name.split('.', 1)[0]
import sys as _sys
if root in _sys.modules and _sys.modules[root] is None:
self.__dict__['_module'] = None
raise ModuleNotFoundError(
f"import of {root!r} halted; None in sys.modules",
name=root,
)
minimum = self.__dict__['_minimum_distribution']
if minimum is not None:
distribution, minimum_version, reason = minimum
try:
current = _distribution_version(distribution)
except Exception:
current = None
if current is not None:
current_release = _release_version(current)
minimum_release = _release_version(minimum_version)
width = max(len(current_release), len(minimum_release))
current_release += (0,) * (width - len(current_release))
minimum_release += (0,) * (width - len(minimum_release))
if current_release < minimum_release:
self.__dict__['_module'] = None
raise OptionalDependencyCompatibilityError(
f"spaCR cannot initialize {distribution} {current}; "
f"version {minimum_version} or newer is required. "
f"{reason} Upgrade with `python -m pip install --upgrade "
f"'{distribution}>={minimum_version},<1.0'`."
)
if module is None:
from importlib import import_module
before = {
key for key in _sys.modules
if key == root or key.startswith(root + '.')
}
if self.__dict__['_block_roots']:
blocker = _BlockTensorFlowFinder()
_sys.meta_path.insert(0, blocker)
try:
module = import_module(name)
except Exception:
self.__dict__['_module'] = None
for key in tuple(_sys.modules):
if ((key == root or key.startswith(root + '.'))
and key not in before):
_sys.modules.pop(key, None)
raise
finally:
try:
_sys.meta_path.remove(blocker)
except ValueError:
pass
else:
try:
module = import_module(name)
except Exception:
self.__dict__['_module'] = None
for key in tuple(_sys.modules):
if ((key == root or key.startswith(root + '.'))
and key not in before):
_sys.modules.pop(key, None)
raise
self.__dict__['_module'] = module
return module
def __getattr__(self, item):
"""Forward attribute reads, importing the module if needed."""
return getattr(self._load(), item)
def __setattr__(self, item, value):
"""Forward attribute writes onto the real module, importing it first."""
setattr(self._load(), item, value)
def __dir__(self):
"""The real module's names, importing it to find them."""
return dir(self._load())
def __repr__(self):
"""Name the module and say whether it has been imported.
Does NOT import it: inspecting a lazy proxy in a debugger must not be the
thing that triggers the import it exists to defer.
"""
loaded = self.__dict__['_module'] is not None
state = 'loaded' if loaded else 'not yet imported'
return f"<lazy module {self.__dict__['_name']!r} ({state})>"
#: ``umap.umap_``, imported on first use and without TensorFlow.
#: ``import umap.umap_`` runs ``umap/__init__.py``, which imports
#: ``umap.parametric_umap`` -> ``tensorflow``. spaCR uses only
#: ``umap.umap_.UMAP`` and never ``ParametricUMAP``, so the TF-backed roots
#: are blocked for that import and umap takes its own documented no-TF path.
#: See :class:`_LazyModule`.
umap = _LazyModule(
'umap.umap_',
block_roots=_TF_BACKED_ROOTS,
minimum_distribution=(
'umap-learn',
'0.5.11',
"Older releases call scikit-learn's removed `force_all_finite` API.",
),
)
from functools import wraps
from skimage.segmentation import watershed
from skimage.feature import peak_local_max
import tifffile
from . import schema, tabular
from .tiff_io import write_tiff
def _load_image(filepath):
"""Load a .tif or .npy image."""
ext = os.path.splitext(filepath)[1].lower()
if ext == '.npy':
return np.load(filepath)
elif ext in ('.tif', '.tiff'):
return tifffile.imread(filepath)
return None
def _save_image(filepath, img):
"""Save image as .tif or .npy matching original format."""
ext = os.path.splitext(filepath)[1].lower()
if ext == '.npy':
np.save(filepath, img)
else:
write_tiff(filepath, img)
def _select_intensity_channel(raw, intensity_channel):
"""Pick one intensity plane out of a raw image, layout-aware.
2-D images (and a ``None`` channel) are returned as-is. A 3-D image is
treated as channel-last when its trailing axis is small (<= 4), else as
channel-first when its leading axis is small, else channel-last.
Shared by the on-disk (:func:`_process_single_fov`) and in-memory
(:func:`_process_single_fov_in_memory`) paths so the two cannot drift:
the on-disk one used to do a bare ``raw[intensity_channel]``, which
silently took a ROW of a 2-D image and the wrong axis of a channel-last
stack.
:param raw: 2-D or 3-D image array.
:param intensity_channel: channel index, or ``None`` to use ``raw`` whole.
:returns: a float32 array.
:raises ValueError: if ``intensity_channel`` is out of bounds.
"""
raw = np.asarray(raw)
if raw.ndim == 2 or intensity_channel is None:
return raw.astype(np.float32)
if raw.ndim == 3:
if raw.shape[-1] <= 4:
if intensity_channel >= raw.shape[-1]:
raise ValueError(
f"intensity_channel={intensity_channel} out of bounds for channel-last image with shape {raw.shape}"
)
return raw[..., intensity_channel].astype(np.float32)
if raw.shape[0] <= 4:
if intensity_channel >= raw.shape[0]:
raise ValueError(
f"intensity_channel={intensity_channel} out of bounds for channel-first image with shape {raw.shape}"
)
return raw[intensity_channel].astype(np.float32)
if intensity_channel >= raw.shape[-1]:
raise ValueError(
f"intensity_channel={intensity_channel} out of bounds for image with shape {raw.shape}"
)
return raw[..., intensity_channel].astype(np.float32)
return raw.astype(np.float32)
def _union_find_root(parent, i):
"""Find a set's representative, compressing the path as it goes.
:param parent: the union-find parent array, modified in place.
:param i: the element to look up.
:returns: its root.
"""
while parent[i] != i:
parent[i] = parent[parent[i]]
i = parent[i]
return i
def _union_find_merge(parent, a, b):
"""Merge the sets holding two elements.
The lower index becomes the root, so the representative of a merged set
is deterministic -- which is what makes labels reproducible between
runs rather than depending on the order objects happened to be visited.
:param parent: the union-find parent array, modified in place.
:param a: one element.
:param b: the other.
"""
ra = _union_find_root(parent, a)
rb = _union_find_root(parent, b)
if ra != rb:
parent[max(ra, rb)] = min(ra, rb)
def _compute_label_perimeters(label_img):
"""Return dict {label: perimeter_pixel_count}."""
boundaries = find_boundaries(label_img, mode='inner')
boundary_labels = label_img[boundaries]
unique, counts = np.unique(boundary_labels[boundary_labels > 0], return_counts=True)
return dict(zip(unique.astype(int), counts.astype(int)))
def _compute_shared_boundaries(label_img):
"""Return dict {(min_label, max_label): shared_pixel_count}."""
shared = {}
for dy, dx in [(-1, 0), (1, 0), (0, -1), (0, 1)]:
shifted = np.roll(np.roll(label_img, dy, axis=0), dx, axis=1)
mask = (label_img > 0) & (shifted > 0) & (label_img != shifted)
if not np.any(mask):
continue
a = label_img[mask].astype(int)
b = shifted[mask].astype(int)
for la, lb in zip(a, b):
pair = (min(la, lb), max(la, lb))
shared[pair] = shared.get(pair, 0) + 1
return shared
def _merge_by_perimeter(label_img, perimeter_fraction, parent):
"""Mark label pairs for merging based on shared perimeter fraction."""
perimeters = _compute_label_perimeters(label_img)
shared = _compute_shared_boundaries(label_img)
for (la, lb), shared_px in shared.items():
perim_a = perimeters.get(la, 1)
perim_b = perimeters.get(lb, 1)
smaller_perim = min(perim_a, perim_b)
if shared_px / smaller_perim >= perimeter_fraction:
_union_find_merge(parent, la, lb)
def _relabel_sequential(label_img):
"""Relabel to sequential uint16 IDs starting at 1."""
present = np.unique(label_img)
present = present[present > 0]
mapping = np.zeros(int(label_img.max()) + 1, dtype=np.uint16)
for new_id, old_id in enumerate(present, start=1):
mapping[int(old_id)] = new_id
return mapping[label_img].astype(np.uint16)
def _apply_union_find(label_img, parent):
"""Apply union-find mapping and relabel sequentially."""
mapping = np.zeros(int(label_img.max()) + 1, dtype=np.int64)
for l in range(1, len(mapping)):
if l in parent:
mapping[l] = _union_find_root(parent, l)
else:
mapping[l] = l
merged = mapping[label_img]
return _relabel_sequential(merged.astype(np.uint16))
def _validated_intensity_bounds(minimum, maximum):
"""Coerce saved bounds without treating nonfinite or negative values as off."""
message = "Intensity bounds must be finite nonnegative numbers"
try:
bounds = tuple(0.0 if value is None else float(value)
for value in (minimum, maximum))
except (TypeError, ValueError) as error:
raise ValueError(message) from error
if any(not np.isfinite(value) or value < 0 for value in bounds):
raise ValueError(message)
return bounds
def _describe_object_filters(filters):
"""The filter list as one line of text, for the run log."""
parts = []
for entry in filters:
low, high = entry["min"], entry["max"]
if low is not None and high is not None:
parts.append(f"{low:g} <= {entry['property']} <= {high:g}")
elif low is not None:
parts.append(f"{entry['property']} >= {low:g}")
elif high is not None:
parts.append(f"{entry['property']} <= {high:g}")
return ", ".join(parts)
def _filter_objects(label_img, intensity_img=None, min_area=0, max_area=0,
remove_border=False, *, min_intensity=0, max_intensity=0,
filters=None):
"""Remove objects by the object filter list and border contact.
ONE FILTER SYSTEM (item 511). The legacy area and absolute mean
intensity bounds are migrated into filter-list entries by
:func:`spacr.qt.mask_engine.legacy_filters` and judged together with
``filters`` -- any scalar scikit-image regionprop the user added -- by
:func:`spacr.qt.mask_engine.filter_removals`, the engine Make Masks
runs, in one ``regionprops_table`` pass per mask.
Intensity properties use the object's own-channel plane, in the units of
that plane. The caller must supply the original pixel values, not a
display-normalized image. These bounds can retain every object or remove
every object in a field.
Parameters
----------
label_img : ndarray (uint16)
Label image.
intensity_img : ndarray or None
Own-channel intensity plane, with exactly the label image's shape.
Required only when an intensity bound or intensity filter is set and
objects exist.
min_area : int
Remove objects with area < min_area. 0 = disabled.
max_area : int
Remove objects with area > max_area. 0 = disabled.
remove_border : bool
Remove objects touching any image edge.
min_intensity, max_intensity : float
Remove objects whose mean is below/above the respective bound.
Equality is retained; 0 disables that side. Object means must be
finite; nonfinite background pixels do not contribute to a mean.
filters : list of dict or None
Filter entries ``{"property", "min", "max"}``; an object is kept
when ``min <= value <= max`` for each, and a None side is off.
Returns
-------
ndarray (uint16)
Filtered and relabelled image.
"""
from .qt.mask_engine import filter_removals, legacy_filters, normalise_filters
min_intensity, max_intensity = _validated_intensity_bounds(
min_intensity, max_intensity)
by_area_rules = legacy_filters(min_area=min_area, max_area=max_area)
by_intensity_rules = legacy_filters(min_intensity=min_intensity,
max_intensity=max_intensity)
listed = normalise_filters(filters)
rules = by_area_rules + by_intensity_rules + listed
labels_present = np.unique(label_img)
labels_present = labels_present[labels_present > 0]
if len(labels_present) == 0:
return label_img
remove = set()
if rules:
removals = filter_removals(label_img, rules, intensity_img,
require_finite_intensity=True)
first_intensity = len(by_area_rules)
first_listed = first_intensity + len(by_intensity_rules)
def _failed_in(low, high):
"""Labels failing an entry whose index is in ``[low, high)``."""
return {removal.label for removal in removals
if any(low <= failed.index < high for failed in removal.failed)}
by_area = _failed_in(0, first_intensity)
if by_area:
print(f" Area filter: removed {len(by_area)}/{len(labels_present)} objects "
f"(min_area={min_area}, max_area={max_area})")
remove.update(by_area)
by_intensity = _failed_in(first_intensity, first_listed)
additional = len(by_intensity - remove)
remove.update(by_intensity)
if additional:
print(f" Intensity filter: removed {additional} additional objects "
f"(min_intensity={min_intensity}, max_intensity={max_intensity})")
by_list = _failed_in(first_listed, len(rules))
additional = len(by_list - remove)
remove.update(by_list)
if additional:
print(f" Object filters: removed {additional} additional objects "
f"({_describe_object_filters(listed)})")
if remove_border:
border_labels = set()
for axis in range(label_img.ndim):
border_labels.update(np.unique(label_img.take(0, axis=axis)).tolist())
border_labels.update(np.unique(label_img.take(-1, axis=axis)).tolist())
border_labels.discard(0)
new_border = border_labels - remove
remove.update(border_labels)
if len(new_border) > 0:
print(f" Border filter: removed {len(new_border)} additional objects")
total_removed = len(remove)
total_original = len(labels_present)
if remove:
mask = np.isin(label_img, list(remove))
label_img[mask] = 0
result = _relabel_sequential(label_img)
remaining_count = len(np.unique(result[result > 0]))
print(f" Filter summary: {total_original} objects → {remaining_count} objects ({total_removed} removed)")
return result
def _process_single_fov_in_memory(mask, intensity_img=None, intensity_channel=None,
do_perimeter_merge=False, perimeter_fraction=0.5,
min_area=0, max_area=0, remove_border_objects=False,
progress_callback=None, fov_index=0, total_fovs=0,
op_name='', *, min_intensity=0, max_intensity=0,
filters=None):
"""Copy one label field, merge by perimeter, then apply shared object filters.
Intensity input is an original own-channel plane or an explicitly
indexed channel-last array. Its dtype is preserved until the shared
filter accumulates object means in float64; layout is never guessed.
"""
start = time.time()
if mask is None:
return None
min_intensity, max_intensity = _validated_intensity_bounds(
min_intensity, max_intensity)
label_img = np.asarray(mask).astype(np.uint16).copy()
n_before = len(np.unique(label_img[label_img > 0]))
if n_before == 0:
print(f" FOV {fov_index}: empty mask, skipping")
return label_img
from .qt.mask_engine import filters_need_intensity
intensity_img_use = None
if ((min_intensity > 0 or max_intensity > 0 or filters_need_intensity(filters))
and intensity_img is not None):
intensity_img_use = np.asarray(intensity_img)
if intensity_img_use.ndim == label_img.ndim + 1:
if intensity_channel is None:
raise ValueError("An explicit own-channel index is required for intensity filtering")
intensity_img_use = intensity_img_use[..., intensity_channel]
all_labels = np.unique(label_img)
all_labels = all_labels[all_labels > 0]
parent = {int(l): int(l) for l in all_labels}
n_before_merge = len(all_labels)
if do_perimeter_merge:
_merge_by_perimeter(label_img, perimeter_fraction, parent)
label_img = _apply_union_find(label_img, parent)
n_after_merge = len(np.unique(label_img[label_img > 0]))
if n_after_merge != n_before_merge:
print(f" FOV {fov_index} merge: {n_before_merge} → {n_after_merge} objects")
label_img = _filter_objects(
label_img,
intensity_img_use,
min_area=min_area,
max_area=max_area,
remove_border=remove_border_objects,
min_intensity=min_intensity,
max_intensity=max_intensity,
filters=filters,
)
duration = time.time() - start
if progress_callback:
progress_callback(fov_index, total_fovs, duration, op_name)
return label_img
[docs]
def merge_split_objects(mask_src, intensity_img_src=None, intensity_channel=None,
perimeter_fraction=0.5,
min_area=0, max_area=0, remove_border_objects=False,
n_jobs=1, progress_callback=None, op_name='', *,
min_intensity=0, max_intensity=0, filters=None):
"""Merge by perimeter and filter labeled objects across a directory of masks.
Runs the shared in-memory merge/filter pipeline on each mask file in
``mask_src`` in parallel, overwriting each mask in place.
:param mask_src: directory containing mask .tif/.tiff/.npy files.
:param intensity_img_src: directory of matched original intensity images,
required when either intensity bound is enabled.
:param intensity_channel: explicit channel-last index for multi-channel
intensity images; unnecessary for single-channel planes.
:param perimeter_fraction: minimum shared-boundary fraction for perimeter-based merging.
:param min_area: remove objects smaller than this (px); 0 disables.
:param max_area: remove objects larger than this (px); 0 disables.
:param remove_border_objects: drop objects touching the image border.
:param n_jobs: parallel worker count.
:param progress_callback: optional callback(fov_index, total, duration, op_name).
:param op_name: label passed to the progress callback.
:param min_intensity: remove objects whose own-channel mean is below this
raw-image value; equality is kept and 0 disables the lower bound.
:param max_intensity: remove objects whose own-channel mean is above this
raw-image value; equality is kept and 0 disables the upper bound.
:param filters: object filter entries (any scalar regionprop with a
minimum and a maximum), judged with the bounds above in one pass.
:returns: None.
"""
valid_ext = ('.tif', '.tiff', '.npy')
mask_files = sorted([f for f in os.listdir(mask_src)
if os.path.splitext(f)[1].lower() in valid_ext])
if not mask_files:
return
do_perimeter_merge = perimeter_fraction > 0
min_intensity, max_intensity = _validated_intensity_bounds(
min_intensity, max_intensity)
mask_paths = [os.path.join(mask_src, f) for f in mask_files]
if intensity_img_src is not None:
intensity_paths = [os.path.join(intensity_img_src, f) for f in mask_files]
else:
intensity_paths = [None] * len(mask_files)
total = len(mask_paths)
from .resource_log import (_array_file_nbytes, _guard_workers,
_parallel_cloudpickle_map)
n_jobs = _guard_workers('merge_split', n_jobs,
_array_file_nbytes(mask_paths[0]))
_parallel_cloudpickle_map((
(_process_single_fov, (mp, ip, intensity_channel,
do_perimeter_merge, perimeter_fraction,
min_area, max_area, remove_border_objects,
progress_callback, idx, total, op_name), {
'min_intensity': min_intensity, 'max_intensity': max_intensity,
'filters': filters})
for idx, (mp, ip) in enumerate(zip(mask_paths, intensity_paths))
), n_jobs)
def _process_single_fov(mask_path, intensity_path, intensity_channel,
do_perimeter_merge, perimeter_fraction,
min_area, max_area, remove_border_objects,
progress_callback=None, fov_index=0, total_fovs=0, op_name='', *,
min_intensity=0, max_intensity=0, filters=None):
"""Load one field and save the result of the same filter used in memory."""
from .qt.mask_engine import filters_need_intensity
start = time.time()
label_img = _load_image(mask_path)
if label_img is None:
return
intensity_img = None
min_intensity, max_intensity = _validated_intensity_bounds(
min_intensity, max_intensity)
if ((min_intensity > 0 or max_intensity > 0 or filters_need_intensity(filters))
and intensity_path is not None):
intensity_img = _load_image(intensity_path)
filtered = _process_single_fov_in_memory(
label_img, intensity_img, intensity_channel,
do_perimeter_merge, perimeter_fraction, min_area, max_area,
remove_border_objects, None, fov_index, total_fovs, op_name,
min_intensity=min_intensity, max_intensity=max_intensity,
filters=filters,
)
import tempfile
import stat
descriptor, temporary = tempfile.mkstemp(
prefix='.spacr-mask-', suffix=os.path.splitext(mask_path)[1],
dir=os.path.dirname(os.path.abspath(mask_path)))
os.close(descriptor)
try:
_save_image(temporary, filtered)
with open(temporary, 'rb') as handle:
os.fsync(handle.fileno())
os.chmod(temporary, stat.S_IMODE(os.stat(mask_path).st_mode))
os.replace(temporary, mask_path)
finally:
if os.path.exists(temporary):
os.unlink(temporary)
if progress_callback:
try:
progress_callback(fov_index, total_fovs, time.time() - start, op_name)
except Exception as error:
from .runctx import _is_overload_failure
if not _is_overload_failure(error):
raise
import logging
logging.getLogger(__name__).warning(
'Mask saved; progress notification overloaded: %s', error)
def _organelle_diagnostic(img, morphology, method, settings):
"""
Generate a diagnostic image for organelle segmentation QC.
Returns the processed intermediate image and a descriptive title,
depending on the morphology mode and method used.
Parameters
----------
img : ndarray
2-D float32 single-channel image.
morphology : str
One of 'spots', 'network', 'irregular', 'ring'.
method : str
Segmentation method used.
settings : dict
Organelle settings.
Returns
-------
diag_img : ndarray
2-D image showing the intermediate processing step.
diag_title : str
Description for the plot title.
"""
img_norm = img.astype(np.float64)
pmin, pmax = np.percentile(img_norm, (1, 99))
if pmax - pmin > 0:
img_norm = np.clip((img_norm - pmin) / (pmax - pmin), 0, 1)
if morphology == 'spots':
if method == 'log':
blobs = blob_log(img_norm,
min_sigma=settings.get('organelle_log_min_sigma', 1),
max_sigma=settings.get('organelle_log_max_sigma', 10),
num_sigma=settings.get('organelle_log_num_sigma', 10),
threshold=settings.get('organelle_log_threshold', 0.01))
diag_img = img_norm.copy()
for y, x, sigma in blobs:
rr, cc = np.ogrid[-int(sigma*2):int(sigma*2)+1, -int(sigma*2):int(sigma*2)+1]
circle = rr**2 + cc**2 <= (sigma * np.sqrt(2))**2
yy = np.clip(int(y) + np.where(circle)[0] - int(sigma*2), 0, img.shape[0]-1)
xx = np.clip(int(x) + np.where(circle)[1] - int(sigma*2), 0, img.shape[1]-1)
diag_img[yy, xx] = 1.0
return diag_img, f'LoG detections ({len(blobs)} blobs)'
elif method == 'dog':
blobs = blob_dog(img_norm,
min_sigma=settings.get('organelle_dog_sigma_low', 1.0),
max_sigma=settings.get('organelle_dog_sigma_high', 3.0),
threshold=settings.get('organelle_log_threshold', 0.01))
diag_img = img_norm.copy()
for y, x, sigma in blobs:
rr, cc = np.ogrid[-int(sigma*2):int(sigma*2)+1, -int(sigma*2):int(sigma*2)+1]
circle = rr**2 + cc**2 <= (sigma * np.sqrt(2))**2
yy = np.clip(int(y) + np.where(circle)[0] - int(sigma*2), 0, img.shape[0]-1)
xx = np.clip(int(x) + np.where(circle)[1] - int(sigma*2), 0, img.shape[1]-1)
diag_img[yy, xx] = 1.0
return diag_img, f'DoG detections ({len(blobs)} blobs)'
else:
radius = settings.get('organelle_tophat_radius', 5)
filtered = white_tophat(img, disk(radius))
return filtered, f'Top-hat filtered (r={radius})'
elif morphology == 'network':
if method == 'ridge':
sigmas = settings.get('organelle_ridge_sigmas', [1, 2, 3])
filter_name = settings.get('organelle_ridge_filter', 'frangi')
ridge_filters = {'frangi': frangi, 'sato': sato, 'meijering': meijering}
enhanced = ridge_filters[filter_name](img_norm, sigmas=sigmas, black_ridges=False)
return enhanced, f'{filter_name} ridge (sigmas={sigmas})'
elif method == 'hysteresis':
low = settings.get('organelle_hysteresis_low', 0.2)
high = settings.get('organelle_hysteresis_high', 0.6)
smooth = gaussian(img, sigma=1)
if low < 1.0:
low_abs = np.percentile(smooth, low * 100)
else:
low_abs = low
if high < 1.0:
high_abs = np.percentile(smooth, high * 100)
else:
high_abs = high
binary = apply_hysteresis_threshold(smooth, low_abs, high_abs)
return binary.astype(np.float64), f'Hysteresis (low={low}, high={high})'
else:
smooth = gaussian(img, sigma=1)
return smooth, 'Gaussian smoothed (σ=1)'
elif morphology == 'irregular':
morph_r = settings.get('organelle_morph_radius', 3)
smooth = gaussian(img, sigma=max(morph_r / 2, 1))
return smooth, f'Gaussian smoothed (σ={max(morph_r/2, 1):.1f})'
elif morphology == 'ring':
sigma_inner = settings.get('organelle_ring_sigma_inner', 1.0)
sigma_outer = settings.get('organelle_ring_sigma_outer', 3.0)
enhanced = np.abs(difference_of_gaussians(img_norm, sigma_inner, sigma_outer))
return enhanced, f'DoG ring enhancement (σ={sigma_inner}/{sigma_outer})'
else:
return img_norm, 'Normalised image'
[docs]
def debug(enabled=True, logger_name = None):
"""Decorator that temporarily sets the given logger to DEBUG for the wrapped call.
:param enabled: no-op when ``False``.
:param logger_name: logger name to tweak; defaults to the function's module logger.
:returns: decorator function.
"""
def decorator(func):
"""Inner decorator that binds the logger for ``func`` and returns the wrapper."""
log = logging.getLogger(logger_name or func.__module__)
@wraps(func)
def wrapper(*args, **kwargs):
"""Temporarily bump the logger to DEBUG while ``func`` runs, then restore its level."""
if not enabled:
return func(*args, **kwargs)
old_level = log.level
try:
log.setLevel(logging.DEBUG)
log.debug(">>> Entering %s", func.__name__)
result = func(*args, **kwargs)
log.debug("<<< Exiting %s", func.__name__)
return result
finally:
log.setLevel(old_level)
return wrapper
return decorator
def _generate_mask_random_cmap(mask):
"""Return a ``ListedColormap`` with a random color per label in ``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
#: The ``png_list`` object-id column each ``crop_mode`` writes, and the object
#: table whose rows that column identifies.
#:
#: One dict rather than a chain of ``if crop_mode ==`` because both directions
#: are needed and they must not drift: :func:`filepaths_to_database` writes the
#: column, and :func:`spacr.io._read_and_join_tables` has to work out, from a
#: database it did not write, which crop mode produced which rows. A database
#: measured with ``crop_mode=['cell','nucleus']`` carries **both** columns, each
#: NULL on the other mode's rows.
PNG_OBJECT_ID_COLUMNS = {
'cell': 'cell_id',
'nucleus': 'nucleus_id',
'pathogen': 'pathogen_id',
'cytoplasm': 'cytoplasm_id',
**{role: f'{role}_id' for role in schema.ORGANELLE_ROLES},
}
#: Reverse of :data:`PNG_OBJECT_ID_COLUMNS`.
PNG_CROP_MODE_BY_ID_COLUMN = {v: k for k, v in PNG_OBJECT_ID_COLUMNS.items()}
[docs]
def object_label_from_png_id(values):
"""Migrate ``png_list``'s ``'o<N>'`` text ids onto the integer object label.
``png_list`` stores an object id as **text** (``'o5'``) because it is the
last component of ``prcfo``; every object table stores the same object as
an **integer** ``object_label``, and the child tables store their parent as
an integer (in practice a float, since ``measure`` writes NaN for "no
overlapping cell") ``cell_id``. Two types for one identity, which is why a
plain SQL ``png_list.cell_id = nucleus.cell_id`` matches **zero rows**
rather than failing: SQLite compares a TEXT value with an INTEGER one by
type class, and text always sorts after numbers. Measured on a database
built by the real writers: 6 crops, 6 nuclei, 0 rows joined.
The integer is canonical — it is what the measurement tables key on — so
this is the one migration, applied on read. It replaces
``series.str[1:].astype(int)``, which crashed on four values the real
writers genuinely produce:
* ``'omulti'`` and ``'onone'`` — :func:`_generate_names` names a crop that
overlaps several cells ``..._multi.png`` and one that overlaps none
``..._none.png``. Both are ordinary outcomes of a real segmentation.
``ValueError: invalid literal for int() with base 10: 'multi'``;
* ``'error'`` — what :func:`_map_wells_png` writes for a name it cannot
parse. ``.str[1:]`` turned it into ``'rror'``, so the exception did not
even name the problem;
* ``NULL`` — every row of a *different* crop mode, in a database measured
with more than one. ``TypeError: int() argument must be ... not
'NoneType'``;
* an already-integer column, from a database whose ids were migrated
elsewhere: ``.str`` raises ``AttributeError`` on a numeric Series.
All four now come back as ``NaN``, which a caller can count and drop —
losing the crop's path for those objects, never the whole read.
:param values: a ``png_list`` object-id column (``cell_id``,
``nucleus_id``, ...), of any dtype.
:returns: a float ``Series`` of object labels, ``NaN`` where the id holds
no integer. Float rather than int because ``NaN`` has no int64.
"""
series = values if isinstance(values, pd.Series) else pd.Series(values)
if series.empty:
return pd.Series([], dtype=float, index=series.index)
return series.map(_one_object_label).astype(float)
def _one_object_label(value):
"""``'o5'`` / ``5`` / ``5.0`` -> ``5.0``; anything else -> ``NaN``.
The scalar half of :func:`object_label_from_png_id`. Numbers are taken
directly rather than routed through :func:`spacr.schema.object_index`,
which reads a *token*: ``str(5.0)`` is ``'5.0'`` and
:func:`spacr.schema.parse_int_token` deliberately refuses that (inventing
``3`` from ``3.7`` is the lie it exists to prevent). A whole-numbered float
in this column is not a fractional label, it is SQLite's REAL affinity, so
it is read as the label it is; a genuinely fractional one is ``NaN``.
"""
if value is None:
return np.nan
if isinstance(value, bool):
return np.nan
if isinstance(value, (int, np.integer)):
return float(value)
if isinstance(value, (float, np.floating)):
number = float(value)
if np.isnan(number) or not number.is_integer():
return np.nan
return number
parsed = schema.object_index(value)
return np.nan if parsed is None else float(parsed)
[docs]
def filepaths_to_database(img_paths, settings, source_folder, crop_mode):
"""Insert cropped PNG filepaths and parsed well/object IDs into the measurements DB.
:param img_paths: iterable of PNG paths for cropped objects.
:param settings: settings dict; ``timelapse`` toggles time_id parsing.
:param source_folder: experiment root; DB is written to ``measurements/measurements.db``.
:param crop_mode: a registered object role, including any organelle slot.
:returns: None.
"""
png_df = pd.DataFrame(img_paths, columns=['png_path'])
png_df['file_name'] = png_df['png_path'].apply(lambda x: os.path.basename(x))
parts = png_df['file_name'].apply(lambda x: pd.Series(_map_wells_png(x, timelapse=settings['timelapse'])))
columns = ['plateID', 'rowID', 'columnID', 'fieldID']
if settings['timelapse']:
columns = columns + ['timeID']
columns = columns + ['prcfo']
if crop_mode in PNG_OBJECT_ID_COLUMNS:
columns = columns + [PNG_OBJECT_ID_COLUMNS[crop_mode]]
png_df[columns] = parts
_append_to_measurements_db(
f'{source_folder}/measurements/measurements.db', 'png_list', png_df,
required=False, store=_measurement_store_for(
f'{source_folder}/measurements/measurements.db', settings))
[docs]
def activation_maps_to_database(img_paths, source_folder, settings):
"""Insert activation-map PNG paths and parsed well IDs into the dataset DB.
:param img_paths: iterable of PNG paths for activation-map images.
:param source_folder: experiment root; DB written to ``measurements/<dataset>.db``.
:param settings: settings dict; must contain ``dataset`` and ``cam_type``.
:returns: None.
"""
from .io import _create_database
png_df = pd.DataFrame(img_paths, columns=['png_path'])
png_df['file_name'] = png_df['png_path'].apply(lambda x: os.path.basename(x))
parts = png_df['file_name'].apply(lambda x: pd.Series(_map_wells_png(x, timelapse=False)))
columns = ['plateID', 'rowID', 'columnID', 'fieldID', 'prcfo', 'object']
png_df[columns] = parts
dataset_name = os.path.splitext(os.path.basename(settings['dataset']))[0]
database_name = f"{source_folder}/measurements/{dataset_name}.db"
if not os.path.exists(database_name):
_create_database(database_name)
try:
conn = sqlite3.connect(database_name, timeout=5)
png_df.to_sql(f"{settings['cam_type']}_list", conn, if_exists='append', index=False)
conn.commit()
except sqlite3.OperationalError as e:
print(f"SQLite error: {e}", flush=True)
traceback.print_exc()
[docs]
def activation_correlations_to_database(df, img_paths, source_folder, settings):
"""Merge per-image correlation stats with parsed well IDs and insert into the dataset DB.
:param df: DataFrame of correlation stats indexed by ``file_name``.
:param img_paths: iterable of PNG paths matching rows of ``df``.
:param source_folder: experiment root; DB written to ``measurements/<dataset>.db``.
:param settings: settings dict; must contain ``dataset`` and ``cam_type``.
:returns: None.
"""
from .io import _create_database
png_df = pd.DataFrame(img_paths, columns=['png_path'])
png_df['file_name'] = png_df['png_path'].apply(lambda x: os.path.basename(x))
parts = png_df['file_name'].apply(lambda x: pd.Series(_map_wells_png(x, timelapse=False)))
columns = ['plateID', 'rowID', 'columnID', 'fieldID', 'prcfo', 'object']
png_df[columns] = parts
png_df.set_index('file_name', inplace=True)
df.set_index('file_name', inplace=True)
merged_df = pd.concat([png_df, df], axis=1)
merged_df.reset_index(inplace=True)
dataset_name = os.path.splitext(os.path.basename(settings['dataset']))[0]
database_name = f"{source_folder}/measurements/{dataset_name}.db"
if not os.path.exists(database_name):
_create_database(database_name)
try:
conn = sqlite3.connect(database_name, timeout=5)
merged_df.to_sql(f"{settings['cam_type']}_correlations", conn, if_exists='append', index=False)
conn.commit()
except sqlite3.OperationalError as e:
print(f"SQLite error: {e}", flush=True)
traceback.print_exc()
[docs]
def calculate_activation_correlations(inputs, activation_maps, file_names, manders_thresholds=None):
"""Compute per-image Pearson and Manders correlations between input and activation channels.
:param inputs: input image batch, tensor of shape ``(B, C, H, W)``.
:param activation_maps: activation-map batch, tensor of shape ``(B, C, H, W)`` or ``(B, H, W)``.
:param file_names: file names corresponding to each image in the batch.
:param manders_thresholds: intensity percentiles used for Manders coefficients. Default ``[15, 50, 75]``.
:returns: DataFrame with one row per image and one column per channel-pair statistic.
"""
if manders_thresholds is None:
manders_thresholds = [15, 50, 75]
inputs = inputs.detach().cpu()
activation_maps = activation_maps.detach().cpu()
batch_size, in_channels, height, width = inputs.shape
if activation_maps.dim() == 3:
activation_maps = activation_maps.unsqueeze(1)
_, act_channels, act_height, act_width = activation_maps.shape
if (height != act_height) or (width != act_width):
activation_maps = torch.nn.functional.interpolate(activation_maps, size=(height, width), mode='bilinear')
correlations_dict = {'file_name': []}
for in_c in range(in_channels):
for act_c in range(act_channels):
correlations_dict[f'channel_{in_c}_activation_{act_c}_pearsons'] = []
for threshold in manders_thresholds:
correlations_dict[f'channel_{in_c}_activation_{act_c}_{threshold}_M1'] = []
correlations_dict[f'channel_{in_c}_activation_{act_c}_{threshold}_M2'] = []
for b in range(batch_size):
input_img = inputs[b]
activation_map = activation_maps[b]
correlations_dict['file_name'].append(file_names[b])
for in_c in range(in_channels):
input_raw = input_img[in_c].flatten().numpy()
for act_c in range(act_channels):
activation_raw = activation_map[act_c].flatten().numpy()
finite = np.isfinite(input_raw) & np.isfinite(activation_raw)
input_channel = input_raw[finite]
activation_channel = activation_raw[finite]
if input_channel.size > 0 and activation_channel.size > 0:
pearson_corr, _ = pearsonr(input_channel, activation_channel)
else:
pearson_corr = np.nan
correlations_dict[f'channel_{in_c}_activation_{act_c}_pearsons'].append(pearson_corr)
for threshold in manders_thresholds:
if input_channel.size > 0 and activation_channel.size > 0:
input_threshold = np.percentile(input_channel, threshold)
activation_threshold = np.percentile(activation_channel, threshold)
mask = (input_channel >= input_threshold) & (activation_channel >= activation_threshold)
if np.sum(mask) > 0:
manders_corr_M1 = np.sum(input_channel[mask] * activation_channel[mask]) / np.sum(input_channel[mask] ** 2)
manders_corr_M2 = np.sum(activation_channel[mask] * input_channel[mask]) / np.sum(activation_channel[mask] ** 2)
else:
manders_corr_M1 = np.nan
manders_corr_M2 = np.nan
else:
manders_corr_M1 = np.nan
manders_corr_M2 = np.nan
correlations_dict[f'channel_{in_c}_activation_{act_c}_{threshold}_M1'].append(manders_corr_M1)
correlations_dict[f'channel_{in_c}_activation_{act_c}_{threshold}_M2'].append(manders_corr_M2)
df_correlations = pd.DataFrame(correlations_dict)
return df_correlations
[docs]
def load_settings(csv_file_path, show=False, setting_key='setting_key', setting_value='setting_value'):
"""Reload a spacr settings CSV (written by :func:`save_settings`) back into a Python dict.
Every spacr pipeline persists its resolved settings alongside its
outputs so that a run can be reproduced. This helper re-parses that
CSV, coercing each value into its original Python type (``bool``,
``int``, ``float``, ``None``, ``list``, ``tuple``, ``dict``,
``str``).
:param csv_file_path: path to the CSV file.
:param show: display the raw DataFrame for debugging. Default ``False``.
:param setting_key: name of the key column. Default
``'setting_key'``; ``'Key'`` is accepted too, see below.
:param setting_value: name of the value column. Default
``'setting_value'``; ``'Value'`` is accepted too.
:returns: dict of parsed settings, ready to pass back into the
original pipeline entry point.
:raises ValueError: if the required key / value columns are missing.
THE TWO SPELLINGS. :func:`save_settings` writes ``Key`` / ``Value``,
while this function's defaults ask for ``setting_key`` /
``setting_value`` -- so the documented inverse pair did not round-trip,
and the example above raised. Callers had each worked around it
separately (``spacr/qt/dnd.py`` tries one spelling and catches the
failure to try the other), which is how it survived: nothing that used
the defaults was reading a file spacr had written.
Either spelling is now read. An explicitly named column still wins, so a
caller that knows its file's header is unaffected.
Example:
.. code-block:: python
from spacr.utils import load_settings
from spacr.core import preprocess_generate_masks
settings = load_settings('/data/plate01/settings/gen_mask_settings.csv')
preprocess_generate_masks(settings)
See Also:
:func:`save_settings` — inverse operation.
"""
df = tabular.read_table(csv_file_path, report=None)
if show:
display(df)
if setting_key not in df.columns or setting_value not in df.columns:
if 'Key' in df.columns and 'Value' in df.columns:
setting_key, setting_value = 'Key', 'Value'
else:
raise ValueError(
f"CSV file must contain {setting_key} and {setting_value} "
f"columns (or the Key/Value pair save_settings writes); "
f"{os.path.basename(str(csv_file_path))} has "
f"{list(df.columns)}.")
def parse_value(value):
"""Parse the string value into the appropriate Python data type."""
if pd.isna(value) or value == '':
return None
if not isinstance(value, str):
return value
if value.strip().lower() == 'true':
return True
if value.strip().lower() == 'false':
return False
if value.startswith(('(', '[', '{')):
try:
parsed_value = ast.literal_eval(value)
if isinstance(parsed_value, dict):
parsed_value = {k: parse_value(v) for k, v in parsed_value.items()}
return parsed_value
except (ValueError, SyntaxError):
pass
try:
if '.' in value:
return float(value)
return int(value)
except ValueError:
pass
return value
result_dict = {key: parse_value(value) for key, value in zip(df[setting_key], df[setting_value])}
return result_dict
[docs]
def console_encoding(stream=None):
"""Return the codec text printed to ``stream`` has to survive.
:param stream: a text stream; defaults to ``sys.stdout``.
:returns: a codec name, ``'utf-8'`` when the stream does not declare one
(a queue-backed GUI console, a StringIO, a captured pipe).
"""
import sys
if stream is None:
stream = getattr(sys, 'stdout', None)
return getattr(stream, 'encoding', None) or 'utf-8'
[docs]
def console_can_encode(text, stream=None):
"""Return ``True`` when ``text`` can be printed to ``stream`` as-is.
:param text: the string about to be printed.
:param stream: text stream to test against; defaults to ``sys.stdout``.
:returns: bool.
"""
try:
text.encode(console_encoding(stream))
except (UnicodeEncodeError, LookupError):
return False
return True
[docs]
def console_safe(text, stream=None):
"""Return ``text`` with anything the console cannot encode replaced by ``?``.
Console decoration must never be able to end a run. No Windows codepage
encodes spaCR's own output set -- ``▸`` (U+25B8) is absent from cp1252,
cp437, cp850, cp932 *and* cp936, and the box-drawing frame is absent from
cp1252 -- and neither does any of them encode the domain vocabulary that
ends up in settings values, such as the parental strain ``Δku80`` or a
``µm`` voxel size. Printing either to a non-UTF-8 stream raises
``UnicodeEncodeError``, and on Windows that is the normal case the moment
stdout is redirected: a batch-queue job, ``spacr-run``, a legacy console.
:param text: the string about to be printed.
:param stream: text stream to encode against; defaults to ``sys.stdout``.
:returns: ``text`` unchanged when it is printable, otherwise a lossy but
printable version of it.
"""
encoding = console_encoding(stream)
try:
text.encode(encoding)
except UnicodeEncodeError:
return text.encode(encoding, errors='replace').decode(encoding,
errors='replace')
except LookupError:
return text.encode('ascii', errors='replace').decode('ascii')
return text
#: Frame glyphs for :func:`pretty_print_settings`: the pretty set, and the
#: ASCII set used when the console cannot encode the pretty one.
_BOX_GLYPHS = {
'unicode': {'tl': '┌', 'tr': '┐', 'bl': '└', 'br': '┘',
'h': '─', 'v': '│', 'bullet': '▸', 'ellipsis': '…'},
'ascii': {'tl': '+', 'tr': '+', 'bl': '+', 'br': '+',
'h': '-', 'v': '|', 'bullet': '>', 'ellipsis': '...'},
}
[docs]
def pretty_print_settings(settings, title="Settings"):
"""Print a settings dict to the console as a tidy, aligned table.
Nicer than dumping a truncated pandas DataFrame: values are grouped by the
spacr settings categories, keys are aligned in a column, long values are
clipped, and the whole thing sits under a boxed title. Purely cosmetic --
used wherever "Saving settings" is shown.
Purely cosmetic, and it stays that way: the frame degrades to ASCII and
every line goes out through :func:`console_safe`, so a console that cannot
encode the decoration prints a plainer table instead of raising
``UnicodeEncodeError``. :func:`spacr.measure.measure_crop` calls this
(through :func:`save_settings`) before it does any work at all, so a
decoration character was enough to end a whole run before the first field
was read.
:param settings: the settings dict to render.
:param title: heading shown in the box.
:returns: None.
"""
try:
from .settings import categories
except Exception:
categories = {}
items = {k: settings[k] for k in settings}
key_w = min(38, max((len(str(k)) for k in items), default=10))
line_w = max(len(title) + 4, key_w + 46)
pretty = ''.join(_BOX_GLYPHS['unicode'].values())
g = _BOX_GLYPHS['unicode'] if console_can_encode(pretty) else _BOX_GLYPHS['ascii']
def _say(line):
"""Print a console-safe form of ``line`` and return ``None``."""
print(console_safe(line))
def _fmt(v):
"""Return ``v`` as text truncated to the table's value width."""
s = str(v)
return s if len(s) <= 44 else s[:41] + g['ellipsis']
def _row(k, v):
"""Return one padded key-and-formatted-value table row."""
return f" {str(k):<{key_w}} {_fmt(v)}"
bar = g['h'] * line_w
_say(f"{g['tl']}{bar}{g['tr']}")
_say(f"{g['v']} {title.ljust(line_w - 1)}{g['v']}")
_say(f"{g['bl']}{bar}{g['br']}")
shown = set()
for cat, keys in categories.items():
rows = [k for k in keys if k in items and k not in shown]
if not rows:
continue
_say(f"{g['bullet']} {cat}")
for k in rows:
_say(_row(k, items[k]))
shown.add(k)
leftover = [k for k in items if k not in shown]
if leftover:
if shown:
_say(f"{g['bullet']} Other")
for k in leftover:
_say(_row(k, items[k]))
print("")
[docs]
def save_settings(settings, name='settings', show=False):
"""Persist a settings dict to ``<src>/settings/<name>.csv`` so a spacr run can be reproduced later.
Called by every pipeline entry point to snapshot the resolved
settings before real work starts. The saved copy has ``test_mode``
and ``plot`` forced to ``False`` so that a downstream
:func:`load_settings` -> re-run produces a full, headless run.
:param settings: settings dict; must contain ``src``.
:param name: base filename (no extension); ``_list`` is appended
when ``src`` is a list. Default ``'settings'``.
:param show: display the DataFrame before writing. Default ``False``.
:returns: None. Writes ``<src>/settings/<name>.csv``.
Example:
.. code-block:: python
from spacr.utils import save_settings
save_settings(my_settings, name='my_experiment', show=True)
See Also:
:func:`load_settings` — inverse operation.
"""
settings_2 = settings.copy()
if isinstance(settings_2['src'], list):
src = settings_2['src'][0]
name = f"{name}_list"
else:
src = settings_2['src']
if 'test_mode' in settings_2.keys():
settings_2['test_mode'] = False
if 'plot' in settings_2.keys():
settings_2['plot'] = False
settings_df = pd.DataFrame(list(settings_2.items()), columns=['Key', 'Value'])
if show:
pretty_print_settings(settings_2, title=name.replace('_', ' ').title())
settings_csv = os.path.join(src,'settings',f'{name}.csv')
try:
os.makedirs(os.path.join(src,'settings'), exist_ok=True)
print(f"Saving settings to {settings_csv}")
settings_df.to_csv(settings_csv, index=False)
_save_settings_json(settings, os.path.splitext(settings_csv)[0] + '.json')
except (OSError, PermissionError) as e:
print(f"Warning: could not save settings to {settings_csv}: {e}. "
f"Continuing without writing the settings copy.")
def _save_settings_json(settings, path):
"""Write the same settings as JSON beside the CSV.
THE CSV LOSES THE TYPES. Every value in it is text, so a list comes back
as ``"[0, 1, 2]"``, ``None`` as ``""`` and ``False`` as ``"False"`` --
which is why loading one needs `ast.literal_eval` and a pile of special
cases, and why a settings file that round-trips through the panel is not
always the file that ran. JSON keeps the shape, so a results folder can
say exactly what produced it.
The CSV stays: it is what every existing loader reads and what a user
opens in a spreadsheet. This is a sibling, not a replacement.
Never raises. A settings copy that cannot be written is a note, not a
reason to lose a finished run.
"""
import json
def plain(value):
"""The value as something JSON can hold, or its repr."""
if value is None or isinstance(value, (bool, int, float, str)):
return value
if isinstance(value, dict):
return {str(k): plain(v) for k, v in value.items()}
if isinstance(value, (list, tuple, set)):
return [plain(v) for v in value]
try:
import numpy as _np
if isinstance(value, _np.generic):
return value.item()
except Exception: # noqa: BLE001
pass
return repr(value)
try:
with open(path, 'w', encoding='utf-8') as out:
json.dump({str(k): plain(v) for k, v in dict(settings).items()},
out, indent=2, sort_keys=True)
except Exception as error: # noqa: BLE001
print(f"Warning: could not write {path}: {error}")
[docs]
def print_progress(files_processed, files_to_process, n_jobs, time_ls=None, batch_size=None, operation_type=""):
"""Print a one-line progress report with an ETA derived from mean step time.
:param files_processed: number of items done (int or list).
:param files_to_process: total items to do (int or list).
:param n_jobs: parallelism used to compute ETA.
:param time_ls: list of per-step durations (seconds) for ETA; ``None`` skips ETA.
:param batch_size: batch size when ``time_ls`` is per batch rather than per image.
:param operation_type: label printed alongside the progress line.
:returns: None.
"""
if isinstance(files_processed, list):
files_processed = len(set(files_processed))
if isinstance(files_to_process, list):
files_to_process = len(set(files_to_process))
if isinstance(batch_size, list):
batch_size = len(batch_size)
if not isinstance(files_processed, int):
try:
files_processed = int(files_processed)
except Exception:
files_processed = 0
if not isinstance(files_to_process, int):
try:
files_to_process = int(files_to_process)
except Exception:
files_to_process = 0
time_info = ""
if time_ls is not None:
average_time = np.mean(time_ls) if len(time_ls) > 0 else 0
try:
effective_jobs = max(1, int(n_jobs))
except (TypeError, ValueError):
effective_jobs = 1
remaining = max(0, files_to_process - files_processed)
time_left = (remaining * average_time / effective_jobs) / 60
if batch_size is None:
time_info = f'Time/image: {average_time:.3f}sec, Time_left: {time_left:.3f} min.'
else:
try:
effective_batch_size = max(1, int(batch_size))
except (TypeError, ValueError):
effective_batch_size = 1
average_time_img = average_time / effective_batch_size
time_info = f'Time/batch: {average_time:.3f}sec, Time/image: {average_time_img:.3f}sec, Time_left: {time_left:.3f} min.'
else:
time_info = None
print(f'Progress: {files_processed}/{files_to_process}, operation_type: {operation_type}, {time_info}')
[docs]
def reset_mp():
"""Set the multiprocessing start method appropriate for the current OS.
Uses ``spawn`` on Windows and ``fork`` on Linux/macOS.
:returns: None.
"""
current_method = get_start_method()
system = platform.system()
if system == 'Windows':
if current_method != 'spawn':
set_start_method('spawn', force=True)
elif system in ('Linux', 'Darwin'):
if current_method != 'fork':
set_start_method('fork', force=True)
[docs]
def is_multiprocessing_process(process):
"""Return ``True`` if ``process`` cmdline contains ``multiprocessing``.
:param process: process object exposing the :mod:`psutil` ``cmdline`` API.
"""
try:
for cmd in process.cmdline():
if 'multiprocessing' in cmd:
return True
except (psutil.NoSuchProcess, psutil.AccessDenied, psutil.ZombieProcess):
pass
return False
[docs]
def close_file_descriptors():
"""Close file descriptors from 3 up to the soft NOFILE limit."""
import resource
soft, hard = resource.getrlimit(resource.RLIMIT_NOFILE)
for fd in range(3, soft):
try:
os.close(fd)
except OSError:
pass
[docs]
def close_multiprocessing_processes():
"""Terminate all detected multiprocessing child processes and close file descriptors."""
current_pid = os.getpid()
for proc in psutil.process_iter(['pid', 'cmdline']):
try:
if proc.info['pid'] == current_pid:
continue
if is_multiprocessing_process(proc):
proc.terminate()
proc.wait(timeout=5)
print(f"Terminated process {proc.info['pid']}")
except (psutil.NoSuchProcess, psutil.AccessDenied, psutil.ZombieProcess) as e:
print(f"Failed to terminate process {proc.info['pid']}: {e}")
close_file_descriptors()
[docs]
def check_mask_folder(src, mask_fldr, resume=False):
"""Return ``True`` if masks in ``src/masks/mask_fldr`` still need generating.
:param src: experiment root containing ``masks/`` and ``stack/`` subfolders.
:param mask_fldr: subfolder name under ``masks/``.
:param resume: accepted for the callers that pass it. Only structurally
complete mask arrays are counted whether or not it is set, so an
empty or truncated array left by an interrupted run is re-queued.
:returns: ``True`` when the mask folder is missing or any expected stack
filename lacks a complete mask. Unrelated masks cannot substitute
for a missing field and do not force complete fields to run again.
"""
from .io import _listdir_visible
from .resume import validate_merged_field
mask_folder = os.path.join(src,'masks',mask_fldr)
stack_folder = os.path.join(src,'stack')
if not os.path.exists(mask_folder):
return True
expected = [file for file in _listdir_visible(stack_folder)
if file.endswith('.npy')]
if all(validate_merged_field(os.path.join(mask_folder, file))[0]
for file in expected):
print(f'All masks have been generated for {mask_fldr}')
return False
else:
return True
[docs]
def smooth_hull_lines(cluster_data):
"""Return the x, y coordinates of a smoothed convex-hull outline of a 2-D point set.
:param cluster_data: 2-D array of point coordinates.
:returns: tuple ``(x, y)`` of spline-interpolated hull coordinates (100 samples).
"""
hull = ConvexHull(cluster_data)
vertices = hull.points[hull.vertices]
vertices = np.vstack([vertices, vertices[0, :]])
tck, u = splprep(vertices.T, u=None, s=0.0)
new_points = splev(np.linspace(0, 1, 100), tck)
return new_points[0], new_points[1]
def _gen_rgb_image(image, channels):
"""Return an ``(H, W, 3)`` RGB image built from selected channels of ``image``."""
rgb_image = np.zeros((image.shape[0], image.shape[1], 3), dtype=np.float32)
for i, chan in enumerate(channels):
if chan < image.shape[2]:
rgb_image[:, :, i] = image[:, :, chan]
return rgb_image
def _outline_and_overlay(image, rgb_image, mask_dims, outline_colors, outline_thickness):
"""Draw each mask's outline over the RGB image.
DRAWN ON THE CALLING THREAD, DELIBERATELY. This used to run in a thread
pool, which aborted the whole process -- SIGABRT, core dumped, no
traceback -- once Qt and Tk had both been initialised in the same
session: cv2 and skimage's contour code are not safe to call off the
main thread with two GUI toolkits resident, and there is nothing to
catch.
Giving the pool up cost almost nothing. There are at most three mask
dimensions and ``find_contours`` holds the GIL throughout, so the
threads bought 3-5% -- measured at 1,257 ms serial against 1,222 ms
threaded on 3x60 objects at 1024 px, and 17.3 s against 16.5 s on 3x200
at 2048 px. A 1.03x speedup is not worth a core dump.
:param image: the merged array, image channels then mask slices.
:param rgb_image: the base to draw over.
:param mask_dims: which slices hold the masks to outline.
:param outline_colors: one colour per mask dimension, cycled.
:param outline_thickness: the outline width, in pixels.
:returns: the overlaid image, the outlines, and the input array.
"""
outlines = []
overlayed_image = rgb_image.copy()
def process_dim(mask_dim):
"""Return a dilated outline image of the labeled mask at ``image[..., mask_dim]``."""
mask = np.take(image, mask_dim, axis=-1)
outline = np.zeros_like(mask, dtype=np.uint8)
for j in np.unique(mask):
if j == 0:
continue
contours = find_contours(mask == j, 0.5)
cv_contours = [np.flip(contour.astype(int), axis=1) for contour in contours]
cv2.drawContours(outline, cv_contours, -1, color=255, thickness=outline_thickness)
return dilation(outline, _square_footprint(outline_thickness))
outlines = [process_dim(mask_dim) for mask_dim in mask_dims]
for i, outline in enumerate(outlines):
color = np.array(outline_colors[i % len(outline_colors)])
for j in np.unique(outline):
if j == 0:
continue
mask = outline == j
overlayed_image[mask] = color
return overlayed_image, outlines, image
def _convert_cq1_well_id(well_id):
"""Convert a linear well index to the CQ1 ``<row_letter><col>`` well format.
24 columns per row is the CQ1's own layout, not an assumption about the
plate, so it stays. What changed is the row letter: ``chr(ord('A') + n)``
walked off the end of the alphabet, so index 1536 came back as
``'\\x8024'`` — a control character where a row label should be.
:func:`spacr.schema.well_id` is bijective base 26 and stays inside the
alphabet however far the index runs.
A token that is not a 1-based index is returned unchanged rather than
converted. The old arithmetic turned index ``0`` into ``'@24'`` — a well
name with a punctuation mark for a row — and, now that
``_extract_filename_metadata`` keeps an unreadable well token instead of
substituting ``'0'``, it would be handed things like ``'1a'``. Keeping the
token leaves two odd wells as two odd wells and never invents a name.
:param well_id: 1-based linear well index.
:returns: the well name, e.g. ``1`` -> ``'A01'``, ``384`` -> ``'P24'``;
or ``str(well_id)`` when it names no well.
"""
index = schema.parse_int_token(well_id, allow_prefix=False)
if index is None or index < 1:
print(f'Not a CQ1 well index: {well_id!r}; keeping it as it is',
flush=True)
return str(well_id)
row, col = divmod(index - 1, 24)
return schema.well_id(row + 1, col + 1)
def _get_cellpose_batch_size():
"""Choose a Cellpose batch size from the GPU's VRAM.
The bounds form an EXHAUSTIVE ladder. The previous ``> 8 and < 12``
style left 8.0, 12.0 and 24.0 GB unmatched, so the batch size was never
assigned, the print below raised ``UnboundLocalError``, and a bare
``except`` silently turned that into a batch size of 8 -- a card with 24
GB quietly running at the smallest batch.
:returns: the batch size, and 8 when there is no CUDA device or its
memory cannot be inspected.
"""
try:
if torch.cuda.is_available():
device_properties = torch.cuda.get_device_properties(0)
vram_gb = device_properties.total_memory / (1024**3)
else:
print("CUDA is not available. Please check your installation and GPU.")
return 8
if vram_gb < 8:
batch_size = 8
elif vram_gb < 12:
batch_size = 16
elif vram_gb < 24:
batch_size = 48
else:
batch_size = 96
print(f"Device {0}: {device_properties.name}, VRAM: {vram_gb:.2f} GB, cellpose batch size: {batch_size}")
return batch_size
except Exception:
LOG.warning(
"Could not inspect CUDA memory; using Cellpose batch size 8",
exc_info=True,
)
return 8
def _extract_filename_metadata(filenames, src, regular_expression, metadata_type='cellvoyager'):
"""Group image paths by the metadata their filenames carry.
Zero padding is undone so ``001`` and ``1`` are one key, through
``_int_or_token``, which KEEPS a token it cannot read rather than
substituting ``0`` -- every unreadable well used to collapse onto well
``0``.
A filename the regex cannot read is reported and skipped rather than
raising, so one odd name does not cost the plate.
:param filenames: the names to parse.
:param src: the folder they are in; also the fallback plate name when
the pattern has no plate group.
:param regular_expression: the compiled pattern.
:param metadata_type: the microscope convention; ``'cq1'`` also converts
the well id, whose scheme differs from the well name it prints.
:returns: ``{(plate, well, field, channel, time, slice): [paths]}``.
"""
images_by_key = defaultdict(list)
for filename in filenames:
match = regular_expression.match(filename)
if match:
try:
try:
plate = match.group('plateID')
except Exception:
plate = os.path.basename(src)
well = match.group('wellID')
if well[0].isdigit():
well = _int_or_token(well)
field = match.group('fieldID')
if field[0].isdigit():
field = _int_or_token(field)
channel = match.group('chanID')
if channel[0].isdigit():
channel = _int_or_token(channel)
if 'timeID' in match.groupdict():
timeID = match.group('timeID')
if timeID[0].isdigit():
timeID = _int_or_token(timeID)
else:
timeID = None
if 'sliceID' in match.groupdict():
sliceID = match.group('sliceID')
if sliceID[0].isdigit():
sliceID = _int_or_token(sliceID)
else:
sliceID = None
if metadata_type =='cq1':
orig_well = well
well = _convert_cq1_well_id(well)
print(f'Converted Well ID: {orig_well} to {well}', end='\r', flush=True)
key = (plate, well, field, channel, timeID, sliceID)
file_path = os.path.join(src, filename)
images_by_key[key].append(file_path)
except IndexError:
print(f"Could not extract information from filename {filename} using provided regex")
else:
print(f"Filename {filename} did not match provided regex: {regular_expression}")
continue
return images_by_key
[docs]
def mask_object_count(mask):
"""Return the number of nonzero labeled objects in ``mask``.
:param mask: integer label image where ``0`` is background. The count is the
number of *distinct* nonzero values, so label IDs need not be contiguous
and gaps left by filtering are not counted. A purely binary mask
therefore reports ``1`` however many blobs it contains.
"""
unique_labels = np.unique(mask)
num_objects = len(unique_labels[unique_labels!=0])
return num_objects
def _update_database_with_merged_info(db_path, df, table='png_list', columns=None):
"""Merge extra columns from ``df`` into ``table`` on ``prcfo`` and rewrite the table."""
if columns is None:
columns = ['pathogen', 'treatment', 'host_cells', 'condition', 'prcfo']
conn = sqlite3.connect(db_path, timeout=30)
try:
existing_df = tabular.read_database(
db_path, [table], report=None, migrate=False)[0]
except Exception as e:
print(f"Failed to read table {table} from database: {e}")
conn.close()
return
if 'prcfo' not in df.columns:
print(f'generating prcfo columns')
try:
df['prcfo'] = df['plateID'].astype(str) + '_' + df['rowID'].astype(str) + '_' + df['columnID'].astype(str) + '_' + df['fieldID'].astype(str) + '_o' + df['object_label'].astype(int).astype(str)
except Exception:
print('Merging on cell failed, trying with cell_id')
try:
df['prcfo'] = df['plateID'].astype(str) + '_' + df['rowID'].astype(str) + '_' + df['columnID'].astype(str) + '_' + df['fieldID'].astype(str) + '_o' + df['cell_id'].astype(int).astype(str)
except Exception as e:
print(e)
try:
merged_df = pd.merge(
existing_df,
df[columns],
on='prcfo',
how='left',
validate='many_to_one',
)
except pd.errors.MergeError:
conn.close()
raise
try:
conn.execute(f"DROP TABLE IF EXISTS {table}")
merged_df.to_sql(table, conn, index=False)
print(f"Table {table} successfully updated in the database.")
except Exception as e:
print(f"Failed to update table {table} in the database: {e}")
finally:
conn.close()
def _generate_representative_images(db_path, cells=None, cell_loc=None, pathogens=None, pathogen_loc=None, treatments=None, treatment_loc=None, channel_of_interest=1, compartments = None, measurement = 'mean_intensity', nr_imgs=16, channel_indices=None, um_per_pixel=0.1, scale_bar_length_um=10, plot=False, fontsize=12, show_filename=True, channel_names=None, update_db=True):
"""Save representative-image grids per condition selected by a compartment measurement ratio."""
if cells is None:
cells = ['HeLa']
if pathogens is None:
pathogens = ['rh']
if treatments is None:
treatments = ['cm']
if compartments is None:
compartments = ['pathogen','cytoplasm']
if channel_indices is None:
channel_indices = [0,1,2]
from .io import _read_and_join_tables, _save_figure
from .plot import _plot_images_on_grid
df = _read_and_join_tables(db_path)
df = annotate_conditions(df, cells, cell_loc, pathogens, pathogen_loc, treatments, treatment_loc)
if update_db:
_update_database_with_merged_info(db_path, df, table='png_list', columns=['pathogen', 'treatment', 'host_cells', 'condition', 'prcfo'])
def _compartment_column(compartment):
"""Return the selected compartment series, or raise for a missing column."""
suffix = f'_channel_{channel_of_interest}_{measurement}'
col = f'{compartment}{suffix}'
if col not in df.columns:
available = sorted({c.split('_channel_')[0] for c in df.columns
if c.endswith(suffix)})
raise KeyError(
f"compartment {compartment!r} has no column {col!r} in the "
f"joined measurement tables. Available compartments for "
f"channel {channel_of_interest} / measurement {measurement!r}: "
f"{', '.join(available) if available else '(none)'}")
return df[col]
if isinstance(compartments, list):
if len(compartments) > 1:
df['new_measurement'] = (_compartment_column(compartments[0])
/ _compartment_column(compartments[1]))
elif len(compartments) == 1:
df['new_measurement'] = _compartment_column(compartments[0])
else:
df['new_measurement'] = df['cell_area']
else:
df['new_measurement'] = df['cell_area']
dfs = {condition: df_group for condition, df_group in df.groupby('condition')}
conditions = df['condition'].dropna().unique().tolist()
for condition in conditions:
df = dfs[condition]
df = _filter_closest_to_stat(df, column='new_measurement', n_rows=nr_imgs, use_median=False)
png_paths_by_condition = df['png_path'].tolist()
fig = _plot_images_on_grid(png_paths_by_condition, channel_indices, um_per_pixel, scale_bar_length_um, fontsize, show_filename, channel_names, plot)
src = os.path.dirname(db_path)
os.makedirs(src, exist_ok=True)
_save_figure(fig=fig, src=src, text=condition)
for channel in channel_indices:
fig = _plot_images_on_grid(png_paths_by_condition, [channel], um_per_pixel, scale_bar_length_um, fontsize, show_filename, channel_names, plot)
_save_figure(fig, src, text=f'channel_{channel}_{condition}')
plt.close()
def _map_values(row, values, locs):
"""Look up the value assigned to the row/column identifier in ``row``."""
if locs:
value_dict = {loc: value for value, loc_list in zip(values, locs) for loc in loc_list}
type_ = 'rowID' if locs[0][0][0] == 'r' else 'columnID'
return value_dict.get(row[type_], None)
return values[0] if values else None
[docs]
def is_list_of_lists(var):
"""Return ``True`` if ``var`` is a list whose every element is also a list.
:param var: value to test, including an empty or nested list.
"""
if isinstance(var, list) and all(isinstance(i, list) for i in var):
return True
return False
[docs]
def normalize_to_dtype(array, p1=2, p2=98, percentile_list=None, new_dtype=None):
"""Percentile-normalize each channel of an image stack into the target dtype range.
:param array: input stack of shape ``(H, W, C)``.
:param p1: lower percentile. Default ``2``.
:param p2: upper percentile. Default ``98``.
:param percentile_list: per-channel ``(low, high)`` pairs; overrides ``p1``/``p2``.
:param new_dtype: target dtype (``np.uint8``/``np.uint16`` or their string forms).
:returns: normalized stack with the same shape as ``array``.
"""
if new_dtype is None:
out_range = (0, np.iinfo(array.dtype).max)
elif new_dtype in [np.uint8, np.uint16]:
out_range = (0, np.iinfo(new_dtype).max)
elif new_dtype in ['uint8', 'uint16']:
new_dtype = np.uint8 if new_dtype == 'uint8' else np.uint16
out_range = (0, np.iinfo(new_dtype).max)
else:
out_range = (0, np.iinfo(array.dtype).max)
nimg = array.shape[2]
new_stack = np.empty_like(array, dtype=array.dtype)
for i in range(nimg):
img = array[:, :, i]
non_zero_img = img[img > 0]
if not percentile_list is None:
percentiles = percentile_list[i]
else:
percentile_1 = p1
percentile_2 = p2
if percentile_list is None:
if non_zero_img.size > 0:
img_min = np.percentile(non_zero_img, percentile_1)
img_max = np.percentile(non_zero_img, percentile_2)
else:
img_min = np.percentile(img, percentile_1)
img_max = np.percentile(img, percentile_2)
else:
img_min = percentiles[0]
img_max = percentiles[1]
img = rescale_intensity(img, in_range=(img_min, img_max), out_range=out_range)
new_stack[:, :, i] = img
return new_stack
def _list_endpoint_subdirectories(base_dir):
"""Return leaf subdirectory paths under ``base_dir``, excluding any named ``figure``."""
endpoint_subdirectories = []
for root, dirs, _ in os.walk(base_dir):
if not dirs:
endpoint_subdirectories.append(root)
endpoint_subdirectories = [path for path in endpoint_subdirectories if os.path.basename(path) != 'figure']
return endpoint_subdirectories
def _generate_names(file_name, cell_id, cell_nucleus_ids, cell_pathogen_ids,
source_folder, crop_mode='cell', timelapse=None,
object_id=None):
"""Build the ``(image_name, folder_path, table_name)`` tuple for a cropped object."""
file_name = schema.escape_field_stem_plate(
file_name, timelapse=bool(timelapse))
non_zero_cell_ids = cell_id[cell_id != 0]
cell_id_str = "multi" if non_zero_cell_ids.size > 1 else str(non_zero_cell_ids[0]) if non_zero_cell_ids.size == 1 else "none"
cell_nucleus_ids = cell_nucleus_ids[cell_nucleus_ids != 0]
cell_nucleus_id_str = "multi" if cell_nucleus_ids.size > 1 else str(cell_nucleus_ids[0]) if cell_nucleus_ids.size == 1 else "none"
cell_pathogen_ids = cell_pathogen_ids[cell_pathogen_ids != 0]
cell_pathogen_id_str = "multi" if cell_pathogen_ids.size > 1 else str(cell_pathogen_ids[0]) if cell_pathogen_ids.size == 1 else "none"
object_ids = np.atleast_1d(object_id)
object_ids = object_ids[
np.array([value is not None and value != 0 for value in object_ids],
dtype=bool)]
object_id_str = ("multi" if object_ids.size > 1 else
str(object_ids[0]) if object_ids.size == 1 else "none")
fldr = f"{source_folder}/data/"
img_name = ""
if crop_mode == 'nucleus':
img_name = f"{file_name}_{cell_id_str}_{cell_nucleus_id_str}.png"
fldr += "single_nucleus/" if cell_nucleus_ids.size == 1 else "multiple_nucleus/" if cell_nucleus_ids.size > 1 else "no_nucleus/"
fldr += "single_pathogen/" if cell_pathogen_ids.size == 1 else "multiple_pathogens/" if cell_pathogen_ids.size > 1 else "uninfected/"
elif crop_mode == 'pathogen':
img_name = f"{file_name}_{cell_id_str}_{cell_pathogen_id_str}.png"
fldr += "single_nucleus/" if cell_nucleus_ids.size == 1 else "multiple_nucleus/" if cell_nucleus_ids.size > 1 else "no_nucleus/"
fldr += "infected/" if cell_pathogen_ids.size >= 1 else "uninfected/"
elif crop_mode in ('cell', 'cytoplasm'):
img_name = f"{file_name}_{cell_id_str}.png"
fldr += "single_nucleus/" if cell_nucleus_ids.size == 1 else "multiple_nucleus/" if cell_nucleus_ids.size > 1 else "no_nucleus/"
fldr += "single_pathogen/" if cell_pathogen_ids.size == 1 else "multiple_pathogens/" if cell_pathogen_ids.size > 1 else "uninfected/"
elif crop_mode in schema.ORGANELLE_ROLES:
img_name = f"{file_name}_{object_id_str}.png"
fldr += "single_nucleus/" if cell_nucleus_ids.size == 1 else "multiple_nucleus/" if cell_nucleus_ids.size > 1 else "no_nucleus/"
fldr += "single_pathogen/" if cell_pathogen_ids.size == 1 else "multiple_pathogens/" if cell_pathogen_ids.size > 1 else "uninfected/"
else:
raise ValueError(
f"_generate_names has no naming rule for crop_mode={crop_mode!r}. "
f"Known crop modes: {', '.join(schema.ALL_ROLES)}.")
parts = file_name.split('_')
plate = parts[0]
well = parts[1]
if timelapse:
timeID = parts[2]
metadata = f'{plate}_{well}_{timeID}'
else:
metadata = f'{plate}_{well}'
fldr = os.path.join(fldr,metadata)
table_name = fldr.replace("/", "_")
return img_name, fldr, table_name
def _find_bounding_box(crop_mask, _id, buffer=10):
"""Return a mask with the padded bounding box of ``_id`` filled with ``_id``."""
object_indices = np.where(crop_mask == _id)
y_min, y_max = object_indices[0].min(), object_indices[0].max()
x_min, x_max = object_indices[1].min(), object_indices[1].max()
y_min = max(y_min - buffer, 0)
y_max = min(y_max + buffer, crop_mask.shape[0] - 1)
x_min = max(x_min - buffer, 0)
x_max = min(x_max + buffer, crop_mask.shape[1] - 1)
new_mask = np.zeros_like(crop_mask)
new_mask[y_min:y_max+1, x_min:x_max+1] = _id
return new_mask
#: Tables whose rows are child objects and therefore carry a parent-cell link.
#: 'organelle' is here because measure._morphological_measurements maps each
#: organelle to its enclosing cell, exactly as it does for nucleus and pathogen.
_CHILD_OBJECT_TABLES = schema.CHILD_OBJECT_TABLES
#: Tables whose rows are parent objects summarised over their organelles. The
#: row IS the parent, so object_label is the only key it needs — the same key
#: set as 'cell'. Written by measure._summarize_organelles_per_parent.
_ORGANELLE_SUMMARY_TABLES = schema.ORGANELLE_SUMMARY_TABLES
#: Tables whose rows are top-level objects with no parent link.
_PARENT_OBJECT_TABLES = schema.PARENT_OBJECT_TABLES
[docs]
class MeasurementUnitsMismatch(ValueError):
"""A measurement frame's units differ from the ones already in the table.
A 2-D field measures areas in px^2; a 3-D field measures volumes, in voxels
or um^3, and writes them into the *same* ``<object>_area`` column, because
that column is read by name by every downstream selector, model and
threshold ever written against a spaCR database and renaming it would break
all of them silently. Appending both into one table would therefore leave a
numeric column that mixes two incompatible quantities with nothing in the
row to tell them apart, which no amount of downstream care could recover
from. So it is refused here instead.
"""
#: What an unstamped row is taken to be. Every spaCR release before 3-D
#: measurement existed could only write 2-D pixel measurements -- a 3-D mask
#: crashed the morphology pass outright -- so a row with no stamp is 2-D/px as
#: a matter of fact, not as an assumption.
_LEGACY_STAMP = (2, 'px')
def _stamp_identity(stamp):
"""Reduce a stamp dict to the ``(ndim, units)`` pair the table is keyed on."""
if not stamp:
return _LEGACY_STAMP
ndim = stamp.get('measurement_ndim')
units = stamp.get('measurement_units')
if ndim is None or units is None:
return _LEGACY_STAMP
return (int(ndim), str(units))
def _existing_measurement_identity(db_path, table):
"""Return the ``(ndim, units)`` pairs already present in ``table``.
Retried like the append it guards: with many workers writing one
``measurements.db`` (above all on a network share) a read can find the
database locked, and failing the field there loses it for nothing.
:param db_path: path to ``measurements.db``.
:param table: object table name.
:returns: see :func:`_existing_measurement_identity_once`.
:raises sqlite3.OperationalError: when every attempt finds it locked.
"""
delay = 0.2
attempt = 1
while True:
try:
return _existing_measurement_identity_once(db_path, table)
except sqlite3.OperationalError as e:
if 'locked' not in str(e).lower() or attempt == DB_WRITE_ATTEMPTS:
raise
time.sleep(delay)
delay *= 2
attempt += 1
def _existing_measurement_identity_once(db_path, table):
"""Return the ``(ndim, units)`` pairs already present in ``table``.
:returns: a set of pairs; empty when the database or table does not exist
yet. Rows written before the stamp existed, and rows whose stamp is
NULL, count as :data:`_LEGACY_STAMP`.
"""
if not os.path.isfile(db_path):
return set()
from .database_concurrency import connect
conn = connect(db_path, readonly=True, timeout=DB_WRITE_TIMEOUT)
try:
exists = conn.execute(
"SELECT 1 FROM sqlite_master WHERE type='table' AND name=?",
(table,)).fetchone()
if not exists:
return set()
have = {row[1] for row in conn.execute(f'PRAGMA table_info("{table}")')}
if not {'measurement_ndim', 'measurement_units'} <= have:
row = conn.execute(f'SELECT 1 FROM "{table}" LIMIT 1').fetchone()
return {_LEGACY_STAMP} if row else set()
rows = conn.execute(
'SELECT DISTINCT measurement_ndim, measurement_units '
f'FROM "{table}"').fetchall()
finally:
conn.close()
found = set()
for ndim, units in rows:
if ndim is None or units is None:
found.add(_LEGACY_STAMP)
else:
found.add((int(ndim), str(units)))
return found
def _assert_measurement_units_compatible(db_path, table, stamp):
"""Refuse to append rows whose units differ from the table's.
:param db_path: path to measurements.db.
:param table: destination table.
:param stamp: the stamp about to be written, or ``None`` for a caller that
supplies none (treated as :data:`_LEGACY_STAMP`, i.e. 2-D pixels).
:raises MeasurementUnitsMismatch: when the table already holds rows in
different units.
"""
incoming = _stamp_identity(stamp)
existing = _existing_measurement_identity(db_path, table)
other = existing - {incoming}
if not other:
return
describe = ', '.join(
f"{n}-D/{u}" for n, u in sorted(other, key=lambda p: (p[0], p[1])))
raise MeasurementUnitsMismatch(
f"refusing to append {incoming[0]}-D/{incoming[1]} rows to the "
f"'{table}' table of {db_path}, which already holds {describe} rows. "
f"A 2-D field writes a px^2 area into <object>_area and a 3-D field "
f"writes a volume into the same column, so one table cannot hold both "
f"without the numbers becoming uncomparable in a way no reader could "
f"detect. Measure the 3-D and 2-D fields into separate output folders, "
f"or re-measure the whole plate one way.")
[docs]
class ImportedCopyNotReleased(ValueError):
"""The table being appended to holds an import's copy of the same field.
``foreign.run_import`` copies the imported frame into the canonical
``cell`` / ``nucleus`` / ``pathogen`` table when the destination is empty,
so a project built purely by import is readable by every spaCR tool. That
copy is a convenience and it stops being one the moment spaCR measures the
same field: :func:`_merge_and_save_to_database` *appends*, so its rows land
beside theirs, in different columns, with nothing in the row marking the
seam -- and every ``count_cell`` downstream becomes the sum of two
populations.
A resume supersedes the copy before measuring (see
:func:`spacr.resume.supersede_imported_copies`) or refuses and says so,
which covers every path *through* a resume. This covers the path around
one: a direct ``measure_crop`` with ``resume`` off.
Raised only when the copy cannot be handed back **provably losslessly** --
no ``foreign_<object>`` to check the rows against, a row in the copy with
no twin in it, a timelapse whose frames the importer never keyed, or a
delete that did not act on the rows the count cleared. In every one of
those cases the field's rows are not written, because a refused write can
be re-run and a mixed table cannot be un-mixed.
"""
def _field_key_predicate(frame, key_columns, alias):
"""A WHERE fragment matching exactly the field identities ``frame`` carries.
One parenthesised ``OR`` of ``AND`` groups over the four key columns, plus
its parameters, so that the count and the delete below can be handed *one*
predicate rather than two statements that could drift apart.
Naming ``rowID`` here is safe, and is worth saying out loud given what that
identifier has cost this project twice: it is used as one of four *key*
columns, always together, quoted and alias-qualified, and it means the
plate row -- which is exactly what it is. The destructive spelling was
``rowid``, the implicit row identity, which a declared ``rowID`` shadows.
:param frame: the rows about to be appended.
:param key_columns: :data:`spacr.schema.FIELD_KEY_COLUMNS`.
:param alias: table alias every column reference is qualified by.
:returns: ``(predicate, params)``.
:raises ImportedCopyNotReleased: when ``frame`` lacks a key column, so the
identity of what is being written cannot be established.
"""
missing = [c for c in key_columns if c not in frame.columns]
if missing:
raise ImportedCopyNotReleased(
f"cannot establish which field these rows belong to: the frame "
f"has no {missing} column(s), and a delete keyed on fewer columns "
f"than the writer used would take other fields with it.")
keys = list(dict.fromkeys(
tuple(row) for row in frame[list(key_columns)].astype(str).itertuples(
index=False, name=None)))
if not keys:
return '0', []
group = '(' + ' AND '.join(
f'{alias}."{c}" = ?' for c in key_columns) + ')'
predicate = '(' + ' OR '.join([group] * len(keys)) + ')'
params = [value for key in keys for value in key]
return predicate, params
def _verified_delete(conn, table, alias, predicate, params, what):
"""Count with a predicate, delete with the *same* predicate, verify.
The shape of :func:`spacr.data_manager._verified_write`, which exists for
the same reason: this project has been destroyed twice by a delete written
against a row identity that was not one. ``DELETE ... WHERE rowid IN (...)``
removed a whole table because every spaCR object table declares a column
called ``rowID`` and SQLite identifiers are case-insensitive; the obvious
repair -- delete by the declared key -- was equally destructive, because an
import's row and a measurement's row for one object share all five key
columns.
So no row identity is named. The caller supplies one predicate string; it
is interpolated once, into both statements, so a later edit cannot change
one without the other, and any difference between the two numbers is a
failure rather than a result.
:param conn: open connection, inside a transaction.
:param table: table to delete from.
:param alias: alias bound to ``table`` in both statements -- the predicate
qualifies its column references by it.
:param predicate: the WHERE clause, without ``WHERE``.
:param params: its parameters, bound to both statements.
:param what: what this delete is, for the error message.
:returns: rows removed, which equals the rows counted.
:raises ImportedCopyNotReleased: on any difference. The caller's
transaction rolls back and nothing is written -- not the delete, and
not the measurements that were to follow it.
"""
counted = int(conn.execute(
f'SELECT COUNT(*) FROM "{table}" AS {alias} WHERE {predicate}',
tuple(params)).fetchone()[0])
removed = int(conn.execute(
f'DELETE FROM "{table}" AS {alias} WHERE {predicate}',
tuple(params)).rowcount or 0)
if removed != counted:
raise ImportedCopyNotReleased(
f"refusing to {what}: the delete removed {removed} row(s) from "
f"'{table}' where the count that gated it said {counted}. The "
f"statement did not act on the rows that were checked, so nothing "
f"about the result can be trusted. The transaction was rolled "
f"back, and this field's measurements were not written.")
return removed
def _release_imported_rows_for_field(db_path, table, frame, timelapse=False):
"""Hand back an import's copy of the field about to be measured into ``table``.
Called immediately before the append, and only for the canonical object
tables an import can have copied into. It asks
:func:`spacr.resume.importer_rows_clause` whether this table holds rows a
foreign import wrote, narrows that to the field ``frame`` is for, and
removes exactly those -- verified against ``foreign_<table>`` row by row
first, and gated on a count taken with the same predicate as the delete.
Scoping to the field is what makes this safe to do at the writer, where a
whole-table release is not. ``resume.supersede_imported_copies`` refuses to
release a table when some field it covers is neither measured nor queued,
because a half-released table leaves that field with no rows at all. Here
the released field's replacement rows are the very next statement, so the
field is never left empty; the import's other fields keep their rows and
their provenance, and are released the same way when their turn comes.
Nothing here can lose a measurement. What is removed is a duplicate of
``foreign_<table>``, which nothing in spaCR may delete from, and the
importer's own numbers stay exactly where they were.
:param db_path: path to ``measurements.db``.
:param table: destination table, one of
:data:`spacr.schema.CANONICAL_OBJECT_TABLES`.
:param frame: the rows about to be appended, carrying the field key.
:param timelapse: True for a timelapse run.
A locked database is retried on the same schedule as the append that
follows it, and for the same reason: ``measure_crop`` writes one field per
worker into a single SQLite file, contention is normal and transient, and a
check that turned a busy database into a lost field would re-introduce the
bug ``_append_to_measurements_db`` exists to prevent. The reads are taken
on a read-only connection so this can never be the thing holding the lock.
:param db_path: path to ``measurements.db``.
:param table: destination table, one of
:data:`spacr.schema.CANONICAL_OBJECT_TABLES`.
:param frame: the rows about to be appended, carrying the field key.
:param timelapse: True for a timelapse run.
:returns: number of imported rows released, ``0`` when the table holds none
for this field -- which is the ordinary case, and costs three reads of
``sqlite_master`` on a project that has never seen an import.
:raises ImportedCopyNotReleased: when the copy is there and cannot be
released provably losslessly. Refusing costs one field's measurements,
which a re-run replaces; mixing costs every count in the project, which
nothing detects.
:raises sqlite3.OperationalError: when the database stays locked for every
attempt, exactly as the append would.
"""
if not os.path.isfile(db_path):
return 0
delay = 0.2
attempt = 1
while True:
try:
return _release_imported_rows_once(db_path, table, frame, timelapse)
except sqlite3.OperationalError as e:
if 'locked' not in str(e).lower() or attempt == DB_WRITE_ATTEMPTS:
raise
print(f"measurements.db busy checking {table} for an imported copy "
f"(attempt {attempt}/{DB_WRITE_ATTEMPTS}): {e}; retrying")
time.sleep(delay)
delay *= 2
attempt += 1
def _release_imported_rows_once(db_path, table, frame, timelapse=False):
"""One attempt of :func:`_release_imported_rows_for_field`.
Split out so the retry above wraps the whole question -- read the
provenance, verify the twins, delete -- rather than any one statement of
it. Retrying a statement inside a transaction could duplicate an earlier
write; retrying the whole thing cannot, because it re-reads the state it
decides on and the delete is gated on a count taken beside it.
"""
from . import resume as _resume
from .database_concurrency import connect, transaction
alias = 's'
conn = connect(db_path, readonly=True, timeout=DB_WRITE_TIMEOUT)
try:
if not conn.execute(
"SELECT 1 FROM sqlite_master WHERE type='table' AND name=?",
(table,)).fetchone():
return 0
importer_clause = _resume.importer_rows_clause(conn, table)
if importer_clause is None:
return 0
total = int(conn.execute(
f'SELECT COUNT(*) FROM "{table}" AS {alias} '
f'WHERE {importer_clause}').fetchone()[0])
if not total:
return 0
have = {row[1] for row in conn.execute(f'PRAGMA table_info("{table}")')}
absent = [c for c in schema.FIELD_KEY_COLUMNS if c not in have]
if absent:
raise ImportedCopyNotReleased(
f"'{table}' in {db_path} holds {total} row(s) a foreign import "
f"copied there, and it has no {absent} column, so which field "
f"they belong to cannot be established. spaCR is about to "
f"measure into the same table and the two populations would be "
f"indistinguishable. Nothing was written. Measure into a "
f"different output folder, or re-run the import.")
key_predicate, params = _field_key_predicate(
frame, schema.FIELD_KEY_COLUMNS, alias)
held = int(conn.execute(
f'SELECT COUNT(*) FROM "{table}" AS {alias} '
f'WHERE {importer_clause} AND {key_predicate}',
tuple(params)).fetchone()[0])
if not held:
return 0
field = str(frame['prcf'].iloc[0]) if 'prcf' in frame.columns else '?'
if timelapse:
raise ImportedCopyNotReleased(
f"'{table}' in {db_path} holds {held} row(s) a foreign import "
f"copied there for field {field}, and this is a timelapse run. "
f"The importer writes no timeID, so which frame a copied row "
f"belongs to cannot be established from the database, and "
f"releasing it on the four field columns alone would be a "
f"delete keyed on fewer columns than the writer used. Nothing "
f"was written. Measure this plate into a different output "
f"folder.")
from .foreign import FOREIGN_PREFIX, _twin_condition
twin = _twin_condition(conn, table)
if twin is None:
raise ImportedCopyNotReleased(
f"'{table}' in {db_path} holds {held} row(s) a foreign import "
f"copied there for field {field}, and "
f"'{FOREIGN_PREFIX}{table}' -- the importer's own copy of "
f"exactly those rows -- is not in this database, so removing "
f"them would destroy the only copy. spaCR will not append "
f"beside them either: the table would then hold two "
f"populations no reader could tell apart. Nothing was written. "
f"Re-run the import to restore it, or measure into a different "
f"output folder.")
orphans = int(conn.execute(
f'SELECT COUNT(*) FROM "{table}" AS {alias} '
f'WHERE {importer_clause} AND {key_predicate} AND NOT {twin}',
tuple(params)).fetchone()[0])
if orphans:
raise ImportedCopyNotReleased(
f"'{table}' in {db_path} holds {held} imported row(s) for "
f"field {field} and {orphans} of them have no matching row in "
f"'{FOREIGN_PREFIX}{table}', so they exist nowhere else and "
f"removing them would lose them. Nothing was written. Re-run "
f"the import into this destination, which rewrites both tables "
f"from the source, and measure again.")
finally:
conn.close()
writer = connect(db_path, timeout=DB_WRITE_TIMEOUT)
try:
with transaction(writer, mode='IMMEDIATE', attempts=6,
busy_timeout=DB_WRITE_TIMEOUT):
return _verified_delete(
writer, table, alias,
f'{importer_clause} AND {twin} AND {key_predicate}', params,
f"release the import's copy of field {field} from '{table}'")
finally:
writer.close()
def _merge_and_save_to_database(morph_df, intensity_df, table_type, source_folder, file_name, experiment, timelapse=False, stamp=None, store=None):
"""Merge morphology and intensity DataFrames and append to the measurements SQLite DB.
``intensity_df`` may be empty: the ``*_organelle_summary`` tables are
morphology-only rollups and have no intensity frame to merge. Requiring
both to be non-empty meant all four summary writes returned silently and
no summary table was ever created.
:param stamp: dict of :data:`MEASUREMENT_STAMP_COLUMNS`, from
:func:`spacr.measure.resolve_measurement_spacing`. Every value is
written onto every row so that a reader can tell whether
``<object>_area`` is a px^2 area or a volume, and in which units,
without guessing. ``None`` writes no stamp columns, which keeps a
direct caller's schema exactly as it was, and is treated as 2-D/px
by the compatibility check.
:raises MeasurementUnitsMismatch: when ``table_type`` already holds
rows measured in other units.
:raises spacr.schema.ObjectTableSchemaError: when a cell, cytoplasm,
nucleus, or pathogen frame violates its canonical identity,
provenance, feature-namespace, or cardinality contract.
"""
from .database_concurrency import _capture_write_operation
if _capture_write_operation('merge', (
morph_df, intensity_df, table_type, source_folder, file_name,
experiment, timelapse, stamp, store)):
return
morph_df = _check_integrity(morph_df)
intensity_df = _check_integrity(intensity_df)
if len(morph_df) == 0:
return
if len(intensity_df) == 0 and table_type not in _ORGANELLE_SUMMARY_TABLES:
print(f"Warning: {table_type} has {len(morph_df)} morphology rows but an "
f"empty intensity frame for {file_name}; nothing written to the "
f"{table_type} table for this field.")
return
_META = ['plateID', 'rowID', 'columnID', 'fieldID', 'prcf', 'file_name', 'path_name']
if table_type in _PARENT_OBJECT_TABLES or table_type in _ORGANELLE_SUMMARY_TABLES:
column_list = ['object_label'] + _META
elif table_type in _CHILD_OBJECT_TABLES:
column_list = ['object_label', 'cell_id'] + _META
else:
raise ValueError(f"Invalid table_type: {table_type}")
if len(intensity_df) > 0:
merged_df = pd.merge(
morph_df,
intensity_df,
on='object_label',
how='outer',
validate='one_to_one',
)
else:
merged_df = morph_df.copy()
merged_df = merged_df.rename(columns={"label_list_x": "label_list_morphology", "label_list_y": "label_list_intensity"})
merged_df['file_name'] = file_name
merged_df['path_name'] = os.path.join(source_folder, file_name + '.npy')
if stamp:
for col in MEASUREMENT_STAMP_COLUMNS:
merged_df[col] = stamp.get(col)
if timelapse:
merged_df[['plateID', 'rowID', 'columnID', 'fieldID', 'timeID', 'prcf']] = merged_df['file_name'].apply(lambda x: pd.Series(_map_wells(x, timelapse)))
else:
merged_df[['plateID', 'rowID', 'columnID', 'fieldID', 'prcf']] = merged_df['file_name'].apply(lambda x: pd.Series(_map_wells(x, timelapse)))
cols = merged_df.columns.tolist()
missing_columns = [col for col in column_list if col not in cols]
if missing_columns == ['cell_id']:
column_list = ['object_label'] + _META
missing_columns = []
if missing_columns:
raise ValueError(f"Columns missing in DataFrame: {missing_columns}")
for i, col in enumerate(column_list):
cols.insert(i, cols.pop(cols.index(col)))
merged_df = merged_df[cols]
if table_type in schema.CANONICAL_OBJECT_TABLES:
merged_df = schema.validate_object_table_frame(
merged_df,
table_type,
timelapse=timelapse,
)
db_path = f'{source_folder}/measurements/measurements.db'
_assert_measurement_units_compatible(db_path, table_type, stamp)
if table_type in schema.CANONICAL_OBJECT_TABLES:
_release_imported_rows_for_field(
db_path, table_type, merged_df, timelapse=timelapse)
_append_to_measurements_db(db_path, table_type, merged_df,
store=store)
#: How many times a locked measurements.db write is retried before it fails.
DB_WRITE_ATTEMPTS = 8
#: Seconds a single connect() waits for the lock, per attempt. Kept at the
#: original 5 s: the retry loop, not a long single wait, is what survives
#: contention, and a long one makes every deliberately-locked-database test
#: pay for it. Raising it to 30 s made the whole suite time out.
DB_WRITE_TIMEOUT = 5.0
def _widen_table_for(conn, table, frame):
"""Add any column ``frame`` has and ``table`` lacks, as NULL for old rows.
Measurement frames legitimately differ field to field — a field with no
pathogen objects produces no pathogen columns, and ``radial_dist`` off
removes a whole block. ``to_sql(if_exists='append')`` refuses the whole
frame in that case with "table X has no column named Y".
:param conn: open sqlite3 connection.
:param table: destination table.
:param frame: the rows about to be appended.
:returns: the list of column names added.
"""
have = {row[1] for row in conn.execute(f'PRAGMA table_info("{table}")')}
if not have:
return []
added = []
for col in frame.columns:
if col in have:
continue
try:
conn.execute(f'ALTER TABLE "{table}" ADD COLUMN "{col}"')
except sqlite3.OperationalError as e:
# PRAGMA above and this ALTER. Its column is indistinguishable from
if 'duplicate column name' not in str(e).lower():
raise
continue
added.append(col)
if added:
conn.commit()
return added
#: How many schema repairs a single append attempts before giving up. Two
#: distinct conditions can each fire once (lost the CREATE TABLE race, then
#: the winner's schema turns out to be narrower than ours), plus slack.
DB_APPEND_REPAIRS = 4
def _sqlite_identifier(value):
"""Quote one SQLite identifier without treating data as SQL.
SQLite accepts double quotes inside an identifier when they are doubled.
NUL cannot occur in an SQLite identifier and is rejected explicitly so a
malformed frame cannot produce a confusing parser error later.
"""
value = str(value)
if '\x00' in value:
raise ValueError("SQLite identifiers cannot contain NUL")
return '"' + value.replace('"', '""') + '"'
def _sqlite_value(value):
"""Return a sqlite3-compatible scalar with pandas nulls normalised."""
if value is None:
return None
missing = pd.isna(value)
if isinstance(missing, (bool, np.bool_)) and bool(missing):
return None
if isinstance(value, pd.Timedelta):
return int(value.value)
if isinstance(value, pd.Timestamp):
return value.to_pydatetime()
if isinstance(value, np.generic):
return value.item()
return value
def _insert_frame(conn, table, frame):
"""Insert ``frame`` without pandas' per-append table-existence query.
``DataFrame.to_sql(if_exists='append')`` performs a ``sqlite_master`` read
before every write. On a shared filesystem that read can hold a SHARED
lock while another worker is trying to commit, creating the writer
starvation reported in issue #15. A parameterised INSERT needs no schema
read and keeps values out of the SQL text.
Table creation and schema repair remain :func:`_append_frame`'s job. The
transaction is committed here to retain ``to_sql``'s all-or-error append
contract; sqlite rolls it back if ``executemany`` or ``commit`` fails.
"""
if frame.empty:
return
columns = [str(column) for column in frame.columns]
if len(columns) != len(set(columns)):
raise ValueError("duplicate names in measurement frame columns")
quoted_table = _sqlite_identifier(table)
quoted_columns = ', '.join(_sqlite_identifier(col) for col in columns)
placeholders = ', '.join('?' for _ in columns)
statement = (
f'INSERT INTO {quoted_table} ({quoted_columns}) '
f'VALUES ({placeholders})'
)
values_by_column = []
for _name, series in frame.items():
if series.dtype.kind == 'm':
if isinstance(series.dtype, getattr(pd, 'ArrowDtype', ())):
values = series.to_numpy(dtype='timedelta64[ns]').view('i8')
else:
values = series.to_numpy().view('i8')
values_by_column.append(values.astype(object))
else:
values_by_column.append(
np.asarray(
[_sqlite_value(value) for value in series],
dtype=object,
)
)
rows = zip(*values_by_column)
try:
conn.executemany(statement, rows)
conn.commit()
except Exception:
conn.rollback()
raise
def _append_frame(conn, table, frame):
"""Append without a per-write schema read, repairing concurrent hazards.
The normal path is a direct parameterised INSERT, avoiding pandas'
``sqlite_master`` existence probe on every field. Only a genuinely absent
table uses ``to_sql`` to create its schema, without rows, then retries the
direct insert. Several first-field workers can race during that one-time
creation; the loser sees ``table already exists`` and retries safely.
The same bounded loop covers schema widening when the worker that won the
creation race used a narrower frame. A failed INSERT is rolled back before
ALTER/retry, so a frame cannot be partly duplicated or silently lost.
:param conn: open sqlite3 connection.
:param table: destination table.
:param frame: rows to append.
:raises sqlite3.OperationalError: the last error, if every repair failed.
"""
last = None
for _ in range(DB_APPEND_REPAIRS):
try:
_insert_frame(conn, table, frame)
return
except sqlite3.OperationalError as e:
last = e
message = str(e).lower()
if 'no such table' in message:
try:
frame.iloc[:0].to_sql(
table, conn, if_exists='append', index=False)
except sqlite3.OperationalError as create_error:
if 'already exists' not in str(create_error).lower():
raise
last = create_error
continue
if 'has no column named' not in message:
raise
added = _widen_table_for(conn, table, frame)
if added:
print(f"measurements.db: added {len(added)} column(s) to "
f"{table} for this field ("
f"{', '.join(added[:6])}"
f"{' ...' if len(added) > 6 else ''}); rows written "
f"earlier are NULL there")
raise last
def _append_to_measurements_db(db_path, table, frame, required=True,
store=None):
"""Append ``frame`` to ``table``, surviving a lock and a widened schema.
This used to be a bare ``except sqlite3.OperationalError: print(...)``,
which hid two different data-loss bugs behind one printed line.
**A locked database dropped the field's rows.** measure_crop writes one
field per worker into a single SQLite file, so contention is normal and
transient — but the rows were discarded while the worker still returned
success and the run reported complete. It reproduced about one run in
twelve on a four-field synthetic set. Now the write is retried with
backoff, and if it still cannot land the error propagates so the caller's
RunLedger records the field as failed and stamps the artifact partial.
**A differing column set dropped the field's rows too.** Measurement
frames legitimately vary between fields, and ``to_sql`` refuses the entire
frame when the table lacks one of its columns. The table is widened
instead, which keeps the rows; columns the frame lacks are simply NULL.
**And losing pandas' CREATE TABLE race dropped them a third time** — see
:func:`_append_frame`, which now retries instead. That was the last
remaining cause of measure_crop measuring only three of four fields.
The connection is closed on every path — leaving it open held the lock
longer and made the contention worse.
:param db_path: path to measurements.db.
:param table: destination table name.
:param frame: rows to append.
:param required: True when losing this table should fail the whole field.
False for side tables such as ``png_list``: a lock there costs the crop
index, and aborting the field over it would throw away the
measurements as well, which are the artifact that matters. Measured -
raising on png_list took the failure rate from 3 in 20 to 8 in 20.
:param store: a DuckDB file, Parquet store or PostgreSQL locator that
also receives the rows once ``measurements.db`` holds them, through
:func:`spacr.tabular.write_database`. ``None`` writes SQLite only.
:raises sqlite3.OperationalError: when every attempt fails and ``required``.
"""
from .database_concurrency import _capture_write_operation, _inside_write_packet
if _capture_write_operation('append', (db_path, table, frame, required, store)):
return
delay = 0.2
attempt = 1
while True:
conn = None
try:
from .database_concurrency import connect
conn = connect(db_path, timeout=DB_WRITE_TIMEOUT)
from .database_schema import migrate_connection
migrate_connection(conn, path=os.path.abspath(db_path))
_append_frame(conn, table, frame)
break
except sqlite3.OperationalError as e:
if 'locked' not in str(e).lower():
if _inside_write_packet(db_path):
raise
print(f"SQLite error writing {table}: {e}")
return
if attempt == DB_WRITE_ATTEMPTS:
if required or _inside_write_packet(db_path):
raise
print(f"giving up writing {table} after "
f"{DB_WRITE_ATTEMPTS} attempts: {e}")
return
print(f"measurements.db busy writing {table} "
f"(attempt {attempt}/{DB_WRITE_ATTEMPTS}): {e}; retrying")
time.sleep(delay)
delay *= 2
attempt += 1
finally:
if conn is not None:
conn.close()
if store is not None:
_append_to_measurement_store(store, table, frame, required)
def _measurement_store_for(db_path, settings):
"""The store ``measurement_backend`` names for ``db_path``, or ``None``.
:param db_path: the run's measurements.db.
:param settings: the run settings.
:returns: ``None`` for sqlite, else the DuckDB, Parquet or PostgreSQL
locator.
"""
backend = str((settings or {}).get('measurement_backend') or 'sqlite')
if backend.lower() == 'sqlite':
return None
from .measure import _measurement_backend_target
return _measurement_backend_target(db_path, settings)
def _append_to_measurement_store(store, table, frame, required):
"""Append ``frame`` to ``table`` in a non-SQLite measurement store.
:param store: the store locator.
:param table: destination table name.
:param frame: rows to append; column names are kept as written.
:param required: re-raise a failed write when True, else report it.
"""
from .tabular import write_database
try:
write_database(frame, store, table, if_exists='append',
canonicalise=False)
except Exception as exc:
if required:
raise
print(f"measurement store: {table} not written ({exc})")
def _safe_int_convert(value, default=0):
"""Return the integer ``value`` denotes, otherwise ``default``.
**This is not a key builder.** It used to be — ``_map_wells`` built
``fieldID`` and ``timeID`` out of it — and because its default is ``0``
and ``0`` is a perfectly good field id, every token it could not read
became field ``f0``: three ImageXpress sites ``s1``/``s2``/``s3`` went in
and one ``prcf`` came out, and a whole ``T0001``/``T0002``/``T0003``
timelapse collapsed onto ``t0``. Nothing said so. Key construction now
goes through :mod:`spacr.schema`, which never invents a number; see
:func:`spacr.schema.field_id` for the graded policy that replaced it.
What is left is the one honest use: undoing zero padding on a regex group
that has already been checked to start with a digit
(``_extract_filename_metadata``, ``io._move_to_chan_folder``). Even there
spaCR no longer relies on ``default`` — those call sites keep the original
token when it holds no integer, because two unreadable wells that both
became ``'0'`` were two wells merged into one.
"Is this an integer?" is answered by :func:`spacr.schema.parse_int_token`,
so that this function and every key in the database agree on the question.
That makes it stricter than the old bare ``int()`` in two inert ways:
``3.7`` and ``True`` now take the default rather than silently becoming
``3`` and ``1``. Inventing ``3`` from ``3.7`` is the same species of lie
as inventing ``0`` from ``'x'``.
:param value: token to convert.
:param default: returned when ``value`` holds no integer. ``None`` takes
it too — the old form raised :class:`TypeError` there while
``resume._safe_int`` returned the default, so ``None`` was a crash in
one code path and field ``f0`` in the other.
:returns: the integer, or ``default``.
"""
parsed = schema.parse_int_token(value, allow_prefix=False)
if parsed is None:
return default
return parsed
def _int_or_token(value):
"""Undo zero padding, or return the token unchanged when it is not a number.
The replacement for ``str(_safe_int_convert(x))`` in the filename-metadata
parsers: ``'001'`` becomes ``'1'`` so that ``'001'`` and ``'1'`` are one
field, while ``'1a'`` stays ``'1a'`` instead of becoming ``'0'`` — which
is the difference between two odd fields staying two fields and every odd
field in the run merging into one.
:param value: a token from a filename regex group.
:returns: the token's integer as a string, or the token unchanged.
"""
parsed = schema.parse_int_token(value, allow_prefix=False)
return str(value) if parsed is None else str(parsed)
def _map_wells(file_name, timelapse=False):
"""Parse a stack file name into ``(plate, row, column, field[, timeid], prcf)``.
A thin adapter over :func:`spacr.schema.parse_field_stem`, which is the
single definition of what those keys are. The tuple shape and the
``'error'`` fallback are unchanged, because callers
(:func:`spacr.predictions.crop_name_metadata`,
:func:`process_vision_results`) read both.
Every difference from the previous hand-rolled body is a case the old one
got wrong, not a change of contract — ``tests/test_schema.py`` pins the
agreement on every name the old one handled:
* ``'AA01'`` (an ordinary 1536-plate well) now parses to ``r27``; it used
to raise inside and destroy the *plate* along with the well.
* a lowercase or whitespace-padded well now parses.
* a vendor-prefixed field parses — ``s3``/``F003``/``T0003`` are field 3,
not field 0 — and a field token holding no integer at all is preserved
(``'xy'`` -> ``'fxy'``) instead of colliding on ``f0``.
* the name is reduced to its basename and stem first, so a full path or a
trailing ``.npy`` no longer leaks a directory into ``plateID`` or turns
``'3.npy'`` into field 0.
:param file_name: stack file name, stem or path.
:param timelapse: parse a fourth component as the timepoint.
:returns: the key tuple, or ``'error'`` in every slot.
"""
try:
field = schema.parse_field_stem(file_name, timelapse=timelapse)
except schema.SchemaError as e:
print(f"Error processing filename: {file_name}")
print(f"Error: {e}")
return ('error',) * (6 if timelapse else 5)
if timelapse:
return (field.plateID, field.rowID, field.columnID, field.fieldID,
field.timeID, field.prcf)
return (field.plateID, field.rowID, field.columnID, field.fieldID,
field.prcf)
def _map_wells_png(file_name, timelapse=False):
"""Parse a cropped-object PNG name into well ids plus ``prcfo`` and object id.
A thin adapter over :func:`spacr.schema.parse_object_stem`; see
:func:`_map_wells` for why. The differences from the previous body, all of
them repairs:
* ``'AA01'`` gave ``('r1', 'c0')`` — the second row letter dropped *and*
a column 0 invented. It now gives ``('r27', 'c1')``, which is what the
object tables carry, so ``png_list`` and ``cell`` join again.
* a well with letters but no column (``'A'``) gave ``'c0'``, which is
indistinguishable from a real column 0; it is now an ``'error'`` row,
the same answer :func:`_map_wells` has always given it.
* an object token holding no integer gave ``'onone'`` whatever it said, so
a nucleus crop overlapping several nuclei (``..._multi.png``) and one
overlapping none (``..._none.png``) shared a ``prcfo``. The token is now
preserved: ``'omulti'`` and ``'onone'``.
* a three-part name (``'plate1_A01_5.png'``) read one token as both the
field *and* the object. ``_generate_names`` never emits one, so it is
now an ``'error'`` row rather than a fabricated identity.
:param file_name: crop PNG name or path.
:param timelapse: parse a timepoint between the field and the object.
:returns: the key tuple, or ``'error'`` in every slot.
"""
try:
obj = schema.parse_object_stem(file_name, timelapse=timelapse)
except schema.SchemaError as e:
print(f"Error processing filename: {file_name}")
print(f"Error: {e}")
return ('error',) * (7 if timelapse else 6)
if timelapse:
return (obj.plateID, obj.rowID, obj.columnID, obj.fieldID, obj.timeID,
obj.prcfo, obj.objectID)
return (obj.plateID, obj.rowID, obj.columnID, obj.fieldID, obj.prcfo,
obj.objectID)
DUPLICATE_COLUMN_SUFFIX = "__dup"
def _check_integrity(df):
"""Deduplicate label columns and collapse them into ``label_list``/``object_label``.
Repeats of a duplicated name are suffixed with their OCCURRENCE index, not
their position in the frame. The previous form used ``enumerate``'s
frame-wide index, so a second ``mean_intensity`` sitting at position 57
became ``mean_intensity_57`` -- a name indistinguishable from a genuinely
parameterised feature like ``homogeneity_distance_8``, and one that moved
whenever an unrelated column was added upstream. ``__dup<n>`` cannot
collide with a feature name. It also left the first occurrence renamed
unless it happened to sit at index 0, even though ``object_label`` is taken
from the first label column.
Counting once rather than re-scanning the column list per column takes this
from O(n^2) to O(n); a measurement frame carries roughly a thousand columns
and this runs twice per field per object type.
:param df: a morphology or intensity measurement frame.
:returns: the frame with label columns collapsed and dropped.
"""
counts = Counter(df.columns)
seen = Counter()
renamed = []
for col in df.columns:
if counts[col] > 1:
n = seen[col]
seen[col] += 1
renamed.append(col if n == 0 else f"{col}{DUPLICATE_COLUMN_SUFFIX}{n}")
else:
renamed.append(col)
df.columns = renamed
label_cols = [col for col in df.columns if 'label' in col]
if len(df) and not label_cols:
raise ValueError(
"_check_integrity: no column containing 'label' in a frame of "
f"{len(df)} rows, so object_label cannot be derived. "
f"Columns: {list(df.columns)[:12]}"
+ (" ..." if len(df.columns) > 12 else ""))
df['label_list'] = df[label_cols].values.tolist()
df['object_label'] = df['label_list'].apply(lambda x: x[0] if x else None)
df = df.drop(columns=label_cols)
df['label_list'] = df['label_list'].astype(str)
return df
def _get_percentiles(array, p1=2, p2=98):
"""Return per-channel ``[p1, p2]`` percentiles from nonzero pixels of an image stack."""
nimg = array.shape[2]
percentiles = []
for v in range(nimg):
img = np.squeeze(array[:, :, v])
non_zero_img = img[img > 0]
if non_zero_img.size > 0:
img_min = np.percentile(non_zero_img, p1)
img_max = np.percentile(non_zero_img, p2)
percentiles.append([img_min, img_max])
else:
img_min = np.percentile(img, p1)
img_max = np.percentile(img, p2)
percentiles.append([img_min, img_max])
return percentiles
def _crop_center(img, cell_mask, new_width, new_height):
"""Crop ``img`` to ``new_width`` x ``new_height`` centered on the mask centroid."""
cell_mask[cell_mask != 0] = 1
mask_3d = np.repeat(cell_mask[:, :, np.newaxis], img.shape[2], axis=2).astype(img.dtype)
img = np.multiply(img, mask_3d).astype(img.dtype)
centroid = np.round(ndi.center_of_mass(cell_mask)).astype(int)
pad_width = max(new_width, new_height)
img = np.pad(img, ((pad_width, pad_width), (pad_width, pad_width), (0, 0)), mode='constant')
cell_mask = np.pad(cell_mask, ((pad_width, pad_width), (pad_width, pad_width)), mode='constant')
centroid += pad_width
start_y = max(0, centroid[0] - new_height // 2)
end_y = min(start_y + new_height, img.shape[0])
start_x = max(0, centroid[1] - new_width // 2)
end_x = min(start_x + new_width, img.shape[1])
img = img[start_y:end_y, start_x:end_x, :]
return img
def _masks_to_masks_stack(masks):
"""Return ``masks`` as a plain list preserving iteration order."""
mask_stack = []
for idx, mask in enumerate(masks):
mask_stack.append(mask)
return mask_stack
def _get_diam(mag, obj):
"""Return an object type's expected diameter at a magnification.
:param mag: the objective magnification.
:param obj: the object type.
:returns: the diameter in pixels.
:raises ValueError: naming the supported types, for anything else --
this used to fall through to an unbound variable and raise
``UnboundLocalError``, which names an implementation detail rather
than the setting the user got wrong.
"""
if obj == 'cell':
diameter = 2 * mag + 80
elif obj == 'cell_large':
diameter = 2 * mag + 120
elif obj == 'nucleus':
diameter = 0.75 * mag + 45
elif obj == 'pathogen':
diameter = mag
else:
raise ValueError(
f"_get_diam: unsupported object type '{obj}'. "
f"Expected one of: cell, cell_large, nucleus, pathogen."
)
return int(diameter)
def _get_object_settings(object_type, settings):
"""Assemble one object type's segmentation settings.
The size bounds are derived from the diameter rather than asked for, so
they scale with the magnification. A pre-SAM Cellpose model name left in
an old settings file is mapped forward HERE, once, rather than carried
into segmentation as if it still selected different weights.
:param object_type: the object being segmented.
:param settings: the run settings.
:returns: the settings for that object.
"""
object_settings = {}
object_settings['diameter'] = _get_diam(settings['magnification'], obj=object_type)
object_settings['minimum_size'] = (object_settings['diameter']**2)/4
object_settings['maximum_size'] = (object_settings['diameter']**2)*10
object_settings['merge'] = False
object_settings['resample'] = True
object_settings['remove_border_objects'] = False
from .settings import normalize_cellpose_model_name
object_settings['model_name'] = normalize_cellpose_model_name(
settings.get(f'{object_type}_model_name'),
object_type=object_type, key=f'{object_type}_model_name')
if object_type == 'cell':
object_settings['filter_size'] = False
object_settings['filter_intensity'] = False
object_settings['restore_type'] = settings.get('cell_restore_type', None)
elif object_type == 'nucleus':
object_settings['filter_size'] = False
object_settings['filter_intensity'] = False
object_settings['restore_type'] = settings.get('nucleus_restore_type', None)
elif object_type == 'pathogen':
object_settings['filter_size'] = False
object_settings['filter_intensity'] = False
object_settings['resample'] = False
object_settings['restore_type'] = settings.get('pathogen_restore_type', None)
object_settings['merge'] = settings['merge_pathogens']
else:
print(f'Object type: {object_type} not supported. Supported object types are : cell, nucleus and pathogen')
if settings['verbose']:
print(object_settings)
return object_settings
def _pivot_counts_table(db_path):
"""Rewrite the object-count table as one row per file, one column per type.
Written to ``pivoted_counts`` rather than over the source table, so the
long form the pipeline appends to is left intact.
:param db_path: the measurements database.
"""
def _read_table_to_dataframe(db_path, table_name='object_counts'):
"""Return the given SQLite table as a DataFrame."""
return tabular.read_database(
db_path, [table_name], report=None, migrate=False)[0]
def _pivot_dataframe(df):
"""Pivot count-type rows into one column per object type, NaNs filled with 0."""
pivoted_df = df.pivot(index='file_name', columns='count_type', values='object_count').reset_index()
pivoted_df = pivoted_df.fillna(0)
return pivoted_df
df = _read_table_to_dataframe(db_path, 'object_counts')
pivoted_df = _pivot_dataframe(df)
conn = sqlite3.connect(db_path, timeout=30)
pivoted_df.to_sql('pivoted_counts', conn, if_exists='replace', index=False)
conn.close()
#: The order the merged stack's channel axis is built in, and the ONLY order
#: that describes it. `io.preprocess_img_data` walks these four settings in
#: exactly this sequence and assigns each new raw channel the next dense
#: position (`seen[ch] = len(mask_channels)`), so the axis is in ROLE order,
#: deduplicated -- not in ascending channel order.
MASK_CHANNEL_ROLE_ORDER = (
"nucleus_channel", "cell_channel", "pathogen_channel",
*(f"{role}_channel" for role in schema.ORGANELLE_ROLES),
)
[docs]
def dense_mask_channel_positions(settings):
"""Map each RAW channel index to its position on the merged stack's axis.
Built the same way `io.preprocess_img_data` builds the stack, because
that is the only thing that makes the answer true: walk the roles in
:data:`MASK_CHANNEL_ROLE_ORDER` and give each newly seen raw channel the
next dense position.
THE TRAP THIS EXISTS TO CLOSE. Several callers computed the position as
``sorted({nucleus, cell, pathogen, organelle})`` instead, which agrees
with role order only when the roles happen to be in ascending channel
order. With ``nucleus_channel=2, cell_channel=0, organelle_channel=1``
the stack is ``[2, 0, 1]`` -- raw channel 1 sits at position 2 -- while
the sorted reading says position 1, which holds the CELL image. Cellpose
then segments organelles on the cell plane, silently, and every count and
intensity downstream is measured from the wrong masks.
:param settings: the run settings, holding the raw ``*_channel`` keys.
:returns: ``{raw_channel: dense_position}``. Channels that are None or
uncoercible are absent, matching the writer's own behaviour.
"""
positions = {}
for key in MASK_CHANNEL_ROLE_ORDER:
raw = settings.get(key)
if raw is None:
continue
try:
raw = int(raw)
except (TypeError, ValueError):
continue
if raw not in positions:
positions[raw] = len(positions)
return positions
def _get_cellpose_channels(settings):
"""Return the channel indices to extract and the per-object-type Cellpose channel remap."""
nucleus_ch = settings.get('cellpose_nucleus_channel')
cell_ch = settings.get('cellpose_cell_channel')
pathogen_ch = settings.get('cellpose_pathogen_channel')
organelle_channels = {
role: settings.get(f'cellpose_{role}_channel')
for role in schema.ORGANELLE_ROLES
}
all_channels = set()
for ch in [nucleus_ch, cell_ch, pathogen_ch,
*organelle_channels.values()]:
if ch is not None:
all_channels.add(ch)
channels_to_extract = sorted(all_channels)
remap = {orig: new for new, orig in enumerate(channels_to_extract)}
cellpose_channels = {}
if nucleus_ch is not None:
cellpose_channels['nucleus'] = [remap[nucleus_ch]]
if cell_ch is not None:
if nucleus_ch is not None:
cellpose_channels['cell'] = [remap[cell_ch], remap[nucleus_ch]]
else:
cellpose_channels['cell'] = [remap[cell_ch]]
if pathogen_ch is not None:
cellpose_channels['pathogen'] = [remap[pathogen_ch]]
for role, channel in organelle_channels.items():
if channel is not None:
cellpose_channels[role] = [remap[channel]]
return channels_to_extract, cellpose_channels
[docs]
def annotate_conditions(df, cells=None, cell_loc=None, pathogens=None, pathogen_loc=None, treatments=None, treatment_loc=None):
"""Annotate ``df`` with host cell, pathogen, treatment, and combined ``condition`` columns.
:param df: DataFrame to annotate; must contain ``rowID``/``columnID``.
:param cells: host cell types (str or list).
:param cell_loc: per-cell-type list-of-lists of row/column identifiers.
:param pathogens: pathogens (str or list).
:param pathogen_loc: per-pathogen list-of-lists of row/column identifiers.
:param treatments: treatments (str or list).
:param treatment_loc: per-treatment list-of-lists of row/column identifiers.
:returns: annotated DataFrame with ``host_cells``, ``pathogen``, ``treatment``, ``condition`` columns.
"""
def _get_type(val):
"""Determine if a value maps to 'rowID' or 'columnID'."""
if isinstance(val, str) and val.startswith('c'):
return 'columnID'
elif isinstance(val, str) and val.startswith('r'):
return 'rowID'
return None
def _map_or_default(column_name, values, loc, df):
"""Assign or map ``values`` into ``column_name`` based on optional row/column ``loc``."""
if isinstance(values, str) and loc is None:
df[column_name] = values
elif isinstance(values, list) and loc is None:
df[column_name] = values[0]
elif values is not None and loc is not None:
value_dict = {val: key for key, loc_list in zip(values, loc) for val in loc_list}
df[column_name] = pd.Series(np.nan, index=df.index, dtype=object)
for val, key in value_dict.items():
loc_type = _get_type(val)
if loc_type:
df.loc[df[loc_type] == val, column_name] = key
_map_or_default('host_cells', cells, cell_loc, df)
_map_or_default('pathogen', pathogens, pathogen_loc, df)
_map_or_default('treatment', treatments, treatment_loc, df)
if pathogens is not None:
df['pathogen'] = df['pathogen'].where(df['pathogen'].notna(), np.nan)
if treatments is not None:
df['treatment'] = df['treatment'].where(df['treatment'].notna(), np.nan)
df['condition'] = df.apply(
lambda x: '_'.join([str(v) for v in [x.get('host_cells'), x.get('pathogen'), x.get('treatment')] if pd.notna(v)]),
axis=1
)
df.loc[df['condition'] == '', 'condition'] = pd.NA
return df
def _split_data(df, group_by, object_type):
"""Group numeric and non-numeric columns of ``df`` separately with per-column aggregation."""
df = df.copy()
time_col = _time_column(df.columns)
if time_col is not None and all(
c in df.columns for c in ('plateID', 'rowID', 'columnID', 'fieldID')):
df['prcft'] = (
df['plateID'].astype(str) + '_' +
df['rowID'].astype(str) + '_' +
df['columnID'].astype(str) + '_' +
df['fieldID'].astype(str) + '_' +
df[time_col].astype(str)
)
try:
prcf = (
df['plateID'].astype(str) + '_' +
df['rowID'].astype(str) + '_' +
df['columnID'].astype(str) + '_' +
df['fieldID'].astype(str)
)
if time_col is not None:
prcf = prcf + '_' + df[time_col].astype(str)
df['prcf'] = prcf
except Exception as e:
print('Exception', e)
df['prcfo'] = df['prcf'].astype(str) + '_' + df[object_type].astype(str)
df = df.set_index(group_by, inplace=False)
df_numeric = df.select_dtypes(include=np.number)
df_non_numeric = df.select_dtypes(exclude=np.number)
from .merge_tables import aggregation_for
agg_dict = {column: aggregation_for(column)
for column in df_numeric.columns}
if len(agg_dict) > 0 and not df_numeric.empty:
grouped_numeric = df_numeric.groupby(df_numeric.index).agg(agg_dict)
else:
grouped_numeric = pd.DataFrame(index=df.index.unique())
if not df_non_numeric.empty:
grouped_non_numeric = df_non_numeric.groupby(df_non_numeric.index).first()
else:
grouped_non_numeric = pd.DataFrame(index=df.index.unique())
return pd.DataFrame(grouped_numeric), pd.DataFrame(grouped_non_numeric)
def _calculate_recruitment(df, channel):
"""Add pathogen-to-compartment recruitment ratio columns for the given intensity channel.
Each output identifies its channel, compartment and numerator statistic,
e.g. ``pathogen_channel_2_cytoplasm_mean_ratio``. Repeated calls preserve
previously computed channels. No spatial slope is inferred or fabricated.
:param df: measurement frame, augmented in place.
:param channel: intensity channel to compare within each compartment.
:returns: the input frame with fifteen channel-specific ratio columns.
The frame is canonicalised first, so a table written before the ring
percentiles were renamed (``outside_75_percentile``) divides correctly
rather than raising ``KeyError`` on the new name. A database read through
``io._read_db`` has already been migrated; a CSV handed in directly has
not.
"""
canonicalize_measurement_columns(df)
statistics = {
'mean': 'mean_intensity', 'q75': 'percentile_75',
'outside_mean': 'outside_mean', 'outside_q75': 'outside_percentile_75',
'periphery_mean': 'periphery_mean',
}
for compartment in ('cell', 'cytoplasm', 'nucleus'):
denominator = df[f'{compartment}_channel_{channel}_mean_intensity']
for name, source in statistics.items():
output = f'pathogen_channel_{channel}_{compartment}_{name}_ratio'
df[output] = df[f'pathogen_channel_{channel}_{source}'] / denominator
return df
def _group_by_well(df):
"""
Group the DataFrame by well coordinates (plate, row, col) and apply mean function to numeric columns
and select the first value for non-numeric columns.
Parameters:
df (DataFrame): The input DataFrame to be grouped.
Returns:
DataFrame: The grouped DataFrame.
"""
numeric_cols = df._get_numeric_data().columns
non_numeric_cols = df.select_dtypes(include=['object']).columns
aggregations = {
**{col: 'mean' for col in numeric_cols},
**{col: 'first' for col in non_numeric_cols},
}
df_grouped = df.groupby(
['plateID', 'rowID', 'columnID'], observed=False
).agg(aggregations)
return df_grouped
[docs]
class Cache:
"""LRU cache with a fixed maximum size.
:param max_size: maximum number of entries retained; oldest is evicted on overflow.
"""
def __init__(self, max_size):
"""Store the size limit and initialize an empty ``OrderedDict``."""
self.cache = OrderedDict()
self.max_size = max_size
[docs]
def get(self, key):
"""Return and refresh ``key``, or ``None`` when it is not cached.
:param key: cache key to look up.
"""
if key in self.cache:
value = self.cache.pop(key)
self.cache[key] = value
return value
return None
[docs]
def put(self, key, value):
"""Insert ``value`` under ``key``, evicting the oldest entry if full.
:param key: cache key under which to store the value.
:param value: object to cache.
"""
if len(self.cache) >= self.max_size:
self.cache.popitem(last=False)
self.cache[key] = value
[docs]
class ScaledDotProductAttention(nn.Module):
"""Standard scaled dot-product attention layer.
:param d_k: dimensionality of key/query vectors used in the scaling factor.
"""
def __init__(self, d_k):
"""Store ``d_k`` used to scale attention logits."""
super(ScaledDotProductAttention, self).__init__()
self.d_k = d_k
[docs]
def forward(self, Q, K, V):
"""Return ``softmax(QK^T / sqrt(d_k)) V``.
:param Q: query tensor.
:param K: key tensor.
:param V: value tensor.
:returns: attention-weighted value tensor.
"""
scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.d_k, dtype=torch.float32))
attention_probs = F.softmax(scores, dim=-1)
output = torch.matmul(attention_probs, V)
return output
[docs]
class SelfAttention(nn.Module):
"""Linear-projected self-attention layer.
:param in_channels: input feature dimension.
:param d_k: projected key/query/value dimension.
"""
def __init__(self, in_channels, d_k):
"""Build the Q/K/V projections and the underlying attention layer."""
super(SelfAttention, self).__init__()
self.W_q = nn.Linear(in_channels, d_k)
self.W_k = nn.Linear(in_channels, d_k)
self.W_v = nn.Linear(in_channels, d_k)
self.attention = ScaledDotProductAttention(d_k)
[docs]
def forward(self, x):
"""Return self-attention over ``x`` of shape ``(B, in_channels)``.
:param x: batch of input feature vectors.
"""
Q = self.W_q(x)
K = self.W_k(x)
V = self.W_v(x)
output = self.attention(Q, K, V)
return output
[docs]
class EarlyFusion(nn.Module):
"""1x1 convolution that fuses input channels down to 64 feature maps.
:param in_channels: number of input channels.
"""
def __init__(self, in_channels):
"""Create the 1x1 fusion convolution."""
super(EarlyFusion, self).__init__()
self.conv1 = nn.Conv2d(in_channels, 64, kernel_size=1, stride=1)
[docs]
def forward(self, x):
"""Return the 64-channel fused feature map.
:param x: image-feature tensor accepted by the 1x1 convolution.
"""
x = self.conv1(x)
return x
[docs]
class SpatialAttention(nn.Module):
"""Spatial attention gate that reweights features by pooled channel statistics.
:param kernel_size: convolution kernel width used to fuse average+max pooled maps.
"""
def __init__(self, kernel_size=7):
"""Build the fusion convolution and sigmoid gate."""
super(SpatialAttention, self).__init__()
self.conv1 = nn.Conv2d(2, 1, kernel_size, padding=kernel_size//2, bias=False)
self.sigmoid = nn.Sigmoid()
[docs]
def forward(self, x):
"""Return the spatial attention map for ``x`` in ``[0, 1]``.
:param x: feature map whose channel statistics define the attention.
"""
avg_out = torch.mean(x, dim=1, keepdim=True)
max_out, _ = torch.max(x, dim=1, keepdim=True)
x = torch.cat([avg_out, max_out], dim=1)
x = self.conv1(x)
return self.sigmoid(x)
[docs]
class MultiScaleBlockWithAttention(nn.Module):
"""Dilated conv block followed by a 1x1 attention convolution.
:param in_channels: input channel count.
:param out_channels: output channel count.
"""
def __init__(self, in_channels, out_channels):
"""Build the dilated convolution and 1x1 spatial-attention convolution."""
super(MultiScaleBlockWithAttention, self).__init__()
self.dilated_conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, dilation=1, padding=1)
self.spatial_attention = nn.Conv2d(out_channels, out_channels, kernel_size=1)
[docs]
def custom_forward(self, x):
"""Apply dilated conv + ReLU followed by the 1x1 spatial attention.
:param x: input feature map for the convolutional block.
"""
x1 = F.relu(self.dilated_conv1(x), inplace=True)
x = self.spatial_attention(x1)
return x
[docs]
def forward(self, x):
"""Forward pass; delegates to :meth:`custom_forward`.
:param x: input feature map for the convolutional block.
"""
return self.custom_forward(x)
[docs]
class CustomCellClassifier(nn.Module):
"""Small classifier stacking :class:`EarlyFusion` and a multi-scale attention block.
:param num_classes: output class count.
:param pathogen_channel: reserved for downstream use; kept for API compatibility.
:param use_attention: reserved for downstream use; kept for API compatibility.
:param use_checkpoint: run the forward pass through ``torch.utils.checkpoint``.
:param dropout_rate: reserved for downstream use; kept for API compatibility.
"""
def __init__(self, num_classes, pathogen_channel, use_attention, use_checkpoint, dropout_rate):
"""Build the fusion, multi-scale, and linear classifier submodules."""
super(CustomCellClassifier, self).__init__()
self.early_fusion = EarlyFusion(in_channels=3)
self.multi_scale_block_1 = MultiScaleBlockWithAttention(in_channels=64, out_channels=64)
self.fc1 = nn.Linear(64, num_classes)
self.use_checkpoint = use_checkpoint
for param in self.parameters():
param.requires_grad = True
[docs]
def custom_forward(self, x):
"""Return the class logits for a batch ``x`` of shape ``(B, 3, H, W)``.
:param x: three-channel image batch to classify.
"""
x = self.early_fusion(x)
x = self.multi_scale_block_1(x)
x = F.adaptive_avg_pool2d(x, (1, 1)).view(x.size(0), -1)
x = F.relu(self.fc1(x), inplace=True)
return x
[docs]
def forward(self, x):
"""Forward pass, optionally through activation checkpointing.
:param x: three-channel image batch to classify.
"""
if self.use_checkpoint:
return _checkpoint_module(self, self.custom_forward, x)
else:
return self.custom_forward(x)
[docs]
class TorchModel(nn.Module):
"""
Thin wrapper around TorchVision classification backbones that:
1) Loads a requested backbone with (optional) pretrained weights
2) Strips its classification head to expose features
3) Adds a simple Linear 'spacr' classifier with `num_classes` outputs
4) Optionally applies dropout before the final classifier
5) Supports gradient checkpointing
Works with most TorchVision **classification** models. Non-classification
(detection/segmentation) models are rejected with a clear error.
"""
def __init__(
self,
model_name: str = "resnet50",
pretrained: bool = True,
dropout_rate: Optional[float] = None,
use_checkpoint: bool = False,
num_classes: int = 2,
multilabel: bool = False,
image_size: int = 224,
):
"""Build the backbone, strip its head, and attach the spaCR linear classifier.
:param model_name: TorchVision classification model to load.
:param pretrained: use ImageNet-pretrained weights when available.
:param dropout_rate: dropout probability applied to backbone and spaCR head; ``None`` disables.
:param use_checkpoint: enable gradient checkpointing through the backbone.
:param num_classes: output class count; ``1`` yields a BCE-style binary head.
:param multilabel: informational flag consumed by external loss/metrics code.
:param image_size: square input resolution used for the dummy forward
pass that infers the backbone's feature width.
:raises ValueError: if ``model_name`` is not a TorchVision model.
"""
super().__init__()
self.model_name = str(model_name)
self.pretrained = bool(pretrained)
self.dropout_rate = (
float(dropout_rate) if dropout_rate is not None else None
)
self.use_checkpoint = bool(use_checkpoint)
self.num_classes = int(num_classes)
self.multilabel = bool(multilabel)
self.image_size = int(image_size) if image_size else 224
self.use_dropout = (dropout_rate is not None)
self.base_model = self._init_base_model(pretrained=bool(pretrained))
if self.model_name == "maxvit_t" and hasattr(self.base_model, "classifier"):
seq = list(self.base_model.classifier.children())
if len(seq) > 0:
self.base_model.classifier = nn.Sequential(*seq[:-1])
if dropout_rate is not None:
self._apply_dropout_rate(self.base_model, float(dropout_rate))
self._remove_head_for_features()
self.num_ftrs = self._infer_feature_dim()
if self.use_dropout:
self.dropout = nn.Dropout(float(dropout_rate))
self.spacr_classifier = nn.Linear(self.num_ftrs, self.num_classes)
def _get_weight_choice(self):
"""
Return the DEFAULT weights enum if available (newer torchvision),
otherwise None to fall back to legacy pretrained=True/False.
"""
enum_attr = f"{self.model_name}_weights"
for attr in dir(models):
if attr.lower() == enum_attr.lower():
enum = getattr(models, attr, None)
if enum is not None and hasattr(enum, "DEFAULT"):
return enum.DEFAULT
return None
def _init_base_model(self, pretrained: bool) -> nn.Module:
"""Build the torchvision backbone this model wraps.
Both weight APIs are supported: the newer ``weights=`` form when
torchvision offers a weight enum for this architecture, and the older
``pretrained=`` flag when it does not.
:param pretrained: load pretrained weights.
:returns: the backbone module.
:raises ValueError: if torchvision has no model of that name.
"""
fn = models.__dict__.get(self.model_name, None)
if fn is None or not callable(fn):
raise ValueError(f"Unknown torchvision model: {self.model_name}")
weights = self._get_weight_choice()
if weights is not None:
return fn(weights=weights if pretrained else None)
else:
return fn(pretrained=pretrained)
def _apply_dropout_rate(self, module: nn.Module, p: float):
"""Set one dropout probability on every dropout layer in a module.
:param module: the subtree to walk.
:param p: the probability to set, on 1-, 2- and 3-D dropout alike.
"""
for m in module.modules():
if isinstance(m, (nn.Dropout, nn.Dropout2d, nn.Dropout3d)):
m.p = p
def _remove_head_for_features(self):
"""
Normalize a wide swath of TorchVision classification heads to Identity.
Also disable auxiliary logits where present (Inception/GoogLeNet).
"""
if hasattr(self.base_model, "aux_logits"):
self.base_model.aux_logits = False
if hasattr(self.base_model, "fc"):
self.base_model.fc = nn.Identity()
return
if hasattr(self.base_model, "classifier"):
if self.model_name != "maxvit_t":
self.base_model.classifier = nn.Identity()
return
if hasattr(self.base_model, "_fc"):
self.base_model._fc = nn.Identity()
return
if hasattr(self.base_model, "heads"):
self.base_model.heads = nn.Identity()
return
if hasattr(self.base_model, "head"):
self.base_model.head = nn.Identity()
return
def _infer_feature_dim(self) -> int:
"""
Forward a dummy tensor through the backbone and determine the flattened
feature size. Uses 224×224 nominal resolution.
"""
self.base_model.eval()
s = int(getattr(self, "image_size", 224)) or 224
with torch.no_grad():
x = torch.zeros(1, 3, s, s)
out = self._run_backbone_raw(x)
if isinstance(out, torch.Tensor) and out.ndim > 2:
out = torch.flatten(out, 1)
if not isinstance(out, torch.Tensor) or out.ndim != 2:
raise RuntimeError(
f"Backbone produced unexpected shape/type for features: {type(out)} / {getattr(out, 'shape', None)}"
)
return int(out.size(1))
def _run_backbone_raw(self, x: torch.Tensor) -> torch.Tensor:
"""
Call the underlying backbone and unwrap common container outputs.
Does NOT apply the new spaCR head.
"""
def forward_fn(t):
"""Run the underlying backbone on ``t`` (used as the checkpoint target)."""
return self.base_model(t)
out = (
_checkpoint_module(self.base_model, forward_fn, x)
if self.use_checkpoint else forward_fn(x)
)
if hasattr(out, "logits"):
out = out.logits
elif isinstance(out, (tuple, list)):
out = out[0]
elif isinstance(out, dict):
raise RuntimeError(
"Selected backbone returned a dict (likely detection/segmentation). "
"Use an image-classification backbone."
)
return out
def _run_backbone(self, x: torch.Tensor) -> torch.Tensor:
"""Run the backbone and flatten its output to ``(N, F)``.
Some backbones return a spatial feature map rather than a vector, so the
trailing dimensions are flattened -- the head expects one row per
sample either way.
:param x: the input batch.
:returns: the features, two-dimensional.
"""
out = self._run_backbone_raw(x)
if isinstance(out, torch.Tensor) and out.ndim > 2:
out = torch.flatten(out, 1)
return out
[docs]
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Return classification logits of shape ``(N, num_classes)``.
:param x: input image batch for the configured TorchVision backbone.
"""
feats = self._run_backbone(x)
if self.use_dropout:
feats = self.dropout(feats)
logits = self.spacr_classifier(feats)
return logits
[docs]
class TorchModel_v2(nn.Module):
"""TorchVision backbone with a spaCR linear head (streamlined variant of :class:`TorchModel`).
:param model_name: TorchVision classification model to load.
:param pretrained: use ImageNet-pretrained weights when available.
:param dropout_rate: dropout probability applied to backbone and spaCR head; ``None`` disables.
:param use_checkpoint: enable gradient checkpointing through the backbone.
:param num_classes: output class count.
:param multilabel: informational flag consumed by external loss/metrics code.
:raises ValueError: if ``model_name`` is not a TorchVision model.
"""
def __init__(
self,
model_name: str = "resnet50",
pretrained: bool = True,
dropout_rate: float = None,
use_checkpoint: bool = False,
num_classes: int = 2,
multilabel: bool = False
):
"""Build the backbone, strip its head, and attach the spaCR classifier."""
super().__init__()
self.model_name = model_name
self.pretrained = bool(pretrained)
self.dropout_rate = (
float(dropout_rate) if dropout_rate is not None else None
)
self.use_checkpoint = bool(use_checkpoint)
self.num_classes = int(num_classes)
self.multilabel = bool(multilabel)
self.base_model = self._init_base_model(pretrained)
if self.model_name == "maxvit_t" and hasattr(self.base_model, "classifier"):
self.base_model.classifier = nn.Sequential(
*list(self.base_model.classifier.children())[:-1]
)
if dropout_rate is not None:
self._apply_dropout_rate(self.base_model, float(dropout_rate))
self.num_ftrs = self._infer_feature_dim()
self._init_spacr_classifier(dropout_rate)
def _apply_dropout_rate(self, module: nn.Module, p: float):
"""Set ``p`` on every dropout layer inside ``module``.
Walks the whole tree, so a backbone with dropout at several depths is set
consistently rather than only at its top level.
"""
for m in module.modules():
if isinstance(m, (nn.Dropout, nn.Dropout2d, nn.Dropout3d)):
m.p = p
def _init_base_model(self, pretrained: bool) -> nn.Module:
"""Build the named torchvision backbone.
An unknown name raises rather than falling back to a default: silently
training a different architecture than the one asked for produces a model
whose results cannot be compared to anything.
"""
fn = models.__dict__.get(self.model_name, None)
if fn is None:
raise ValueError(f"Unknown torchvision model: {self.model_name}")
weights = self._get_weight_choice()
if weights is not None:
return fn(weights=weights if pretrained else None)
else:
return fn(pretrained=bool(pretrained))
def _get_weight_choice(self):
"""The torchvision ``DEFAULT`` weights enum for this model, or ``None``.
``None`` means torchvision ships no pretrained weights under that name, in
which case the backbone starts from random initialisation.
"""
for attr in dir(models):
if attr.lower() == f"{self.model_name}_weights":
return getattr(models, attr).DEFAULT
return None
def _remove_head_for_features(self):
"""Replace the classifier head with identity so the backbone returns features.
``maxvit_t`` IS EXCLUDED: its classifier holds the pooling the forward
pass needs, so replacing it removes more than the head.
"""
if hasattr(self.base_model, "fc"):
self.base_model.fc = nn.Identity()
elif hasattr(self.base_model, "classifier"):
if self.model_name != "maxvit_t":
self.base_model.classifier = nn.Identity()
def _infer_feature_dim(self) -> int:
"""The backbone's feature width, measured by running one dummy image.
MEASURED RATHER THAN TABULATED, so a torchvision version that changes a
backbone's width does not silently mismatch the classifier. Costs one
224x224 forward pass at construction.
"""
self._remove_head_for_features()
self.base_model.eval()
with torch.no_grad():
out = self.base_model(torch.randn(1, 3, 224, 224))
if out.ndim > 2:
out = torch.flatten(out, 1)
return int(out.size(1))
def _init_spacr_classifier(self, dropout_rate: float):
"""Attach the linear head, and dropout before it when a rate was given.
``dropout_rate=None`` means no dropout layer at all rather than a layer
with ``p=0``.
"""
self.use_dropout = dropout_rate is not None
if self.use_dropout:
self.dropout = nn.Dropout(float(dropout_rate))
self.spacr_classifier = nn.Linear(self.num_ftrs, self.num_classes)
def _run_backbone(self, x: torch.Tensor) -> torch.Tensor:
"""Run the backbone, through gradient checkpointing when enabled.
Checkpointing recomputes activations in the backward pass instead of
storing them: less memory, more compute, and the same output.
"""
if self.use_checkpoint:
return _checkpoint_module(
self.base_model, lambda t: self.base_model(t), x)
return self.base_model(x)
[docs]
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Return classification logits of shape ``(N, num_classes)``.
:param x: input image batch for the configured TorchVision backbone.
"""
feats = self._run_backbone(x)
if feats.ndim > 2:
feats = torch.flatten(feats, 1)
if self.use_dropout:
feats = self.dropout(feats)
logits = self.spacr_classifier(feats)
return logits
[docs]
class FocalLossWithLogits(nn.Module):
"""Focal loss for binary, multiclass, and multilabel targets.
Auto-selects the BCE or cross-entropy branch based on the shapes of
``logits`` and ``target``:
- binary: logits ``(N,)`` or ``(N,1)``; target float ``(N,)`` in ``{0,1}``.
- multiclass: logits ``(N,C)``; target long ``(N,)`` in ``[0..C-1]``.
- multilabel: logits ``(N,C)``; target float ``(N,C)`` in ``{0,1}``.
:param alpha: class-balancing factor (float or 1-D tensor of shape ``(C,)``).
:param gamma: focusing parameter.
:param reduction: one of ``'mean'``, ``'sum'``, ``'none'``.
"""
def __init__(self, alpha=1.0, gamma=2.0, reduction="mean"):
"""Store the focal-loss hyperparameters."""
super().__init__()
self.gamma = float(gamma)
self.reduction = reduction
self.alpha = alpha
[docs]
def forward(self, logits, target):
"""Return the focal loss value for the chosen ``reduction`` mode.
:param logits: unnormalized binary, multiclass, or multilabel scores.
:param target: labels shaped for the corresponding logits branch.
"""
if logits.ndim == 1 or logits.size(-1) == 1 or (
logits.ndim == 2 and target.ndim == 2 and target.size(1) == logits.size(1)
):
logits = logits.view_as(target)
bce = F.binary_cross_entropy_with_logits(logits, target, reduction="none")
p = torch.sigmoid(logits)
pt = target * p + (1 - target) * (1 - p)
loss = (self.alpha * (1 - pt).pow(self.gamma) * bce)
else:
if target.dtype != torch.long:
target = target.long()
logp = F.log_softmax(logits, dim=1)
p = torch.exp(logp)
pt = p.gather(1, target.unsqueeze(1)).squeeze(1)
ce = F.nll_loss(logp, target, reduction="none")
if isinstance(self.alpha, torch.Tensor):
alpha = self.alpha.to(logits.device)[target]
else:
alpha = float(self.alpha)
loss = alpha * (1 - pt).pow(self.gamma) * ce
if self.reduction == "mean":
return loss.mean()
if self.reduction == "sum":
return loss.sum()
return loss
[docs]
class ResNet(nn.Module):
"""ResNet backbone with a two-layer spaCR binary-classification head.
:param resnet_type: one of ``'resnet18'``/``'resnet34'``/``'resnet50'``/``'resnet101'``/``'resnet152'``.
:param dropout_rate: dropout probability before the final linear layer; ``None`` disables.
:param use_checkpoint: enable gradient checkpointing through the ResNet backbone.
:param init_weights: ``'imagenet'`` for pretrained weights or ``'none'`` for random init.
:raises ValueError: if ``resnet_type`` is unsupported, or if
``init_weights`` is neither ``'imagenet'`` nor ``'none'``.
"""
def __init__(self, resnet_type='resnet50', dropout_rate=None, use_checkpoint=False, init_weights='imagenet'):
"""Select the backbone and delegate head construction to :meth:`initialize_base`."""
super(ResNet, self).__init__()
resnet_map = {
'resnet18': {'func': models.resnet18, 'weights': ResNet18_Weights.IMAGENET1K_V1},
'resnet34': {'func': models.resnet34, 'weights': ResNet34_Weights.IMAGENET1K_V1},
'resnet50': {'func': models.resnet50, 'weights': ResNet50_Weights.IMAGENET1K_V1},
'resnet101': {'func': models.resnet101, 'weights': ResNet101_Weights.IMAGENET1K_V1},
'resnet152': {'func': models.resnet152, 'weights': ResNet152_Weights.IMAGENET1K_V1}
}
if resnet_type not in resnet_map:
raise ValueError(f"Invalid resnet_type. Choose from {list(resnet_map.keys())}")
self.initialize_base(resnet_map[resnet_type], dropout_rate, use_checkpoint, init_weights)
[docs]
def initialize_base(self, base_model_dict, dropout_rate, use_checkpoint, init_weights):
"""Build the backbone (with or without pretrained weights) and the two-layer head.
:param base_model_dict: dict with keys ``func`` (model constructor) and ``weights``.
:param dropout_rate: dropout probability applied between the two linear layers.
:param use_checkpoint: enable gradient checkpointing through the backbone.
:param init_weights: ``'imagenet'`` or ``'none'``.
:raises ValueError: if ``init_weights`` is neither ``'imagenet'`` nor ``'none'``.
"""
if init_weights == 'imagenet':
self.resnet = base_model_dict['func'](weights=base_model_dict['weights'])
elif init_weights == 'none':
self.resnet = base_model_dict['func'](weights=None)
else:
raise ValueError("init_weights should be either 'imagenet' or 'none'")
self.fc1 = nn.Linear(1000, 500)
self.use_dropout = dropout_rate != None
self.use_checkpoint = use_checkpoint
if self.use_dropout:
self.dropout = nn.Dropout(dropout_rate)
self.fc2 = nn.Linear(500, 1)
[docs]
def forward(self, x):
"""Return the flattened single-logit prediction for input batch ``x``.
:param x: image batch for the configured ResNet backbone.
"""
if self.use_checkpoint:
x = _checkpoint_module(self.resnet, self.resnet, x)
else:
x = self.resnet(x)
x = F.relu(self.fc1(x))
if self.use_dropout:
x = self.dropout(x)
logits = self.fc2(x).flatten()
return logits
[docs]
def split_my_dataset(dataset, split_ratio=0.1):
"""Randomly split ``dataset`` into ``(train, val)`` subsets.
:param dataset: source dataset.
:param split_ratio: fraction of samples reserved for validation.
:returns: ``(train_subset, val_subset)``.
"""
num_samples = len(dataset)
indices = list(range(num_samples))
split_idx = int((1 - split_ratio) * num_samples)
random.shuffle(indices)
train_indices, val_indices = indices[:split_idx], indices[split_idx:]
train_dataset = Subset(dataset, train_indices)
val_dataset = Subset(dataset, val_indices)
return train_dataset, val_dataset
[docs]
def classification_metrics(all_labels, prediction_pos_probs, loss, epoch):
"""Return a one-row DataFrame of accuracy, PR-AUC, and optimal-threshold stats.
:param all_labels: ground-truth binary labels.
:param prediction_pos_probs: predicted positive-class probabilities.
:param loss: loss tensor for the epoch (``.item()`` is called).
:param epoch: epoch number used as the row index.
:returns: DataFrame indexed by epoch with accuracy, per-class accuracy, loss,
PR-AUC, and optimal threshold columns.
:raises ValueError: if ``all_labels`` and ``prediction_pos_probs`` have different lengths.
"""
if len(all_labels) != len(prediction_pos_probs):
raise ValueError(f"all_labels ({len(all_labels)}) and pred_labels ({len(prediction_pos_probs)}) have different lengths")
unique_labels = np.unique(all_labels)
if len(unique_labels) >= 2:
pr_labels = np.array(all_labels).astype(int)
precision, recall, thresholds = precision_recall_curve(pr_labels, prediction_pos_probs, pos_label=1)
pr_auc = auc(recall, precision)
thresholds = np.append(thresholds, 0.0)
f1_scores = 2 * (precision * recall) / (precision + recall)
optimal_idx = np.nanargmax(f1_scores)
optimal_threshold = thresholds[optimal_idx]
pred_labels = [int(p > 0.5) for p in prediction_pos_probs]
if len(unique_labels) < 2:
optimal_threshold = 0.5
pred_labels = [int(p > optimal_threshold) for p in prediction_pos_probs]
pr_auc = np.nan
data = {'label': all_labels, 'pred': pred_labels}
df = pd.DataFrame(data)
pc_df = df[df['label'] == 1.0]
nc_df = df[df['label'] == 0.0]
correct = df[df['label'] == df['pred']]
acc_all = len(correct) / len(df)
if len(pc_df) > 0:
correct_pc = pc_df[pc_df['label'] == pc_df['pred']]
acc_pc = len(correct_pc) / len(pc_df)
else:
acc_pc = np.nan
if len(nc_df) > 0:
correct_nc = nc_df[nc_df['label'] == nc_df['pred']]
acc_nc = len(correct_nc) / len(nc_df)
else:
acc_nc = np.nan
data_dict = {'accuracy': acc_all, 'neg_accuracy': acc_nc, 'pos_accuracy': acc_pc, 'loss':loss.item(),'prauc':pr_auc, 'optimal_threshold':optimal_threshold}
data_df = pd.DataFrame(data_dict, index=[str(epoch)])
return data_df
[docs]
def compute_irm_penalty(losses, dummy_w, device):
"""Return the IRM penalty as the sum of squared gradient dot-products across environments.
:param losses: per-environment loss tensors.
:param dummy_w: scalar dummy weight used for gradient computation.
:param device: torch device on which to compute the penalty.
:returns: scalar IRM penalty value.
"""
weighted_losses = [loss.clone().detach().requires_grad_(True).to(device) * dummy_w for loss in losses]
gradients = [grad(w_loss, dummy_w, create_graph=True)[0] for w_loss in weighted_losses]
irm_penalty = 0.0
for g1, g2 in combinations(gradients, 2):
irm_penalty += (g1.dot(g2))**2
return irm_penalty
def _list_torchvision_model_names() -> set[str]:
"""Every torchvision classification FACTORY, and nothing else.
THE FALLBACK USED TO POLLUTE THE ANSWER. It added every public callable
in ``torchvision.models``, which is the factories plus the classes they
build (``AlexNet``, ``ResNet``) plus every weights enum
(``AlexNet_Weights``). Two consequences, both real: the names offered
to a user who mistyped began "AlexNet, AlexNet_Weights, ConvNeXt,
ConvNeXt_Base_Weights" -- twenty entries and not one of them a name
that works -- and `choose_model('AlexNet_Weights')` passed the name
check and failed inside the wrapper instead.
The modern API answers exactly this question, so the fallback runs only
when it answers nothing, and then keeps only lower-case factory names.
"""
try:
names = set(tv_models.list_models(module=tv_models))
except Exception: # noqa: BLE001
names = set()
if names:
return names
return {
name for name, value in tv_models.__dict__.items()
if not name.startswith("_") and callable(value)
and name.islower() and not name.endswith("_weights")
}
[docs]
def choose_model(model_type: str,
device: torch.device,
init_weights: bool = True,
dropout_rate: float = 0.0,
use_checkpoint: bool = False,
channels: int = 3,
height: int = 224,
width: int = 224,
chan_dict: Optional[dict[str, Any]] = None,
num_classes: int = 2,
verbose: bool = False) -> Optional[nn.Module]:
"""Instantiate a classification model by name for binary or multiclass problems.
:param model_type: TorchVision model name (e.g. ``'resnet50'``, ``'vit_b_16'``).
``'custom'`` passes the name check but then raises ``NotImplementedError``,
as no custom builder is wired up.
:param device: unused; the model is built on the CPU and the caller moves it.
:param init_weights: load pretrained weights when available.
:param dropout_rate: dropout probability before the classifier head (``None``/``0`` disables).
:param use_checkpoint: enable gradient checkpointing for the backbone.
:param channels: unused; the forward sanity check always feeds 3 channels.
:param height: square input resolution, forwarded as ``TorchModel(image_size=...)``;
it therefore fixes the dummy-forward size used to infer the backbone feature
dimension (and so the size of the classifier head) as well as both dimensions
of the square forward sanity check (falls back to ``224`` when falsy).
:param width: unused; ``height`` sets both dimensions.
:param chan_dict: unused; reserved for the unimplemented custom branch.
:param num_classes: output class count; ``1`` yields a single-logit BCE head.
:param verbose: print the model structure when ``True``.
:returns: The instantiated ``nn.Module``.
:raises ValueError: ``model_type`` names no backbone, or the built model
does not produce logits of the requested shape.
Unsupported names raise immediately and include close TorchVision matches
when available, so configuration errors are reported before training.
"""
import difflib
tv_names = _list_torchvision_model_names()
valid_names = set(tv_names) | {"custom"}
if model_type not in valid_names:
close = difflib.get_close_matches(str(model_type), sorted(tv_names),
n=5, cutoff=0.6)
suggestion = (f" Did you mean {close}?" if close else
f" Names spaCR can build include "
f"{sorted(tv_names)[:8]} and {len(tv_names) - 8} more.")
raise ValueError(
f"model_type={model_type!r} names no classification backbone "
f"torchvision provides.{suggestion}")
print(
f"Model parameters: Architecture: {model_type} "
f"init_weights: {init_weights} dropout_rate: {dropout_rate} "
f"use_checkpoint: {use_checkpoint}", flush=True
)
if model_type == "custom":
raise NotImplementedError(
"Model type 'custom' selected but no CustomCellClassifier is wired. "
"Provide your implementation or use a TorchVision backbone."
)
head_dim = max(1, int(num_classes))
img_size = int(height) if height else 224
base_model = TorchModel(
model_name=model_type,
pretrained=bool(init_weights),
dropout_rate=(dropout_rate if (dropout_rate and dropout_rate > 0) else None),
use_checkpoint=use_checkpoint,
num_classes=head_dim,
image_size=img_size,
)
try:
base_model.eval()
with torch.no_grad():
dummy = torch.randn(1, 3, img_size, img_size)
z = base_model(dummy)
if isinstance(z, dict):
raise RuntimeError("Selected model returned a dict, not logits.")
if not isinstance(z, torch.Tensor) or z.ndim != 2 or z.size(1) != head_dim:
raise RuntimeError(
f"Expected logits of shape (1,{head_dim}); got {type(z)} / {getattr(z, 'shape', None)}"
)
except Exception as error: # noqa: BLE001
raise ValueError(
f"model_type={model_type!r} built, but its forward pass does not "
f"produce {head_dim} logit(s) at {img_size}x{img_size}: {error}"
) from error
if verbose:
print("\n", base_model)
return base_model
[docs]
def calculate_loss(output, target, prefer_focal=False, gamma=2.0, alpha=1.0, reduction="mean"):
"""Auto-select and return a loss for binary, multiclass, or multilabel problems.
Dispatches based on the shapes/dtypes of ``output`` and ``target``:
- binary: logits ``(N,1)``, float targets in ``{0,1}`` -> BCE / focal-BCE.
- multiclass: logits ``(N,C)``, long targets ``(N,)`` -> CE / focal-CE.
- multilabel: logits ``(N,C)``, float targets ``(N,C)`` -> BCE / focal-BCE.
:param output: model logits.
:param target: ground-truth labels.
:param prefer_focal: use the focal-loss variant instead of plain CE/BCE.
:param gamma: focal-loss focusing parameter.
:param alpha: focal-loss class-balancing factor.
:param reduction: one of ``'mean'``, ``'sum'``, ``'none'``.
:returns: scalar loss tensor (or per-sample tensor when ``reduction='none'``).
"""
def _focal_bce_with_logits(logits, y, alpha=1.0, gamma=2.0, reduction="mean"):
"""Return focal binary cross-entropy for ``logits`` and targets ``y``."""
p = torch.sigmoid(logits)
ce = F.binary_cross_entropy_with_logits(logits, y, reduction="none")
p_t = p * y + (1 - p) * (1 - y)
loss = alpha * (1 - p_t).pow(gamma) * ce
if reduction == "mean":
return loss.mean()
elif reduction == "sum":
return loss.sum()
return loss
def _focal_cross_entropy(logits, y_idx, alpha=1.0, gamma=2.0, reduction="mean"):
"""Return focal cross-entropy for logits and class indices ``y_idx``."""
log_p = F.log_softmax(logits, dim=1)
p = log_p.exp()
log_p_t = log_p.gather(1, y_idx.view(-1,1)).squeeze(1)
p_t = p.gather(1, y_idx.view(-1,1)).squeeze(1)
loss = -alpha * (1 - p_t).pow(gamma) * log_p_t
if reduction == "mean":
return loss.mean()
elif reduction == "sum":
return loss.sum()
return loss
if output.ndim == 1:
output = output.unsqueeze(1)
N, C = output.shape[0], output.shape[1]
if C == 1:
target = target.float().view(N, 1)
if prefer_focal:
return _focal_bce_with_logits(output, target, alpha=alpha, gamma=gamma, reduction=reduction)
return F.binary_cross_entropy_with_logits(output, target, reduction=reduction)
if target.dtype == torch.long and target.ndim == 1:
if prefer_focal:
return _focal_cross_entropy(output, target, alpha=alpha, gamma=gamma, reduction=reduction)
return F.cross_entropy(output, target, reduction=reduction)
if target.ndim == 1:
target = torch.nn.functional.one_hot(target.long(), num_classes=C).float()
else:
target = target.float().view(N, C)
if prefer_focal:
return _focal_bce_with_logits(output, target, alpha=alpha, gamma=gamma, reduction=reduction)
return F.binary_cross_entropy_with_logits(output, target, reduction=reduction)
[docs]
def pick_best_model(src):
"""Return the strongest checkpoint anywhere below ``src``.
Current artifacts are ranked by their stored validation metric and role;
legacy files fall back to their ``_acc_``/``_epoch_`` filename fields.
:param src: model directory or a checkpoint path.
:returns: absolute path to the top-ranked checkpoint.
"""
if os.path.isfile(src):
return os.path.abspath(src)
if not os.path.isdir(src):
raise FileNotFoundError(f"Model directory does not exist: {src}")
pth_files = sorted(glob.glob(os.path.join(src, "**", "*.pth"),
recursive=True))
if not pth_files:
raise FileNotFoundError(f"No .pth model checkpoints found below {src}")
pattern = re.compile(r'_epoch_(\d+)_acc_(\d+(?:\.\d+)?)')
epoch_pattern = re.compile(r'_epoch_(\d+)')
def sort_key(x):
"""Return ``(role, accuracy, epoch)`` from metadata or legacy name."""
role_rank = 2 if "_best_" in os.path.basename(x) else 0
accuracy = float("-inf")
epoch = 0
try:
payload = torch.load(x, map_location="cpu", weights_only=False)
if isinstance(payload, dict):
role = payload.get("artifact_role")
if role == "best":
role_rank = 2
elif role == "milestone":
role_rank = 1
metrics = payload.get("metrics") or {}
value = metrics.get("accuracy")
if value is not None and np.isfinite(float(value)):
accuracy = float(value)
training = payload.get("training_state") or {}
epoch = int(training.get("epoch") or 0)
except Exception:
pass
match = pattern.search(os.path.basename(x))
if match and not np.isfinite(accuracy):
epoch = int(match.group(1))
accuracy = float(match.group(2)) / 100.0
elif epoch == 0:
match = epoch_pattern.search(os.path.basename(x))
if match:
epoch = int(match.group(1))
return role_rank, accuracy, epoch
return max(pth_files, key=sort_key)
[docs]
def get_paths_from_db(df, png_df, image_type='cell_png'):
"""Return rows of ``png_df`` whose path contains ``image_type`` and whose ``prcfo`` is in ``df``.
:param df: DataFrame indexed by ``prcfo`` identifiers.
:param png_df: DataFrame of PNG metadata with ``png_path`` and ``prcfo`` columns.
:param image_type: substring that must appear in ``png_path``.
:returns: filtered subset of ``png_df``.
"""
objects = df.index.tolist()
filtered_df = png_df[png_df['png_path'].str.contains(image_type) & png_df['prcfo'].isin(objects)]
return filtered_df
[docs]
def save_file_lists(dst, data_set, ls):
"""Write ``ls`` as a single-column CSV named ``<data_set>.csv`` under ``dst``.
:param dst: destination directory.
:param data_set: column name and file stem.
:param ls: iterable of values to persist.
:returns: None.
"""
df = pd.DataFrame(ls, columns=[data_set])
tabular.write_table(df, f'{dst}/{data_set}.csv', canonicalise=False)
return
[docs]
def augment_single_image(args):
"""Save six augmentations of one image (original, 90/180/270 rotations, H/V flips).
:param args: ``(img_path, dst)`` tuple.
:returns: None.
"""
img_path, dst = args
img = read_image_rgb(img_path, cv2.IMREAD_UNCHANGED)
if img is None:
raise ValueError(f"Could not read image: {img_path}")
filename = os.path.basename(img_path).split('.')[0]
write_image_rgb(os.path.join(dst, f"{filename}_original.png"), img)
img_rot_90 = cv2.rotate(img, cv2.ROTATE_90_CLOCKWISE)
write_image_rgb(os.path.join(dst, f"{filename}_rot_90.png"), img_rot_90)
img_rot_180 = cv2.rotate(img, cv2.ROTATE_180)
write_image_rgb(os.path.join(dst, f"{filename}_rot_180.png"), img_rot_180)
img_rot_270 = cv2.rotate(img, cv2.ROTATE_90_COUNTERCLOCKWISE)
write_image_rgb(os.path.join(dst, f"{filename}_rot_270.png"), img_rot_270)
img_flip_hor = cv2.flip(img, 1)
write_image_rgb(os.path.join(dst, f"{filename}_flip_hor.png"), img_flip_hor)
img_flip_ver = cv2.flip(img, 0)
write_image_rgb(os.path.join(dst, f"{filename}_flip_ver.png"), img_flip_ver)
[docs]
def augment_images(file_paths, dst):
"""Run :func:`augment_single_image` in parallel over ``file_paths``.
Spawn workers so they inherit no locks held by other threads in this
process. Close and join the pool after mapping; terminating forked workers
can hang while their exit handlers wait on inherited locks.
:param file_paths: iterable of source image paths.
:param dst: destination folder (created if missing).
:returns: None.
"""
if not os.path.exists(dst):
os.makedirs(dst)
args_list = [(img_path, dst) for img_path in file_paths]
if not args_list:
return
from .resource_log import _array_file_nbytes, _guard_workers
workers = _guard_workers('augment', cpu_count(),
_array_file_nbytes(args_list[0][0]))
workers = max(1, min(int(workers), len(args_list)))
pool = Pool(workers, context=_augment_pool_context())
try:
pool.map(augment_single_image, args_list)
finally:
pool.close()
pool.join()
def _augment_pool_context():
"""Return the multiprocessing context :func:`augment_images` uses."""
import multiprocessing
return multiprocessing.get_context('spawn')
[docs]
def suggest_training_changes(
dst,
train_csv=None,
val_csv=None,
last_k=25,
min_epochs=10,
gap_threshold_acc=0.05,
plateau_eps=1e-3,
noisy_var_ratio=0.03,
):
"""Inspect saved training/validation progress CSVs and propose concrete training changes.
:param dst: folder where progress CSVs were saved.
:param train_csv: explicit train-CSV path; auto-detected in ``dst`` if ``None``.
:param val_csv: explicit val-CSV path; auto-detected in ``dst`` if ``None``.
:param last_k: number of recent epochs used for trend and plateau checks.
:param min_epochs: minimum epochs before most suggestions are issued.
:param gap_threshold_acc: accuracy generalization-gap threshold (train - val).
:param plateau_eps: absolute slope threshold used to declare a plateau.
:param noisy_var_ratio: instability flag threshold on ``stdev/mean`` of recent val loss.
:returns: dict with ``summary`` (key scalars), ``flags`` (short codes),
and ``suggestions`` (ordered suggestion strings).
"""
import os, glob
import numpy as np
import pandas as pd
def _scalar(val):
"""Ensure a single float even if a Series sneaks through.
The Series branch is currently unreachable: every call site passes
``<Series>.iloc[<int>]``, which yields a numpy scalar. It is kept as a
deliberate guard because the label-based ``.loc`` lookups this
function used to rely on returned a Series whenever the progress CSV
had a duplicated index, and that is an easy regression to reintroduce.
"""
if isinstance(val, pd.Series):
return float(val.iloc[0])
return float(val)
def _find_csv(root, hint):
"""Return the lexically last matching CSV in ``root``, if any."""
cs = sorted(glob.glob(os.path.join(root, f"*{hint}*.csv")))
return cs[-1] if cs else None
def _normalize_cols(df):
"""Return ``df`` with normalized, aliased, first-occurrence columns."""
m = {c: c.strip().lower() for c in df.columns}
df = df.rename(columns=m)
df = df.loc[:, ~df.columns.duplicated(keep='first')]
aliases = {
"accuracy": ["acc", "accuracy", "train_acc", "val_acc"],
"loss": ["loss", "train_loss", "val_loss"],
"f1_macro": ["f1_macro", "macro_f1", "f1macro", "f1"],
"epoch": ["epoch", "epochs", "step"],
"lr": ["lr", "learning_rate"],
}
name_map = {}
for canon, opts in aliases.items():
for o in opts:
if o in df.columns:
name_map[o] = canon
df = df.rename(columns=name_map)
df = df.loc[:, ~df.columns.duplicated(keep='first')]
return df
def _poly_slope(y):
"""Return the finite linear slope of ``y``, or zero when undefined."""
if len(y) < 2 or np.allclose(y, y[0]):
return 0.0
x = np.arange(len(y), dtype=float)
mask = np.isfinite(y)
if mask.sum() < 2:
return 0.0
coef = np.polyfit(x[mask], y[mask], 1)
return float(coef[0])
def _last_seq(series, k):
"""Return at most the final ``k`` values as a floating-point array."""
s = np.asarray(series, dtype=float)
return s[-min(k, len(s)):] if len(s) else np.array([])
train_csv = train_csv or _find_csv(dst, "train")
val_csv = val_csv or _find_csv(dst, "val")
out = {"summary": {}, "flags": [], "suggestions": []}
if not train_csv or not os.path.exists(train_csv):
out["flags"].append("missing_train_csv")
out["suggestions"].append("Could not locate train CSV; ensure _save_progress writes a train CSV in dst.")
return out
if not val_csv or not os.path.exists(val_csv):
out["flags"].append("missing_val_csv")
out["suggestions"].append("Could not locate val CSV; enable validation logging in _save_progress.")
return out
tr = tabular.read_table(train_csv, report=None)
va = tabular.read_table(val_csv, report=None)
tr = _normalize_cols(tr)
va = _normalize_cols(va)
for col in ("epoch", "loss"):
if col not in tr.columns or col not in va.columns:
out["flags"].append(f"missing_required_col:{col}")
out["suggestions"].append(f"Progress CSVs lack '{col}'. Ensure _save_progress writes epoch and loss.")
return out
best_pos = int(va["loss"].argmin())
best_val_loss = _scalar(va["loss"].iloc[best_pos])
best_epoch = int(_scalar(va["epoch"].iloc[best_pos])) if "epoch" in va.columns else (best_pos + 1)
final = {
"train_loss": float(tr["loss"].iloc[-1]),
"val_loss": float(va["loss"].iloc[-1]),
}
if "accuracy" in tr.columns:
final["train_accuracy"] = _scalar(tr["accuracy"].iloc[-1])
if "accuracy" in va.columns:
final["val_accuracy"] = _scalar(va["accuracy"].iloc[-1])
if "f1_macro" in tr.columns:
final["train_f1_macro"] = _scalar(tr["f1_macro"].iloc[-1])
if "f1_macro" in va.columns:
final["val_f1_macro"] = _scalar(va["f1_macro"].iloc[-1])
tr_last = _last_seq(tr["loss"], last_k)
va_last = _last_seq(va["loss"], last_k)
slope_tr = _poly_slope(tr_last)
slope_va = _poly_slope(va_last)
val_mean = float(np.nanmean(va_last)) if len(va_last) else np.nan
val_std = float(np.nanstd(va_last)) if len(va_last) else np.nan
unstable = (len(va_last) >= max(5, last_k//2)) and np.isfinite(val_mean) and (val_std > noisy_var_ratio * max(val_mean, 1e-8))
gen_gap = None
if "accuracy" in tr.columns and "accuracy" in va.columns:
gen_gap = _scalar(tr["accuracy"].iloc[-1]) - _scalar(va["accuracy"].iloc[-1])
f1_nan_train = "f1_macro" in tr.columns and np.isnan(tr["f1_macro"]).mean() > 0.2
f1_nan_val = "f1_macro" in va.columns and np.isnan(va["f1_macro"]).mean() > 0.2
since_best = int(tr.shape[0] - (best_pos + 1))
val_loss_delta_from_best = float(va["loss"].iloc[-1] - best_val_loss)
out["summary"].update(
dict(
best_epoch=best_epoch,
best_val_loss=best_val_loss,
final_metrics=final,
slope_train_loss_last_k=slope_tr,
slope_val_loss_last_k=slope_va,
val_loss_std_last_k=val_std,
epochs=len(tr),
since_best=since_best,
gen_gap_acc=gen_gap,
)
)
E = len(tr)
if E < min_epochs:
out["flags"].append("few_epochs")
out["suggestions"].append(f"Only {E} epochs logged (<{min_epochs}). Consider training longer or using a warmer LR schedule.")
if len(va_last) >= max(5, last_k//2) and abs(slope_va) < plateau_eps:
out["flags"].append("val_plateau")
out["suggestions"].extend([
"Validation loss plateau detected: try ReduceLROnPlateau (factor=0.1, patience=5–10) or cosine annealing with warm restarts.",
"Add/strengthen data augmentation; if already heavy, try stochastic depth/label smoothing=0.05–0.1.",
"If capacity may be limiting, consider a larger backbone or unfreezing more layers after a warmup.",
])
overfit_like = False
if slope_tr < -plateau_eps and slope_va > plateau_eps:
overfit_like = True
if gen_gap is not None and gen_gap > gap_threshold_acc:
overfit_like = True
if overfit_like:
out["flags"].append("overfitting")
out["suggestions"].extend([
"Overfitting signs: increase regularization (weight_decay e.g. 0.05→0.1), enable/raise dropout (e.g. 0.2–0.5).",
"Increase augmentation (color jitter, random crops, flips, CutMix/MixUp).",
"Use early stopping on val loss; keep the best checkpoint (epoch with min val loss).",
"Consider smaller head or freeze more backbone layers for longer warmup.",
])
train_acc_low = ("accuracy" in tr.columns and final.get("train_accuracy", 0.0) < 0.70)
losses_not_decreasing = (slope_tr > -plateau_eps and slope_va > -plateau_eps)
if train_acc_low and losses_not_decreasing:
out["flags"].append("underfitting")
out["suggestions"].extend([
"Underfitting signs: increase learning rate 2–4× or use a longer schedule (more epochs with decay).",
"Reduce regularization (lower weight_decay), or increase model capacity (bigger backbone).",
"Verify labels and channel order/normalization; large label noise or wrong preprocessing can cap accuracy.",
])
if unstable:
out["flags"].append("unstable_training")
out["suggestions"].extend([
"Validation loss is noisy: lower LR (e.g., ×0.5), increase batch size, or enable gradient clipping (clip_norm=1.0).",
"Ensure deterministic preprocessing and consistent image normalization.",
])
if f1_nan_train or f1_nan_val:
out["flags"].append("f1_nan_detected")
out["suggestions"].extend([
"F1(macro) shows NaN—ensure each split has ≥2 classes and use stratified sampling.",
"If highly imbalanced, prefer class weights or focal loss (you already use focal—verify label distribution).",
])
if since_best >= max(5, last_k//2) and val_loss_delta_from_best > plateau_eps:
out["flags"].append("past_best_regression")
out["suggestions"].extend([
f"Validation loss has worsened by +{val_loss_delta_from_best:.4f} since best epoch {best_epoch}: adopt early stopping and keep best checkpoint.",
"Also try ReduceLROnPlateau triggered on val loss.",
])
if ("accuracy" in va.columns and "f1_macro" in va.columns
and np.isfinite(final.get("val_accuracy", np.nan))
and np.isfinite(final.get("val_f1_macro", np.nan))
and (final["val_accuracy"] - final["val_f1_macro"] > 0.10)):
out["flags"].append("class_imbalance_suspected")
out["suggestions"].extend([
"Accuracy ≫ macro-F1 suggests imbalance: use class weights, oversampling, or stronger focal loss (gamma 2–3, tune alpha).",
"Track per-class metrics/confusion matrices to verify rare classes.",
])
out["suggestions"] = list(dict.fromkeys(out["suggestions"]))
return out
def _infer_indices(target: torch.Tensor, num_classes: int) -> torch.Tensor:
"""Return class indices (N,) from target that may be long or one-hot/float."""
if target.dtype == torch.long:
return target.view(-1)
if target.ndim == 2 and target.size(1) == num_classes:
return target.argmax(dim=1).long()
return (target.view(-1) > 0.5).long()
[docs]
def estimate_class_counts(loader, num_classes: int, src=None, classes=None) -> torch.Tensor:
"""Return per-class sample counts as a ``LongTensor`` of length ``num_classes``.
When ``src`` and ``classes`` are provided the counts are taken from the file
listings under ``src/<class>``, avoiding a slow DataLoader iteration on NAS.
:param loader: fallback DataLoader iterated only when folder info is missing.
:param num_classes: number of output classes.
:param src: parent folder containing per-class subfolders.
:param classes: ordered class-folder names matching ``src``.
:returns: ``LongTensor`` of per-class counts.
"""
if src is not None and classes is not None:
counts = torch.zeros(num_classes, dtype=torch.long)
for i, cls in enumerate(classes):
cls_dir = os.path.join(src, cls)
if os.path.isdir(cls_dir):
n = sum(1 for f in os.listdir(cls_dir) if os.path.isfile(os.path.join(cls_dir, f)))
counts[i] = n
print(f"Class counts (from folders): {dict(zip(classes, counts.tolist()))}")
return counts
print("Warning: counting classes by iterating DataLoader (slow on NAS). "
"Pass src and classes to avoid this.")
counts = torch.zeros(num_classes, dtype=torch.long)
for _, y, _ in loader:
y = y.detach()
idx = _infer_indices(y, num_classes)
binc = torch.bincount(idx, minlength=num_classes)
counts[:num_classes] += binc[:num_classes]
return counts
[docs]
def build_loss(loss_type: str = "ce",
num_classes: int = 2,
class_counts: Optional[torch.Tensor] = None,
label_smoothing: float = 0.0,
focal_gamma: float = 2.0,
focal_alpha: Optional[float] = None,
logit_adjust_tau: float = 0.0,
asl_gamma_pos: float = 0.0,
asl_gamma_neg: float = 4.0,
asl_clip: float = 0.05):
"""Return a closure ``loss_fn(logits, target)`` implementing the requested loss.
Supported ``loss_type`` values: ``'ce'``, ``'ce_smooth'``, ``'ce_weighted'``,
``'focal_ce'``, ``'bce'``, ``'focal_bce'``, ``'logit_adjust_ce'``, ``'asl'``, ``'auto'``.
``num_classes==1`` selects binary (BCE variants); ``>=2`` selects multiclass (CE variants).
:param loss_type: loss identifier (see above).
:param num_classes: output class count.
:param class_counts: per-class sample counts used to derive weights or logit adjustment.
:param label_smoothing: label-smoothing epsilon for ``ce_smooth``.
:param focal_gamma: focal-loss focusing parameter.
:param focal_alpha: focal-loss class-balancing factor (float or per-class tensor).
:param logit_adjust_tau: strength of the Menon-et-al. logit adjustment; 0 disables.
:param asl_gamma_pos: asymmetric-loss gamma for positives.
:param asl_gamma_neg: asymmetric-loss gamma for negatives.
:param asl_clip: asymmetric-loss negative-probability clip.
:returns: ``loss_fn(logits, target)`` callable returning a scalar tensor.
:raises ValueError: if ``loss_type`` is unknown or incompatible with ``num_classes``.
"""
lt = (loss_type or "ce").lower()
def _infer_indices(target: torch.Tensor, C: int) -> torch.Tensor:
"""Return class indices from an index vector or 2-D target matrix."""
if target.ndim == 2:
return target.argmax(dim=1).long()
return target.long().view(-1)
class_weights = None
logit_adjust = None
if class_counts is not None:
counts = class_counts.to(dtype=torch.float)
counts = torch.clamp(counts, min=1.0)
priors = counts / counts.sum()
inv = 1.0 / priors
class_weights = (inv / inv.mean()).to(dtype=torch.float)
if logit_adjust_tau > 0:
logit_adjust = (float(logit_adjust_tau) * priors.log()).to(dtype=torch.float)
def _focal_bce(logits, y, alpha, gamma):
"""Return mean focal binary cross-entropy for ``logits`` and ``y``."""
p = torch.sigmoid(logits)
ce = F.binary_cross_entropy_with_logits(logits, y, reduction="none")
pt = p * y + (1 - p) * (1 - y)
w = (1 - pt).pow(gamma)
if alpha is not None:
w = w * (alpha * y + (1 - alpha) * (1 - y))
return (w * ce).mean()
def _focal_ce(logits, y_idx, alpha, gamma):
"""Return mean focal cross-entropy for logits and class indices."""
log_p = F.log_softmax(logits, dim=1)
p = log_p.exp()
log_p_t = log_p.gather(1, y_idx.view(-1, 1)).squeeze(1)
p_t = p.gather(1, y_idx.view(-1, 1)).squeeze(1)
w = (1 - p_t).pow(gamma)
if alpha is not None:
if torch.is_tensor(alpha) and alpha.numel() > 1:
a = alpha.to(logits.device)[y_idx]
else:
a = float(alpha)
loss = -a * w * log_p_t
else:
loss = -w * log_p_t
return loss.mean()
def _asl(logits, y, gpos, gneg, clip):
"""Return mean asymmetric multilabel loss for logits and targets."""
x_sigmoid = torch.sigmoid(logits)
xs_pos = x_sigmoid
xs_neg = 1 - x_sigmoid
if clip and clip > 0:
xs_neg = torch.clamp(xs_neg + clip, max=1.0)
loss = y * torch.log(xs_pos.clamp_min(1e-8)) + (1 - y) * torch.log(xs_neg.clamp_min(1e-8))
pt = xs_pos * y + xs_neg * (1 - y)
one_sided = (1 - pt).pow(gpos * y + gneg * (1 - y))
return -(one_sided * loss).mean()
def _auto_choice() -> str:
"""Return the default loss name from class count and imbalance."""
if num_classes >= 2:
if class_counts is not None:
props = (class_counts.float() / class_counts.sum().clamp_min(1))
if props.min() < 0.10:
return "logit_adjust_ce"
return "ce"
else:
return "bce"
if lt == "auto":
lt = _auto_choice()
if num_classes == 1:
if lt in ("bce", "binary_cross_entropy_with_logits"):
def loss_fn(logits, target):
"""Closure: compute the selected per-batch loss from ``(logits, target)``."""
y = target.float().view(-1, 1)
return F.binary_cross_entropy_with_logits(logits, y)
elif lt in ("focal_bce", "focal", "focal_loss"):
def loss_fn(logits, target):
"""Closure: compute the selected per-batch loss from ``(logits, target)``."""
y = target.float().view(-1, 1)
return _focal_bce(logits, y, focal_alpha, focal_gamma)
else:
raise ValueError(f"loss_type '{loss_type}' not valid for binary (num_classes=1)")
return loss_fn
if lt in ("ce", "cross_entropy"):
def loss_fn(logits, target):
"""Closure: compute the selected per-batch loss from ``(logits, target)``."""
y = _infer_indices(target, num_classes)
w = class_weights.to(logits.device) if class_weights is not None else None
return F.cross_entropy(logits, y, weight=w)
elif lt in ("ce_smooth", "label_smoothing"):
def loss_fn(logits, target):
"""Closure: compute the selected per-batch loss from ``(logits, target)``."""
y = _infer_indices(target, num_classes)
w = class_weights.to(logits.device) if class_weights is not None else None
return F.cross_entropy(logits, y, weight=w, label_smoothing=float(label_smoothing))
elif lt in ("ce_weighted",):
if class_weights is None:
raise ValueError("ce_weighted requires class_counts (to derive weights).")
def loss_fn(logits, target):
"""Closure: compute the selected per-batch loss from ``(logits, target)``."""
y = _infer_indices(target, num_classes)
return F.cross_entropy(logits, y, weight=class_weights.to(logits.device))
elif lt in ("focal_ce", "focal", "focal_loss"):
alpha = None
if focal_alpha is not None:
alpha = focal_alpha if torch.is_tensor(focal_alpha) else float(focal_alpha)
if torch.is_tensor(alpha) and alpha.numel() == num_classes:
alpha = alpha.to(torch.float)
def loss_fn(logits, target):
"""Closure: compute the selected per-batch loss from ``(logits, target)``."""
y = _infer_indices(target, num_classes)
return _focal_ce(logits, y, alpha, focal_gamma)
elif lt in ("logit_adjust_ce", "la_ce"):
if class_counts is None:
raise ValueError("logit_adjust_ce requires class_counts.")
adjust = logit_adjust.to(torch.float) if logit_adjust is not None else None
def loss_fn(logits, target):
"""Closure: compute the selected per-batch loss from ``(logits, target)``."""
y = _infer_indices(target, num_classes)
z = logits if adjust is None else (logits + adjust.to(logits.device))
return F.cross_entropy(z, y)
elif lt in ("asl", "asymmetric_loss"):
def loss_fn(logits, target):
"""Closure: compute the selected per-batch loss from ``(logits, target)``."""
if target.ndim == 1:
y = F.one_hot(target.long(), num_classes=num_classes).float()
else:
y = target.float().view(-1, num_classes)
return _asl(logits, y, asl_gamma_pos, asl_gamma_neg, asl_clip)
else:
raise ValueError(f"Unknown loss_type '{loss_type}'")
return loss_fn
[docs]
def augment_classes(dst, nc, pc, generate=True, move=True,
group_by='well', test_size=0.1):
"""Augment negative and positive class images and split them into train/test folders.
:param dst: destination root; augmented images land under ``aug_nc``/``aug_pc`` and
move into ``aug/{train,test}/{nc,pc}``.
:param nc: negative-class source image paths.
:param pc: positive-class source image paths.
:param generate: run augmentation before moving files.
:param move: split augmented images into train/test folders.
:param group_by: source identity held intact across train/test. Default
``'well'``; ``'cell'`` permits sibling objects on both sides.
:param test_size: requested test fraction; whole groups make it approximate.
:returns: None.
"""
aug_nc = os.path.join(dst,'aug_nc')
aug_pc = os.path.join(dst,'aug_pc')
all_ = len(nc)+len(pc)
if generate == True:
os.makedirs(aug_nc, exist_ok=True)
if __name__ == '__main__':
augment_images(file_paths=nc, dst=aug_nc)
os.makedirs(aug_pc, exist_ok=True)
if __name__ == '__main__':
augment_images(file_paths=pc, dst=aug_pc)
if move == True:
aug = os.path.join(dst,'aug')
aug_train_nc = os.path.join(aug,'train/nc')
aug_train_pc = os.path.join(aug,'train/pc')
aug_test_nc = os.path.join(aug,'test/nc')
aug_test_pc = os.path.join(aug,'test/pc')
os.makedirs(aug_train_nc, exist_ok=True)
os.makedirs(aug_train_pc, exist_ok=True)
os.makedirs(aug_test_nc, exist_ok=True)
os.makedirs(aug_test_pc, exist_ok=True)
aug_nc_list = [os.path.join(aug_nc, file) for file in os.listdir(aug_nc)]
aug_pc_list = [os.path.join(aug_pc, file) for file in os.listdir(aug_pc)]
from .classifier_evaluation import grouped_split, split_group_values
all_paths = aug_nc_list + aug_pc_list
labels = ([0] * len(aug_nc_list)) + ([1] * len(aug_pc_list))
level, groups = split_group_values(
group_by=group_by, paths=all_paths,
table='augmented crop dataset')
train_idx, test_idx, split_report = grouped_split(
groups, labels, test_size, seed=_run_random_state(42),
group_by=level)
train_set, test_set = set(train_idx.tolist()), set(test_idx.tolist())
nc_train_data = [path for i, path in enumerate(all_paths)
if i in train_set and labels[i] == 0]
nc_test_data = [path for i, path in enumerate(all_paths)
if i in test_set and labels[i] == 0]
pc_train_data = [path for i, path in enumerate(all_paths)
if i in train_set and labels[i] == 1]
pc_test_data = [path for i, path in enumerate(all_paths)
if i in test_set and labels[i] == 1]
print(split_report.summary())
i=0
for path in nc_train_data:
i+=1
shutil.move(path, os.path.join(aug_train_nc, os.path.basename(path)))
print(f'{i}/{all_}', end='\r', flush=True)
for path in nc_test_data:
i+=1
shutil.move(path, os.path.join(aug_test_nc, os.path.basename(path)))
print(f'{i}/{all_}', end='\r', flush=True)
for path in pc_train_data:
i+=1
shutil.move(path, os.path.join(aug_train_pc, os.path.basename(path)))
print(f'{i}/{all_}', end='\r', flush=True)
for path in pc_test_data:
i+=1
shutil.move(path, os.path.join(aug_test_pc, os.path.basename(path)))
print(f'{i}/{all_}', end='\r', flush=True)
print(f'Train nc: {len(os.listdir(aug_train_nc))}, Train pc:{len(os.listdir(aug_train_pc))}, Test nc:{len(os.listdir(aug_test_nc))}, Test pc:{len(os.listdir(aug_test_pc))}')
return
[docs]
def annotate_predictions(csv_loc):
"""Read prediction CSV and add plate/well/field/object columns plus a ``cond`` label.
:param csv_loc: path to a predictions CSV with a ``path`` column of PNG paths.
:returns: DataFrame enriched with parsed metadata and a ``cond`` column
(``'screen'``/``'pc'``/``'nc'`` from the plate/well convention).
"""
df = tabular.read_table(csv_loc, report=None)
df['filename'] = df['path'].apply(lambda x: x.split('/')[-1])
df[['plateID', 'well', 'fieldID', 'object']] = df['filename'].str.split('_', expand=True)
df['object'] = df['object'].str.replace('.png', '')
def assign_condition(row):
"""Return the condition label (``'screen'``/``'pc'``/``'nc'`` or ``''``) for a metadata row."""
plate = int(row['plateID'])
col = int(row['well'][1:])
if col > 3:
if plate in [1, 2, 3, 4]:
return 'screen'
elif plate in [5, 6, 7, 8]:
return 'pc'
elif col in [1, 2, 3]:
return 'nc'
else:
return ''
df['cond'] = pd.Series(
(assign_condition(row) for _, row in df.iterrows()),
index=df.index,
dtype=object,
)
return df
[docs]
def initiate_counter(counter_, lock_):
"""Initialize shared multiprocessing ``counter`` and ``lock`` globals.
:param counter_: shared ``multiprocessing.Value`` counter.
:param lock_: shared ``multiprocessing.Lock`` guarding the counter.
:returns: None.
"""
global counter, lock
counter = counter_
lock = lock_
[docs]
def add_images_to_tar(paths_chunk, tar_path, total_images):
"""Add ``paths_chunk`` images to ``tar_path``, updating the shared counter for progress.
:param paths_chunk: list of image paths to add.
:param tar_path: destination tar archive path.
:param total_images: overall image count used to render progress.
:returns: None.
"""
with tarfile.open(tar_path, 'w') as tar:
for i, img_path in enumerate(paths_chunk):
arcname = os.path.basename(img_path)
try:
tar.add(img_path, arcname=arcname)
with lock:
counter.value += 1
if counter.value % 10 == 0:
print_progress(counter.value, total_images, n_jobs=1, time_ls=None, batch_size=None, operation_type="generating .tar dataset")
except FileNotFoundError:
print(f"File not found: {img_path}")
[docs]
def generate_fraction_map(df, gene_column, min_frequency=0.0):
"""Return a wells-by-genes fraction matrix, dropping columns below ``min_frequency``.
:param df: long-format DataFrame with ``prc``, ``count``, ``well_read_sum`` columns.
:param gene_column: column identifying the gene/guide.
:param min_frequency: drop columns whose maximum fraction is below this cutoff.
:returns: DataFrame indexed by ``prc`` with per-gene fractions.
"""
df['fraction'] = df['count']/df['well_read_sum']
genes = df[gene_column].unique().tolist()
wells = df['prc'].unique().tolist()
print(len(genes),len(wells))
independent_variables = pd.DataFrame(
np.nan, columns=genes, index=wells, dtype=float)
for index, row in df.iterrows():
prc = row['prc']
gene = row[gene_column]
fraction = row['fraction']
independent_variables.loc[prc,gene]=fraction
independent_variables = independent_variables.dropna(axis=1, how='all')
independent_variables = independent_variables.dropna(axis=0, how='all')
independent_variables['sum'] = independent_variables.sum(axis=1)
independent_variables = independent_variables.fillna(0.0)
independent_variables = independent_variables.drop(columns=[col for col in independent_variables.columns if independent_variables[col].max() < min_frequency])
independent_variables = independent_variables.drop('sum', axis=1)
independent_variables.index.name = 'prc'
return independent_variables
[docs]
def fishers_odds(df, threshold=0.5, phenotyp_col='mean_pred'):
"""Fisher's exact test per mutant column against a binarized phenotype label.
:param df: DataFrame with per-mutant presence columns plus ``phenotyp_col``.
:param threshold: cutoff below which ``phenotyp_col`` is called "high phenotype".
:param phenotyp_col: name of the phenotype column.
:returns: DataFrame with columns ``Mutant``, ``OddsRatio``, ``PValue``, ``AdjustedPValue``.
"""
df['high_phenotype'] = df[phenotyp_col] < threshold
results = []
mutants = df.columns[:-2]
mutants = [item for item in mutants if item not in ['count_prc','mean_pathogen_area']]
print(f'fishers df')
display(df)
for mutant in mutants:
contingency_table = pd.crosstab(df[mutant] > 0, df['high_phenotype'])
if contingency_table.shape == (2, 2):
odds_ratio, p_value = fisher_exact(contingency_table)
results.append((mutant, odds_ratio, p_value))
else:
results.append((mutant, float('nan'), float('nan')))
results_df = pd.DataFrame(results, columns=['Mutant', 'OddsRatio', 'PValue'])
filtered_results_df = results_df.dropna(
subset=['OddsRatio', 'PValue']).copy()
pvalues = filtered_results_df['PValue'].values
if len(pvalues) > 0:
adjusted_pvalues = multipletests(pvalues, method='fdr_bh')[1]
filtered_results_df.loc[:, 'AdjustedPValue'] = adjusted_pvalues
else:
print("No p-values to adjust. Check your data filtering steps.")
return filtered_results_df
[docs]
def model_metrics(model):
"""Print RMSE/MAE/Durbin-Watson and show residual/QQ/scale-location diagnostic plots.
:param model: fitted statsmodels regression result.
:returns: None.
"""
rmse = np.sqrt(model.mse_resid)
mae = np.mean(np.abs(model.resid))
durbin_w_value = durbin_watson(model.resid)
print("\nAdditional Metrics:")
print(f"Root Mean Squared Error (RMSE): {rmse}")
print(f"Mean Absolute Error (MAE): {mae}")
print(f"Durbin-Watson: {durbin_w_value}")
with figure_style(theme_target()):
fig, ax = plt.subplots(2, 2, figsize=(15, 12))
from .figures.bundle import _register_figure_data
_register_figure_data(fig, lambda: pd.DataFrame({"fitted": np.asarray(model.fittedvalues), "residual": np.asarray(model.resid)}), x="fitted", y="residual", kind="scatter")
ax[0, 0].scatter(model.fittedvalues, model.resid, edgecolors = 'k', facecolors = 'none')
ax[0, 0].set_title('Residuals vs Fitted')
ax[0, 0].set_xlabel('Fitted values')
ax[0, 0].set_ylabel('Residuals')
sns.histplot(model.resid, kde=True, ax=ax[0, 1])
ax[0, 1].set_title('Histogram of Residuals')
ax[0, 1].set_xlabel('Residuals')
sm.qqplot(model.resid, fit=True, line='45', ax=ax[1, 0])
ax[1, 0].set_title('QQ Plot')
standardized_resid = model.get_influence().resid_studentized_internal
ax[1, 1].scatter(model.fittedvalues, np.sqrt(np.abs(standardized_resid)), edgecolors = 'k', facecolors = 'none')
ax[1, 1].set_title('Scale-Location')
ax[1, 1].set_xlabel('Fitted values')
ax[1, 1].set_ylabel(r'$\sqrt{|Standardized Residuals|}$')
plt.tight_layout()
plt.show()
[docs]
def check_multicollinearity(x):
"""Checks multicollinearity of the predictors by computing the VIF.
:param x: DataFrame of the design matrix -- one row per observation, one
column per predictor, and *no* response column. Every column is fed to
``variance_inflation_factor`` via ``x.values``, so all columns must be
numeric; categorical predictors have to be one-hot encoded first. Add an
explicit constant column if you want the intercept accounted for, since
without one the VIFs are inflated by the shared mean. Perfectly
collinear columns yield ``inf``.
"""
vif_data = pd.DataFrame()
vif_data["Variable"] = x.columns
vif_data["VIF"] = [variance_inflation_factor(x.values, i) for i in range(x.shape[1])]
return vif_data
[docs]
def lasso_reg(merged_df, alpha_value=0.01, reg_type='lasso'):
"""Fit Lasso or Ridge on one-hot-encoded gene/grna/plate/row/column predictors.
:param merged_df: DataFrame with ``gene``, ``grna``, ``plateID``, ``rowID``, ``columnID``, ``pred``.
:param alpha_value: regularization strength.
:param reg_type: ``'lasso'`` or ``'ridge'``.
:returns: DataFrame with ``Feature`` and ``Coefficient`` columns.
"""
X = merged_df[['gene', 'grna', 'plateID', 'rowID', 'columnID']]
y = merged_df['pred']
encoder = OneHotEncoder(drop='first')
X_encoded = encoder.fit_transform(X).toarray()
feature_names = encoder.get_feature_names_out(input_features=X.columns)
reg_type = str(reg_type).strip().lower()
if reg_type == 'ridge':
ridge = Ridge(alpha=alpha_value)
ridge.fit(X_encoded, y)
coeff_dict = dict(zip(feature_names, ridge.coef_))
elif reg_type == 'lasso':
lasso = Lasso(alpha=alpha_value)
lasso.fit(X_encoded, y)
coeff_dict = dict(zip(feature_names, lasso.coef_))
else:
raise ValueError(
f"Unsupported reg_type {reg_type!r}; expected 'lasso' or 'ridge'."
)
coeff_df = pd.DataFrame(list(coeff_dict.items()), columns=['Feature', 'Coefficient'])
return coeff_df
[docs]
def MLR(merged_df, refine_model):
"""Fit a multiple-linear regression on gene:grna interactions plus plate/row/column terms.
:param merged_df: DataFrame with ``gene``, ``grna``, ``plate``, ``row``, ``column``, ``pred`` columns.
:param refine_model: refit after removing outliers by residuals and Cook's distance.
:returns: tuple ``(max_effects, max_effects_pvalues, model, df)``.
"""
from .plot import _reg_v_plot
model = smf.ols("pred ~ gene + grna + gene:grna + plate + row + column", merged_df).fit()
model_metrics(model)
if refine_model:
std_resid = model.get_influence().resid_studentized_internal
outliers_resid = np.where(np.abs(std_resid) > 3)[0]
(c, p) = model.get_influence().cooks_distance
outliers_cooks = np.where(c > 4/(len(merged_df)-merged_df.shape[1]-1))[0]
outliers = reduce(np.union1d, (outliers_resid, outliers_cooks))
merged_df_filtered = merged_df.drop(merged_df.index[outliers])
display(merged_df_filtered)
model = smf.ols("pred ~ gene + grna + gene:grna + row + column", merged_df_filtered).fit()
print("Number of outliers detected by standardized residuals:", len(outliers_resid))
print("Number of outliers detected by Cook's distance:", len(outliers_cooks))
model_metrics(model)
print(model.summary())
interaction_coeffs = {key: val for key, val in model.params.items() if "gene[T." in key and ":grna[T." in key}
interaction_pvalues = {key: val for key, val in model.pvalues.items() if "gene[T." in key and ":grna[T." in key}
max_effects = {}
max_effects_pvalues = {}
for key, val in interaction_coeffs.items():
gene_name = key.split(":")[0].replace("gene[T.", "").replace("]", "")
if gene_name not in max_effects or abs(max_effects[gene_name]) < abs(val):
max_effects[gene_name] = val
max_effects_pvalues[gene_name] = interaction_pvalues[key]
for key in max_effects:
print(f"Key: {key}: {max_effects[key]}, p:{max_effects_pvalues[key]}")
df = pd.DataFrame([max_effects, max_effects_pvalues])
df = df.transpose()
df = df.rename(columns={df.columns[0]: 'effect', df.columns[1]: 'p'})
df = df.sort_values(by=['effect', 'p'], ascending=[False, True])
_reg_v_plot(df)
return max_effects, max_effects_pvalues, model, df
[docs]
def get_files_from_dir(dir_path, file_extension="*"):
"""Return glob matches for ``dir_path/file_extension``.
:param dir_path: directory to list. It is joined to the pattern rather than
walked, so subdirectories are never searched and a nonexistent path
yields an empty list instead of an error.
:param file_extension: despite the name this is a full glob *pattern*, not a
suffix -- pass ``'*.tif'``, not ``'.tif'`` or ``'tif'``, or nothing
matches. Matching follows the filesystem's case sensitivity, and dotfiles
are excluded by ``glob`` semantics. Default ``'*'`` returns every
non-hidden entry, directories included.
"""
return glob.glob(os.path.join(dir_path, file_extension))
[docs]
def create_circular_mask(h, w, center=None, radius=None):
"""Return a boolean circular mask of shape ``(h, w)`` centered on ``center``.
:param h: image height.
:param w: image width.
:param center: ``(x, y)`` center; defaults to the image middle.
:param radius: circle radius; defaults to the largest circle fitting inside.
:returns: boolean ndarray where ``True`` marks pixels within ``radius``.
"""
if center is None:
center = (int(w/2), int(h/2))
if radius is None:
radius = min(center[0], center[1], w-center[0], h-center[1])
Y, X = np.ogrid[:h, :w]
dist_from_center = np.sqrt((X - center[0])**2 + (Y-center[1])**2)
mask = dist_from_center <= radius
return mask
[docs]
def apply_mask(image, output_value=0):
"""Zero out (or set to ``output_value``) pixels outside a circular mask fit to ``image``.
The circle is not a parameter: :func:`create_circular_mask` is called with no
center or radius, so it is always the largest circle inscribed in the frame.
On a non-square image that means the *short* side sets the radius.
:param image: 2-D ``(H, W)`` or 3-D ``(H, W, C)`` array. The mask is built
from the first two axes and broadcast across every channel, so all
channels are cropped identically.
:param output_value: fill written outside the circle. It goes through
``np.where``, so a value the input dtype cannot hold promotes the whole
result -- passing ``np.nan`` to a ``uint16`` image returns a float array,
not a masked integer one. Default ``0``.
"""
h, w = image.shape[:2]
mask = create_circular_mask(h, w)
if len(image.shape) > 2:
mask = np.repeat(mask[:, :, np.newaxis], image.shape[2], axis=2)
masked_image = np.where(mask, image, output_value)
return masked_image
[docs]
def invert_image(image):
"""Return the intensity-inverted image, reflected through the dtype range.
The pivot is ``iinfo.min + iinfo.max``, which is the dtype maximum for
every unsigned dtype -- so ``uint8`` and ``uint16`` invert exactly as they
always did -- and ``-1`` for a signed one, the same convention
:func:`skimage.util.invert` uses. Reflecting through the range instead of
subtracting from the ceiling is what keeps a signed image in range: under
``int8``, ``-100`` inverts to ``99`` rather than to ``227``, which used to
wrap silently to ``-29``.
:param image: array with an *integer* dtype. A float or boolean image
raises ``ValueError`` rather than inverting -- convert or rescale to an
integer dtype first. The pivot is the dtype range, not the image
range, so a dim ``uint16`` image inverts against 65535 and comes back
near-white; normalize to the dtype range first if you want a
contrast-preserving inversion.
:returns: the inverted image, in the input dtype. Every value stays in
range, so nothing wraps.
:raises ValueError: ``image`` does not have an integer dtype.
"""
image = np.asarray(image)
if not np.issubdtype(image.dtype, np.integer):
raise ValueError(
f"invert_image needs an integer dtype to know what to invert "
f"against; got {image.dtype}. Rescale to uint8/uint16 first.")
info = np.iinfo(image.dtype)
pivot = np.asarray(info.min + info.max, dtype=image.dtype)
inverted_image = pivot - image
return inverted_image.astype(image.dtype, copy=False)
[docs]
def resize_images_and_labels(images, labels, target_height, target_width, show_example=True):
"""Resize aligned image/label lists to ``target_height`` x ``target_width``.
:param images: iterable of source images (2-D or 3-D).
:param labels: matching iterable of label masks, or ``None``.
:param target_height: output height in pixels.
:param target_width: output width in pixels.
:param show_example: display an example of the resized pair when ``True``.
:returns: ``(resized_images, resized_labels)`` lists.
"""
from .plot import plot_resize
resized_images = []
resized_labels = []
if not images is None and not labels is None:
for image, label in zip(images, labels):
if image.ndim == 2:
image_shape = (target_height, target_width)
elif image.ndim == 3:
image_shape = (target_height, target_width, image.shape[-1])
resized_image = resizescikit(image, image_shape, preserve_range=True, anti_aliasing=True).astype(image.dtype)
resized_label = resizescikit(label, (target_height, target_width), order=0, preserve_range=True, anti_aliasing=False).astype(label.dtype)
if resized_image.shape[-1] == 1:
resized_image = np.squeeze(resized_image)
resized_images.append(resized_image)
resized_labels.append(resized_label)
elif not images is None:
for image in images:
if image.ndim == 2:
image_shape = (target_height, target_width)
elif image.ndim == 3:
image_shape = (target_height, target_width, image.shape[-1])
resized_image = resizescikit(image, image_shape, preserve_range=True, anti_aliasing=True).astype(image.dtype)
if resized_image.shape[-1] == 1:
resized_image = np.squeeze(resized_image)
resized_images.append(resized_image)
elif not labels is None:
for label in labels:
resized_label = resizescikit(label, (target_height, target_width), order=0, preserve_range=True, anti_aliasing=False).astype(label.dtype)
resized_labels.append(resized_label)
if show_example:
if not images is None and not labels is None:
plot_resize(images, resized_images, labels, resized_labels)
elif not images is None:
plot_resize(images, resized_images, images, resized_images)
elif not labels is None:
plot_resize(labels, resized_labels, labels, resized_labels)
return resized_images, resized_labels
[docs]
def resize_labels_back(labels, orig_dims):
"""Resize a list of label masks back to their original ``(width, height)``.
:param labels: iterable of label masks.
:param orig_dims: matching iterable of ``(width, height)`` tuples.
:returns: list of resized label masks.
:raises ValueError: if lengths differ or ``orig_dims`` entries are malformed.
"""
resized_labels = []
if len(labels) != len(orig_dims):
raise ValueError("The length of labels and orig_dims must match.")
for label, dims in zip(labels, orig_dims):
if not isinstance(dims, tuple) or len(dims) != 2:
raise ValueError("Each element in orig_dims must be a tuple of two integers representing the original dimensions (width, height)")
resized_label = resizescikit(label, dims, order=0, preserve_range=True, anti_aliasing=False).astype(label.dtype)
resized_labels.append(resized_label)
return resized_labels
[docs]
def calculate_iou(mask1, mask2):
"""Return the intersection-over-union of two binary masks after zero-padding to a common shape.
Unlike :func:`jaccard_index`, this returns ``0`` rather than ``nan`` when both
masks are empty, which is what makes it safe to call inside the matching loop.
:param mask1: 2-D array. Any nonzero value counts as foreground, so a
multi-label crop is treated as one merged object -- pass a single
object's mask if you want a per-object IoU.
:param mask2: 2-D array compared against ``mask1``. Shapes may differ;
:func:`pad_to_same_shape` zero-pads both at the bottom and right, which
assumes the two masks share a top-left origin. Two crops taken from
different offsets in the same image will score meaninglessly low.
"""
mask1, mask2 = pad_to_same_shape(mask1, mask2)
intersection = np.logical_and(mask1, mask2).sum()
union = np.logical_or(mask1, mask2).sum()
return intersection / union if union != 0 else 0
[docs]
def match_masks(true_masks, pred_masks, iou_threshold):
"""Greedy match each predicted mask to a still-unmatched true mask above ``iou_threshold``.
:param true_masks: iterable of ground-truth masks.
:param pred_masks: iterable of predicted masks.
:param iou_threshold: minimum IoU to count as a match.
:returns: list of ``(true_mask, pred_mask)`` matched pairs.
"""
matches = []
matched_true_masks_indices = set()
for pred_mask in pred_masks:
for true_mask_index, true_mask in enumerate(true_masks):
if true_mask_index not in matched_true_masks_indices:
iou = calculate_iou(true_mask, pred_mask)
if iou >= iou_threshold:
matches.append((true_mask, pred_mask))
matched_true_masks_indices.add(true_mask_index)
break
return matches
[docs]
def compute_average_precision(matches, num_true_masks, num_pred_masks):
"""Return ``(precision, recall)`` given match count, true count, and predicted count.
Despite the name this computes a single precision/recall point, not an
averaged precision; :func:`compute_ap_over_iou_thresholds` is what integrates
those points into an AP.
:param matches: the pair list from :func:`match_masks`. Only its length is
used, so any sized container works, but it must be the *matched* pairs
rather than all candidate pairs -- matching is greedy and one-to-one, so
the length is the true-positive count.
:param num_true_masks: total ground-truth objects; drives false negatives as
``num_true_masks - len(matches)``. Passing a count smaller than the match
count silently yields a recall above 1, which
:func:`compute_ap_over_iou_thresholds` rejects with ``ValueError``.
:param num_pred_masks: total predicted objects, used the same way for false
positives. Both counts are the totals for the whole field, not per class.
Zero denominators return ``0`` instead of raising.
"""
TP = len(matches)
FP = num_pred_masks - TP
FN = num_true_masks - TP
precision = TP / (TP + FP) if TP + FP > 0 else 0
recall = TP / (TP + FN) if TP + FN > 0 else 0
return precision, recall
[docs]
def pad_to_same_shape(mask1, mask2):
"""Zero-pad ``mask1`` and ``mask2`` to their element-wise maximum shape.
Padding is appended at the bottom and right only, so the two masks are aligned
on their top-left corner. This is an alignment *assumption*, not a
registration: crops taken from different offsets are not brought into
correspondence by padding them.
:param mask1: 2-D array. Only axes 0 and 1 are considered, so a 3-D stack is
padded on its first two axes and left ragged on the third.
:param mask2: 2-D array padded to the same element-wise maximum shape. Each
mask is padded independently, so the larger one along a given axis is
returned untouched on that axis.
"""
shape_diff = np.array([max(mask1.shape[0], mask2.shape[0]) - mask1.shape[0],
max(mask1.shape[1], mask2.shape[1]) - mask1.shape[1]])
pad_mask1 = ((0, shape_diff[0]), (0, shape_diff[1]))
shape_diff = np.array([max(mask1.shape[0], mask2.shape[0]) - mask2.shape[0],
max(mask1.shape[1], mask2.shape[1]) - mask2.shape[1]])
pad_mask2 = ((0, shape_diff[0]), (0, shape_diff[1]))
padded_mask1 = np.pad(mask1, pad_mask1, mode='constant', constant_values=0)
padded_mask2 = np.pad(mask2, pad_mask2, mode='constant', constant_values=0)
return padded_mask1, padded_mask2
[docs]
def compute_ap_over_iou_thresholds(true_masks, pred_masks, iou_thresholds):
"""Return the area under the precision-recall curve swept over ``iou_thresholds``.
:param true_masks: sequence of *per-object* ground-truth masks, one entry per
object -- not a single label image. Its length is the ground-truth count
used for recall, so filtering objects out changes the denominator.
:param pred_masks: sequence of per-object predicted masks, matched greedily
against ``true_masks`` at each threshold. Matching walks predictions in
the order given and claims the first free true mask that clears the
threshold, so the ordering can change which pairs form.
:param iou_thresholds: iterable of IoU cutoffs to sweep. The curve is the
trapezoid over the resulting points sorted by recall, so a single
threshold gives an area of ``0`` -- pass at least two (COCO convention is
``np.linspace(0.5, 0.95, 10)``). Duplicate thresholds contribute
zero-width segments and do not count.
:raises ValueError: if a computed precision or recall falls outside
``[0, 1]``, which indicates the mask counts disagree with the matches.
"""
precision_recall_pairs = []
for iou_threshold in iou_thresholds:
matches = match_masks(true_masks, pred_masks, iou_threshold)
precision, recall = compute_average_precision(matches, len(true_masks), len(pred_masks))
if not 0 <= precision <= 1 or not 0 <= recall <= 1:
raise ValueError(f'Precision or recall out of bounds. Precision: {precision}, Recall: {recall}')
precision_recall_pairs.append((precision, recall))
precision_recall_pairs = sorted(precision_recall_pairs, key=lambda x: x[1])
sorted_precisions = [p[0] for p in precision_recall_pairs]
sorted_recalls = [p[1] for p in precision_recall_pairs]
return _trapezoid(sorted_precisions, x=sorted_recalls)
[docs]
def compute_segmentation_ap(true_masks, pred_masks, iou_thresholds=np.linspace(0.5, 0.95, 10)):
"""Return the COCO-style segmentation AP by matching connected components across IoU thresholds.
This is the whole-image entry point: unlike
:func:`compute_ap_over_iou_thresholds` it takes label images and splits them
into objects itself.
:param true_masks: ground-truth label or binary image for one field. It is
re-run through ``label()``, so existing IDs are discarded and touching
objects that share a border merge into one component -- the AP is
computed on connected components, not on the IDs you supply.
:param pred_masks: predicted mask for the same field, treated identically.
Each object is reduced to its bounding-box crop by ``regionprops``, so
objects are compared shape-to-shape with their positions dropped; two
identically shaped cells in different corners score as a perfect match.
:param iou_thresholds: IoU cutoffs to sweep. Default
``np.linspace(0.5, 0.95, 10)`` is the COCO sweep. This default array is
evaluated once at import and shared by every call, so do not mutate it
in place.
"""
true_mask_labels = label(true_masks)
pred_mask_labels = label(pred_masks)
true_mask_regions = [region.image for region in regionprops(true_mask_labels)]
pred_mask_regions = [region.image for region in regionprops(pred_mask_labels)]
return compute_ap_over_iou_thresholds(true_mask_regions, pred_mask_regions, iou_thresholds)
[docs]
def jaccard_index(mask1, mask2):
"""Return the Jaccard/IoU index of two binary masks.
:param mask1: array of any shape; nonzero is foreground, so a multi-label
mask collapses to one merged object.
:param mask2: array that must already have the *same shape* as ``mask1`` --
there is no padding step here, so mismatched shapes either raise a
broadcast error or, worse, broadcast silently against a length-1 axis.
Use :func:`calculate_iou` when the shapes can differ; it also returns
``0`` for two empty masks, whereas this divides by zero and returns
``nan`` with a runtime warning.
"""
intersection = np.logical_and(mask1, mask2)
union = np.logical_or(mask1, mask2)
return np.sum(intersection) / np.sum(union)
[docs]
def dice_coefficient(mask1, mask2):
"""Return the Dice similarity of two masks, treating any nonzero value as foreground.
:param mask1: array binarized with ``> 0``, so negative values are counted as
*background* -- a signed difference image will not behave as expected.
:param mask2: array of the same shape as ``mask1``; like
:func:`jaccard_index` there is no padding step, so shapes must already
agree. Two empty masks return ``1.0`` here (defined as perfect
agreement) rather than ``nan``.
"""
mask1 = np.where(mask1 > 0, 1, 0)
mask2 = np.where(mask2 > 0, 1, 0)
intersection = np.sum(mask1 & mask2)
total = np.sum(mask1) + np.sum(mask2)
if total == 0:
return 1.0
return 2.0 * intersection / total
[docs]
def boundary_f1_score(mask_true, mask_pred, dilation_radius=1):
"""Return the boundary F1 score between two masks with tolerance ``dilation_radius``.
Both masks are binarized before the boundary is taken, so this scores the
outline of the *foreground as a whole*: boundaries where two labeled objects
abut are interior to that foreground and do not appear. Split/merge errors
between touching cells are therefore invisible to this metric.
:param mask_true: reference label or binary mask, reduced to its boundary by
:func:`extract_boundaries`.
:param mask_pred: predicted mask of the same shape; the two boundary images
are intersected element-wise, so the masks must be pixel-registered.
:param dilation_radius: half-width of the square structuring element, giving a
band ``2 * dilation_radius + 1`` pixels wide. This is the matching
tolerance -- raising it forgives localization error but also thickens both
boundaries, so scores rise for every model and stop being comparable
across different radii. Default ``1``.
"""
boundary_true = extract_boundaries(mask_true, dilation_radius)
boundary_pred = extract_boundaries(mask_pred, dilation_radius)
intersection = np.logical_and(boundary_true, boundary_pred)
precision = np.sum(intersection) / (np.sum(boundary_pred) + 1e-6)
recall = np.sum(intersection) / (np.sum(boundary_true) + 1e-6)
f1 = 2 * (precision * recall) / (precision + recall + 1e-6)
return f1
def _remove_noninfected(stack, cell_dim, nucleus_dim, pathogen_dim):
"""Zero out cells (and their nuclei) that contain no pathogen labels."""
if not cell_dim is None:
cell_mask = stack[:, :, cell_dim]
else:
cell_mask = np.zeros_like(stack)
if not nucleus_dim is None:
nucleus_mask = stack[:, :, nucleus_dim]
else:
nucleus_mask = np.zeros_like(stack)
if not pathogen_dim is None:
pathogen_mask = stack[:, :, pathogen_dim]
else:
pathogen_mask = np.zeros_like(stack)
for cell_label in np.unique(cell_mask)[1:]:
cell_region = cell_mask == cell_label
labels_in_cell = np.unique(pathogen_mask[cell_region])
labels_in_cell = labels_in_cell[labels_in_cell != 0]
if len(labels_in_cell) == 0:
cell_mask[cell_region] = 0
nucleus_mask[cell_region] = 0
if not cell_dim is None:
stack[:, :, cell_dim] = cell_mask
if not nucleus_dim is None:
stack[:, :, nucleus_dim] = nucleus_mask
return stack
def _remove_outside_objects(stack, cell_dim, nucleus_dim, pathogen_dim):
"""Zero out pathogens (and their nuclei) that do not overlap any cell."""
if not cell_dim is None:
cell_mask = stack[:, :, cell_dim]
else:
return stack
if pathogen_dim is None:
return stack
pathogen_mask = stack[:, :, pathogen_dim]
nucleus_mask = None if nucleus_dim is None else stack[:, :, nucleus_dim]
pathogen_labels = np.unique(pathogen_mask)[1:]
for pathogen_label in pathogen_labels:
pathogen_region = pathogen_mask == pathogen_label
cell_in_pathogen_region = np.unique(cell_mask[pathogen_region])
cell_in_pathogen_region = cell_in_pathogen_region[cell_in_pathogen_region != 0]
if len(cell_in_pathogen_region) == 0:
pathogen_mask[pathogen_region] = 0
if nucleus_mask is not None:
nuclei_in_pathogen = np.unique(nucleus_mask[pathogen_region])
nuclei_in_pathogen = nuclei_in_pathogen[
nuclei_in_pathogen != 0]
for nucleus_label in nuclei_in_pathogen:
nucleus_mask[nucleus_mask == nucleus_label] = 0
stack[:, :, cell_dim] = cell_mask
if nucleus_dim is not None:
stack[:, :, nucleus_dim] = nucleus_mask
stack[:, :, pathogen_dim] = pathogen_mask
return stack
def _remove_multiobject_cells(stack, mask_dim, cell_dim, nucleus_dim, pathogen_dim, object_dim):
"""Zero out cells containing more than one object in ``object_dim``."""
if mask_dim is None or object_dim is None:
return stack
cell_mask = stack[:, :, mask_dim]
object_mask = stack[:, :, object_dim]
nucleus_mask = None if nucleus_dim is None else stack[:, :, nucleus_dim]
pathogen_mask = None if pathogen_dim is None else stack[:, :, pathogen_dim]
for cell_label in np.unique(cell_mask)[1:]:
cell_region = cell_mask == cell_label
labels_in_cell = np.unique(object_mask[cell_region])
labels_in_cell = labels_in_cell[labels_in_cell != 0]
if len(labels_in_cell) > 1:
cell_mask[cell_region] = 0
if nucleus_mask is not None:
nucleus_mask[cell_region] = 0
if pathogen_mask is not None:
pathogens_in_cell = np.unique(pathogen_mask[cell_region])
pathogens_in_cell = pathogens_in_cell[pathogens_in_cell != 0]
for pathogen_label in pathogens_in_cell:
pathogen_mask[pathogen_mask == pathogen_label] = 0
if cell_dim is not None:
stack[:, :, cell_dim] = cell_mask
if nucleus_dim is not None:
stack[:, :, nucleus_dim] = nucleus_mask
if pathogen_dim is not None:
stack[:, :, pathogen_dim] = pathogen_mask
return stack
[docs]
def merge_touching_objects(mask, threshold=0.25):
"""Merge touching labeled objects whose shared boundary exceeds ``threshold`` of the smaller perimeter.
:param mask: labeled mask.
:param threshold: fraction of the smaller perimeter required to merge.
:returns: merged label mask.
"""
perimeters = {}
labels = np.unique(mask)
for label in labels:
if label != 0:
edges = morphology.erosion(mask == label) ^ (mask == label)
perimeters[label] = np.sum(edges)
shared_perimeters = {}
dilated = morphology.dilation(mask > 0)
for label in labels:
if label != 0:
dilated_label = morphology.dilation(mask == label)
touching_labels = np.unique(mask[dilated & (dilated_label != 0) & (mask != 0)])
for touching_label in touching_labels:
if touching_label != label:
shared_boundary = dilated_label & morphology.dilation(mask == touching_label)
shared_perimeters[(label, touching_label)] = np.sum(shared_boundary)
for (label1, label2), shared_perimeter in shared_perimeters.items():
if shared_perimeter > threshold * min(perimeters[label1], perimeters[label2]):
mask[mask == label2] = label1
return mask
[docs]
def remove_intensity_objects(image, mask, intensity_threshold, mode):
"""Drop labeled objects whose mean intensity is on the wrong side of ``intensity_threshold``.
:param image: intensity image.
:param mask: labeled mask aligned to ``image``.
:param intensity_threshold: cutoff value.
:param mode: ``'low'`` removes below-threshold objects, ``'high'`` removes above.
:returns: filtered label mask.
"""
props = regionprops_table(mask, image, properties=('label', 'mean_intensity'))
if mode == 'low':
labels_to_remove = props['label'][props['mean_intensity'] < intensity_threshold]
if mode == 'high':
labels_to_remove = props['label'][props['mean_intensity'] > intensity_threshold]
mask[np.isin(mask, labels_to_remove)] = 0
return mask
def _filter_closest_to_stat(df, column, n_rows, use_median=False):
"""Return the ``n_rows`` rows of ``df`` closest to the mean or median of ``column``."""
if use_median:
target_value = df[column].median()
else:
target_value = df[column].mean()
df['diff'] = (df[column] - target_value).abs()
result_df = df.sort_values(by='diff').head(n_rows)
result_df = result_df.drop(columns=['diff'])
return result_df
def _find_similar_sized_images(file_list):
"""Return the largest group of image paths sharing the same cropped size/aspect ratio."""
size_to_paths = defaultdict(list)
for path in file_list:
img = read_image_rgb(path, cv2.IMREAD_UNCHANGED)
if img is not None:
if img.ndim == 3:
mask = np.any(img != 0, axis=2)
else:
mask = img != 0
coords = np.argwhere(mask)
if coords.size == 0:
continue
y0, x0 = coords.min(axis=0)
y1, x1 = coords.max(axis=0) + 1
cropped_img = img[y0:y1, x0:x1]
height, width = cropped_img.shape[:2]
aspect_ratio = width / height
size_key = (width, height, round(aspect_ratio, 2))
size_to_paths[size_key].append(path)
largest_group = max(size_to_paths.values(), key=len)
return largest_group
def _relabel_parent_with_child_labels(parent_mask, child_mask):
"""Relabel parent objects to match their overlapping child labels."""
parent_labels = label(parent_mask, background=0)
child_labels = child_mask
parent_mask_new = np.zeros_like(parent_mask)
unique_child_labels = np.unique(child_labels)[1:]
for child_label in unique_child_labels:
child_area_mask = (child_labels == child_label)
overlapping_parent_label = np.unique(parent_labels[child_area_mask])
for parent_label in overlapping_parent_label:
if parent_label != 0:
parent_mask_new[parent_labels == parent_label] = child_label
for parent_label in np.unique(parent_mask_new)[1:]:
parent_area_mask = (parent_mask_new == parent_label)
child_labels_in_parent = np.unique(child_mask[parent_area_mask])
child_labels_in_parent = child_labels_in_parent[child_labels_in_parent != 0]
if len(child_labels_in_parent) > 1:
first_child_label = child_labels_in_parent[0]
for child_label in child_labels_in_parent:
child_mask[child_mask == child_label] = first_child_label
return parent_mask_new, child_mask
def _exclude_objects(cell_mask, nucleus_mask, pathogen_mask, cytoplasm_mask, uninfected=True):
"""Drop cells missing required companion objects and clear other masks outside kept cells."""
filtered_cells = np.zeros_like(cell_mask)
for cell_label in np.unique(cell_mask):
if cell_label == 0:
continue
cell_region = cell_mask == cell_label
has_nucleus = np.any(nucleus_mask[cell_region])
has_cytoplasm = np.any(cytoplasm_mask[cell_region])
has_pathogen = np.any(pathogen_mask[cell_region])
if uninfected:
if has_nucleus and has_cytoplasm:
filtered_cells[cell_region] = cell_label
else:
if has_nucleus and has_cytoplasm and has_pathogen:
filtered_cells[cell_region] = cell_label
nucleus_mask = nucleus_mask * (filtered_cells > 0)
pathogen_mask = pathogen_mask * (filtered_cells > 0)
cytoplasm_mask = cytoplasm_mask * (filtered_cells > 0)
return filtered_cells, nucleus_mask, pathogen_mask, cytoplasm_mask
def _merge_overlapping_objects(mask1, mask2):
"""Merge overlapping objects across two masks using a 90% overlap heuristic."""
labeled_1 = label(mask1)
num_1 = np.max(labeled_1)
for m1_id in range(1, num_1 + 1):
current_1_mask = labeled_1 == m1_id
overlapping_2_labels = np.unique(mask2[current_1_mask])
overlapping_2_labels = overlapping_2_labels[overlapping_2_labels != 0]
if len(overlapping_2_labels) > 1:
overlap_percentages = [np.sum(current_1_mask & (mask2 == m2_label)) / np.sum(current_1_mask) * 100 for m2_label in overlapping_2_labels]
max_overlap_label = overlapping_2_labels[np.argmax(overlap_percentages)]
max_overlap_percentage = max(overlap_percentages)
if max_overlap_percentage >= 90:
for m2_label in overlapping_2_labels:
if m2_label != max_overlap_label:
mask1[(current_1_mask) & (mask2 == m2_label)] = 0
else:
for m2_label in overlapping_2_labels[1:]:
mask2[mask2 == m2_label] = overlapping_2_labels[0]
return mask1, mask2
def _filter_object(mask, min_value, max_value=None):
"""Zero out label values outside the allowed pixel-count range.
:param min_value: drop objects smaller than this. 0 or None disables it.
:param max_value: drop objects LARGER than this. None -- the default,
and what every existing run does -- disables it.
:returns: ``mask``, filtered in place.
THE UPPER BOUND EXISTS BECAUSE THE LOWER ONE IS NOT ENOUGH. A
segmentation blow-up -- one "cell" covering a quarter of the field, two
cells merged by a bright bridge -- passes every minimum there is, gets
measured, and carries its area into the classifier and the regression.
A minimum can only remove debris.
:returns: the number of objects removed is NOT returned; the caller
counts them, because it is the caller that knows which object type
this is and can say so.
"""
count = np.bincount(mask.ravel())
too_small = count < (min_value or 0)
if max_value:
too_big = count > max_value
else:
too_big = np.zeros_like(too_small)
remove = np.where(too_small | too_big)[0]
remove = remove[remove != 0]
mask[np.isin(mask, remove)] = 0
return mask
def _filter_cp_masks(masks, flows, filter_size, filter_intensity, minimum_size, maximum_size, remove_border_objects, merge, batch, plot, figuresize):
"""Post-process Cellpose masks: optional merge, size filter, intensity filter, border removal."""
from .plot import plot_masks
mask_stack = []
for idx, (mask, flow, image) in enumerate(zip(masks, flows[0], batch)):
if plot and idx == 0:
num_objects = mask_object_count(mask)
print(f'Number of objects before filtration: {num_objects}')
plot_masks(batch=image, masks=mask, flows=flow, cmap='inferno', figuresize=figuresize, nr=1, file_type='.npz', print_object_number=True)
if merge:
mask = merge_touching_objects(mask, threshold=0.66)
if plot and idx == 0:
num_objects = mask_object_count(mask)
print(f'Number of objects after merging adjacent objects, : {num_objects}')
plot_masks(batch=image, masks=mask, flows=flow, cmap='inferno', figuresize=figuresize, nr=1, file_type='.npz', print_object_number=True)
if filter_size:
props = measure.regionprops_table(mask, properties=['label', 'area'])
valid_labels = props['label'][np.logical_and(props['area'] > minimum_size, props['area'] < maximum_size)]
mask = np.isin(mask, valid_labels) * mask
if plot and idx == 0:
num_objects = mask_object_count(mask)
print(f'Number of objects after size filtration >{minimum_size} and <{maximum_size} : {num_objects}')
plot_masks(batch=image, masks=mask, flows=flow, cmap='inferno', figuresize=figuresize, nr=1, file_type='.npz', print_object_number=True)
if filter_intensity:
intensity_image = image[:, :, 1]
props = measure.regionprops_table(mask, intensity_image=intensity_image, properties=['label', 'mean_intensity'])
mean_intensities = np.array(props['mean_intensity']).reshape(-1, 1)
if mean_intensities.shape[0] >= 2:
kmeans = KMeans(n_clusters=2, random_state=0).fit(mean_intensities)
centroids = kmeans.cluster_centers_
dist_between_centroids = distance.euclidean(centroids[0], centroids[1])
distance_threshold = 0.25
if dist_between_centroids > distance_threshold:
high_intensity_cluster = np.argmax(centroids)
valid_labels = np.array(props['label'])[kmeans.labels_ == high_intensity_cluster]
mask = np.isin(mask, valid_labels) * mask
if plot and idx == 0:
num_objects = mask_object_count(mask)
props_after = measure.regionprops_table(mask, intensity_image=intensity_image, properties=['label', 'mean_intensity'])
mean_intensities_after = np.mean(np.array(props_after['mean_intensity']))
average_intensity_before = np.mean(mean_intensities)
print(f'Number of objects after potential intensity clustering: {num_objects}. Mean intensity before:{average_intensity_before:.4f}. After:{mean_intensities_after:.4f}.')
plot_masks(batch=image, masks=mask, flows=flow, cmap='inferno', figuresize=figuresize, nr=1, file_type='.npz', print_object_number=True)
if remove_border_objects:
mask = clear_border(mask)
if plot and idx == 0:
num_objects = mask_object_count(mask)
print(f'Number of objects after removing border objects, : {num_objects}')
plot_masks(batch=image, masks=mask, flows=flow, cmap='inferno', figuresize=figuresize, nr=1, file_type='.npz', print_object_number=True)
mask_stack.append(mask)
return mask_stack
def _object_filter(df, object_type, size_range, intensity_range, mask_chans, mask_chan):
"""
Filter the DataFrame based on object type, size range, and intensity range.
Args:
df (pandas.DataFrame): The DataFrame to filter.
object_type (str): The type of object to filter.
size_range (list or None): The range of object sizes to filter.
intensity_range (list or None): The range of object intensities to filter.
mask_chans (list): The list of mask channels.
mask_chan (int): The index of the mask channel to use.
Returns:
pandas.DataFrame: The filtered DataFrame.
"""
if not size_range is None:
if isinstance(size_range, list):
if isinstance(size_range[0], int):
df = df[df[f'{object_type}_area'] > size_range[0]]
print(f'After {object_type} minimum area filter: {len(df)}')
if isinstance(size_range[1], int):
df = df[df[f'{object_type}_area'] < size_range[1]]
print(f'After {object_type} maximum area filter: {len(df)}')
if not intensity_range is None:
if isinstance(intensity_range, list):
if isinstance(intensity_range[0], int):
df = df[df[f'{object_type}_channel_{mask_chans[mask_chan]}_mean_intensity'] > intensity_range[0]]
print(f'After {object_type} minimum mean intensity filter: {len(df)}')
if isinstance(intensity_range[1], int):
df = df[df[f'{object_type}_channel_{mask_chans[mask_chan]}_mean_intensity'] < intensity_range[1]]
print(f'After {object_type} maximum mean intensity filter: {len(df)}')
return df
def _get_regex(metadata_type, img_format, custom_regex=None):
"""Return the filename pattern for a microscope convention.
THE VOCABULARY IS A TABLE, not an if/elif chain: every convention is one
record in ``spacr.regex_infer._METADATA_CONVENTIONS``, carrying its
vendor, its instrument family, real example filenames, what each named
group means, where it was sourced, and whether it is confirmed or
provisional. Adding a microscope is adding a record; nothing here
changes.
THE IMPORT IS ABSOLUTE ON PURPOSE AND MUST STAY THAT WAY.
:func:`spacr.qt.widgets.preview_controls._get_regex_callable` lifts THIS
FUNCTION ALONE out of the source file with ``ast`` and executes it in an
empty namespace, so that a dropdown can learn a filename pattern without
paying the 3.2 s and ~900 MB that importing ``spacr.utils`` costs. A
relative ``from .regex_infer import ...`` has no package to resolve
against there, raises ImportError, and is swallowed by that caller's
``except Exception`` -- so the previews would quietly stop grouping
files and nothing would say so. ``spacr.regex_infer`` imports nothing
outside the standard library, which is what makes this affordable.
:param metadata_type: the convention. The four spaCR has always had are
``'cellvoyager'``, ``'cq1'``, ``'auto'`` and ``'custom'``; the rest
are in the table. Matched EXACTLY -- ``'CellVoyager'`` is a typo and
is refused, because silently correcting it would also silently
correct a name that meant something else.
:param img_format: the file extension the pattern should end on;
``None`` means ``tif``.
:param custom_regex: the pattern, for ``'custom'``.
:returns: the pattern.
:raises ValueError: NAMING THE VOCABULARY, for an unrecognised type.
Falling through left the variable unbound and raised "cannot access
local variable 'regex'" -- an error about an implementation detail
rather than about the setting that was wrong.
"""
from spacr.regex_infer import _METADATA_CONVENTIONS, _metadata_pattern
print(f"Image_format: {img_format}")
if img_format == None:
img_format = 'tif'
try:
regex = _metadata_pattern(metadata_type, img_format, custom_regex)
except KeyError:
known = ", ".join(repr(record["key"])
for record in _METADATA_CONVENTIONS)
raise ValueError(
f"metadata_type={metadata_type!r} is not one of {known}. "
f"Choose one of those, or use 'custom' with a regular "
f"expression of your own.")
print(f'regex mode:{metadata_type} regex:{regex}')
return regex
def _run_test_mode(src, regex, timelapse=False, test_images=10, random_test=True):
"""Copy a small sample of the source into a test folder.
A timelapse is cut to ONE image set rather than the requested number:
the point of a test run there is a complete sequence, and ten partial
sequences test nothing.
Raw images are sampled from ``orig/`` AND from the plate folder itself,
which is where the full pipeline reads them from. A plate that a killed
run left half moved into ``orig/``, or one given new images after an
earlier run, holds raw images in both, and sampling only ``orig/`` tested
a plate the real run would not see. A name present in both is taken once,
from ``orig/``.
:param src: the folder to sample from.
:param regex: the filename pattern.
:param timelapse: treat the source as a timelapse.
:param test_images: how many image sets to take.
:param random_test: sample at random rather than taking the first.
:returns: the test folder.
"""
from .io import _listdir_visible
if timelapse:
test_images = 1
test_folder_path = os.path.join(src, 'test')
os.makedirs(test_folder_path, exist_ok=True)
regular_expression = re.compile(regex)
folders = [src]
if os.path.isdir(os.path.join(src, 'orig')):
folders = [os.path.join(src, 'orig'), src]
found_in = {}
for folder in folders:
listed = [filename for filename in _listdir_visible(folder) if regular_expression.match(filename)]
for filename in listed:
if (filename not in found_in
and os.path.isfile(os.path.join(folder, filename))):
found_in[filename] = folder
all_filenames = list(found_in)
print(f'Found {len(all_filenames)} files')
images_by_set = defaultdict(list)
fallback_plate = os.path.basename(folders[0])
for filename in all_filenames:
match = regular_expression.match(filename)
plate = match.group('plateID') if 'plateID' in match.groupdict() else fallback_plate
well = match.group('wellID')
field = match.group('fieldID')
set_identifier = (plate, well, field)
images_by_set[set_identifier].append(filename)
set_identifiers = list(images_by_set.keys())
if random_test:
random.seed(42)
random.shuffle(set_identifiers)
selected_sets = set_identifiers[:test_images]
print(f'Using {len(selected_sets)} random image set(s) for test model')
for set_identifier in selected_sets:
for filename in images_by_set[set_identifier]:
shutil.copy(os.path.join(found_in[filename], filename),
test_folder_path)
return test_folder_path
#: The only stock Cellpose weights that exist from Cellpose 4 (SAM) onward.
CPSAM_MODEL = 'cpsam'
#: Pre-SAM Cellpose model names that older settings files may still carry.
#: Cellpose 4 removed every one of them — ``models.MODEL_NAMES == ['cpsam']``
#: — and silently resolves an unknown name to cpsam, so honouring them would
#: only mislead. They are ACCEPTED-BUT-MAPPED aliases: a settings CSV written
#: against Cellpose 3 still loads and still runs, it just runs the model that
#: actually exists and says so. They are deliberately NOT offered anywhere in
#: the UI — see ``spacr.settings.normalize_cellpose_model_name``.
LEGACY_CELLPOSE_MODELS = ('cyto', 'cyto2', 'cyto3', 'cyto_2', 'cyto_3',
'nuclei', 'nucleus', 'toxo_pv_lumen', 'toxo_cyto')
#: Notices already printed by :func:`_resolve_cellpose_pretrained` this run.
#: A plate is segmented field by field but the model choice is made from the
#: same settings every time, so the substitution notice is worth exactly one
#: line per (message, object type) — not one per field, which on a 1000-field
#: plate buried the run log under thousands of identical warnings.
_REPORTED_CELLPOSE_NOTICES = set()
def _installed_cellpose_models():
"""The stock model names the INSTALLED Cellpose advertises.
``()`` when Cellpose is not importable or says nothing, which makes
every caller fall through to the behaviour that shipped. Read here
rather than hard-coded because the list grew between 4.0 and 4.2 and
will grow again.
Only ``MODEL_NAMES`` -- the stock weights. A user-registered checkpoint
is already handled by the file branch of
:func:`_resolve_cellpose_pretrained`, and is a path rather than a name.
"""
try:
return tuple(getattr(cp_models, "MODEL_NAMES", ()) or ())
except Exception:
return ()
[docs]
def reset_cellpose_model_reports():
"""Forget which Cellpose model notices have already been printed.
Call this at the start of a run so a second run in the same process (a
GUI session segmenting a second plate) reports its model choice again
instead of inheriting the first run's silence.
"""
_REPORTED_CELLPOSE_NOTICES.clear()
def _report_cellpose_once(key, message):
"""Print ``message`` the first time ``key`` is seen this run.
:param key: hashable identity of the notice; repeats are dropped.
:param message: text to print.
:returns: True if it was printed, False if it was suppressed as a repeat.
"""
if key in _REPORTED_CELLPOSE_NOTICES:
return False
_REPORTED_CELLPOSE_NOTICES.add(key)
print(message)
return True
def _for_object(object_type):
"""Return ``' for <object_type>'``, or ``''`` when the caller did not say.
``_choose_model`` used to default ``object_type='cell'``, so a call that
never named an object type still announced one: asking for the nucleus
model printed "using 'cpsam' for cell". An unnamed object type is now
left unnamed rather than guessed.
"""
return f" for {object_type}" if object_type else ""
def _resolve_cellpose_pretrained(model_name, object_type=None, restore_type=None):
"""Return the ``pretrained_model`` string Cellpose 4 should actually load.
Cellpose 4 ships exactly one model, ``cpsam``. ``model_type=`` and
``diam_mean=`` are accepted-and-ignored by ``CellposeModel`` (it logs
"not used in v4.0.1+"), and an unrecognised ``pretrained_model`` resolves
to cpsam with only a log warning — so the pre-SAM names were never
actually loading the model they named. ``diameter``, by contrast, is
still honoured: ``CellposeModel.eval`` rescales the image by
``30. / diameter``. Every legacy name is therefore mapped to cpsam
explicitly, and said out loud once, rather than pretending.
A ``model_name`` that names an existing FILE is treated as a fine-tuned
checkpoint and returned as-is. ``pretrained_model`` used to be hard-coded
to 'cpsam', so every model produced by spaCR's own Train Cellpose module
was silently discarded and the stock weights used instead — the trained
model could never actually be applied to anything.
:param model_name: 'cpsam', a legacy pre-SAM name (mapped to cpsam), or a
path to a fine-tuned checkpoint.
:param object_type: 'cell' / 'nucleus' / 'pathogen' / 'organelle', or None
when the caller genuinely has no object type to name.
:param restore_type: unsupported under Cellpose 4; reported and ignored.
:returns: the string to pass as ``pretrained_model``.
:raises FileNotFoundError: if ``model_name`` looks like a path but no file
is there. Falling back to cpsam would silently segment with the wrong
weights, which is worse than stopping.
"""
clause = _for_object(object_type)
if restore_type is not None:
_report_cellpose_once(
('restore', restore_type, object_type),
f"restore_type={restore_type!r} is not supported on Cellpose 4 "
f"(the denoise/deblur/upsample checkpoints are pre-SAM). Ignoring it.")
name = str(model_name).strip() if model_name else ''
if name and name in _installed_cellpose_models():
if name != CPSAM_MODEL:
_report_cellpose_once(
('stock', name, object_type),
f"Using Cellpose model {name!r}{clause}.")
return name
if name and name not in LEGACY_CELLPOSE_MODELS and name != CPSAM_MODEL:
if os.path.isfile(name):
_report_cellpose_once(
('checkpoint', name, object_type),
f"Loading fine-tuned Cellpose checkpoint{clause}: {name}")
return name
from .model_zoo import _ensure_model_file
downloaded = _ensure_model_file(name, kinds=("cellpose",))
if downloaded is not None:
_report_cellpose_once(
('checkpoint', name, object_type),
f"Loading model-zoo Cellpose checkpoint{clause}: {downloaded}")
return str(downloaded)
if os.sep in name or name.endswith(('.pth', '.pt')):
raise FileNotFoundError(
f"Cellpose model {name!r}{clause} looks like a "
f"checkpoint path but no file is there. Cellpose would quietly "
f"fall back to the stock cpsam weights, so this stops instead. "
f"Check the path, or use 'cpsam' for the stock model.")
_report_cellpose_once(
('unknown', name, object_type),
f"Unknown Cellpose model {name!r}; using 'cpsam'{clause}.")
elif name in LEGACY_CELLPOSE_MODELS:
_report_cellpose_once(
('legacy', name, object_type),
f"Cellpose model {name!r} predates Cellpose-SAM and is no longer "
f"available; using 'cpsam'{clause}.")
return CPSAM_MODEL
def _choose_model(model_name, device, object_type=None, restore_type=None, object_settings=None):
"""Return the Cellpose model to segment ``object_type`` with.
Thin wrapper over :func:`_resolve_cellpose_pretrained` — see there for
what Cellpose 4 does and does not still honour.
:param model_name: 'cpsam', a legacy pre-SAM name (mapped to cpsam), or a
path to a fine-tuned checkpoint.
:param device: torch device passed through to Cellpose.
:param object_type: 'cell' / 'nucleus' / 'pathogen' / 'organelle'. Left
unset it is reported as unset rather than guessed as 'cell'.
:param restore_type: unsupported under Cellpose 4; reported and ignored.
:param object_settings: unused, kept for call-site compatibility.
:returns: a ``CellposeModel``.
:raises FileNotFoundError: if ``model_name`` looks like a path but no file
is there.
"""
if object_settings is None:
object_settings = {}
pretrained = _resolve_cellpose_pretrained(
model_name, object_type=object_type, restore_type=restore_type)
from .accelerator import cellpose_kwargs
kwargs = cellpose_kwargs()
if device is not None:
resolved_cpu = str(kwargs.get("device")).split(":", 1)[0] == "cpu"
kwargs["device"] = device
if str(device).split(":", 1)[0] == "cpu":
kwargs.update(gpu=False, use_bfloat16=False)
elif resolved_cpu:
kwargs["gpu"] = True
kwargs.pop("use_bfloat16", None)
return cp_models.CellposeModel(pretrained_model=pretrained, **kwargs)
[docs]
class SelectChannels:
"""Callable transform that zeroes out image channels not present in ``channels``.
:param channels: iterable of 1-based channel indices to keep (1=red, 2=green, 3=blue).
"""
def __init__(self, channels):
"""Store the list of channels to preserve."""
self.channels = channels
[docs]
def __call__(self, img):
"""Return ``img`` with unselected RGB channels zeroed."""
img = img.clone()
if 1 not in self.channels:
img[0, :, :] = 0
if 2 not in self.channels:
img[1, :, :] = 0
if 3 not in self.channels:
img[2, :, :] = 0
return img
def _activation_map_to_2d(activation_map):
"""Return an activation map shaped for ``imshow``.
The two ``plot_activation_grid`` implementations disagreed about this:
the saliency one transposed a leading 3 to channels-last while the
Grad-CAM one did nothing, so ``(3, H, W)`` rendered in one and raised
``TypeError`` in the other, and ``(1, H, W)`` raised in both. Same name,
same signature, incompatible inputs.
Accepts ``(H, W)``, ``(1, H, W)`` and ``(3, H, W)``; anything else is
handed back untouched so ``imshow`` still raises rather than this
silently reshaping a map it does not understand.
"""
if getattr(activation_map, "ndim", 0) == 3:
if activation_map.shape[0] == 1:
return activation_map[0]
if activation_map.shape[0] == 3:
return np.transpose(activation_map, (1, 2, 0))
return activation_map
[docs]
class SaliencyMapGenerator:
"""Generate saliency maps and predictions for a binary classifier.
:param model: trained PyTorch model with a single-logit binary output.
"""
def __init__(self, model):
"""Store the model to be probed."""
self.model = model
[docs]
def compute_saliency_maps(self, X, y):
"""Return absolute-gradient saliency maps for inputs ``X``.
:param X: differentiable input image batch to probe.
:param y: binary labels selecting the signed output scores.
"""
self.model.eval()
X.requires_grad_()
scores = self.model(X).squeeze()
target_scores = scores * (2 * y - 1)
self.model.zero_grad()
target_scores.backward(torch.ones_like(target_scores))
saliency = X.grad.abs()
return saliency
[docs]
def compute_saliency_and_predictions(self, X):
"""Return saliency maps and the model's own predicted classes.
:param X: differentiable input image batch to classify and probe.
"""
self.model.eval()
X.requires_grad_()
raw = self.model(X)
if raw.ndim > 1 and raw.shape[-1] > 1:
predictions = raw.argmax(dim=-1).long()
scores = raw.gather(-1, predictions.unsqueeze(-1)).squeeze(-1)
target_scores = scores
else:
scores = raw.squeeze()
predictions = (scores > 0).long()
target_scores = scores * (2 * predictions - 1)
self.model.zero_grad()
target_scores.backward(torch.ones_like(target_scores))
saliency = X.grad.abs()
return saliency, predictions
[docs]
def plot_activation_grid(self, X, saliency, predictions, overlay=True, normalize=False):
"""Render a grid overlaying saliency maps on inputs with predicted-class labels.
The grid is always eight columns wide with ``ceil(N / 8)`` rows, and
unused panels in an incomplete last row are hidden. The figure is
returned, never shown.
:param X: batch tensor shaped ``(N, C, H, W)``; ``N`` fixes the grid
size, and an empty batch raises ``ValueError`` from ``subplots``.
The pixels are read only under ``overlay``, where the sample is
permuted to ``(H, W, C)`` -- ``C`` of 1, 3 or 4 renders, ``C`` of 2
reaches ``imshow`` as an invalid shape and raises ``TypeError``.
:param saliency: torch tensor of at least ``N`` entries. It is indexed
and moved to the CPU on every iteration even when ``overlay`` is
false, so a numpy array raises ``AttributeError`` either way. Each
entry may be ``(H, W)``, ``(1, H, W)`` or ``(3, H, W)``; a leading
singleton is removed and a leading RGB dimension is transposed to
channels-last. Other three-dimensional shapes reach ``imshow`` and
raise ``TypeError``.
:param predictions: sequence supporting ``predictions[i].item()``,
whose scalar is stamped in each panel's corner. A plain Python list
of ints raises ``AttributeError``.
:param overlay: true draws the input beneath a translucent map; false
draws the map alone. Default ``True``.
:param normalize: percentile-stretch the input image only; the saliency
map is always drawn raw. Has no effect unless ``overlay`` is true,
and a channel that is flat between its 2nd and 98th percentiles
divides by zero and comes out ``NaN``. Default ``False``.
:returns: the Matplotlib ``Figure``.
"""
N = X.shape[0]
rows = (N + 7) // 8
with figure_style(theme_target()):
fig, axs = plt.subplots(rows, 8, figsize=(16, rows * 2), squeeze=False)
from .figures.bundle import _register_figure_data
_register_figure_data(fig, None, kind="image", title="Saliency")
for ax in axs.flat:
ax.axis('off')
for i in range(N):
ax = axs[i // 8, i % 8]
saliency_map = _activation_map_to_2d(saliency[i].cpu().numpy())
if overlay:
img_np = X[i].permute(1, 2, 0).detach().cpu().numpy()
if normalize:
img_np = self.percentile_normalize(img_np)
ax.imshow(img_np)
ax.imshow(saliency_map, cmap='jet', alpha=0.5)
else:
ax.imshow(saliency_map, cmap='jet')
ax.text(5, 25, str(predictions[i].item()), fontsize=12, color='white', weight='bold',
bbox=dict(facecolor='black', alpha=0.7, boxstyle='round,pad=0.2'))
ax.axis('off')
plt.tight_layout(pad=0)
return fig
[docs]
def percentile_normalize(self, img, lower_percentile=2, upper_percentile=98):
"""Per-channel percentile-normalize ``img`` into ``[0, 1]``.
:param img: channels-last image array to normalize.
"""
img_normalized = np.zeros_like(img)
for c in range(img.shape[2]):
low = np.percentile(img[:, :, c], lower_percentile)
high = np.percentile(img[:, :, c], upper_percentile)
img_normalized[:, :, c] = np.clip((img[:, :, c] - low) / (high - low), 0, 1)
return img_normalized
[docs]
class GradCAMGenerator:
"""Grad-CAM (and variants) map generator for binary classifiers.
:param model: trained model to inspect.
:param target_layer: dotted attribute path to the convolutional layer to probe.
:param cam_type: variant identifier, e.g. ``'gradcam'``.
"""
def __init__(self, model, target_layer, cam_type='gradcam'):
"""Store the model, resolve the target layer, and register activation/gradient hooks."""
self.model = model
self.model.eval()
self.target_layer = target_layer
self.cam_type = cam_type
self.gradients = None
self.activations = None
self.target_layer_module = self.get_layer(self.model, self.target_layer)
self.hook_layers()
[docs]
def hook_layers(self):
"""Register forward/backward hooks that capture activations and gradients."""
def forward_hook(module, input, output):
"""Forward hook: cache the target layer's output activations."""
self.activations = output
def backward_hook(module, grad_input, grad_output):
"""Backward hook: cache the gradient flowing into the target layer's output."""
self.gradients = grad_output[0]
self.target_layer_module.register_forward_hook(forward_hook)
self.target_layer_module.register_full_backward_hook(backward_hook)
[docs]
def get_layer(self, model, target_layer):
"""Resolve a dotted attribute path into the referenced submodule.
:param model: root model from which attribute traversal starts.
:param target_layer: dot-separated submodule attribute path.
"""
modules = target_layer.split('.')
layer = model
for module in modules:
layer = getattr(layer, module)
return layer
[docs]
def compute_gradcam_maps(self, X, y):
"""Return a normalized Grad-CAM map for one sample.
:param X: single-sample differentiable input batch to probe.
:param y: binary label selecting the signed output score.
"""
X.requires_grad_()
scores = self.model(X).squeeze()
target_scores = scores * (2 * y - 1)
self.model.zero_grad()
target_scores.backward(torch.ones_like(target_scores))
pooled_gradients = torch.mean(self.gradients, dim=[0, 2, 3])
for i in range(self.activations.size(1)):
self.activations[:, i, :, :] *= pooled_gradients[i]
gradcam = torch.mean(self.activations, dim=1, keepdim=True)
gradcam = F.relu(gradcam)
gradcam = F.interpolate(gradcam, size=X.shape[2:], mode='bilinear')
gradcam = gradcam.squeeze().cpu().detach().numpy()
gradcam -= gradcam.min()
peak = gradcam.max()
if peak > 0:
gradcam /= peak
else:
gradcam.fill(0.0)
return gradcam
[docs]
def compute_gradcam_and_predictions(self, X):
"""Return Grad-CAM maps and predictions for every sample.
:param X: differentiable input image batch to classify and probe.
"""
self.model.eval()
X.requires_grad_()
raw = self.model(X)
if raw.ndim > 1 and raw.shape[-1] > 1:
predictions = raw.argmax(dim=-1).long()
else:
predictions = (raw.squeeze() > 0).long()
gradcam_maps = []
for i in range(X.size(0)):
gradcam_map = self.compute_gradcam_maps(X[i].unsqueeze(0), predictions[i])
gradcam_maps.append(gradcam_map)
return torch.from_numpy(np.stack(gradcam_maps)), predictions
[docs]
def plot_activation_grid(self, X, gradcam, predictions, overlay=True, normalize=False):
"""Render a grid overlaying Grad-CAM maps on inputs with predicted-class labels.
The grid is always eight columns wide with ``ceil(N / 8)`` rows, and
unused panels in an incomplete last row are hidden. The figure is
returned, never shown.
:param X: batch tensor shaped ``(N, C, H, W)``; ``N`` fixes the grid
size, and an empty batch raises ``ValueError`` from ``subplots``.
The pixels are read only under ``overlay``, where the sample is
permuted to ``(H, W, C)`` -- ``C`` of 1, 3 or 4 renders, ``C`` of 2
raises ``TypeError`` from ``imshow``.
:param gradcam: torch tensor with at least ``N`` entries. Each map may
be ``(H, W)``, ``(1, H, W)`` or ``(3, H, W)``, with the same shape
normalization used by the saliency twin. It is indexed and moved
to the CPU on every iteration, so it must be a tensor either way.
:param predictions: sequence supporting ``predictions[i].item()``,
whose scalar is stamped in each panel's corner. A plain Python list
of ints raises ``AttributeError``.
:param overlay: true draws the input beneath a translucent map; false
draws the map alone. Default ``True``.
:param normalize: percentile-stretch the input image only; the Grad-CAM
map is always drawn raw. Has no effect unless ``overlay`` is true,
and a channel that is flat between its 2nd and 98th percentiles
divides by zero and comes out ``NaN``. Default ``False``.
:returns: the Matplotlib ``Figure``.
"""
N = X.shape[0]
rows = (N + 7) // 8
with figure_style(theme_target()):
fig, axs = plt.subplots(rows, 8, figsize=(16, rows * 2), squeeze=False)
from .figures.bundle import _register_figure_data
_register_figure_data(fig, None, kind="image", title="Grad-CAM")
for ax in axs.flat:
ax.axis('off')
for i in range(N):
ax = axs[i // 8, i % 8]
gradcam_map = _activation_map_to_2d(gradcam[i].cpu().numpy())
if overlay:
img_np = X[i].permute(1, 2, 0).detach().cpu().numpy()
if normalize:
img_np = self.percentile_normalize(img_np)
ax.imshow(img_np)
ax.imshow(gradcam_map, cmap='jet', alpha=0.5)
else:
ax.imshow(gradcam_map, cmap='jet')
ax.text(5, 25, str(predictions[i].item()), fontsize=12, color='white', weight='bold',
bbox=dict(facecolor='black', alpha=0.7, boxstyle='round,pad=0.2'))
ax.axis('off')
plt.tight_layout(pad=0)
return fig
[docs]
def percentile_normalize(self, img, lower_percentile=2, upper_percentile=98):
"""Per-channel percentile-normalize ``img`` into ``[0, 1]``.
:param img: channels-last image array to normalize.
"""
img_normalized = np.zeros_like(img)
for c in range(img.shape[2]):
low = np.percentile(img[:, :, c], lower_percentile)
high = np.percentile(img[:, :, c], upper_percentile)
img_normalized[:, :, c] = np.clip((img[:, :, c] - low) / (high - low), 0, 1)
return img_normalized
[docs]
def preprocess_image(image_path, normalize=True, image_size=224, channels=None):
"""Load and preprocess ``image_path`` into a batched tensor ready for classification.
:param image_path: path to the source image.
:param normalize: apply ImageNet mean/std normalization.
:param image_size: square resize dimension.
:param channels: reserved for downstream use; kept for API compatibility.
:returns: ``(pil_image, input_tensor)`` where the tensor has shape ``(1, 3, H, W)``.
"""
if channels is None:
channels = [1,2,3]
preprocess = transforms.Compose([
transforms.Resize((image_size, image_size)),
transforms.ToTensor(),
])
image = Image.open(image_path).convert('RGB')
input_tensor = preprocess(image)
if normalize:
input_tensor = transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])(input_tensor)
input_tensor = input_tensor.unsqueeze(0)
return image, input_tensor
[docs]
def class_visualization(target_y, model_path, dtype, img_size=224, channels=None, l2_reg=1e-3, learning_rate=25, num_iterations=100, blur_every=10, max_jitter=16, show_every=25, class_names = None):
"""Synthesize an input image that maximizes the classifier score for ``target_y``.
:param target_y: target class index.
:param model_path: path to the trained model checkpoint.
:param dtype: tensor dtype; overridden internally based on CUDA availability.
:param img_size: square image size (pixels).
:param channels: input channels (defaults to ``[0, 1, 2]``).
:param l2_reg: L2 regularization weight on the pixel norm.
:param learning_rate: gradient-ascent step size.
:param num_iterations: optimization iteration count.
:param blur_every: interval (iterations) between periodic Gaussian blurs.
:param max_jitter: maximum pixel jitter applied per iteration.
:param show_every: interval (iterations) between preview plots.
:param class_names: display names for the classes (defaults to ``['nc', 'pc']``).
:returns: deprocessed image as a numpy array.
"""
if channels is None:
channels = [0,1,2]
if class_names is None:
class_names = ['nc', 'pc']
def jitter(img, ox, oy):
"""Return ``img`` shifted (rolled) by ``ox`` and ``oy`` pixels along the spatial axes."""
return torch.roll(torch.roll(img, ox, dims=2), oy, dims=3)
def blur_image(img, sigma=1):
"""In-place Gaussian blur of each channel of ``img`` with standard deviation ``sigma``."""
img_np = img.cpu().numpy()
for i in range(img_np.shape[1]):
img_np[:, i] = gaussian_filter(img_np[:, i], sigma=sigma)
img.copy_(torch.tensor(img_np).to(img.device))
def deprocess(img_tensor):
"""Undo ImageNet normalization and return an ``(H, W, 3)`` numpy image in ``[0, 1]``."""
img_tensor = img_tensor.clone()
for c in range(3):
img_tensor[:, c] = img_tensor[:, c] * SQUEEZENET_STD[c] + SQUEEZENET_MEAN[c]
img_tensor = img_tensor.clamp(0, 1)
return img_tensor.squeeze().permute(1, 2, 0).cpu().numpy()
SQUEEZENET_MEAN = [0.485, 0.456, 0.406]
SQUEEZENET_STD = [0.229, 0.224, 0.225]
model = torch.load(model_path, weights_only=False)
from .accelerator import is_cuda
dtype = torch.cuda.FloatTensor if is_cuda() else torch.FloatTensor
len_chans = len(channels)
model.type(dtype)
img = torch.randn(1, len_chans, img_size, img_size).mul_(1.0).type(dtype).requires_grad_()
for t in range(num_iterations):
ox, oy = random.randint(0, max_jitter), random.randint(0, max_jitter)
img.data.copy_(jitter(img.data, ox, oy))
score = model(img)
if target_y == 0:
target_score = -score
else:
target_score = score
target_score = target_score - l2_reg * torch.norm(img)
target_score.backward()
with torch.no_grad():
img += learning_rate * img.grad / torch.norm(img.grad)
img.grad.zero_()
img.data.copy_(jitter(img.data, -ox, -oy))
for c in range(3):
lo = float(-SQUEEZENET_MEAN[c] / SQUEEZENET_STD[c])
hi = float((1.0 - SQUEEZENET_MEAN[c]) / SQUEEZENET_STD[c])
img.data[:, c].clamp_(min=lo, max=hi)
if t % blur_every == 0:
blur_image(img.data, sigma=0.5)
if t == 0 or (t + 1) % show_every == 0 or t == num_iterations - 1:
plt.imshow(deprocess(img.data.clone().cpu()))
class_name = class_names[target_y]
plt.title('%s\nIteration %d / %d' % (class_name, t + 1, num_iterations))
plt.gcf().set_size_inches(4, 4)
plt.axis('off')
plt.show()
return deprocess(img.data.cpu())
[docs]
def get_submodules(model, prefix=''):
"""Return all dotted submodule names of ``model`` in traversal order.
:param model: PyTorch module to walk.
:param prefix: optional prefix prepended to returned names.
:returns: list of dotted submodule names.
"""
submodules = []
for name, module in model.named_children():
full_name = prefix + ('.' if prefix else '') + name
submodules.append(full_name)
submodules.extend(get_submodules(module, full_name))
return submodules
[docs]
class GradCAM:
"""Named-hook Grad-CAM implementation for arbitrary target layers.
:param model: trained model to inspect.
:param target_layers: list of dotted layer names to hook.
:param use_cuda: if true, move the model and inputs to CUDA
unconditionally; the caller must ensure CUDA is available.
"""
def __init__(self, model, target_layers=None, use_cuda=True):
"""Store the model and move it to CUDA if requested."""
self.model = model
self.model.eval()
self.target_layers = target_layers
self.cuda = use_cuda
if self.cuda:
self.model = model.cuda()
[docs]
def forward(self, input):
"""Return the model output for ``input``.
:param input: tensor passed directly to the wrapped model.
"""
return self.model(input)
[docs]
def __call__(self, x, index=None):
"""Return the normalized CAM heatmap for input ``x``, targeting class ``index``."""
if self.cuda:
x = x.cuda()
features = []
def hook(module, input, output):
"""Forward hook: append the target layer's output to ``features``.
``retain_grad()`` is required: PyTorch only populates ``.grad`` on
leaf tensors, so without it ``features[0].grad`` is None below and
GradCAM died with "'NoneType' object has no attribute 'cpu'".
"""
if output.requires_grad:
output.retain_grad()
features.append(output)
handles = []
for name, module in self.model.named_modules():
if name in self.target_layers:
handles.append(module.register_forward_hook(hook))
output = self.forward(x)
if index is None:
index = np.argmax(output.data.cpu().numpy())
one_hot = np.zeros((1, output.size()[-1]), dtype=np.float32)
one_hot[0][index] = 1
one_hot = torch.from_numpy(one_hot).requires_grad_(True)
if self.cuda:
one_hot = one_hot.cuda()
one_hot = torch.sum(one_hot * output)
self.model.zero_grad()
one_hot.backward(retain_graph=True)
grads_val = features[0].grad.cpu().data.numpy()
target = features[0].cpu().data.numpy()[0, :]
weights = np.mean(grads_val, axis=(2, 3))[0, :]
cam = np.zeros(target.shape[1:], dtype=np.float32)
for i, w in enumerate(weights):
cam += w * target[i, :, :]
cam = np.maximum(cam, 0)
cam = cv2.resize(np.atleast_2d(cam), (x.size(2), x.size(3)))
cam = cam - np.min(cam)
peak = np.max(cam)
if peak > 0:
cam = cam / peak
else:
cam.fill(0.0)
for handle in handles:
handle.remove()
return cam
[docs]
def show_cam_on_image(img, mask):
"""Return ``img`` overlaid with a jet colormap of ``mask`` as an 8-bit RGB image.
The sum of heatmap and image is renormalized by its own peak, so the output
brightness is relative to the single hottest pixel -- two images overlaid
separately are not comparable to each other on absolute intensity.
:param img: 3-channel ``(H, W, 3)`` image already scaled to ``[0, 1]``. It is
added to the colormap rather than blended, so a ``[0, 255]`` image
swamps the heatmap and the result is a near-uniform wash. A 2-D
grayscale array fails to broadcast against the 3-channel heatmap.
An image negative enough that the blend has no positive pixel left
raises rather than returning a black frame -- a black attribution map
is indistinguishable from "the model looked nowhere", which is a
claim this function must never make on the strength of bad input.
:param mask: ``(H, W)`` activation map in ``[0, 1]``, matching ``img`` in
height and width. Values outside that range are CLIPPED to it, with a
:class:`RuntimeWarning`, so an un-normalized CAM saturates at the hot
end instead of wrapping the ``np.uint8`` cast: before this was clipped,
``1.1`` landed at the cold end of jet, ``1.5`` in the middle and ``2.0``
back at the top, which could render the hottest region of a map as the
coldest colour. An all-zero mask does not produce a black overlay: jet
maps 0 to a non-zero color, so a zero mask over a zero image
renormalizes to a saturated flat field.
:raises ValueError: ``img`` or ``mask`` contains NaN or infinity, or the
blend has no positive pixel to normalize against.
"""
import warnings
mask = np.asarray(mask, dtype=np.float32)
img = np.asarray(img, dtype=np.float32)
if not np.isfinite(mask).all():
raise ValueError(
"the activation map contains NaN or infinity, so it cannot be "
"coloured; a CAM that came out non-finite means the attribution "
"itself failed and the overlay would hide that")
if not np.isfinite(img).all():
raise ValueError(
"the image contains NaN or infinity, so the overlay cannot be "
"normalized")
low, high = float(mask.min()), float(mask.max())
if low < 0.0 or high > 1.0:
warnings.warn(
f"activation map spans [{low:g}, {high:g}]; show_cam_on_image "
f"expects [0, 1] and is clipping to it. Normalize the CAM to "
f"avoid saturating the colormap.",
RuntimeWarning, stacklevel=2)
scaled = np.clip(mask, 0.0, 1.0)
heatmap = cv2.applyColorMap(np.uint8(255 * scaled), cv2.COLORMAP_JET)
heatmap = np.float32(heatmap) / 255
cam = heatmap + img
peak = float(np.max(cam))
if not peak > 0:
raise ValueError(
f"the heatmap blend peaks at {peak:g}, so there is nothing to "
f"normalize against; show_cam_on_image needs an image scaled to "
f"[0, 1] (this one spans [{float(img.min()):g}, "
f"{float(img.max()):g}])")
cam = np.clip(cam / peak, 0.0, 1.0)
return np.uint8(255 * cam)
[docs]
def recommend_target_layers(model):
"""Return ``([last_conv_layer], all_conv_layers)`` from ``model``.
:param model: PyTorch module to scan for ``Conv2d`` layers.
:returns: tuple ``(recommended, all)`` of layer-name lists.
:raises ValueError: if the model contains no convolutional layers.
"""
target_layers = []
for name, module in model.named_modules():
if isinstance(module, torch.nn.Conv2d):
target_layers.append(name)
if target_layers:
return [target_layers[-1]], target_layers
else:
raise ValueError("No convolutional layers found in the model.")
[docs]
class IntegratedGradients:
"""Compute integrated-gradients attributions for a classifier.
:param model: trained PyTorch model.
"""
def __init__(self, model):
"""Store the model and switch it to eval mode."""
self.model = model
self.model.eval()
[docs]
def generate_integrated_gradients(self, input_tensor, target_label_idx, baseline=None, num_steps=50):
"""Return integrated gradients from ``baseline`` to ``input_tensor`` for ``target_label_idx``.
:param input_tensor: input sample tensor.
:param target_label_idx: target class index whose logit is attributed.
:param baseline: reference tensor (defaults to zeros of the same shape).
:param num_steps: number of Riemann-sum interpolation steps.
:returns: attribution ndarray with the shape of ``input_tensor``.
"""
if baseline is None:
baseline = torch.zeros_like(input_tensor)
assert baseline.shape == input_tensor.shape
scaled_inputs = [(baseline + (float(i) / num_steps) * (input_tensor - baseline)).requires_grad_(True) for i in range(0, num_steps + 1)]
grads = []
for scaled_input in scaled_inputs:
out = self.model(scaled_input)
self.model.zero_grad()
out[0, target_label_idx].backward(retain_graph=True)
grads.append(scaled_input.grad.data.cpu().numpy())
avg_grads = np.mean(grads[:-1], axis=0)
integrated_grads = (input_tensor.cpu().data.numpy() - baseline.cpu().data.numpy()) * avg_grads
return integrated_grads
[docs]
def get_db_paths(src):
"""Return the standard ``measurements/measurements.db`` paths for one or more source roots.
:param src: plate folder, or list of plate folders for a multi-plate run. A
bare string is wrapped, so the return type is always a list and a
single-plate caller still has to index ``[0]``. These are the run roots
that Measure wrote into, not the ``measurements`` folder itself -- the
``measurements/measurements.db`` suffix is appended here. Nothing is
checked for existence, so a typo produces a path that only fails later at
connect time.
"""
if isinstance(src, str):
src = [src]
db_paths = [os.path.join(source, 'measurements/measurements.db') for source in src]
return db_paths
[docs]
def get_sequencing_paths(src):
"""Return the standard ``sequencing/sequencing_data.csv`` paths for one or more source roots.
:param src: plate folder, or list of plate folders, given in the *same order*
as the corresponding :func:`get_db_paths` call -- the two lists are
zipped positionally when measurements are joined to barcode counts, so a
reordered list silently pairs a plate's images with another plate's
reads. A bare string is wrapped, so the result is always a list, and no
path is checked for existence.
"""
if isinstance(src, str):
src = [src]
seq_paths = [os.path.join(source, 'sequencing/sequencing_data.csv') for source in src]
return seq_paths
[docs]
def load_image_paths(c, visualize):
"""Load the ``png_list`` table into a DataFrame indexed by ``prcfo`` and optionally filter by object.
:param c: open sqlite3 cursor.
:param visualize: object-type prefix (``'cell'``/``'nucleus'``/...) or falsy to keep all rows.
:returns: DataFrame of PNG metadata indexed by ``prcfo``.
"""
c.execute(f'SELECT * FROM png_list')
data = c.fetchall()
columns_info = c.execute(f'PRAGMA table_info(png_list)').fetchall()
column_names = [col_info[1] for col_info in columns_info]
image_paths_df = pd.DataFrame(data, columns=column_names)
if visualize:
object_visualize = visualize + '_png'
image_paths_df = image_paths_df[image_paths_df['png_path'].str.contains(object_visualize)]
image_paths_df = image_paths_df.set_index('prcfo')
return image_paths_df
[docs]
def merge_dataframes(df, image_paths_df, verbose):
"""Merge ``df`` into ``image_paths_df`` on the shared ``prcfo`` index.
:param df: feature DataFrame with a ``prcfo`` column.
:param image_paths_df: DataFrame indexed by ``prcfo``.
:param verbose: display the merged DataFrame.
:returns: merged DataFrame.
"""
df.set_index('prcfo', inplace=True)
df = image_paths_df.merge(
df,
left_index=True,
right_index=True,
validate='many_to_one',
)
if verbose:
display(df)
return df
[docs]
def filter_columns(df, filter_by):
"""Return ``df`` restricted to columns matching ``filter_by`` (or morphology columns).
:param df: source DataFrame.
:param filter_by: substring required in column names, or ``'morphology'`` to drop channel columns.
:returns: column-filtered DataFrame.
"""
if filter_by != 'morphology':
cols_to_include = [col for col in df.columns if filter_by in str(col)]
else:
cols_to_include = [col for col in df.columns if 'channel' not in str(col)]
df = df[cols_to_include]
return df
[docs]
def reduction_and_clustering(
numeric_data, n_neighbors, min_dist, metric, eps, min_samples,
clustering, reduction_method='umap', verbose=False, embedding=None,
n_jobs=-1, mode='fit', model=False, reducer_options=None,
prefer_gpu=False, random_seed=42):
"""Reduce ``numeric_data`` to 2-D and cluster the embedding.
Supported reducers are UMAP, t-SNE, PCA, Isomap and Spectral Embedding.
``reducer_options`` carries only method-specific settings; irrelevant
options are never forwarded. RAPIDS is opt-in and applies to UMAP, t-SNE
and PCA, with the actual backend retained on the fitted reducer.
:param numeric_data: rows of numeric features to embed and cluster.
:param n_neighbors: reducer neighborhood size, or a row fraction as a
float; also supplies the default t-SNE perplexity.
:param min_dist: minimum embedding distance used by UMAP.
:param metric: distance metric used by the reducer and DBSCAN.
:param eps: DBSCAN neighborhood radius.
:param min_samples: DBSCAN minimum neighborhood size, or KMeans cluster
count when ``clustering='kmeans'``.
:param clustering: clustering algorithm, ``'dbscan'`` or ``'kmeans'``.
"""
from .resource_log import _guard_workers, _table_nbytes
n_jobs = _guard_workers('umap', n_jobs, _table_nbytes(numeric_data))
values = np.asarray(numeric_data)
options = dict(reducer_options or {})
aliases = {
't-sne': 'tsne', 't_sne': 'tsne',
'spectral_embedding': 'spectral', 'spectral-embedding': 'spectral',
}
method = aliases.get(
str(reduction_method or 'umap').strip().lower(),
str(reduction_method or 'umap').strip().lower(),
)
supported = ('umap', 'tsne', 'pca', 'isomap', 'spectral')
if method not in supported:
raise ValueError(
f"Unsupported reduction method: {reduction_method}. Supported "
f"methods are {', '.join(supported)}")
gpu_supported = ('umap', 'tsne', 'pca')
if prefer_gpu and method not in gpu_supported:
raise ValueError(
f"GPU acceleration is not available for {method}. Turn GPU off "
f"or choose one of {', '.join(gpu_supported)}.")
if isinstance(n_neighbors, float):
n_neighbors = int(n_neighbors * len(values))
n_neighbors = max(2, int(n_neighbors))
seed = _run_random_state(int(random_seed))
if mode == 'fit':
backend = 'cpu'
if method == 'umap':
kwargs = dict(
n_neighbors=n_neighbors, n_components=2, metric=metric,
min_dist=float(min_dist), random_state=seed,
transform_seed=seed, n_jobs=n_jobs, verbose=bool(verbose),
)
from .gpu_reduce import make_reducer
reducer, backend = make_reducer(
'umap', prefer_gpu=bool(prefer_gpu), **kwargs)
elif method == 'tsne':
requested_perplexity = float(
options.get('perplexity', n_neighbors))
if requested_perplexity <= 0 or len(values) < 2:
raise ValueError(
"t-SNE perplexity must be greater than 0 and smaller "
f"than the {len(values)} input rows; got "
f"{requested_perplexity}.")
perplexity = min(requested_perplexity, float(len(values) - 1))
if verbose and perplexity != requested_perplexity:
print(f'Adjusted t-SNE perplexity from '
f'{requested_perplexity:g} to {perplexity:g} for '
f'{len(values)} rows')
kwargs = dict(
n_components=2, perplexity=perplexity,
early_exaggeration=float(
options.get('early_exaggeration', 12.0)),
learning_rate=float(options.get('learning_rate', 200.0)),
max_iter=int(options.get('max_iter', 1000)), metric=metric,
init='random', verbose=int(bool(verbose)), random_state=seed,
)
if not prefer_gpu:
kwargs['n_jobs'] = n_jobs
from .gpu_reduce import make_reducer
reducer, backend = make_reducer(
'tsne', prefer_gpu=bool(prefer_gpu), **kwargs)
elif method == 'pca':
kwargs = dict(
n_components=2, whiten=bool(options.get('whiten', False)),
svd_solver=str(options.get('svd_solver', 'auto')),
random_state=seed,
)
from .gpu_reduce import make_reducer
reducer, backend = make_reducer(
'pca', prefer_gpu=bool(prefer_gpu), **kwargs)
elif method == 'isomap':
graph_neighbors = min(
max(1, int(options.get('n_neighbors', n_neighbors))),
max(1, len(values) - 1),
)
reducer = Isomap(
n_neighbors=graph_neighbors,
n_components=2, metric=metric,
path_method=str(options.get('path_method', 'auto')),
n_jobs=n_jobs,
)
else:
affinity = str(options.get('affinity', 'nearest_neighbors'))
kwargs = dict(
n_components=2, affinity=affinity, random_state=seed,
n_jobs=n_jobs,
)
if affinity == 'nearest_neighbors':
kwargs['n_neighbors'] = min(
max(1, int(options.get('n_neighbors', n_neighbors))),
max(1, len(values) - 1),
)
reducer = SpectralEmbedding(**kwargs)
if prefer_gpu and backend != 'cuml':
raise RuntimeError(
f"GPU was requested for {method}, but cuML could not build "
"the reducer. No CPU fallback was run; turn GPU off to use "
"the CPU backend.")
embedding = reducer.fit_transform(values)
if hasattr(embedding, 'get'):
embedding = embedding.get()
embedding = np.asarray(embedding)
try:
reducer._spacr_backend = backend
reducer._spacr_reduction_method = method
except Exception:
pass
if verbose:
print(f'Trained and fit reducer: {method} on {backend}')
else:
if model is None or model is False:
raise ValueError("Model is None. Please provide a model for transform.")
transform = getattr(model, 'transform', None)
if not callable(transform):
raise ValueError(
f"{method} cannot transform new rows after fitting. Turn off "
"embedding_by_controls or choose UMAP, PCA, or Isomap.")
embedding = transform(values)
if hasattr(embedding, 'get'):
embedding = embedding.get()
embedding = np.asarray(embedding)
reducer = model
if verbose:
print('Fit data to reducer')
if clustering == 'dbscan':
clustering_model = DBSCAN(eps=eps, min_samples=min_samples, metric=metric, n_jobs=n_jobs)
elif clustering == 'kmeans':
clustering_model = KMeans(n_clusters=min_samples, random_state=_run_random_state(42))
else:
raise ValueError(f"Unsupported clustering method: {clustering}. Supported methods are 'dbscan' and 'kmeans'")
clustering_model.fit(embedding)
labels = clustering_model.labels_ if clustering == 'dbscan' else clustering_model.predict(embedding)
if verbose:
print(f'Embedding shape: {embedding.shape}')
return embedding, labels, reducer
[docs]
def remove_noise(embedding, labels):
"""Drop rows of ``embedding`` (and ``labels``) whose label is DBSCAN noise (``-1``).
Rows are removed, not renumbered, so positional indices into the original
data (an ``image_paths`` list, a DataFrame row order) no longer line up with
the returned arrays -- filter those alongside, using the same mask, or keep
the identities before calling.
:param embedding: ``(N, D)`` ndarray of points. It is filtered by boolean
mask, so a Python list or a DataFrame will not index correctly; pass a
numpy array.
:param labels: length-``N`` ndarray of cluster labels, aligned row-for-row
with ``embedding``. Only ``-1`` is treated as noise, which is the DBSCAN
convention -- KMeans labels contain no ``-1`` and pass through unchanged,
making this a no-op rather than an error on KMeans output.
"""
non_noise_indices = labels != -1
embedding = embedding[non_noise_indices]
labels = labels[non_noise_indices]
return embedding, labels
[docs]
def plot_embedding(embedding, image_paths, labels, image_nr, img_zoom, colors,
plot_by_cluster, plot_outlines, plot_points, plot_images,
smooth_lines, black_background, figuresize, dot_size,
remove_image_canvas, verbose, interactive_payload=None,
theme_colors=None, point_color='cluster',
point_alpha=0.65, outline_width=1.0):
"""Plot a 2-D embedding with cluster outlines, points, and optional image overlays.
:param embedding: ``(N, 2)`` array of 2-D points (e.g. UMAP output).
:param image_paths: length-``N`` image paths used for the overlays; ``None`` skips them.
:param labels: length-``N`` cluster labels; ``-1`` denotes noise.
:param image_nr: number of images to overlay (per cluster when ``plot_by_cluster``).
:param img_zoom: zoom factor applied to each overlaid thumbnail.
:param colors: palette of per-cluster colors, one entry per unique label.
:param plot_by_cluster: sample the overlaid images per cluster instead of at random.
:param plot_outlines: draw a hull/smoothed outline around each cluster.
:param plot_points: render the scatter points (otherwise plotted invisibly).
:param plot_images: overlay the images from ``image_paths``.
:param smooth_lines: use a smoothed hull polyline instead of the convex hull edges.
:param black_background: use the white-on-black default theme instead of
black-on-white; entries in ``theme_colors`` override it per role.
:param figuresize: figure side length in inches; also scales label and tick fonts.
:param dot_size: scatter marker size in points.
:param remove_image_canvas: mask out zero-valued pixels of each overlaid image.
:param verbose: forwarded to the cluster and image helpers, which ignore it.
:param interactive_payload: optional object stashed on the figure as
``_spacr_umap_payload`` so the Qt bridge can keep point/image identities.
Default ``None``.
:param theme_colors: dict with ``background``/``foreground``/``border`` colors.
Default ``None``.
:param point_color: ``'cluster'``/``'viridis'`` colors points per cluster; any other
Matplotlib color is applied to every point. Default ``'cluster'``.
:param point_alpha: scatter opacity, clamped to ``[0, 1]``. Default ``0.65``.
:param outline_width: hull line width in points, floored at ``0.1``. Default ``1.0``.
:returns: matplotlib ``Figure``.
"""
with plt.rc_context():
unique_labels = np.unique(labels)
colors, label_to_color_index = assign_colors(unique_labels, colors)
cluster_centers = [np.mean(embedding[labels == cluster_label], axis=0) for cluster_label in unique_labels]
fig, ax = setup_plot(
figuresize, black_background, theme_colors=theme_colors)
from .figures.bundle import _register_figure_data
_register_figure_data(fig, lambda: pd.DataFrame({"embedding_1": np.asarray(embedding)[:, 0], "embedding_2": np.asarray(embedding)[:, 1], "cluster": np.asarray(labels).astype(str)}), x="embedding_1", y="embedding_2", hue="cluster", kind="scatter")
plot_clusters(
ax, embedding, labels, colors, cluster_centers, plot_outlines,
plot_points, smooth_lines, figuresize, dot_size, verbose,
point_color=point_color, point_alpha=point_alpha,
outline_width=outline_width,
)
if not image_paths is None and plot_images:
plot_umap_images(ax, image_paths, embedding, labels, image_nr, img_zoom, colors, plot_by_cluster, remove_image_canvas, verbose)
if interactive_payload is not None:
fig._spacr_umap_payload = interactive_payload
plt.show()
return fig
[docs]
def generate_colors(num_clusters, black_background):
"""Return a deterministic Viridis RGBA palette for cluster points.
:param num_clusters: how many colors to sample, evenly spaced across the
Viridis range ``0.08``-``0.92`` (the extremes are trimmed so the darkest
cluster stays visible on black and the lightest on white). Coerced with
``int()`` and floored at ``1``, so ``0`` or a negative still yields a
one-entry palette rather than an empty one. Count the clusters you will
actually plot: DBSCAN's ``-1`` noise label is not drawn, so including it
shifts every real cluster's color.
:param black_background: accepted for call compatibility with the plotting
helpers but **not used** -- the palette is Viridis either way. Set the
background through ``setup_plot``/``theme_colors`` instead; changing this
flag will not change the colors you get back.
"""
count = max(int(num_clusters), 1)
positions = np.linspace(0.08, 0.92, count)
return mpl.colormaps['viridis'](positions)
[docs]
def assign_colors(unique_labels, random_colors):
"""Return colors and their positional mapping for the unique labels.
:param unique_labels: cluster labels in the order assigned palette indices.
:param random_colors: iterable of color values converted to tuples.
"""
colors = [tuple(color) for color in random_colors]
label_to_color_index = {label: index for index, label in enumerate(unique_labels)}
return colors, label_to_color_index
def _plot_theme_colors(black_background, theme_colors=None):
"""Resolve serializable GUI colors or the historical CLI fallback."""
fallback = {
'background': 'black' if black_background else 'white',
'foreground': 'white' if black_background else 'black',
'border': 'white' if black_background else 'black',
}
if not isinstance(theme_colors, dict):
return fallback
resolved = fallback.copy()
for role in resolved:
value = theme_colors.get(role)
if value:
try:
mpl.colors.to_rgba(value)
except (TypeError, ValueError):
continue
resolved[role] = value
return resolved
def _style_plot_axes(fig, ax, colors):
"""Apply one theme to a Matplotlib figure, axes, and axis lines."""
background = colors['background']
foreground = colors['foreground']
border = colors['border']
fig.patch.set_facecolor(background)
ax.set_facecolor(background)
ax.tick_params(axis='both', colors=foreground)
ax.xaxis.label.set_color(foreground)
ax.yaxis.label.set_color(foreground)
ax.title.set_color(foreground)
for spine in ax.spines.values():
spine.set_color(border)
[docs]
def setup_plot(figuresize, black_background, theme_colors=None):
"""Create a square Matplotlib figure using scoped theme colors.
Parameters
----------
figuresize : float
Figure width and height in inches.
black_background : bool
Use the legacy dark or light fallback when ``theme_colors`` is not
supplied.
theme_colors : mapping, optional
``background``, ``foreground``, and ``border`` colors. Missing or
invalid entries use the fallback palette.
Returns
-------
tuple
The ``(figure, axes)`` pair.
Notes
-----
Theme values are applied inside :func:`matplotlib.rc_context` and then to
the created artists. Global Matplotlib settings are not modified.
"""
import matplotlib as mpl
colors = _plot_theme_colors(black_background, theme_colors)
with mpl.rc_context({
'figure.facecolor': colors['background'],
'axes.facecolor': colors['background'],
'axes.edgecolor': colors['border'],
'text.color': colors['foreground'],
'xtick.color': colors['foreground'],
'ytick.color': colors['foreground'],
'axes.labelcolor': colors['foreground'],
}):
fig, ax = plt.subplots(1, 1, figsize=(figuresize, figuresize))
_style_plot_axes(fig, ax, colors)
return fig, ax
[docs]
def plot_clusters(ax, embedding, labels, colors, cluster_centers,
plot_outlines, plot_points, smooth_lines, figuresize=10,
dot_size=50, verbose=False, point_color='cluster',
point_alpha=0.65, outline_width=1.0):
"""Draw cluster outlines, points, and centroid labels onto ``ax`` for a 2-D embedding.
:param ax: Matplotlib axes to draw into.
:param embedding: ``(N, 2)`` array of 2-D points (e.g. UMAP output).
:param labels: length-``N`` cluster labels; ``-1`` denotes noise.
:param colors: iterable of per-cluster colors, one per unique label.
:param cluster_centers: iterable of ``(x, y)`` centroids, one per unique label.
:param plot_outlines: draw a hull/smoothed outline around each cluster.
:param plot_points: render the scatter points (otherwise plotted invisibly).
:param smooth_lines: use a smoothed hull polyline instead of the convex hull edges.
:param figuresize: base size in inches used to scale axis label and tick fonts. Default ``10``.
:param dot_size: scatter marker size in points. Default ``50``.
:param verbose: unused placeholder kept for API compatibility. Default ``False``.
:param point_color: ``'cluster'``/``'viridis'`` (or empty) colors points per cluster;
any other Matplotlib color is applied to every point. Default ``'cluster'``.
:param point_alpha: scatter opacity, clamped to ``[0, 1]``; ignored when
``plot_points`` is ``False``. Default ``0.65``.
:param outline_width: hull line width in points, floored at ``0.1``. Default ``1.0``.
:returns: None.
"""
unique_labels = np.unique(labels)
alpha = max(0.0, min(1.0, float(point_alpha)))
width = max(0.1, float(outline_width))
fixed_color = None
if str(point_color).strip().lower() not in {"", "cluster", "viridis"}:
try:
fixed_color = mpl.colors.to_rgba(point_color)
except (TypeError, ValueError):
fixed_color = None
for cluster_label, color, center in zip(unique_labels, colors, cluster_centers):
cluster_data = embedding[labels == cluster_label]
marker_color = fixed_color or color
if smooth_lines:
if cluster_data.shape[0] > 2:
try:
x_smooth, y_smooth = smooth_hull_lines(cluster_data)
if plot_outlines:
ax.plot(
x_smooth, y_smooth, color=color, linewidth=width)
except Exception:
LOG.debug(
"Could not draw a smoothed hull for cluster %r",
cluster_label,
exc_info=True,
)
else:
if cluster_data.shape[0] > 2:
try:
hull = ConvexHull(cluster_data)
for simplex in hull.simplices:
if plot_outlines:
ax.plot(
hull.points[simplex, 0],
hull.points[simplex, 1],
color=color, linewidth=width,
)
except Exception:
LOG.debug(
"Could not draw a convex hull for cluster %r",
cluster_label,
exc_info=True,
)
if plot_points:
ax.scatter(cluster_data[:, 0], cluster_data[:, 1], s=dot_size, c=[marker_color], alpha=alpha, label=f'Cluster {cluster_label if cluster_label != -1 else "Noise"}')
else:
ax.scatter(cluster_data[:, 0], cluster_data[:, 1], s=dot_size, c=[marker_color], alpha=0, label=f'Cluster {cluster_label if cluster_label != -1 else "Noise"}')
ax.text(
center[0], center[1], str(cluster_label), fontsize=12,
ha='center', va='center',
color=ax.xaxis.label.get_color(),
bbox={
'facecolor': ax.get_facecolor(),
'edgecolor': 'none',
'alpha': 0.8,
'pad': 1.5,
},
)
legend = ax.legend(loc='best', fontsize=int(figuresize * 0.75))
if legend is not None:
legend.get_frame().set_facecolor(ax.get_facecolor())
legend.get_frame().set_edgecolor(ax.spines['left'].get_edgecolor())
for text in legend.get_texts():
text.set_color(ax.xaxis.label.get_color())
ax.set_xlabel('UMAP Dimension 1', fontsize=int(figuresize * 0.75))
ax.set_ylabel('UMAP Dimension 2', fontsize=int(figuresize * 0.75))
ax.tick_params(
axis='both', which='major', labelsize=int(figuresize * 0.75))
[docs]
def plot_umap_images(ax, image_paths, embedding, labels, image_nr, img_zoom, colors, plot_by_cluster, remove_image_canvas, verbose):
"""Overlay sample images from ``image_paths`` on the UMAP embedding in ``ax``.
:param ax: axes the thumbnails are added to, as frameless annotation boxes.
:param image_paths: paths addressed by the same positional index as
``embedding``, so the two must share a row order; a short list raises
``IndexError``.
:param embedding: ``(N, 2)`` array whose selected rows give each thumbnail
its data-space position.
:param labels: cluster labels aligned with ``embedding``. Read only when
``plot_by_cluster`` is true; ``None`` is accepted otherwise.
:param image_nr: with ``plot_by_cluster`` false, the exact number of rows
sampled at random from the whole embedding, so a value above ``N``
raises ``ValueError`` from ``random.sample``. With it true, a
per-cluster cap -- a cluster no larger than this contributes every
member, unsampled.
:param img_zoom: scale factor handed to ``OffsetImage``. It sizes the
thumbnail from the file's own pixel dimensions in display space, so
rescaling the axes does not change how big the image is drawn.
:param colors: only zipped against ``np.unique(labels)`` to drive the
iteration; the color itself is never drawn. Its LENGTH is therefore a
silent limit, and because ``np.unique`` includes the ``-1`` noise label
a palette sized to the real clusters leaves the last cluster with no
images. Unused (``None`` is fine) when ``plot_by_cluster`` is false.
:param plot_by_cluster: true samples per cluster and skips label ``-1``;
false ignores ``labels`` and ``colors`` entirely and samples globally.
:param remove_image_canvas: forwarded to :func:`plot_image`; true masks
zero-valued pixels out and restricts the inputs to PIL modes ``L``,
``I`` and ``RGB``.
:param verbose: accepted and ignored, here and in the helper it is passed
to; any value at all is tolerated.
:returns: None.
"""
if plot_by_cluster:
cluster_indices = {label: np.where(labels == label)[0] for label in np.unique(labels) if label != -1}
plot_images_by_cluster(ax, image_paths, embedding, labels, image_nr, img_zoom, colors, cluster_indices, remove_image_canvas, verbose)
else:
indices = random.sample(range(len(embedding)), image_nr)
for i, index in enumerate(indices):
x, y = embedding[index]
img = Image.open(image_paths[index])
plot_image(ax, x, y, img, img_zoom, remove_image_canvas)
[docs]
def plot_images_by_cluster(ax, image_paths, embedding, labels, image_nr, img_zoom, colors, cluster_indices, remove_image_canvas, verbose):
"""Overlay up to ``image_nr`` images per cluster on the embedding in ``ax``.
:param ax: axes the thumbnails are added to, as frameless annotation boxes.
:param image_paths: paths addressed by the indices held in
``cluster_indices``, so they must be in the embedding's row order.
:param embedding: ``(N, 2)`` array supplying each thumbnail's position.
:param labels: only ``np.unique(labels)`` is used, to decide which clusters
to visit; ``-1`` is skipped as noise.
:param image_nr: per-cluster cap. A cluster no larger than this contributes
all of its members -- no sampling happens.
:param img_zoom: scale factor handed to ``OffsetImage``, applied to the
file's own pixel dimensions rather than to data units.
:param colors: accepted for caller compatibility but not read. Thumbnail
overlays do not use a cluster color, and palette length no longer
limits how many labels are visited.
:param cluster_indices: mapping of label to the row indices to draw from.
Looked up with ``.get(label, [])``, so a label present in ``labels``
but absent here plots nothing instead of raising.
:param remove_image_canvas: forwarded to :func:`plot_image`.
:param verbose: accepted and ignored; any value at all is tolerated.
:returns: None.
"""
for cluster_label in np.unique(labels):
if cluster_label == -1:
continue
indices = cluster_indices.get(cluster_label, [])
if len(indices) > image_nr:
indices = random.sample(list(indices), image_nr)
for index in indices:
x, y = embedding[index]
img = Image.open(image_paths[index])
plot_image(ax, x, y, img, img_zoom, remove_image_canvas)
[docs]
def plot_image(ax, x, y, img, img_zoom, remove_image_canvas=True):
"""Place a zoomed thumbnail of ``img`` at ``(x, y)`` on ``ax``.
:param ax: axes the thumbnail is added to, as a frameless annotation box.
:param x: data-space x coordinate the thumbnail is anchored at.
:param y: data-space y coordinate the thumbnail is anchored at.
:param img: PIL image when ``remove_image_canvas`` is true, since
``img.mode`` is read; any array-like otherwise.
:param img_zoom: scale factor handed to ``OffsetImage``. It sizes the
thumbnail from the source's pixel dimensions in display space, so the
drawn size is unchanged by the axis limits.
:param remove_image_canvas: true swaps the image for an RGBA array whose
alpha channel hides zero-valued pixels, which accepts only PIL modes
``L``, ``I`` and ``RGB`` -- ``RGBA`` and ``P`` raise ``ValueError``, and
a numpy array raises ``AttributeError`` because it has no ``mode``. An
all-zero ``L`` image divides by its own zero maximum and comes out
``NaN`` rather than raising. False just calls ``np.array``.
Default ``True``.
:returns: None.
"""
if remove_image_canvas:
img = remove_canvas(img)
else:
img = np.array(img)
imagebox = OffsetImage(img, zoom=img_zoom)
ab = AnnotationBbox(imagebox, (x, y), frameon=False)
ax.add_artist(ab)
[docs]
def remove_canvas(img):
"""Return ``img`` as RGBA with zero-valued pixels made transparent.
:param img: PIL image in ``L``, ``I``, or ``RGB`` mode.
"""
if img.mode in ['L', 'I']:
img_data = np.array(img)
img_data = img_data / np.max(img_data)
alpha_channel = (img_data > 0).astype(float)
img_data_rgb = np.stack([img_data] * 3, axis=-1)
img_data_with_alpha = np.dstack([img_data_rgb, alpha_channel])
elif img.mode == 'RGB':
img_data = np.array(img)
img_data = img_data / 255.0
alpha_channel = (np.sum(img_data, axis=-1) > 0).astype(float)
img_data_with_alpha = np.dstack([img_data, alpha_channel])
else:
raise ValueError(f"Unsupported image mode: {img.mode}")
return img_data_with_alpha
[docs]
def plot_clusters_grid(embedding, labels, image_nr, image_paths, colors, figuresize, black_background, verbose, theme_colors=None):
"""Plot a grid of example images per cluster label discovered in ``labels``.
:param embedding: accepted and never read -- the panels are built from
``labels`` and ``image_paths`` alone, so ``None`` works.
:param labels: cluster labels in the row order of ``image_paths``. ``-1`` is
dropped as noise, and if nothing else remains the function prints
``No clusters found.`` and returns ``None`` instead of a figure.
:param image_nr: per-cluster cap on how many images are opened. A cluster no
larger than this contributes all of its members.
:param image_paths: paths addressed positionally by the label array; a list
shorter than ``labels`` raises ``IndexError``.
:param colors: palette indexed downstream by the cluster LABEL itself rather
than by its rank, so the palette has to be long enough to reach the
largest label -- labels ``0`` and ``5`` against a two-color palette
raise ``IndexError``. Entries need at least three components.
:param figuresize: per-cluster panel size in inches, shrunk downstream so
the whole row never exceeds 200 inches.
:param black_background: picks the white-on-black fallback theme instead of
black-on-white; ``theme_colors`` overrides it per role.
:param verbose: only ever prints for STRING cluster labels; silent for the
integer labels DBSCAN and KMeans produce.
:param theme_colors: dict with ``background``/``foreground``/``border``
colors. Entries Matplotlib cannot parse are dropped silently and fall
back to the ``black_background`` choice. Default ``None``.
:returns: the Matplotlib ``Figure``, or ``None`` when every label is ``-1``.
"""
unique_labels = np.unique(labels)
num_clusters = len(unique_labels[unique_labels != -1])
if num_clusters == 0:
print("No clusters found.")
return
cluster_images = {label: [] for label in unique_labels if label != -1}
cluster_indices = {label: np.where(labels == label)[0] for label in unique_labels if label != -1}
for cluster_label, indices in cluster_indices.items():
if len(indices) > image_nr:
indices = random.sample(list(indices), image_nr)
for index in indices:
img_path = image_paths[index]
img_array = Image.open(img_path)
img = np.array(img_array)
cluster_images[cluster_label].append(img)
fig = plot_grid(
cluster_images, colors, figuresize, black_background, verbose,
theme_colors=theme_colors)
return fig
[docs]
def plot_grid(cluster_images, colors, figuresize, black_background, verbose, theme_colors=None):
"""Render one column per cluster of representative images with colored borders and labels.
:param cluster_images: ordered mapping of cluster label to that cluster's
list of image arrays; one column per key, and an empty mapping raises
``ValueError`` from ``subplots``.
:param colors: palette used consistently for both panel borders and legend
swatches. Integer labels index it by label (wrapping when necessary),
while string labels use their position in ``cluster_images``. An empty
palette falls back to neutral grey, and entries need at least three
components.
:param figuresize: figure height in inches and the label font size; the
width is this times the cluster count. It is shrunk to
``200 / n_clusters`` when that product would exceed 200 inches, which
silently caps the font size too.
:param black_background: picks the white-on-black fallback theme instead of
black-on-white.
:param verbose: prints the label and its index for STRING cluster labels
only; integer labels never print anything.
:param theme_colors: dict with ``background``/``foreground``/``border``
colors overriding the ``black_background`` fallback; values Matplotlib
cannot parse are ignored. Default ``None``.
:returns: the Matplotlib ``Figure``, which is also passed to ``plt.show``.
"""
num_clusters = len(cluster_images)
max_figsize = 200
if figuresize * num_clusters > max_figsize:
figuresize = max_figsize / num_clusters
plot_colors = _plot_theme_colors(black_background, theme_colors)
with figure_style(theme_target()):
grid_fig, grid_axes = plt.subplots(1, num_clusters, figsize=(figuresize * num_clusters, figuresize), gridspec_kw={'wspace': 0.2, 'hspace': 0})
from .figures.bundle import _register_figure_data
_register_figure_data(grid_fig, None, kind="montage", title="Cluster images")
grid_fig.patch.set_facecolor(plot_colors['background'])
if num_clusters == 1:
grid_axes = [grid_axes]
cluster_labels = list(cluster_images.keys())
def cluster_color(cluster_label):
"""Resolve one cluster's color once for panels and legend."""
if isinstance(cluster_label, str):
idx = cluster_labels.index(cluster_label)
else:
idx = int(cluster_label)
return (colors[idx % len(colors)] if len(colors)
else (0.5, 0.5, 0.5))
for cluster_label, axes in zip(cluster_labels, grid_axes):
axes.set_facecolor(plot_colors['background'])
images = cluster_images[cluster_label]
num_images = len(images)
grid_size = int(np.ceil(np.sqrt(num_images)))
image_size = 0.9 / grid_size
whitespace = (1 - grid_size * image_size) / (grid_size + 1)
if isinstance(cluster_label, str) and verbose:
print(
f'Lable: {cluster_label} '
f'index: {cluster_labels.index(cluster_label)}')
color = cluster_color(cluster_label)
axes.add_patch(plt.Rectangle((0, 0), 1, 1, transform=axes.transAxes, color=color[:3]))
axes.axis('off')
for i, img in enumerate(images):
row = i // grid_size
col = i % grid_size
x_pos = (col + 1) * whitespace + col * image_size
y_pos = 1 - ((row + 1) * whitespace + (row + 1) * image_size)
ax_img = axes.inset_axes([x_pos, y_pos, image_size, image_size], transform=axes.transAxes)
ax_img.imshow(img, cmap='gray', aspect='auto')
ax_img.axis('off')
ax_img.set_aspect('equal')
ax_img.set_facecolor(color[:3])
spacing_factor = 0.5
for i, cluster_label in enumerate(cluster_labels):
color = cluster_color(cluster_label)
label_y = 1 - (i + 1) * (spacing_factor / num_clusters)
grid_fig.text(
1.05, label_y, f'Cluster {cluster_label}',
verticalalignment='center', fontsize=figuresize,
color=plot_colors['foreground'])
grid_fig.patches.append(plt.Rectangle((1, label_y - 0.02), 0.03, 0.03, transform=grid_fig.transFigure, color=color[:3], clip_on=False))
plt.show()
return grid_fig
[docs]
def generate_path_list_from_db(db_path, file_metadata):
"""Return all ``png_path`` values from ``db_path`` optionally filtered by ``file_metadata`` substrings.
:param db_path: path to the measurements SQLite DB.
:param file_metadata: substring or list of substrings to LIKE-match against ``png_path``.
:returns: list of PNG paths.
"""
all_paths = []
print(f"Reading DataBase: {db_path}")
try:
with sqlite3.connect(db_path, timeout=30) as conn:
cursor = conn.cursor()
if file_metadata:
if isinstance(file_metadata, str):
cursor.execute("SELECT png_path FROM png_list WHERE png_path LIKE ?", (f"%{file_metadata}%",))
elif isinstance(file_metadata, list):
query = "SELECT png_path FROM png_list WHERE " + " OR ".join(
["png_path LIKE ?" for _ in file_metadata])
params = [f"%{meta}%" for meta in file_metadata]
cursor.execute(query, params)
else:
cursor.execute("SELECT png_path FROM png_list")
while True:
rows = cursor.fetchmany(1000)
if not rows:
break
all_paths.extend([row[0] for row in rows])
except sqlite3.Error as e:
print(f"Database error: {e}")
return
except Exception as e:
print(f"Error: {e}")
return
return all_paths
[docs]
def correct_paths(df, base_path, folder='data'):
"""Rewrite PNG paths (in a DataFrame or list) so they live under ``base_path/folder``.
A non-string entry is passed through untouched. ``png_list`` is LEFT-joined
onto the object tables, so any object whose crop was never written arrives
here with ``png_path`` = NaN -- a state
:func:`spacr.io._read_and_join_tables` documents as healthy
(``len(merged) == len(cell) > len(png_list)``: ``save_png`` off for a field,
a crop that failed to write, an interrupted run, or a ``cell_id`` that could
not be migrated). Testing ``base_path not in path`` on that NaN raised
``TypeError: argument of type 'float' is not iterable`` and took the whole
embedding down over one missing thumbnail. There is no path to re-anchor for
such a row, and it has to keep its position so the rewritten column still
aligns with ``df``.
Delegate rewriting to :func:`spacr.crops.reanchor_path`, which handles
same-platform moves, Windows paths read on Linux, and old absolute paths.
It finds the rightmost anchor component after normalizing separators and
checks existing roots component by component. Paths with no matching
``folder`` component pass through unchanged; their count and one example
are printed so unresolved paths remain visible.
:param df: DataFrame with a ``png_path`` column, or a list of paths.
:param base_path: destination root to prepend.
:param folder: intermediate folder name that anchors the rewrite.
:returns: DataFrame + list, or list, mirroring the input type.
"""
from .crops import NO_ANCHOR, REANCHORED, reanchor_path
if isinstance(df, pd.DataFrame):
if 'png_path' not in df.columns:
print("No 'png_path' column found in the dataframe.")
return df, None
else:
image_paths = df['png_path'].to_list()
elif isinstance(df, list):
image_paths = df
adjusted_image_paths = []
unanchored = []
n_paths = 0
for path in image_paths:
if not isinstance(path, str) or not path:
adjusted_image_paths.append(path)
continue
n_paths += 1
new_path, outcome = reanchor_path(path, base_path, anchors=(folder,))
if outcome == NO_ANCHOR:
unanchored.append(path)
adjusted_image_paths.append(new_path)
if unanchored:
print(f"{len(unanchored):,} of {n_paths:,} recorded paths could not be "
f"re-anchored under {base_path}: they contain no '{folder}' "
f"component. The first is {unanchored[0]}")
if isinstance(df, pd.DataFrame):
df['png_path'] = adjusted_image_paths
return df, adjusted_image_paths
else:
return adjusted_image_paths
[docs]
def delete_folder(folder_path):
"""Recursively delete ``folder_path`` if it exists (files and subdirectories included).
:param folder_path: directory to remove, contents and all. A missing path or
a plain file is reported on stdout and ignored -- the function never
raises for those, so it cannot be used to confirm that a delete
happened; check with ``os.path.isdir`` afterwards if that matters.
Deletion is unconditional and unprompted, with no trash or dry-run, so a
wrong path is not recoverable. Symlinked subdirectories are not descended
into but are still handed to ``os.rmdir``, which raises on a symlink, so
a tree containing one aborts part-way through.
"""
if os.path.exists(folder_path) and os.path.isdir(folder_path):
for root, dirs, files in os.walk(folder_path, topdown=False):
for name in files:
os.remove(os.path.join(root, name))
for name in dirs:
os.rmdir(os.path.join(root, name))
os.rmdir(folder_path)
print(f"Folder '{folder_path}' has been deleted.")
else:
print(f"Folder '{folder_path}' does not exist or is not a directory.")
[docs]
def measure_test_mode(settings):
"""Copy a random subset of source files into a ``test/merged`` folder when ``test_mode`` is on.
Fewer files than ``test_nr`` is not an error. test_mode is the setting a
user reaches for on a SMALL plate, and ``random.sample`` raised
``ValueError: Sample larger than population or is negative`` on exactly
that case -- so the one folder you most want to smoke-test first was the
one folder test_mode refused to run on.
Only visible ``.npy`` arrays are sampled, so a macOS ``._`` sidecar is never
measured in place of a field. The folder's ``.spacr_plane_layout.json`` is
copied across as well when it exists: it is what says which plane is which,
and a ``test/merged`` without it is read as a legacy folder, against the
default plane order.
:param settings: settings dict; must contain ``src``, ``test_mode``, ``test_nr``.
:returns: settings dict with ``src`` optionally redirected to the test folder.
:raises ValueError: if there is nothing to sample -- an empty ``src``, or a
``test_nr`` below 1. Sampling zero files would point ``src`` at an
empty ``test/merged`` and the run would report "no fields found",
blaming the wrong thing.
"""
if settings['test_mode']:
if not os.path.basename(settings['src']) == 'test':
all_files = [f for f in os.listdir(settings['src'])
if f.endswith('.npy') and not f.startswith('.')
and os.path.isfile(os.path.join(settings['src'], f))]
n_test = min(int(settings['test_nr']), len(all_files))
if n_test < 1:
raise ValueError(
f"test_mode is on but nothing can be sampled from "
f"{settings['src']}: it holds {len(all_files)} file(s) and "
f"test_nr is {settings['test_nr']}. Point src at a folder "
f"with merged arrays in it, and set test_nr to at least 1.")
if n_test < int(settings['test_nr']):
print(f"test_mode: {settings['src']} holds {len(all_files)} "
f"file(s), fewer than test_nr={settings['test_nr']}; "
f"measuring all {n_test}.")
random_files = random.sample(all_files, n_test)
src = os.path.join(os.path.dirname(settings['src']),'test', 'merged')
if os.path.exists(src):
delete_folder(src)
os.makedirs(src, exist_ok=True)
for file in random_files:
shutil.copy(os.path.join(settings['src'], file), os.path.join(src,file))
from .crops import MERGED_LAYOUT_SIDECAR
layout = os.path.join(settings['src'], MERGED_LAYOUT_SIDECAR)
if os.path.isfile(layout):
shutil.copy(layout, os.path.join(src, MERGED_LAYOUT_SIDECAR))
settings['src'] = src
print(f'Changed source folder to {src} for test mode')
else:
print(f'Test mode enabled, using source folder {settings["src"]}')
return settings
[docs]
def normalize_feature_filter(filter_by):
"""Normalize text representations of an unfiltered feature selection.
Settings imported from CSV files and older Qt sessions can contain the
literal string ``"None"``. Treating that as a feature-name substring
removes every measurement column, although the UI means "all channels".
:param filter_by: the raw setting value. A string is stripped and, if it
case-insensitively matches one of the "no filter" spellings
(``""``, ``"none"``, ``"null"``, ``"all"``, ``"all_channels"``,
``"all channels"``, ``"*"``), collapsed to ``None`` -- otherwise the
stripped string is returned as a feature-name substring. Anything that is
not a string (a real ``None``, or a list of channels) passes through
untouched, so this is safe to apply unconditionally to whatever the
settings dict holds. Note ``"*"`` means *no filter*, not a glob: real
patterns are not supported, and the surviving string is matched as a
plain substring of the column name.
"""
if isinstance(filter_by, str):
value = filter_by.strip()
if value.lower() in {
"", "none", "null", "all", "all_channels", "all channels", "*",
}:
return None
return value
return filter_by
def _available_feature_filters(columns):
"""Return useful channel/filter choices represented by feature names."""
options = {
match
for column in columns
for match in re.findall(r"channel_\d+", str(column))
}
morphology_tokens = (
"area", "major_axis_length", "minor_axis_length", "eccentricity",
"extent", "perimeter", "solidity", "zernike_",
)
if any(any(token in str(column) for token in morphology_tokens)
for column in columns):
options.add("morphology")
return sorted(options)
def _feature_filter_matches(columns, filter_by):
"""Return columns selected by the same public filter forms as the UI."""
if filter_by == "morphology":
morphology_tokens = (
"area", "area_bbox", "major_axis_length", "minor_axis_length",
"eccentricity", "extent", "perimeter", "euler_number", "solidity",
"zernike_", "area_filled", "convex_area",
"equivalent_diameter_area", "feret_diameter_max",
)
return [
column for column in columns
if any(token in str(column) for token in morphology_tokens)
]
if isinstance(filter_by, list):
terms = [f"channel_{channel}" for channel in filter_by]
elif isinstance(filter_by, int):
terms = [f"channel_{filter_by}"]
else:
terms = [str(filter_by)]
return [
column for column in columns
if any(term in str(column) for term in terms)
]
[docs]
def preprocess_data(
df,
filter_by,
remove_highly_correlated,
log_data,
exclude,
column_list=False,
*,
batch_correction="none",
batch_column="plateID",
batch_control_column=None,
batch_control_values=None,
batch_covariate_column=None,
batch_combat_mean_only=False,
batch_min_samples=3,
batch_missing_control="error",
):
"""Prepare a feature matrix by filtering, decorrelating, log-transforming, and scaling ``df``.
:param df: input DataFrame.
:param filter_by: channel of interest passed to
:func:`filter_dataframe_features`; ``None`` and its text forms disable
filtering.
:param remove_highly_correlated: correlation cutoff (float) or ``True`` to use ``0.95``; ``False`` disables.
:param log_data: apply ``log(x + 1e-6)`` to numeric columns.
:param exclude: features to exclude from filtering.
:param column_list: optional explicit column subset applied before selecting numeric columns.
:param batch_correction: ``none``, ``center``, ``zscore``,
``robust_zscore``, ``control_center`` or ``combat``.
:param batch_column: metadata column identifying acquisition batches.
:param batch_control_column: metadata column selecting reference controls.
:param batch_control_values: reference value(s) for ``control_center``.
:param batch_covariate_column: metadata column(s) naming the biology
``combat`` must preserve. Required by ``combat``, ignored by every
other method — and left blank, ``combat`` refuses to run rather than
removing the contrast along with the plate effect.
:param batch_combat_mean_only: correct only ``combat``'s additive shift
and leave each batch's scale alone.
:param batch_min_samples: minimum rows/reference controls per batch.
:param batch_missing_control: ``error`` or ``skip`` when a batch lacks
enough controls.
:returns: standard-scaled ``ndarray`` of numeric features.
:raises ValueError: if no numeric columns remain after filtering.
"""
metadata_df = df
filter_by = normalize_feature_filter(filter_by)
explicit_features = column_list or ()
excluded_features = (
[exclude] if isinstance(exclude, str) else (exclude or ())
)
allow_unknown = not bool(filter_by or column_list)
df = schema.coerce_model_feature_types(
df,
extra_features=explicit_features,
exclude=excluded_features,
allow_unknown=allow_unknown,
)
available_features = schema.model_feature_columns(
df,
extra_features=explicit_features,
exclude=excluded_features,
allow_unknown=allow_unknown,
)
if filter_by is not None:
if not _feature_filter_matches(available_features, filter_by):
choices = _available_feature_filters(available_features)
choices_text = ", ".join(choices) if choices else "none"
raise ValueError(
f"filter_by={filter_by!r} matched no measurement features. "
f"Available feature filters: {choices_text}. Set filter_by "
f"to None to use every declared measurement feature."
)
df, _ = filter_dataframe_features(df, channel_of_interest=filter_by, exclude=exclude)
if column_list:
df = df[column_list]
numeric_data = schema.model_feature_frame(
df,
extra_features=explicit_features,
exclude=excluded_features,
allow_unknown=allow_unknown,
)
numeric_data = _resolve_missing_model_features(numeric_data)
if numeric_data.empty:
if filter_by is not None:
raise ValueError(
f"filter_by={filter_by!r} initially matched measurement "
f"features, but none remained after removing excluded, "
f"constant, correlated, or incomplete columns. Choose another "
f"filter or set filter_by to None."
)
raise ValueError(
"No numeric measurement columns are available. Check the selected "
"tables and excluded features."
)
if not remove_highly_correlated is False:
if isinstance(remove_highly_correlated, float):
numeric_data = remove_highly_correlated_columns(numeric_data, remove_highly_correlated)
else:
numeric_data = remove_highly_correlated_columns(numeric_data, 0.95)
if log_data:
numeric_data = np.log(numeric_data + 1e-6)
if str(batch_correction or "none").strip().lower() not in {
"none", "off", "false",
}:
from .batch_correction import correct_from_metadata
numeric_data, correction_report = correct_from_metadata(
numeric_data,
metadata_df.loc[numeric_data.index],
batch_correction=batch_correction,
batch_column=batch_column,
batch_control_column=batch_control_column,
batch_control_values=batch_control_values,
batch_covariate_column=batch_covariate_column,
batch_combat_mean_only=batch_combat_mean_only,
batch_min_samples=batch_min_samples,
batch_missing_control=batch_missing_control,
)
print(
"Batch correction "
f"{correction_report.method}: {len(correction_report.batches)} "
f"batch(es), centroid spread "
f"{correction_report.centroid_spread_before} -> "
f"{correction_report.centroid_spread_after}."
)
for note in correction_report.warnings:
print(f"Warning: batch correction: {note}")
numeric_data = numeric_data.fillna(numeric_data.mean())
scaler = StandardScaler(copy=True, with_mean=True, with_std=True)
numeric_data = scaler.fit_transform(numeric_data)
return numeric_data
[docs]
def remove_low_variance_columns(df, threshold=0.01, verbose=False):
"""Drop numeric columns whose variance is below ``threshold``.
:param df: input DataFrame.
:param threshold: variance cutoff.
:param verbose: print the dropped column names.
:returns: filtered DataFrame.
"""
numerical_cols = df.select_dtypes(include=[np.number])
low_variance_cols = numerical_cols.var()[numerical_cols.var() < threshold].index.tolist()
if verbose:
print(f"Removed columns due to low variance: {low_variance_cols}")
df = df.drop(columns=low_variance_cols)
return df
#: The morphology measurements, by their bare skimage names. A column
#: belongs to this group when one of these appears anywhere in its name, so
#: `cell_area`, `nucleus_solidity` and `pathogen_zernike_7` are all
#: morphology whichever object they were measured on.
MORPHOLOGY_FEATURES = (
'area', 'area_bbox', 'major_axis_length', 'minor_axis_length',
'eccentricity', 'extent', 'perimeter', 'euler_number', 'solidity',
'area_filled', 'convex_area', 'equivalent_diameter_area',
'feret_diameter_max',
) + tuple(f'zernike_{i}' for i in range(25))
#: The feature group meaning "the shape of the object, whatever it was
#: stained with". Spelled out because it is a value a user picks, not an
#: implementation detail. Its canonical reader lives in the lightweight
#: settings module so merely opening Classify does not import this module's
#: torch/cv2/matplotlib stack.
from .settings import (
FEATURE_SELECTION_MORPHOLOGY as MORPHOLOGY,
canonical_feature_selection as feature_selection,
)
#: Every group the panel offers, in the order it offers them.
FEATURE_GROUPS = (0, 1, 2, 3, MORPHOLOGY)
[docs]
def feature_columns(columns, selection):
"""Which of ``columns`` the selection keeps. Order preserved.
The union over the selection's members, so ``[1, 'morphology']`` is
channel 1's intensities AND the shapes rather than the empty
intersection of the two.
COLOCALISATION BELONGS TO BOTH CHANNELS IT MEASURES. A
``cell_channel_1_channel_2_pearsons`` column names two channels and
survives a request for either -- which is what makes "localization"
reachable without a separate setting: ask for the channel and its
relationships come with it.
:param columns: ordered column names available for selection.
:param selection: channel, morphology group, text filter, mixture, or
``None`` as accepted by :func:`feature_selection`.
"""
canonical = feature_selection(selection)
if canonical is None:
return list(columns)
members = canonical if isinstance(canonical, list) else [canonical]
keep = []
for column in columns:
name = str(column)
for member in members:
if member == MORPHOLOGY:
if any(base in name for base in MORPHOLOGY_FEATURES):
keep.append(column)
break
elif isinstance(member, int):
if f"channel_{member}" in name:
keep.append(column)
break
elif str(member) in name:
keep.append(column)
break
return keep
def _resolve_missing_model_features(df):
"""Make non-finite measurements fit-ready without inventing zero signal.
A missing value and either sign of infinity carry the same information at
this boundary: no finite measurement exists for that object. A measurement
absent from every object is removed. A measurement available for at least
one object is retained, with its non-finite rows filled by that feature's
median. Median imputation gives an object with no measurement the typical
observed value, so the absence itself cannot masquerade as unusually
strong or weak signal; it is also robust to the long-tailed intensity
distributions common here.
:param df: numeric model-feature frame.
:returns: a frame with no missing values.
"""
df = df.replace([np.inf, -np.inf], np.nan)
missing = df.isna()
all_missing = missing.all(axis=0)
all_missing_columns = all_missing[all_missing].index.tolist()
if all_missing_columns:
df = df.drop(columns=all_missing_columns)
partially_missing = df.isna().any(axis=0)
partial_columns = partially_missing[partially_missing].index.tolist()
missing_values = (
int(df[partial_columns].isna().sum().sum())
if partial_columns else 0
)
if partial_columns:
medians = df[partial_columns].median(axis=0, skipna=True)
df = df.copy()
df[partial_columns] = df[partial_columns].fillna(medians)
print(
f"Dropped {len(all_missing_columns)} columns with NaN values "
"(all values were missing)"
)
if partial_columns:
print(
f"Median-imputed {missing_values} missing value(s) in "
f"{len(partial_columns)} partially observed feature column(s)"
)
return df
[docs]
def filter_dataframe_features(df, channel_of_interest, exclude=None, remove_low_variance_features=True, remove_highly_correlated_features=True, verbose=False):
"""Restrict a features DataFrame to a channel of interest and clean up correlated/low-variance columns.
:param df: input DataFrame.
:param channel_of_interest: int, str, list, or ``'morphology'`` to select feature groups.
:param exclude: feature(s) to drop from the final list. A single name or
any number of them; the Qt 'Exclude' field collects a list.
:param remove_low_variance_features: apply :func:`remove_low_variance_columns`.
:param remove_highly_correlated_features: apply :func:`remove_highly_correlated_columns`.
:param verbose: print filter details.
:returns: ``(filtered_df, features)``.
"""
excluded_features = (
[exclude] if isinstance(exclude, str) else (exclude or ())
)
missing_exclusions = [
feature for feature in excluded_features if feature not in df.columns
]
if missing_exclusions:
raise ValueError(
"Requested feature exclusions are not present in the input "
f"table: {missing_exclusions}. Available columns: "
f"{sorted(map(str, df.columns))}.")
df = schema.coerce_model_feature_types(df, exclude=excluded_features)
declared_features = schema.model_feature_columns(
df, exclude=excluded_features)
legacy_non_features = {
col for col in df.columns
if '_id' in col or 'count' in col
}
count_and_id_columns = [
col for col in df.columns
if col not in declared_features or col in legacy_non_features
]
declared_features = [
col for col in declared_features if col not in legacy_non_features]
if verbose:
print("Columns to remove:", count_and_id_columns)
df = df[declared_features].copy()
selection = feature_selection(channel_of_interest)
if selection is not None:
keep = feature_columns(df.columns, selection)
columns_to_drop = [col for col in df.columns if col not in set(keep)]
df = df.drop(columns=columns_to_drop)
if verbose:
print(f"Removed columns: {columns_to_drop}")
df = _resolve_missing_model_features(df)
if remove_low_variance_features:
df = remove_low_variance_columns(df, threshold=0.01, verbose=verbose)
if remove_highly_correlated_features:
df = remove_highly_correlated_columns(df, threshold=0.95, verbose=verbose)
features = schema.model_feature_columns(df)
if isinstance(exclude, list):
features = [feature for feature in features if feature not in exclude]
elif isinstance(exclude, str):
features = [feature for feature in features if feature != exclude]
filtered_df = df[features]
return filtered_df, features
[docs]
def check_overlap(current_position, other_positions, threshold):
"""Return ``True`` if ``current_position`` is within ``threshold`` of any point in ``other_positions``.
:param current_position: candidate point as a sequence of coordinates. Any
dimensionality works as long as it matches the entries of
``other_positions``; a genuine length mismatch raises ``ValueError`` from
the subtraction, while a length-1 entry broadcasts silently and yields a
meaningless distance.
:param other_positions: already-placed points to test against. Scanned
linearly with an early return, so cost grows with the number of placed
items -- this is the inner loop of the image-scatter layout. An empty
sequence returns ``False``, so the first placement always succeeds.
:param threshold: minimum center-to-center Euclidean separation, in the same
units as the coordinates (data units for an embedding, not pixels or
points). The comparison is strict ``<``, so a distance exactly equal to
``threshold`` counts as *not* overlapping. Because it measures centers,
set it to roughly the thumbnail width; half of that still lets images
overlap visually.
"""
for other_position in other_positions:
distance = np.linalg.norm(np.array(current_position) - np.array(other_position))
if distance < threshold:
return True
return False
[docs]
def find_non_overlapping_position(x, y, image_positions, threshold, max_attempts=100):
"""Return a nearby ``(x, y)`` jittered position that does not collide with ``image_positions``.
:param x: original x.
:param y: original y.
:param image_positions: previously placed points.
:param threshold: minimum allowed spacing.
:param max_attempts: retry budget before giving up.
:returns: ``(x, y)`` tuple; original position if no non-overlapping spot is found.
"""
offset_range = 10
attempts = 0
while attempts < max_attempts:
random_offset_x = random.uniform(-offset_range, offset_range)
random_offset_y = random.uniform(-offset_range, offset_range)
new_x = x + random_offset_x
new_y = y + random_offset_y
if not check_overlap((new_x, new_y), image_positions, threshold):
return new_x, new_y
attempts += 1
return x, y
[docs]
def search_reduction_and_clustering(numeric_data, n_neighbors, min_dist, metric, eps, min_samples, clustering, reduction_method, verbose, reduction_param=None, embedding=None, n_jobs=-1):
"""Variant of :func:`reduction_and_clustering` accepting extra reducer kwargs via ``reduction_param``.
:param numeric_data: numeric data matrix.
:param n_neighbors: UMAP ``n_neighbors`` or t-SNE perplexity (int or fraction).
:param min_dist: UMAP ``min_dist``.
:param metric: distance metric.
:param eps: DBSCAN ``eps``.
:param min_samples: DBSCAN ``min_samples`` or KMeans cluster count.
:param clustering: ``'dbscan'`` or ``'kmeans'``.
:param reduction_method: ``'umap'`` or ``'tsne'``.
:param verbose: print progress.
:param reduction_param: extra kwargs forwarded to the reducer.
:param embedding: precomputed embedding to skip fitting.
:param n_jobs: parallel worker count.
:returns: ``(embedding, labels)``.
:raises ValueError: on unsupported ``reduction_method`` or ``clustering``.
"""
from .resource_log import _guard_workers, _table_nbytes
n_jobs = _guard_workers('umap', n_jobs, _table_nbytes(numeric_data))
if isinstance(n_neighbors, float):
n_neighbors = int(n_neighbors * len(numeric_data))
if n_neighbors <= 1:
n_neighbors = 2
print(f'n_neighbors cannota be less than 2. Setting n_neighbors to {n_neighbors}')
reduction_param = reduction_param or {}
reduction_param = {k: v for k, v in reduction_param.items() if k not in ['perplexity', 'n_neighbors', 'min_dist', 'metric', 'method']}
if reduction_method == 'umap':
reducer = umap.UMAP(n_neighbors=n_neighbors, min_dist=min_dist, metric=metric, n_jobs=n_jobs, **reduction_param)
elif reduction_method == 'tsne':
reducer = TSNE(n_components=2, perplexity=n_neighbors, metric=metric, n_jobs=n_jobs, **reduction_param)
else:
raise ValueError(f"Unsupported reduction method: {reduction_method}. Supported methods are 'umap' and 'tsne'")
if embedding is None:
embedding = reducer.fit_transform(numeric_data)
if clustering == 'dbscan':
clustering_model = DBSCAN(eps=eps, min_samples=min_samples, metric=metric)
elif clustering == 'kmeans':
from sklearn.cluster import KMeans
clustering_model = KMeans(n_clusters=min_samples, random_state=_run_random_state(42))
else:
raise ValueError(f"Unsupported clustering method: {clustering}. Supported methods are 'dbscan' and 'kmeans'")
clustering_model.fit(embedding)
labels = clustering_model.labels_ if clustering == 'dbscan' else clustering_model.predict(embedding)
if verbose:
print(f'Embedding shape: {embedding.shape}')
return embedding, labels
[docs]
def load_image(image_path):
"""Load and preprocess an image.
The preprocessing is fixed to the ImageNet recipe used by
:func:`extract_features`: resize to 224x224, then normalize with the ImageNet
channel means and standard deviations. None of it is configurable.
:param image_path: path to any file PIL can open. It is forced through
``convert('RGB')``, so a 16-bit or float microscopy TIFF is downcast to
8-bit and a single-channel image is replicated across three channels
rather than rejected -- the dynamic range of a raw scientific image is
lost here, so rescale to 8-bit yourself if that matters. Aspect ratio is
not preserved: ``Resize((224, 224))`` takes both dimensions, so
non-square crops are stretched, not letterboxed. Returns a
``(1, 3, 224, 224)`` tensor with the batch axis already added, so it can
be fed to a model directly but must be concatenated, not stacked, to
batch several images.
"""
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
image = Image.open(image_path).convert('RGB')
image = transform(image).unsqueeze(0)
return image
[docs]
def check_normality(series):
"""Helper function to check if a feature is normally distributed.
This is a *failure to reject* at alpha 0.05, not evidence of normality: the
answer is ``True`` whenever the D'Agostino-Pearson test does not find
significant skew or kurtosis. Small samples therefore look normal for want of
power, and very large ones fail on deviations too small to matter -- which is
what decides whether :func:`perform_statistical_tests` sends a feature to
ANOVA or to Kruskal-Wallis.
:param series: one numeric column of observations, pooled across all groups.
``NaN`` propagates and makes the p-value ``NaN``, which compares ``False``
against alpha and so is reported as normal -- drop missing values first.
Under 8 observations ``scipy`` cannot run the skew test, returns ``NaN``,
and the feature is likewise reported as normal; a constant column
behaves the same way. Values are treated as one sample, so a strongly
bimodal feature whose groups are each normal is judged on the mixture.
"""
k2, p = stats.normaltest(series)
alpha = 0.05
if p < alpha:
return False
return True
[docs]
def random_forest_feature_importance(all_df, cluster_col='cluster'):
"""Rank features by how well they predict the cluster label.
Z-scales the numeric feature columns and fits a 100-tree
``RandomForestClassifier`` against ``cluster_col``.
:param all_df: DataFrame with the numeric feature columns and the cluster column.
:param cluster_col: Column holding the cluster label, excluded from the
features. Default ``'cluster'``.
:returns: DataFrame with ``Feature`` and ``Importance`` columns, sorted by
descending importance.
"""
numeric_features = schema.model_feature_columns(
all_df,
allow_unknown=True,
exclude=[cluster_col],
)
X = all_df[numeric_features]
y = all_df[cluster_col]
scaler = StandardScaler()
X_scaled = scaler.fit_transform(X)
model = RandomForestClassifier(n_estimators=100, random_state=_run_random_state(42))
model.fit(X_scaled, y)
feature_importances = model.feature_importances_
importance_df = pd.DataFrame({
'Feature': numeric_features,
'Importance': feature_importances
}).sort_values(by='Importance', ascending=False)
return importance_df
[docs]
def combine_results(rf_df, anova_df, kruskal_df):
"""Combine the results into a single DataFrame.
All three frames are keyed on ``Feature`` and carry exactly one row per
feature: ``rf_df`` is built from the feature list, and
:func:`perform_statistical_tests` sends each feature to *either* ANOVA or
Kruskal-Wallis, never both. Hence ``one_to_one``. A repeated ``Feature``
-- the signature of a frame with duplicated column names, or of two runs'
results concatenated by mistake -- would multiply the importance rows and
report the same feature several times as if independently ranked.
:param rf_df: random-forest results keyed uniquely by ``Feature``.
:param anova_df: ANOVA results keyed uniquely by ``Feature``.
:param kruskal_df: Kruskal-Wallis results keyed uniquely by ``Feature``.
"""
combined_df = rf_df.merge(anova_df, on='Feature', how='left',
validate='one_to_one')
combined_df = combined_df.merge(kruskal_df, on='Feature', how='left',
validate='one_to_one')
return combined_df
[docs]
def cluster_feature_analysis(all_df, cluster_col='cluster'):
"""
Perform Random Forest feature importance, ANOVA for normally distributed features,
and Kruskal-Wallis for non-normally distributed features. Combine results into a single DataFrame.
:param all_df: DataFrame holding the numeric feature columns *and*
``cluster_col``, one row per object. The same frame is passed to the
Random Forest and to the statistical tests, so the two rankings describe
the same rows. The features are selected by
``schema.model_feature_columns(..., allow_unknown=True)``, which means
stray numeric bookkeeping columns (row/column indices, object IDs) are
picked up as features unless you drop them first. Because each feature is
routed to *either* ANOVA or Kruskal-Wallis, the merged output has exactly
one of the two p-value pairs filled per row and ``NaN`` in the other.
:param cluster_col: column holding the group label, excluded from the
features and used as the Random Forest target and the grouping variable
for the tests. It needs at least two distinct labels, and every group
needs enough rows for the test to run. DBSCAN's ``-1`` noise label is not
special-cased here, so it is analysed as if it were a real cluster --
strip it with :func:`remove_noise` first if you do not want that.
Default ``'cluster'``.
"""
rf_df = random_forest_feature_importance(all_df, cluster_col)
anova_df, kruskal_df = perform_statistical_tests(all_df, cluster_col)
combined_df = combine_results(rf_df, anova_df, kruskal_df)
return combined_df
def _merge_cells_without_nucleus(adj_cell_mask: np.ndarray, nuclei_mask: np.ndarray):
"""
Relabel any cell that lacks a nucleus to the ID of an adjacent
cell that *does* contain a nucleus.
Parameters
----------
adj_cell_mask : np.ndarray
Labelled (0 = background) cell mask after all other merging steps.
nuclei_mask : np.ndarray
Labelled (0 = background) nuclei mask.
Returns
-------
np.ndarray
Updated cell mask with nucleus-free cells merged into
neighbouring nucleus-bearing cells.
"""
out = adj_cell_mask.copy()
nuc_labels = np.unique(nuclei_mask[nuclei_mask > 0])
cells_with_nuc = set()
for nuc_id in nuc_labels:
labels, counts = np.unique(adj_cell_mask[nuclei_mask == nuc_id],
return_counts=True)
keep = labels > 0
labels = labels[keep]
counts = counts[keep]
if labels.size:
cells_with_nuc.add(labels[np.argmax(counts)])
boundaries = find_boundaries(adj_cell_mask, mode="thick")
adj_map = defaultdict(set)
ys, xs = np.where(boundaries)
h, w = adj_cell_mask.shape
for y, x in zip(ys, xs):
src = adj_cell_mask[y, x]
if src == 0:
continue
for dy in (-1, 0, 1):
for dx in (-1, 0, 1):
ny, nx = y + dy, x + dx
if 0 <= ny < h and 0 <= nx < w:
dst = adj_cell_mask[ny, nx]
if dst != 0 and dst != src:
adj_map[src].add(dst)
cells_no_nuc = set(np.unique(adj_cell_mask)) - {0} - cells_with_nuc
for cell_id in cells_no_nuc:
neighbours = adj_map.get(cell_id, set()) & cells_with_nuc
if neighbours:
target = sorted(neighbours)[0]
out[out == cell_id] = target
return out.astype(np.uint16)
def _merge_cells_based_on_parasite_overlap(parasite_mask, cell_mask, nuclei_mask, organelle_mask, overlap_threshold=5, perimeter_threshold=30):
"""Merge cells that share a parasite/nucleus or a large fraction of perimeter.
Overlap and perimeter decisions use the mask's current object IDs, including
nonconsecutive labels. Cell IDs are compacted only for the returned mask;
intermediate component IDs must never be used to index the original mask.
"""
labeled_cells = cell_mask
labeled_parasites = label(parasite_mask)
labeled_nuclei = label(nuclei_mask)
num_parasites = np.max(labeled_parasites)
num_nuclei = np.max(labeled_nuclei)
for parasite_id in range(1, num_parasites + 1):
current_parasite_mask = labeled_parasites == parasite_id
overlapping_cell_labels = np.unique(labeled_cells[current_parasite_mask])
overlapping_cell_labels = overlapping_cell_labels[overlapping_cell_labels != 0]
if len(overlapping_cell_labels) > 1:
overlap_percentages = [
np.sum(current_parasite_mask & (labeled_cells == cell_label)) / np.sum(current_parasite_mask) * 100
for cell_label in overlapping_cell_labels
]
for cell_label, overlap_percentage in zip(overlapping_cell_labels, overlap_percentages):
if overlap_percentage > overlap_threshold:
first_label = overlapping_cell_labels[0]
for other_label in overlapping_cell_labels[1:]:
cell_mask[cell_mask == other_label] = first_label
for nucleus_id in range(1, num_nuclei + 1):
current_nucleus_mask = labeled_nuclei == nucleus_id
overlapping_cell_labels = np.unique(labeled_cells[current_nucleus_mask])
overlapping_cell_labels = overlapping_cell_labels[overlapping_cell_labels != 0]
if len(overlapping_cell_labels) > 1:
overlap_percentages = [
np.sum(current_nucleus_mask & (labeled_cells == cell_label)) / np.sum(current_nucleus_mask) * 100
for cell_label in overlapping_cell_labels
]
if all(overlap_percentage > overlap_threshold for overlap_percentage in overlap_percentages):
first_label = overlapping_cell_labels[0]
for other_label in overlapping_cell_labels[1:]:
cell_mask[cell_mask == other_label] = first_label
labeled_cells = cell_mask.copy()
cell_regions = regionprops(labeled_cells)
for region in cell_regions:
cell_label = region.label
cell_mask_binary = labeled_cells == cell_label
overlapping_nuclei = np.unique(nuclei_mask[cell_mask_binary])
overlapping_nuclei = overlapping_nuclei[overlapping_nuclei != 0]
if len(overlapping_nuclei) == 0:
perimeter = region.perimeter
dilated_cell = binary_dilation(
cell_mask_binary, structure=_square_footprint(3))
neighbor_cells = np.unique(labeled_cells[dilated_cell])
neighbor_cells = neighbor_cells[(neighbor_cells != 0) & (neighbor_cells != cell_label)]
shared_borders = [
np.sum((labeled_cells == neighbor_label) & dilated_cell) for neighbor_label in neighbor_cells
]
shared_border_percentages = [shared_border / perimeter * 100 for shared_border in shared_borders]
if shared_borders:
max_shared_border_index = np.argmax(shared_border_percentages)
max_shared_border_percentage = shared_border_percentages[max_shared_border_index]
if max_shared_border_percentage > perimeter_threshold:
cell_mask[labeled_cells == cell_label] = neighbor_cells[max_shared_border_index]
relabeled_cell_mask, _ = label(cell_mask, return_num=True)
return relabeled_cell_mask.astype(np.uint16)
[docs]
def process_mask_file_adjust_cell(file_name, parasite_folder, cell_folder, nuclei_folder, organelle_folder=None, overlap_threshold=5, perimeter_threshold=30, *, output_folder=None):
"""Load one triple of parasite/cell/nuclei masks, merge cells in place, and return the elapsed time.
:param file_name: mask file name (must exist in all folders).
:param parasite_folder: folder of parasite masks.
:param cell_folder: folder of cell masks (overwritten in place). The
adjusted mask replaces the old one atomically, so a run killed
during the write leaves the previous whole mask, never a truncated
one.
:param nuclei_folder: folder of nuclei masks.
:param organelle_folder: optional folder of organelle masks.
:param overlap_threshold: fractional overlap threshold used by the merger.
:param perimeter_threshold: shared-perimeter threshold used by the merger.
:param output_folder: optional separate destination for adjusted masks.
None retains in-place adjustment. An explicit destination must differ
from every source mask folder, including through directory symlinks.
:returns: elapsed seconds.
:raises ValueError: if the matching cell or nuclei mask file is missing,
or a mask file holds pickled objects: masks are plain arrays, and
nothing is unpickled.
"""
start = time.perf_counter()
parasite_path = os.path.join(parasite_folder, file_name)
cell_path = os.path.join(cell_folder, file_name)
nuclei_path = os.path.join(nuclei_folder, file_name)
if output_folder is not None and any(
os.path.realpath(output_folder) == os.path.realpath(folder)
for folder in (parasite_folder, cell_folder, nuclei_folder, organelle_folder)
if folder is not None):
raise ValueError('The adjusted-mask output folder must differ from all source folders')
if not (os.path.exists(cell_path) and os.path.exists(nuclei_path)):
raise ValueError(f"Corresponding cell or nuclei mask file for {file_name} not found.")
parasite_mask = np.load(parasite_path, allow_pickle=False)
cell_mask = np.load(cell_path, allow_pickle=False)
nuclei_mask = np.load(nuclei_path, allow_pickle=False)
organelle_mask = None
if organelle_folder is not None:
organelle_path = os.path.join(organelle_folder, file_name)
if os.path.exists(organelle_path):
organelle_mask = np.load(organelle_path, allow_pickle=False)
merged_cell_mask = _merge_cells_based_on_parasite_overlap(parasite_mask, cell_mask, nuclei_mask, organelle_mask, overlap_threshold, perimeter_threshold)
from .io import _save_array_atomic
output_path = cell_path
if output_folder is not None:
os.makedirs(output_folder, exist_ok=True)
output_path = os.path.join(output_folder, file_name)
_save_array_atomic(output_path, merged_cell_mask)
end = time.perf_counter()
return end - start
#: The record, beside the cell masks, of which ones :func:`adjust_cell_masks`
#: already adjusted in place, as ``{file name: sha256 of the adjusted file}``.
ADJUSTED_CELLS_LEDGER = '.cell_masks_adjusted.json'
def _file_sha256(path):
"""The SHA-256 of a file's bytes, or None when it cannot be read.
:param path: the file.
:returns: the hex digest, or None.
"""
import hashlib
digest = hashlib.sha256()
try:
with open(path, 'rb') as handle:
for block in iter(lambda: handle.read(1 << 20), b''):
digest.update(block)
except OSError:
return None
return digest.hexdigest()
def _read_adjusted_cells(cell_folder):
"""The in-place adjustment record for ``cell_folder``.
:param cell_folder: the folder of cell masks.
:returns: ``{file name: sha256}``; empty when there is no record or it
cannot be read, which means every mask is adjusted, as before the
record existed.
"""
import json
path = os.path.join(cell_folder, ADJUSTED_CELLS_LEDGER)
try:
with open(path, encoding='utf-8') as handle:
record = json.load(handle)
except (OSError, ValueError):
return {}
if not isinstance(record, dict):
return {}
return {str(name): str(digest) for name, digest in record.items()}
def _write_adjusted_cells(cell_folder, record):
"""Replace the in-place adjustment record atomically.
:param cell_folder: the folder of cell masks.
:param record: ``{file name: sha256}`` to write.
"""
import json
import tempfile
path = os.path.join(cell_folder, ADJUSTED_CELLS_LEDGER)
fd, temporary = tempfile.mkstemp(prefix='.spacr_tmp_', suffix='.json',
dir=cell_folder)
try:
with os.fdopen(fd, 'w', encoding='utf-8') as handle:
json.dump(record, handle, indent=0, sort_keys=True)
os.replace(temporary, path)
finally:
if os.path.exists(temporary):
os.remove(temporary)
[docs]
def adjust_cell_masks(parasite_folder, cell_folder, nuclei_folder, organelle_folder=None, overlap_threshold=5, perimeter_threshold=30, n_jobs=None, *, output_folder=None):
"""Run :func:`process_mask_file_adjust_cell` in parallel across matching mask files.
:param parasite_folder: folder of parasite masks.
:param cell_folder: folder of cell masks (overwritten in place).
:param nuclei_folder: folder of nuclei masks.
:param organelle_folder: optional folder of organelle masks.
:param overlap_threshold: fractional overlap threshold used by the merger.
:param perimeter_threshold: shared-perimeter threshold used by the merger.
:param n_jobs: worker count; ``None`` defaults to ``cpu_count() - 2`` and
values below two run inline without starting a child process.
:param output_folder: optional separate folder for adjusted masks. None
preserves the historical in-place behavior. A separate folder keeps
all source masks byte-identical, and every selected field is rebuilt
from its source on each invocation, including after interrupted work.
In place, each adjusted mask's SHA-256 is recorded in
:data:`ADJUSTED_CELLS_LEDGER` as it lands, and a mask whose bytes
still match its record is left alone on the next run: adjusting an
adjusted mask merges it again, so a re-run used to change the result
every time. A cell mask segmented again no longer matches and is
adjusted afresh.
:returns: None.
:raises ValueError: if the three folders contain different numbers of files
or mismatched filenames, or a mask is truncated, nonnumeric, empty,
not two-dimensional or has different dimensions from its partners.
Available organelle masks are checked too. Header-only validation of
every field finishes before any mask is changed or workers are started.
An explicit output folder must differ from every source folder and
must not contain masks outside the selected field set.
"""
from .io import _listdir_visible
parasite_files = sorted([f for f in _listdir_visible(parasite_folder) if f.endswith('.npy')])
cell_files = sorted([f for f in _listdir_visible(cell_folder) if f.endswith('.npy')])
nuclei_files = sorted([f for f in _listdir_visible(nuclei_folder) if f.endswith('.npy')])
if not (len(parasite_files) == len(cell_files) == len(nuclei_files)):
raise ValueError("The number of files in the folders do not match.")
if parasite_files != cell_files or parasite_files != nuclei_files:
groups = [set(parasite_files), set(cell_files), set(nuclei_files)]
unmatched = set.union(*groups) - set.intersection(*groups)
raise ValueError(
"Mask filenames do not match across parasite, cell and nuclei folders: "
f"{', '.join(sorted(unmatched)[:5])}. No cell masks were changed.")
if organelle_folder is not None and os.path.exists(organelle_folder):
organelle_files = sorted([f for f in _listdir_visible(organelle_folder) if f.endswith('.npy')])
if len(organelle_files) != len(parasite_files):
print(f'Warning: organelle mask count ({len(organelle_files)}) does not match other masks ({len(parasite_files)}). Organelle masks will be loaded per-file where available.')
else:
organelle_folder = None
if output_folder is not None:
if any(os.path.realpath(output_folder) == os.path.realpath(folder)
for folder in (parasite_folder, cell_folder, nuclei_folder, organelle_folder)
if folder is not None):
raise ValueError('The adjusted-mask output folder must differ from all source folders')
if os.path.isdir(output_folder):
extra = {name for name in _listdir_visible(output_folder)
if name.endswith('.npy')} - set(parasite_files)
if extra:
raise ValueError(f'Adjusted-mask output folder contains unrelated fields: '
f'{", ".join(sorted(extra)[:5])}')
from .cancellation import checkpoint
from .resume import read_npy_header
for name in parasite_files:
checkpoint()
paths = [os.path.join(folder, name)
for folder in (parasite_folder, cell_folder, nuclei_folder)]
if organelle_folder is not None:
candidate = os.path.join(organelle_folder, name)
if os.path.exists(candidate):
paths.append(candidate)
expected_shape = None
for path in paths:
header = read_npy_header(path)
shape = header['shape']
if (len(shape) != 2 or min(shape) <= 0
or header['expected_bytes'] is None
or header['actual_bytes'] < header['expected_bytes']):
raise ValueError(f'Invalid or incomplete two-dimensional mask: {path}')
if expected_shape is not None and shape != expected_shape:
raise ValueError(f'Mask dimensions do not match for {name}: '
f'{path} has {shape}, expected {expected_shape}')
expected_shape = shape
record = None
if output_folder is None:
record = _read_adjusted_cells(cell_folder)
already = {name for name in parasite_files
if name in record and record[name] == _file_sha256(
os.path.join(cell_folder, name))}
if already:
print(f'{len(already)} of {len(parasite_files)} cell masks were '
'already adjusted by an earlier run and are left as they are.')
record = {name: digest for name, digest in record.items()
if name in already}
parasite_files = [name for name in parasite_files if name not in already]
if not parasite_files:
return
if n_jobs is None:
n_jobs = max(1, cpu_count() - 2)
else:
n_jobs = max(1, int(n_jobs))
from .resource_log import _array_file_nbytes, _guard_workers
n_jobs = _guard_workers('adjust_masks', n_jobs, _array_file_nbytes(
os.path.join(cell_folder, parasite_files[0])))
time_ls = []
files_to_process = len(parasite_files)
process_fn = partial(process_mask_file_adjust_cell,
parasite_folder=parasite_folder,
cell_folder=cell_folder,
nuclei_folder=nuclei_folder,
organelle_folder=organelle_folder,
overlap_threshold=overlap_threshold,
perimeter_threshold=perimeter_threshold,
**({'output_folder': output_folder} if output_folder is not None else {}))
if n_jobs == 1:
durations = map(process_fn, parasite_files)
for i, (name, duration) in enumerate(zip(parasite_files, durations), 1):
time_ls.append(duration)
if record is not None:
record[name] = _file_sha256(os.path.join(cell_folder, name))
_write_adjusted_cells(cell_folder, record)
print_progress(i, files_to_process, n_jobs=n_jobs, time_ls=time_ls, batch_size=None, operation_type='adjust_cell_masks')
return
with Pool(n_jobs) as pool:
for i, (name, duration) in enumerate(
zip(parasite_files, pool.imap(process_fn, parasite_files)), 1):
time_ls.append(duration)
if record is not None:
record[name] = _file_sha256(os.path.join(cell_folder, name))
_write_adjusted_cells(cell_folder, record)
print_progress(i, files_to_process, n_jobs=n_jobs, time_ls=time_ls,
batch_size=None,
operation_type='adjust_cell_masks')
[docs]
def process_masks(mask_folder, image_folder, channel, batch_size=50, n_clusters=2, plot=False):
"""Cluster object morphology/intensity across a mask folder and keep the largest cluster in place.
:param mask_folder: folder of ``.npy`` masks.
:param image_folder: matching folder of ``.npy`` intensity images.
:param channel: channel index used for intensity measurements.
:param batch_size: number of files to load per batch.
:param n_clusters: number of KMeans clusters.
:param plot: show a PCA scatter of the clustered objects.
:returns: None.
"""
def read_files_in_batches(folder, batch_size=50):
"""Yield sorted lists of ``.npy`` filenames from ``folder`` in chunks of ``batch_size``."""
files = [f for f in os.listdir(folder) if f.endswith('.npy')]
files.sort()
for i in range(0, len(files), batch_size):
yield files[i:i + batch_size]
def measure_morphology_and_intensity(mask, image):
"""Return a list of dicts with area/mean_intensity/perimeter/eccentricity per labeled region."""
properties = measure.regionprops(mask, intensity_image=image)
properties_list = [{'area': p.area, 'mean_intensity': p.intensity_mean, 'perimeter': p.perimeter, 'eccentricity': p.eccentricity} for p in properties]
return properties_list
def cluster_objects(properties, n_clusters=2):
"""Return a fitted ``KMeans`` object clustering the property dicts into ``n_clusters`` groups."""
data = np.array([[p['area'], p['mean_intensity'], p['perimeter'], p['eccentricity']] for p in properties])
kmeans = KMeans(n_clusters=n_clusters, random_state=0).fit(data)
return kmeans
def remove_objects_not_in_largest_cluster(mask, labels, largest_cluster_label):
"""Return ``mask`` with all labeled regions removed except those in ``largest_cluster_label``."""
cleaned_mask = np.zeros_like(mask)
for idx, region in enumerate(measure.regionprops(mask)):
if labels[idx] == largest_cluster_label:
cleaned_mask[mask == region.label] = region.label
return cleaned_mask
def plot_clusters(properties, labels):
"""Show a 2-D PCA scatter of the property vectors colored by cluster label."""
data = np.array([[p['area'], p['mean_intensity'], p['perimeter'], p['eccentricity']] for p in properties])
pca = PCA(n_components=2)
data_2d = pca.fit_transform(data)
plt.scatter(data_2d[:, 0], data_2d[:, 1], c=labels, cmap='viridis')
plt.xlabel('PCA Component 1')
plt.ylabel('PCA Component 2')
plt.title('Object Clustering')
plt.show()
all_properties = []
for batch in read_files_in_batches(mask_folder, batch_size):
mask_files = [os.path.join(mask_folder, file) for file in batch]
image_files = [os.path.join(image_folder, file) for file in batch]
masks = [np.load(file) for file in mask_files]
images = [np.load(file)[:, :, channel] for file in image_files]
for i, mask in enumerate(masks):
image = images[i]
properties = measure_morphology_and_intensity(mask, image)
all_properties.extend(properties)
kmeans = cluster_objects(all_properties, n_clusters)
labels = kmeans.labels_
if plot:
plot_clusters(all_properties, labels)
label_index = 0
for batch in read_files_in_batches(mask_folder, batch_size):
mask_files = [os.path.join(mask_folder, file) for file in batch]
masks = [np.load(file) for file in mask_files]
for i, mask in enumerate(masks):
batch_properties = measure_morphology_and_intensity(mask, mask)
if not batch_properties:
continue
batch_labels = labels[label_index:label_index + len(batch_properties)]
largest_cluster_label = np.bincount(batch_labels).argmax()
cleaned_mask = remove_objects_not_in_largest_cluster(mask, batch_labels, largest_cluster_label)
np.save(mask_files[i], cleaned_mask)
label_index += len(batch_properties)
[docs]
def process_vision_results(df, threshold=0.5):
"""Split image paths into well identifiers and binarize the ``pred`` column.
:param df: DataFrame with ``path`` and ``pred`` columns.
:param threshold: cutoff used to derive ``cv_predictions``.
:returns: enriched DataFrame with ``plateID``, ``rowID``, ``columnID``, ``fieldID``, ``prc``, ``cv_predictions``.
"""
mapped_values = df['path'].apply(lambda x: _map_wells_png(x))
df['plateID'] = mapped_values.apply(lambda x: x[0])
df['rowID'] = mapped_values.apply(lambda x: x[1])
df['columnID'] = mapped_values.apply(lambda x: x[2])
df['fieldID'] = mapped_values.apply(lambda x: x[3])
df['object'] = (df['path'].str.rsplit('/', n=1).str[-1]
.str.split('.').str[0].str.rsplit('_', n=1).str[-1])
df['prc'] = schema.compose_prc_column(df)
df['cv_predictions'] = (df['pred'] >= threshold).astype(int)
return df
[docs]
def feature_folder_name(channel_of_interest) -> str:
"""A folder name for one feature selection. Safe on every filesystem.
``None`` is ``all_features``; a channel is ``channel_1``; several are
``channels_1_2``; ``morphology`` is itself; a mixture joins them in the
order given; and a free-text filter is slugified, because a user may
reasonably filter on ``mean_intensity`` and a column fragment can carry
anything.
:param channel_of_interest: feature selection accepted by
:func:`feature_selection`.
"""
selection = feature_selection(channel_of_interest)
if selection is None:
return 'all_features'
def _one(member):
"""Return one filesystem-safe selection-member slug."""
return re.sub(r'[^0-9A-Za-z]+', '_', str(member)).strip('_') or 'x'
if isinstance(selection, int):
return f"channel_{selection}"
if not isinstance(selection, list):
return _one(selection)
if all(isinstance(member, int) for member in selection):
return "channels_" + "_".join(str(m) for m in selection)
return "_".join(f"channel_{m}" if isinstance(m, int) else _one(m)
for m in selection)
[docs]
def get_ml_results_paths(src, model_type='xgboost', channel_of_interest=1):
"""Return the standard set of ML output paths for the given model and channel selection.
:param src: experiment root.
:param model_type: model identifier (used in the results folder name).
:param channel_of_interest: int, list, ``'morphology'``, or ``None`` (aliased to ``all_features``).
:returns: 10-tuple of paths ``(data, permutation, feature_importance, model_metrics,
permutation_fig, feature_importance_fig, shap_fig, plate_heatmap, settings, ml_features)``.
:raises ValueError: if ``channel_of_interest`` has an unsupported type.
"""
feature_string = feature_folder_name(channel_of_interest)
res_fldr = os.path.join(src, 'results', model_type, feature_string)
print(f'Saving results to {res_fldr}')
os.makedirs(res_fldr, exist_ok=True)
data_path = os.path.join(res_fldr, 'results.csv')
permutation_path = os.path.join(res_fldr, 'permutation.csv')
feature_importance_path = os.path.join(res_fldr, 'feature_importance.csv')
model_metricks_path = os.path.join(res_fldr, f'{model_type}_model.csv')
permutation_fig_path = os.path.join(res_fldr, 'permutation.pdf')
feature_importance_fig_path = os.path.join(res_fldr, 'feature_importance.pdf')
shap_fig_path = os.path.join(res_fldr, 'shap.pdf')
plate_heatmap_path = os.path.join(res_fldr, 'plate_heatmap.pdf')
settings_csv = os.path.join(res_fldr, 'ml_settings.csv')
ml_features = os.path.join(res_fldr, 'ml_features.csv')
return data_path, permutation_path, feature_importance_path, model_metricks_path, permutation_fig_path, feature_importance_fig_path, shap_fig_path, plate_heatmap_path, settings_csv, ml_features
[docs]
def augment_image(image):
"""Return a list of PIL images covering 4 rotations x 2 horizontal reflections of ``image``.
The 8 outputs are the dihedral group of the square and include the unmodified
original as element 0, so the list is an 8x expansion, not 8 *extra* images.
Ordering is rotation-major -- ``[0deg, 0deg flipped, 90deg, 90deg flipped, ...]``
-- which matters if you are keeping a parallel list of labels.
:param image: a PIL image or a numpy array. Arrays are used as-is; PIL images
are converted first. A 2-D grayscale input is expanded to 3 channels via
``cv2.cvtColor``, so every result is RGB even when the input was not --
channel count is not preserved. The rotations go through ``cv2``, which
expects ``uint8`` or another OpenCV-supported dtype; the final
``Image.fromarray`` likewise rejects the float or 16-bit arrays typical
of raw microscopy, so convert to 8-bit before calling. Because 90-degree
rotations swap height and width, a non-square input yields images of two
different shapes in the same list.
"""
augmented_images = []
if isinstance(image, Image.Image):
image = np.array(image)
if len(image.shape) == 2:
image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)
transformations = [
None,
cv2.ROTATE_90_CLOCKWISE,
cv2.ROTATE_180,
cv2.ROTATE_90_COUNTERCLOCKWISE
]
for transform in transformations:
if transform is not None:
rotated = cv2.rotate(image, transform)
else:
rotated = image
augmented_images.append(rotated)
flipped = cv2.flip(rotated, 1)
augmented_images.append(flipped)
augmented_images = [Image.fromarray(img) for img in augmented_images]
return augmented_images
[docs]
def augment_dataset(dataset, is_grayscale=False):
"""Expand ``dataset`` by 8x through rotation and horizontal reflection of every image tensor.
:param dataset: iterable of ``(tensor, label, filename)``.
:param is_grayscale: informational flag (retained for API compatibility).
:returns: list of augmented ``(tensor, label, filename)`` tuples.
:raises TypeError: if an image is not a ``torch.Tensor``.
"""
augmented_dataset = []
for img, label, filename in dataset:
augmented_images = []
if not isinstance(img, torch.Tensor):
raise TypeError(f"Expected torch.Tensor, got {type(img)}")
angles = [0, 90, 180, 270]
for angle in angles:
rotated = torchvision.transforms.functional.rotate(img, angle)
augmented_images.append(rotated)
flipped = torchvision.transforms.functional.hflip(rotated)
augmented_images.append(flipped)
for aug_img in augmented_images:
augmented_dataset.append((aug_img, label, filename))
return augmented_dataset
[docs]
def convert_and_relabel_masks(folder_path):
"""
Converts all int64 npy masks in a folder to uint16 with relabeling to ensure all labels are retained.
Parameters:
- folder_path (str): The path to the folder containing int64 npy mask files.
Returns:
- None
:param folder_path: directory containing ``.npy`` masks to inspect and
convert in place.
"""
files = [f for f in os.listdir(folder_path) if f.endswith('.npy')]
for file in files:
file_path = os.path.join(folder_path, file)
mask = np.load(file_path)
if mask.dtype != np.int64:
print(f"Skipping {file} as it is not int64.")
continue
unique_labels = np.unique(mask)
if unique_labels.max() > 65535:
print(f"Warning: The mask in {file} contains values that exceed the uint16 range and will be relabeled.")
relabeled_mask = measure.label(mask, background=0)
unique_relabeled = np.unique(relabeled_mask)
if unique_relabeled.max() > 65535:
print(f"Error: Relabeling failed for {file} as it still contains values that exceed the uint16 range.")
continue
relabeled_mask = relabeled_mask.astype(np.uint16)
np.save(file_path, relabeled_mask)
print(f"Converted {file} and saved as uint16_{file}")
[docs]
def correct_masks(src):
"""Convert cell masks under ``src/masks/cell_mask_stack`` to uint16 and re-stack arrays.
Relabels masks so they fit in ``uint16`` and then re-concatenates the four
array folders under ``src`` in the layout expected downstream.
:param src: Root folder of a spacr run containing a ``masks/`` subfolder.
:returns: None.
"""
from .io import _load_and_concatenate_arrays
cell_path = os.path.join(src,'masks', 'cell_mask_stack')
convert_and_relabel_masks(cell_path)
_load_and_concatenate_arrays(src, [0,1,2,3], 1, 0, 2)
[docs]
def count_reads_in_fastq(fastq_file):
"""Return the number of reads in a gzipped FASTQ file.
Counts total lines and divides by four (the FASTQ record length).
:param fastq_file: Path to a ``.fastq.gz`` file.
:returns: Integer read count.
"""
count = 0
with gzip.open(fastq_file, "rt") as f:
for _ in f:
count += 1
return count // 4
[docs]
def get_cuda_version():
"""Return the installed CUDA toolkit version as a digit-only string, or ``None``.
Parses the ``nvcc --version`` output; the dots are stripped so ``11.8`` becomes ``"118"``.
:returns: Version string without dots, or ``None`` if ``nvcc`` is missing or fails.
"""
try:
output = subprocess.check_output(['nvcc', '--version'], stderr=subprocess.STDOUT).decode('utf-8')
if 'release' in output:
return output.split('release ')[1].split(',')[0].replace('.', '')
except (subprocess.CalledProcessError, FileNotFoundError):
return None
[docs]
def all_elements_match(list1, list2):
"""Return ``True`` if every element of ``list1`` is contained in ``list2``.
:param list1: iterable of items to test.
:param list2: iterable acting as the reference set.
:returns: ``True`` when ``list1`` is a subset of ``list2``, else ``False``.
"""
return all(element in list2 for element in list1)
[docs]
def prepare_batch_for_segmentation(batch):
"""Cast a batch to ``float32`` and per-image max-normalize any image whose max exceeds 1.
:param batch: ``(N, ...)`` numpy array of images.
:returns: The same array cast to ``float32`` with each image scaled to ``[0, 1]``.
"""
if batch.dtype != np.float32:
batch = batch.astype(np.float32)
for i in range(batch.shape[0]):
if batch[i].max() > 1:
batch[i] = batch[i] / batch[i].max()
return batch
[docs]
def check_index(df, elements=5, split_char='_'):
"""Validate that every index label in ``df`` splits into ``elements`` parts on ``split_char``.
:param df: DataFrame whose index labels are compound identifiers.
:param elements: Expected number of parts after splitting. Default ``5``.
:param split_char: Delimiter used to split each index label. Default ``'_'``.
:returns: None.
:raises ValueError: if any index label does not split into ``elements`` parts.
"""
problematic_indices = []
for idx in df.index:
parts = str(idx).split(split_char)
if len(parts) != elements:
problematic_indices.append(idx)
if problematic_indices:
print("Indices that cannot be separated into 5 parts:")
for idx in problematic_indices:
print(idx)
raise ValueError(f"Found {len(problematic_indices)} problematic indices that do not split into {elements} parts.")
[docs]
def map_condition(col_value, neg='c1', pos='c2', mix='c3'):
"""Map a column-ID value to one of ``'neg'``, ``'pos'``, ``'mix'``, or ``'screen'``.
:param col_value: Column identifier from the plate metadata.
:param neg: Column ID that corresponds to negative controls. Default ``'c1'``.
:param pos: Column ID that corresponds to positive controls. Default ``'c2'``.
:param mix: Column ID that corresponds to mixed controls. Default ``'c3'``.
:returns: Condition label; any unlisted column returns ``'screen'``.
"""
if col_value == neg:
return 'neg'
elif col_value == pos:
return 'pos'
elif col_value == mix:
return 'mix'
else:
return 'screen'
[docs]
def download_models(repo_id="einarolafsson/models", retries=5, delay=5):
"""
Downloads all model files from Hugging Face and stores them in the `resources/models` directory
within the installed `spacr` package.
Args:
repo_id (str): The repository ID on Hugging Face (default is 'einarolafsson/models').
retries (int): Number of retry attempts in case of failure.
delay (int): Delay in seconds between retries.
Returns:
str: The local path to the downloaded models.
"""
package_dir = os.path.dirname(spacr_path)
local_dir = os.path.join(package_dir, 'resources', 'models')
if not os.path.exists(local_dir):
os.makedirs(local_dir)
elif len(os.listdir(local_dir)) > 0:
return local_dir
attempt = 0
while attempt < retries:
try:
files = list_repo_files(repo_id, repo_type="dataset")
print(f"Files in repository: {files}")
for file_name in files:
for download_attempt in range(retries):
try:
url = f"https://huggingface.co/datasets/{repo_id}/resolve/main/{file_name}?download=true"
print(f"Downloading file from: {url}")
response = requests.get(url, stream=True)
print(f"HTTP response status: {response.status_code}")
response.raise_for_status()
local_file_path = os.path.join(local_dir, os.path.basename(file_name))
with open(local_file_path, 'wb') as file:
for chunk in response.iter_content(chunk_size=8192):
file.write(chunk)
print(f"Downloaded model file: {file_name} to {local_file_path}")
break
except (requests.HTTPError, requests.Timeout) as e:
print(f"Error downloading {file_name}: {e}. Retrying in {delay} seconds...")
time.sleep(delay)
else:
raise Exception(f"Failed to download {file_name} after multiple attempts.")
return local_dir
except (requests.HTTPError, requests.Timeout) as e:
print(f"Error downloading files: {e}. Retrying in {delay} seconds...")
attempt += 1
time.sleep(delay)
raise Exception("Failed to download model files after multiple attempts.")
[docs]
def generate_cytoplasm_mask(nucleus_mask, cell_mask):
"""
Generates a cytoplasm mask from nucleus and cell masks.
Parameters:
- nucleus_mask (np.array): Binary or segmented mask of the nucleus (non-zero values represent nucleus).
- cell_mask (np.array): Binary or segmented mask of the whole cell (non-zero values represent cell).
Returns:
- cytoplasm_mask (np.array): Copy of cell_mask with nucleus pixels set to 0, keeping the cell labels elsewhere (pathogens are not considered).
:param nucleus_mask: nucleus mask whose nonzero pixels are excluded.
:param cell_mask: labeled cell mask copied into the cytoplasm result.
"""
nucleus_mask = np.array(nucleus_mask)
cell_mask = np.array(cell_mask)
cytoplasm_mask = np.where(nucleus_mask != 0, 0, cell_mask)
return cytoplasm_mask
[docs]
def add_column_to_database(settings):
"""
Adds a new column to the database table by matching on a common column from the DataFrame.
If the column already exists in the database, it adds the column with a suffix.
NaN values will remain as NULL in the database.
Parameters:
settings (dict): A dictionary containing the following keys:
csv_path (str): Path to the CSV file with the data to be added.
db_path (str): Path to the SQLite database (or connection string for other databases).
table_name (str): The name of the table in the database.
update_column (str): The name of the new column in the DataFrame to add to the database.
match_column (str): The common column used to match rows.
Returns:
None
"""
df = tabular.read_table(settings['csv_path'], report=None)
if (df[settings['update_column']] == 0).any():
print("Replacing all 0 values with 2 in the update column.")
df[settings['update_column']] = df[settings['update_column']].replace(0, 2)
conn = sqlite3.connect(settings['db_path'], timeout=30)
cursor = conn.cursor()
cursor.execute(f"PRAGMA table_info({settings['table_name']})")
columns_in_db = [col[1] for col in cursor.fetchall()]
if settings['update_column'] in columns_in_db:
suffix = 1
new_column_name = f"{settings['update_column']}_{suffix}"
while new_column_name in columns_in_db:
suffix += 1
new_column_name = f"{settings['update_column']}_{suffix}"
print(f"Column '{settings['update_column']}' already exists. Using new column name: '{new_column_name}'")
else:
new_column_name = settings['update_column']
cursor.execute(f"ALTER TABLE {settings['table_name']} ADD COLUMN {new_column_name} INTEGER")
print(f"Added new column '{new_column_name}' to the table '{settings['table_name']}'.")
for index, row in df.iterrows():
value_to_update = row[settings['update_column']]
match_value = row[settings['match_column']]
if pd.isna(value_to_update):
value_to_update = None
query = f"""
UPDATE {settings['table_name']}
SET {new_column_name} = ?
WHERE {settings['match_column']} = ?
"""
cursor.execute(query, (value_to_update, match_value))
conn.commit()
conn.close()
print(f"Updated '{new_column_name}' in '{settings['table_name']}' using '{settings['match_column']}'.")
[docs]
def fill_holes_in_mask(mask):
"""Fill the holes inside each object of a label mask, keeping every id.
Delegates to :func:`spacr.qt.mask_engine.fill_label_holes`, the one hole
filler for label images. This used to run ``ndimage.label`` over the
mask first, which made every pair of touching objects one object: with
Cellpose ``fill_in`` on (its default in Apply), a field of 74 adjacent
cells was saved as 8. Now no object is merged or renumbered,
and a hole takes the id of the object that encloses it.
Args:
mask (np.ndarray): A labeled mask where each object has a unique integer value.
A boolean mask is labelled by connectivity first.
Returns:
np.ndarray: The mask with holes filled and the original labels preserved,
in the input's dtype.
"""
from .qt.mask_engine import fill_label_holes
return fill_label_holes(mask)
[docs]
def control_filelist(folder, mode='columnID', values=None):
"""Return filenames in ``folder`` whose row or column ID matches one of ``values``.
The filename is split on ``_`` and the second token is inspected: characters
after the first (``mode='columnID'``) or the leading character
(``mode='rowID'``) are matched against ``values``.
:param folder: Directory to scan.
:param mode: ``'columnID'`` matches trailing digits, ``'rowID'`` matches leading letter.
Default ``'columnID'``.
:param values: Iterable of allowed ID strings. Defaults to ``['01', '02']``.
:returns: List of matching filenames.
"""
if values is None:
values = ['01','02']
files = os.listdir(folder)
if mode == 'columnID':
filtered_files = [file for file in files if file.split('_')[1][1:] in values]
if mode == 'rowID':
filtered_files = [file for file in files if file.split('_')[1][:1] in values]
return filtered_files
from .database_schema import (
DB_COLUMN_RENAMES,
DB_COLUMN_RENAME_PATTERNS as _DB_COLUMN_RENAME_PATTERNS,
canonical_column_name,
)
DB_COLUMN_RENAME_PATTERNS = _DB_COLUMN_RENAME_PATTERNS
[docs]
def canonicalize_measurement_columns(df):
"""Rename legacy column spellings on an in-memory measurement frame.
The DataFrame counterpart of :func:`rename_columns_in_db`, for frames that
did not come from a spaCR database and so never passed through it — a CSV
exported by an older release, or a frame a user assembled themselves.
Follows the same never-destructive rule: a rename whose target is already
present is skipped, so a frame carrying both spellings keeps both rather
than losing one to a silently dropped duplicate. The rule itself lives in
:func:`spacr.schema.canonical_rename_plan`, which this and
``schema.canonicalise_columns`` both call so the two frame canonicalisers
cannot drift apart again — and which folds case, because these frames are
written with ``to_sql`` and SQLite compares identifiers
case-insensitively.
:param df: A measurement DataFrame.
:returns: ``df`` with legacy column names replaced (a copy is not made;
the frame is renamed in place and returned).
"""
from .schema import canonical_rename_plan
mapping = canonical_rename_plan(df.columns)
if mapping:
df.columns = [mapping.get(name, name) for name in df.columns]
return df
[docs]
def rename_columns_in_db(db_path):
"""Rename legacy column spellings across every table in a SQLite database.
Applies :data:`DB_COLUMN_RENAMES` — the plate-metadata names — and then
:data:`DB_COLUMN_RENAME_PATTERNS` — the two feature families that were
spelled inconsistently — to every user table. A rename is skipped when the
target name already exists in that table, which gives three properties
worth relying on:
* **Idempotent.** After a rename the legacy name is gone, so a second run
finds nothing to do. Running it on every read is therefore free after the
first.
* **Never destructive.** A table that somehow carries *both* spellings —
say ``time_id`` and ``timeID`` — keeps both, untouched. Neither column is
dropped and nothing raises; the readers accept either spelling, so the
data stays reachable and a human can decide which one is authoritative.
Dropping or overwriting one of them here would destroy data to tidy a
name, which is never the right trade.
* **All or nothing.** SQLite's DDL *is* transactional, but Python's sqlite3
driver only opens an implicit transaction for DML (INSERT/UPDATE/DELETE/
REPLACE) — an ``ALTER TABLE`` runs in autocommit and lands immediately.
So the previous version, which relied on a trailing ``con.commit()``,
left a database half-migrated when a later rename raised. The
transaction is opened explicitly here and rolled back on any error, and
the connection is closed in a ``finally``.
A partial migration would not corrupt anything — each rename is
independently valid and the next read finishes the job — but "the schema
changed and then the call raised" is not a state a user should have to
reason about.
:param db_path: Path to the SQLite database file to update in place.
:returns: The list of ``(table, old, new)`` renames performed.
"""
from .database_schema import repair_legacy_columns
renamed = list(repair_legacy_columns(db_path))
metadata = [entry for entry in renamed if entry[1] in DB_COLUMN_RENAMES]
features = [entry for entry in renamed if entry[1] not in DB_COLUMN_RENAMES]
for table, old, new in metadata:
print(f"Renamed `{table}`.`{old}` → `{new}`")
if features:
by_table = {}
for table, old, new in features:
by_table.setdefault(table, []).append((old, new))
for table, pairs in by_table.items():
old, new = pairs[0]
print(f"Renamed {len(pairs)} legacy feature column(s) in `{table}` "
f"to the canonical spelling, e.g. `{old}` → `{new}`")
return renamed
#: Both spellings of the timepoint column. ``timeID`` is canonical; ``time_id``
#: is what ``filepaths_to_database`` wrote into ``png_list`` before the two were
#: unified, and survives in databases written by those releases until
#: :func:`rename_columns_in_db` migrates them.
TIME_COLUMN_ALIASES = ('timeID', 'time_id')
def _time_column(columns):
"""Return whichever timepoint spelling ``columns`` carries, or ``None``."""
columns = set(columns)
for name in TIME_COLUMN_ALIASES:
if name in columns:
return name
return None
[docs]
def group_feature_class(df, feature_groups=None, name='compartment'):
"""Add a column tagging each feature with its compartment (or other group) label.
Matches feature names against the tokens in ``feature_groups`` and stores the
result in a new column ``name``. When ``name == 'channel'``, unmatched
features are relabeled ``'morphology'``.
:param df: DataFrame with a ``feature`` column.
:param feature_groups: Iterable of substrings/regex tokens to look for in each
feature name. Defaults to ``['cell', 'cytoplasm', 'nucleus', 'pathogen']``.
:param name: Name of the column added to ``df``. Default ``'compartment'``.
:returns: ``df`` with the new group column populated.
"""
if feature_groups is None:
feature_groups = ['cell', 'cytoplasm', 'nucleus', 'pathogen']
def find_feature_class(feature, compartments):
"""Return the group label(s) matched in ``feature`` — joined with '-' when more than one hits."""
matches = [compartment for compartment in compartments if re.search(compartment, feature)]
if len(matches) > 1:
return '-'.join(matches)
elif matches:
return matches[0]
else:
return None
df[name] = pd.Series(
(find_feature_class(feature, feature_groups)
for feature in df['feature']),
index=df.index,
dtype=object,
)
if name == 'channel':
df['channel'] = df['channel'].fillna('morphology')
return df
[docs]
def cleanup_pipeline_folders(src, keep_intermediate=False, keep_original=False,
verbose=True):
"""Delete the intermediate mask-pipeline folders once ``merged/`` is built.
By default spaCR keeps only ``merged/`` (the concatenated image+mask arrays
that Measure reads). This removes ``stack/`` + ``masks/`` (their data is
embedded in ``merged/`` and object labels are recorded in the database) and
the raw ``orig/`` backup, unless the caller opts to keep them.
Heavily guarded so it never destroys un-merged data: ``stack/`` + ``masks/``
are only removed when ``merged/`` is non-empty AND every ``stack/*.npy`` has
a matching ``merged/*.npy`` (i.e. every field of view was merged).
:param src: run root folder (holds ``merged/``, ``stack/``, ``masks/``, ``orig/``).
:param keep_intermediate: keep ``stack/`` + ``masks/`` when True.
:param keep_original: keep the raw ``orig/`` backup when True.
:returns: list of folder paths that were deleted.
"""
import os
import shutil
from .io import _listdir_visible
merged = os.path.join(src, 'merged')
stack = os.path.join(src, 'stack')
masks = os.path.join(src, 'masks')
orig = os.path.join(src, 'orig')
deleted = []
if not os.path.isdir(merged):
if verbose:
print("cleanup skipped: no merged/ folder — nothing removed")
return deleted
merged_files = {f for f in _listdir_visible(merged) if f.endswith('.npy')}
if not merged_files:
if verbose:
print("cleanup skipped: merged/ is empty — keeping intermediates")
return deleted
if not keep_intermediate:
stack_files = set()
if os.path.isdir(stack):
stack_files = {f for f in _listdir_visible(stack) if f.endswith('.npy')}
if stack_files and not stack_files.issubset(merged_files):
missing = len(stack_files - merged_files)
if verbose:
print(f"cleanup: keeping stack/ + masks/ — {missing} field(s) "
"not present in merged/")
else:
for folder in (stack, masks):
if os.path.isdir(folder):
shutil.rmtree(folder, ignore_errors=True)
deleted.append(folder)
for d in _listdir_visible(src):
p = os.path.join(src, d)
if os.path.isdir(p) and d.isdigit():
shutil.rmtree(p, ignore_errors=True)
deleted.append(p)
if not keep_original and os.path.isdir(orig):
shutil.rmtree(orig, ignore_errors=True)
deleted.append(orig)
if verbose and deleted:
print(f"cleanup: removed {', '.join(os.path.basename(d) for d in deleted)} "
"(kept merged/)")
return deleted
[docs]
def delete_intermedeate_files(settings):
"""Remove intermediate per-channel and stack folders under ``settings['src']``.
Safeguarded to only run when a ``merged/`` folder is present and the ``orig/``
backup folder exists, so raw inputs are preserved.
:param settings: Dict with an ``'src'`` key naming the run's root folder.
:returns: None.
"""
path_orig = os.path.join(settings['src'], 'orig')
path_stack = os.path.join(settings['src'], 'stack')
merged_stack = os.path.join(settings['src'], 'merged')
path_norm_chan_stack = os.path.join(settings['src'], 'masks')
path_1 = os.path.join(settings['src'], '1')
path_2 = os.path.join(settings['src'], '2')
path_3 = os.path.join(settings['src'], '3')
path_4 = os.path.join(settings['src'], '4')
path_5 = os.path.join(settings['src'], '5')
path_6 = os.path.join(settings['src'], '6')
path_7 = os.path.join(settings['src'], '7')
path_8 = os.path.join(settings['src'], '8')
path_9 = os.path.join(settings['src'], '9')
path_10 = os.path.join(settings['src'], '10')
paths = [path_stack, path_norm_chan_stack, path_1, path_2, path_3, path_4, path_5, path_6, path_7, path_8, path_9, path_10]
if 'src' not in settings:
print("No 'src' key in settings dictionary.")
return
if not os.path.exists(settings['src']):
print(f"{settings['src']} does not exist.")
return
if not os.path.exists(path_orig):
print(f"{path_orig} does not exist.")
return
merged_len = len(os.listdir(merged_stack)) if os.path.isdir(merged_stack) else 0
stack_len = len(os.listdir(path_stack)) if os.path.isdir(path_stack) else 0
if stack_len == 0 or merged_len < stack_len:
return
for path in paths:
if os.path.exists(path):
try:
shutil.rmtree(path)
print(f"Deleted {path}")
except OSError as e:
print(f"{path} could not be deleted: {e}. Delete manually.")
[docs]
def filter_and_save_csv(input_csv, output_csv, column_name, upper_threshold, lower_threshold):
"""
Reads a CSV into a DataFrame, keeps the rows whose column value falls OUTSIDE
the two thresholds, and saves the filtered DataFrame to a new CSV file.
The two tests are combined with OR, not AND, so this is a two-tailed
selection that keeps the extremes and discards the middle. Both comparisons
are strict, so a value exactly equal to either threshold is dropped, and so
is ``NaN``. Passing an ``upper_threshold`` below ``lower_threshold`` makes
the two conditions cover the whole line and nothing is filtered out at all.
Parameters:
input_csv (str): Path to the input CSV file, read with ``pd.read_csv``.
output_csv (str): Path to save the filtered CSV file, written without
the index. Its parent directory must already exist -- pandas raises
``OSError`` rather than creating it.
column_name (str): Column the two comparisons are applied to. A name
that is not in the frame raises ``KeyError``, and a text column
raises ``TypeError`` when compared against a numeric threshold.
upper_threshold (float): Rows strictly greater than this are retained.
lower_threshold (float): Rows strictly less than this are retained too;
everything between the two bounds is discarded.
Returns:
None. The filtered frame is written to ``output_csv``, shown with
``display`` for notebook users, and the destination is printed.
"""
df = tabular.read_table(input_csv, report=None)
filtered_df = df[(df[column_name] > upper_threshold) | (df[column_name] < lower_threshold)]
tabular.write_table(filtered_df, output_csv)
display(filtered_df)
print(f"Filtered DataFrame saved to {output_csv}")
[docs]
def calculate_shortest_distance(df, object1, object2):
"""
Calculate the shortest edge-to-edge distance between two objects (e.g., pathogen and nucleus).
Parameters:
- df: Pandas DataFrame containing measurements
- object1: String, name of the first object (e.g., "pathogen")
- object2: String, name of the second object (e.g., "nucleus")
Returns:
- df: Pandas DataFrame with a new column for shortest edge-to-edge distance.
:param df: measurement frame containing centroid and Feret-diameter
columns for both objects.
:param object1: prefix of the first object's measurement columns.
:param object2: prefix of the second object's measurement columns.
"""
centroid_distance = np.sqrt(
(df[f'{object1}_channel_0_centroid_weighted-0'] - df[f'{object2}_channel_0_centroid_weighted-0'])**2 +
(df[f'{object1}_channel_0_centroid_weighted-1'] - df[f'{object2}_channel_0_centroid_weighted-1'])**2
)
object1_radius = df[f'{object1}_feret_diameter_max'] / 2
object2_radius = df[f'{object2}_feret_diameter_max'] / 2
shortest_distance = centroid_distance - (object1_radius + object2_radius)
df[f'{object1}_{object2}_shortest_distance'] = np.maximum(shortest_distance, 0)
return df
[docs]
def normalize_src_path(src):
"""
Ensures that the 'src' value is properly formatted as either a list of strings or a single string.
Args:
src (str or list): The input source path(s).
Returns:
list or str: A correctly formatted list if the input was a list (or string representation of a list),
otherwise a single string.
"""
if isinstance(src, list):
return src
if isinstance(src, str):
try:
evaluated_src = ast.literal_eval(src)
if isinstance(evaluated_src, list) and all(isinstance(item, str) for item in evaluated_src):
return evaluated_src
except (SyntaxError, ValueError):
pass
return src
raise ValueError(f"Invalid type for 'src': {type(src).__name__}, expected str or list")
[docs]
def generate_image_path_map(root_folder, valid_extensions=("tif", "tiff", "png", "jpg", "jpeg", "bmp", "czi", "nd2", "lif")):
"""
Recursively scans a folder and its subfolders for images, then creates a mapping of:
{original_image_path: new_image_path}, where the new path includes all subfolder names.
Args:
root_folder (str): The root directory to scan for images.
valid_extensions (tuple): Tuple of valid image file extensions.
Returns:
dict: A dictionary mapping original image paths to their new paths.
"""
image_path_map = {}
for dirpath, dirnames, filenames in os.walk(root_folder):
dirnames[:] = [name for name in dirnames
if name != "consolidated" and not name.startswith('.')]
for file in filenames:
if file.startswith('.'):
continue
ext = file.lower().split('.')[-1]
if ext in valid_extensions:
relative_path = os.path.relpath(dirpath, root_folder)
if relative_path == os.curdir:
folder_parts = []
else:
folder_parts = relative_path.split(os.sep)
folder_info = "_".join(folder_parts)
new_filename = f"{folder_info}_{file}" if folder_info else file
original_path = os.path.join(dirpath, file)
new_path = os.path.join(root_folder, new_filename)
image_path_map[original_path] = new_path
return image_path_map
[docs]
def copy_images_to_consolidated(image_path_map, root_folder):
"""
Copies images from their original locations to a 'consolidated' folder,
renaming them according to the generated dictionary.
Args:
image_path_map (dict): Dictionary mapping {original_path: new_path}.
root_folder (str): The root directory where the 'consolidated' folder will be created.
"""
consolidated_folder = os.path.join(root_folder, "consolidated")
os.makedirs(consolidated_folder, exist_ok=True)
files_processed = 0
files_to_process = len(image_path_map)
time_ls= []
for original_path, new_path in image_path_map.items():
start = time.time()
new_filename = os.path.basename(new_path)
new_file_path = os.path.join(consolidated_folder, new_filename)
shutil.copy2(original_path, new_file_path)
files_processed += 1
stop = time.time()
duration = (stop - start)
time_ls.append(duration)
print_progress(files_processed, files_to_process, n_jobs=1, time_ls=time_ls, batch_size=None, operation_type=f'Consolidating images')
[docs]
def remove_outliers_by_group(df, group_col, value_col, method='iqr', threshold=1.5):
"""
Removes outliers from `value_col` within each group defined by `group_col`.
Rows are selected, never modified: the original index is preserved and a new
frame is returned. A row whose value is ``NaN`` fails the comparison and is
always dropped, whichever method is used.
Parameters:
df (pd.DataFrame): The input DataFrame.
group_col (str): Column name to group by, or a list of column names.
Grouping passes ``observed=False``, so unused categories of a
Categorical are kept. Rows whose group key is missing are discarded,
because pandas drops ``NaN`` group keys and the per-row bound then
comes back ``NaN``.
value_col (str): Column containing values to check for outliers. A name
that is not in the frame raises ``KeyError``.
method (str): 'iqr' or 'zscore'. Anything else raises ``ValueError``.
The two now agree on tiny groups: a one-row group has an undefined
standard deviation, and since one row cannot be an outlier within
its own group it is KEPT under both. It used to be dropped by
'zscore' and kept by 'iqr'.
threshold (float): Multiplier on the IQR (default 1.5), or the z-score
cutoff. Must be >= 0; a negative value inverts the keep-band and is
refused, because under 'iqr' it silently emptied every group with a
nonzero IQR. Note ``0`` under 'zscore' still keeps only rows sitting
exactly on the group mean, which is what a zero cutoff means.
Under 'zscore' an outlier inflates its own group's standard
deviation, so the usual cutoffs keep far more than 'iqr' does on the
same data -- that is the statistic, not a defect.
Returns:
pd.DataFrame: A DataFrame with outliers removed.
"""
if threshold < 0:
raise ValueError(
f"threshold must be >= 0, not {threshold!r}: a negative value "
"inverts the keep-band and silently deletes whole groups.")
grouped = df.groupby(group_col, observed=False)[value_col]
if method == 'iqr':
q1 = grouped.transform(lambda values: values.quantile(0.25))
q3 = grouped.transform(lambda values: values.quantile(0.75))
iqr = q3 - q1
keep = df[value_col].between(
q1 - threshold * iqr,
q3 + threshold * iqr,
)
elif method == 'zscore':
mean = grouped.transform('mean')
std = grouped.transform('std')
keep = (df[value_col] - mean).abs() <= threshold * std
keep = keep | std.isna()
else:
raise ValueError("method must be 'iqr' or 'zscore'")
return df.loc[keep]