"""Image, dataset, and SQLite input/output helpers used across spaCR."""
import readlif.reader
import os, re, json, sqlite3, torch, time, random, shutil, cv2, tarfile, glob, queue, threading, tifffile, czifile, atexit, readlif, tempfile, logging, warnings
from . import _gc as gc
import numpy as np
import pandas as pd
from PIL import Image
from collections import defaultdict, Counter
from contextlib import contextmanager
from .resource_log import _parallel_thread_executor as ThreadPoolExecutor
from pathlib import Path
from matplotlib.animation import FuncAnimation
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 skimage.util import img_as_uint
from skimage.exposure import rescale_intensity
import skimage.measure as measure
from skimage import exposure
import imageio.v2 as imageio2
import matplotlib.pyplot as plt
from io import BytesIO
from multiprocessing import cpu_count
from .resource_log import _parallel_pool as Pool
from torch.utils.data import Dataset, DataLoader as _TorchDataLoader, random_split, Subset, WeightedRandomSampler
from .resource_log import _parallel_data_loader as DataLoader
from torchvision.transforms import ToTensor
import seaborn as sns
from nd2reader import ND2Reader
from torchvision import transforms
pyczi = None
from .errors import RunLedger
from .image_colors import read_image_rgb
from .classification_pixels import DECLARED_UINT8, read_classification_image, validate_policy
from .tiff_io import write_tiff
LOG = logging.getLogger(__name__)
from . import convert as _cv
from . import crop_source as _crop_source
from .object_roles import (CHILD_ROLES, ORGANELLE_ROLES,
enabled_organelle_roles, join_how)
from .png_list import (PNG_LIST_ID_COLUMNS, _merged_field_paths,
_object_id_int, crop_rows_from_png_list)
from .crops import MERGED_LAYOUT_SIDECAR
from .merge_tables import reconcile_duplicates
from .figures.style import figure_style, theme_target
def _escaped_field_stem(plate, well, field, time):
"""Compose a merged-stack stem with its free-text plate component escaped.
Every stack this module writes is ``plate_well_field_time``, whatever
``timelapse`` is set to. The well, the field and the timepoint are drawn
from a bounded vocabulary; the PLATE is free text, taken from a regex group
or -- far more often -- from ``os.path.basename(src)``, a folder name. A
plate folder called ``exp_1`` therefore produced ``exp_1_A01_1_1.npy``:
five separator-delimited components for a four-component grammar.
_map_wells('exp_1_A01_1_1.npy') -> ('error',) * 5
The whole plate could not be measured. Before the identity keys were
escaped it was worse rather than better -- the same name parsed as plate
``exp``, well ``1``, field ``A01`` and was wrong QUIETLY.
The plate is escaped here, at the writer, rather than through
:func:`spacr.schema.escape_field_stem_plate`, because the writer holds the
four components separately and so has no splitting to get wrong. That
helper is the same rule for a caller that holds only a joined name, and it
has to guess where the plate ends; escaping before the join means there is
nothing to guess. :func:`spacr.schema.parse_field_stem` reads the result
back as ``exp_1`` character for character.
:param plate: free-text plate id. ``None`` keeps its historical spelling,
the literal ``'None'``, rather than silently becoming an empty
identity.
:param time: the timepoint, or ``''`` for the channel-folder layout, which
has always written a trailing empty component.
:returns: the stem, without an extension.
:raises spacr.schema.KeyParseError: when the plate is empty. An empty plate
is not an identity, and every field written under one would merge with
every other.
"""
from .schema import escape_filename_component
return f'{escape_filename_component(str(plate))}_{well}_{field}_{time}'
#: Subfolders of a plate source folder whose file stems are field ids written
#: by this module, and are therefore what :func:`migrate_unescaped_plate_names`
#: renames. ``orig/`` and the raw drop are deliberately absent: those names are
#: the vendor's, not spaCR's, and are not field stems.
FIELD_STEM_FOLDERS = ('stack', 'norm_channel_stack', 'merged', 'masks')
#: Extensions those folders hold.
FIELD_STEM_SUFFIXES = ('.npy', '.npz', '.tif', '.tiff')
[docs]
def migrate_unescaped_plate_names(src, dry_run=False):
"""Escape the plate component of arrays a previous release wrote raw.
A plate folder whose name holds an underscore -- ``exp_1`` -- used to
produce ``exp_1_A01_1_1.npy``, five separator-delimited components for a
four-component grammar, and ``utils._map_wells`` answered ``('error',) * 5``
for every field of it. The plate could not be measured at all.
Nothing in the measurement database needs migrating, because there is
none: every frame of such a plate was refused. What DOES need moving is
everything upstream of the measurement -- ``stack/``, ``norm_channel_stack/``,
``merged/`` and, above all, ``masks/``, which is hours to days of
segmentation. Renaming those makes the plate measurable without
re-segmenting it, which is the whole reason this exists rather than a note
saying "re-run the plate".
THE NEW NAME IS NOT GUESSED. Every stem this module writes ends in three
fixed tokens -- well, field, timepoint -- whatever ``timelapse`` is set to,
so everything before them is the plate however many underscores it holds.
A stem that is already escaped, or whose plate holds no separator, is left
alone: the rename is a no-op for every ordinary plate, which is what makes
it safe to run over a folder that does not need it.
Crops under ``data/`` are not touched. A plate that could not be measured
has none, and the ``nightly``-only crop names that ``_generate_names``
mis-escaped came with a ``png_list`` table whose identities are wrong too,
so there is nothing there to salvage by renaming -- re-run ``measure_crop``.
PUBLIC, because the person who needs it is a user with an ``exp_1``
folder full of masks, and a recovery tool they have to reach past a
leading underscore to call is a recovery tool most people will not find.
:param src: the plate source folder, the one holding ``merged/``.
:param dry_run: report the renames without performing them.
:returns: list of ``(old_path, new_path)`` pairs, renamed unless
``dry_run``.
:raises FileExistsError: if a destination is already occupied, before
anything is moved. A half-applied rename is worse than none.
Example:
.. code-block:: python
>>> from spacr.io import migrate_unescaped_plate_names
>>> migrate_unescaped_plate_names('/data/exp_1', dry_run=True)
[('/data/exp_1/merged/exp_1_A01_1_1.npy',
'/data/exp_1/merged/exp%5F1_A01_1_1.npy')]
"""
from .schema import KEY_SEPARATOR, KeyParseError, escape_field_stem_plate
planned = []
for folder in FIELD_STEM_FOLDERS:
root = os.path.join(src, folder)
if not os.path.isdir(root):
continue
for base, _dirs, files in os.walk(root):
for name in sorted(files):
stem, suffix = os.path.splitext(name)
if suffix.lower() not in FIELD_STEM_SUFFIXES:
continue
parts = stem.split(KEY_SEPARATOR)
if len(parts) <= 4:
continue
try:
safe = escape_field_stem_plate(stem, timelapse=True)
except KeyParseError:
continue
planned.append((os.path.join(base, name),
os.path.join(base, safe + suffix)))
occupied = [new for _old, new in planned if os.path.exists(new)]
if occupied:
raise FileExistsError(
f'refusing to migrate {src!r}: {len(occupied)} destination(s) '
f'already exist, starting with {occupied[0]!r}. Move or delete '
f'them first — a half-applied rename leaves the plate in a state '
f'neither the old reader nor the new one can read.')
if not dry_run:
for old, new in planned:
os.rename(old, new)
return planned
def _load_pylibczi():
"""Load the optional high-performance CZI reader when it is needed."""
try:
from pylibCZIrw import czi
except (ImportError, OSError) as exc:
raise ImportError(
"High-performance CZI conversion requires pylibCZIrw. "
"Install it with `pip install 'spacr[czi]'`. "
"Python 3.14 users can continue to use spaCR's czifile-based "
"CZI readers until pylibCZIrw publishes a CPython 3.14 wheel."
) from exc
return czi
[docs]
def process_non_tif_non_2D_images(folder):
"""Split multi-dimensional or non-TIFF images in ``folder`` into per-channel TIFFs.
Grayscale non-TIFF images are converted to TIFF in place. Multi-
dimensional images (3D/4D/5D) are split into one grayscale TIFF per
``(channel, Z, T)`` combination. Bit depth is preserved.
A file that cannot be read is recorded on a
:class:`spacr.errors.RunLedger` and skipped, so one corrupt image
does not abort the folder — but the ledger prints a loud summary of
everything that was skipped once the folder is done.
:param folder: Directory containing the input images.
:returns: the :class:`spacr.errors.RunLedger` for the conversion, so
callers can check ``ledger.is_complete`` before trusting the
folder's contents.
"""
def save_grayscale_images(image, base_name, folder, dtype, channel=None, z=None, t=None):
"""Save grayscale images with appropriate suffix based on channel, z, and t, preserving bit depth.
:param image: A single 2D plane already sliced out of the stack.
:param base_name: Stem of the source file. Every plane cut from one
source shares it, so the ``_C``/``_Z``/``_T`` suffix built below is
the only thing keeping those planes from overwriting each other.
:param folder: Directory the TIFF is written to. The converter passes
the folder it is scanning, so planes land beside the
multi-dimensional file they came from.
:param dtype: NumPy dtype the plane is cast to before writing. This
cast is the whole of the "bit depth is preserved" promise — pass
the dtype ``load_image`` reported for the source rather than a
convenient default, or 16-bit data is silently rewritten.
:param channel: 1-based channel index appended as ``_C``. ``None``
leaves that part of the suffix off.
:param z: 1-based Z index appended as ``_Z``; ``None`` omits it.
:param t: 1-based time index appended as ``_T``; ``None`` omits it.
"""
suffix = f"_C{channel}"
if z is not None:
suffix += f"_Z{z}"
if t is not None:
suffix += f"_T{t}"
output_filename = os.path.join(folder, f"{base_name}{suffix}.tif")
write_tiff(output_filename, image.astype(dtype))
def split_channels(image, folder, base_name, dtype):
"""Splits the image into channels and handles 3D, 4D, and 5D image cases.
:param image: Array whose axis order is assumed to be
``(height, width, channel[, Z[, T]])``. A 2D array returns without
writing anything (the caller handles those), and anything with
more than five axes falls through silently, writing nothing. The
axis order is never checked, so a channel-first array from some
other reader is sliced along width instead of being rejected.
:param folder: Directory the per-plane TIFFs are written into.
:param base_name: Stem shared by every plane written from this image;
the ``_C``/``_Z``/``_T`` indices are appended to it, all 1-based.
:param dtype: NumPy dtype passed straight through to each write, so it
should be the source image's own dtype to keep its bit depth.
"""
if image.ndim == 2:
return
elif image.ndim == 3:
for c in range(image.shape[2]):
save_grayscale_images(image[..., c], base_name, folder, dtype, channel=c+1)
elif image.ndim == 4:
for z in range(image.shape[3]):
for c in range(image.shape[2]):
save_grayscale_images(image[..., c, z], base_name, folder, dtype, channel=c+1, z=z+1)
elif image.ndim == 5:
for t in range(image.shape[4]):
for z in range(image.shape[3]):
for c in range(image.shape[2]):
save_grayscale_images(image[..., c, z, t], base_name, folder, dtype, channel=c+1, z=z+1, t=t+1)
def load_image(file_path):
"""Loads image from various formats and returns it as a numpy array along with its dtype.
:param file_path: Path to the image. Only the extension selects the
reader: ``.tif`` and ``.tiff`` go to tifffile, ``.png``, ``.jpg``
and ``.jpeg`` to PIL, ``.czi`` to czifile, ``.nd2`` to ND2Reader.
Content is never sniffed, so a mislabelled file is read with the
wrong reader, and any other extension raises ``ValueError``.
"""
ext = os.path.splitext(file_path)[1].lower()
if ext in ['.tif', '.tiff']:
image = tifffile.imread(file_path)
return image, image.dtype
elif ext in ['.png', '.jpg', '.jpeg']:
image = np.array(Image.open(file_path))
return image, image.dtype
elif ext == '.czi':
with czifile.CziFile(file_path) as czi:
image = czi.asarray()
return image, image.dtype
elif ext == '.nd2':
with ND2Reader(file_path) as nd2:
image = np.array(nd2)
return image, image.dtype
else:
raise ValueError(f"Unsupported file extension: {ext}")
def convert_grayscale_to_tiff(image, filename, folder, dtype):
"""Convert grayscale images that are not in TIFF format to TIFF, preserving bit depth.
:param image: The decoded 2D plane to write.
:param filename: Base name of the source file, not a path. Only the
extension is stripped, so an absolute path here would be joined
with ``folder`` and win, writing the TIFF back next to the
original instead of into ``folder``.
:param folder: Directory the ``.tif`` is written into; the original
file is left in place next to it rather than replaced.
:param dtype: NumPy dtype the plane is cast to, which is what carries
the source bit depth into the TIFF.
"""
base_name = os.path.splitext(filename)[0]
output_filename = os.path.join(folder, f"{base_name}.tif")
write_tiff(output_filename, image.astype(dtype))
print(f"Converted grayscale image {filename} to TIFF with bit depth {dtype}.")
supported_formats = ['.tif', '.tiff', '.png', '.jpg', '.jpeg', '.czi', '.nd2']
ledger = RunLedger('process_non_tif_non_2D_images')
for filename in os.listdir(folder):
file_path = os.path.join(folder, filename)
ext = os.path.splitext(file_path)[1].lower()
if ext in supported_formats:
print(f"Processing {filename}")
with ledger.item(filename, stage='split_channels',
echo=f"Error processing {filename}"):
image, dtype = load_image(file_path)
if image.ndim == 2:
if ext not in ['.tif', '.tiff']:
convert_grayscale_to_tiff(image, filename, folder, dtype)
else:
print(f"Image {filename} is already grayscale and in TIFF format, skipping.")
continue
base_name = os.path.splitext(filename)[0]
split_channels(image, folder, base_name, dtype)
ledger.finalize()
return ledger
def _load_images_and_labels(image_files, label_files, invert=False):
"""Load a Cellpose training set, keeping each name beside its pixels.
THE NAMES ARE BUILT BESIDE THE PIXELS, one append each. They used to be
``sorted(basename(f) for f in image_files)`` while the arrays were filled
in the CALLER's order -- and the caller shuffles before calling here, so
``identify_masks_finetune`` wrote every mask under a DIFFERENT image's
filename: a whole plate of segmentations silently attributed to the
wrong wells.
Sorting was only half of it. Each loop skips a file that will not read,
which shortens the arrays while a precomputed name list keeps every
entry -- so one unreadable file misnamed every mask after it even when
the input was already in order.
:param image_files: image paths.
:param label_files: label paths, positionally matched to the images.
:param invert: invert each image after loading.
:returns: ``(images, labels, image_names, label_names)``, every list in
the same order and the names always matching the arrays beside them.
"""
from cellpose import io as cellpose_io
from .utils import invert_image
images = []
labels = []
image_names = []
label_names = []
if image_files and label_files:
for img_file, lbl_file in zip(image_files, label_files):
image = cellpose_io.imread(img_file)
if image is None:
print(f"WARNING: Could not load image: {img_file}")
continue
if invert:
image = invert_image(image)
if image.max() > 1:
image = image / image.max()
label = cellpose_io.imread(lbl_file)
if label is None:
print(f"WARNING: Could not load label: {lbl_file}")
continue
images.append(image)
labels.append(label)
image_names.append(os.path.basename(img_file))
label_names.append(os.path.basename(lbl_file))
elif image_files:
for img_file in image_files:
image = cellpose_io.imread(img_file)
if image is None:
print(f"WARNING: Could not load image: {img_file}")
continue
if invert:
image = invert_image(image)
if image.max() > 1:
image = image / image.max()
images.append(image)
image_names.append(os.path.basename(img_file))
elif label_files:
for lbl_file in label_files:
label = cellpose_io.imread(lbl_file)
if label is None:
print(f"WARNING: Could not load label: {lbl_file}")
continue
labels.append(label)
label_names.append(os.path.basename(lbl_file))
image_dir = os.path.dirname(image_files[0]) if image_files else None
label_dir = os.path.dirname(label_files[0]) if label_files else None
print(f'Loaded {len(images)} images and {len(labels)} labels from {image_dir} and {label_dir}')
if images and labels:
print(f'image shape: {images[0].shape}, image type: {images[0].dtype}; '
f'label shape: {labels[0].shape}, label type: {labels[0].dtype}')
return images, labels, image_names, label_names
def _load_normalized_images_and_labels(image_files, label_files, channels=None, percentiles=None,
invert=False, visualize=False, remove_background=False,
background=0, Signal_to_noise=10, target_height=None, target_width=None,
rescale=True):
"""Load a Cellpose training set, percentile-normalised and optionally resized.
With no explicit percentiles, the upper one is chosen per channel as the
first of 98, 99, 99.9, 99.99 and 99.999 that clears the signal
threshold, then averaged across the set -- so a channel whose signal
lives in a thin tail is not flattened by a fixed 99th, and a channel
that is mostly background does not have its noise stretched to full
scale.
:param image_files: image paths.
:param label_files: label paths, or ``None`` for images alone.
:param channels: channel indices to keep from a multi-channel image.
:param percentiles: an explicit ``[low, high]`` pair; anything that is
not a two-element list is ignored in favour of the per-channel
search.
:param invert: invert each image after loading.
:param visualize: plot the before/after of the resize.
:param remove_background: zero everything below ``background``.
:param background: the background level.
:param Signal_to_noise: how far above ``background`` a percentile must
sit to count as signal.
:param target_height: resize height, or ``None`` to keep the original.
:param target_width: resize width, or ``None`` to keep the original.
:param rescale: ``False`` returns each image as loaded -- channels
picked, inverted, background removed and resized, but NOT rescaled --
for a caller that lets Cellpose normalise each image itself, as the
live preview does. ``percentiles``, ``Signal_to_noise`` and the
percentile search are then unused.
:returns: ``(images, labels, image_names, label_names, orig_dims)``.
Labels are resized with nearest-neighbour and no anti-aliasing,
because interpolating a label array invents object ids.
"""
from cellpose import io as cellpose_io
from .plot import plot_resize
from .utils import invert_image
from skimage.transform import resize as resizescikit
if isinstance(percentiles, list) and len(percentiles) == 2:
try:
percentiles = [int(percentiles[0]), int(percentiles[1])]
except ValueError:
percentiles = None
else:
percentiles = None
signal_thresholds = float(background) * float(Signal_to_noise)
lower_percentile = 2
images, labels, orig_dims = [], [], []
num_channels = 4
percentiles_1 = [[] for _ in range(num_channels)]
percentiles_99 = [[] for _ in range(num_channels)]
image_names = [os.path.basename(f) for f in image_files]
image_dir = os.path.dirname(image_files[0])
if label_files is not None:
label_names = [os.path.basename(f) for f in label_files]
label_dir = os.path.dirname(label_files[0])
else:
label_names, label_dir = [], None
for i, img_file in enumerate(image_files):
image = cellpose_io.imread(img_file)
orig_dims.append((image.shape[0], image.shape[1]))
if invert:
image = invert_image(image)
if channels is not None and image.ndim == 3:
image = image[..., channels]
if remove_background:
image = np.where(image < background, 0, image)
if image.ndim < 3:
image = np.expand_dims(image, axis=-1)
if rescale and percentiles is None:
for c in range(image.shape[-1]):
p1 = np.percentile(image[..., c], lower_percentile)
percentiles_1[c].append(p1)
for percentile in [98, 99, 99.9, 99.99, 99.999]:
p = np.percentile(image[..., c], percentile)
if float(p) > signal_thresholds:
percentiles_99[c].append(p)
break
if target_height and target_width:
image_shape = (target_height, target_width) if image.ndim == 2 else (target_height, target_width, image.shape[-1])
image = resizescikit(image, image_shape, preserve_range=True, anti_aliasing=True).astype(image.dtype)
images.append(image)
if not rescale:
normalized_images = images
elif percentiles is None:
used = [c for c, p in enumerate(percentiles_1) if p]
avg_p1 = [np.mean(percentiles_1[c]) for c in used]
avg_p99 = [np.mean(percentiles_99[c]) if percentiles_99[c] else avg_p1[i]
for i, c in enumerate(used)]
print(f'Average 1st percentiles: {avg_p1}, Average 99th percentiles: {avg_p99}')
normalized_images = [
np.stack([rescale_intensity(img[..., c], in_range=(avg_p1[c], avg_p99[c]), out_range=(0, 1))
for c in range(img.shape[-1])], axis=-1) for img in images
]
else:
normalized_images = [
np.stack([rescale_intensity(img[..., c],
in_range=(np.percentile(img[..., c], percentiles[0]),
np.percentile(img[..., c], percentiles[1])),
out_range=(0, 1)) for c in range(img.shape[-1])], axis=-1)
for img in images
]
if label_files is not None:
labels = [resizescikit(cellpose_io.imread(lbl_file),
(target_height, target_width) if target_height and target_width else orig_dims[i],
order=0, preserve_range=True, anti_aliasing=False).astype(np.uint8)
for i, lbl_file in enumerate(label_files)]
print(f'Loaded and normalized {len(normalized_images)} images and {len(labels)} labels from {image_dir} and {label_dir}')
if visualize and images and labels:
plot_resize(images, normalized_images, labels, labels)
return normalized_images, labels, image_names, label_names, orig_dims
[docs]
class CombineLoaders:
"""Randomized interleaving of multiple live DataLoaders.
Each step shuffles the loaders that have not been exhausted, probes them
in that random order, and yields the first available ``(loader_index,
batch)`` pair. Exhausted loaders are removed, so every batch from every
input loader is yielded once even when the loaders have different lengths.
:param train_loaders: DataLoaders to combine.
:raises StopIteration: when every wrapped loader is exhausted.
"""
def __init__(self, train_loaders):
"""Store loaders and initialise per-loader iterators."""
self.train_loaders = train_loaders
self.loader_iters = [(i, iter(loader))
for i, loader in enumerate(train_loaders)]
[docs]
def __iter__(self):
"""Return self — this object is its own iterator."""
return self
[docs]
def __next__(self):
"""Return ``(loader_index, batch)`` from a randomly-chosen live loader."""
while self.loader_iters:
random.shuffle(self.loader_iters)
for pos, (idx, loader_iter) in enumerate(self.loader_iters):
try:
batch = next(loader_iter)
except StopIteration:
continue
if pos:
self.loader_iters = self.loader_iters[pos:]
return idx, batch
self.loader_iters = []
raise StopIteration
[docs]
class CombinedDataset(Dataset):
"""Concatenation of multiple ``Dataset`` objects behind a single index space.
:param datasets: Datasets to concatenate; their samples must be
index-compatible.
:param shuffle: If True, index lookups are permuted once at
construction time. Default ``True``.
"""
def __init__(self, datasets, shuffle=True):
"""Precompute per-dataset lengths and optionally shuffle indices."""
self.datasets = datasets
self.lengths = [len(dataset) for dataset in datasets]
self.total_length = sum(self.lengths)
self.shuffle = shuffle
if shuffle:
self.indices = list(range(self.total_length))
random.shuffle(self.indices)
else:
self.indices = None
[docs]
def __getitem__(self, index):
"""Return the sample at ``index`` from the appropriate sub-dataset."""
if self.shuffle:
index = self.indices[index]
for dataset, length in zip(self.datasets, self.lengths):
if index < length:
return dataset[index]
index -= length
[docs]
def __len__(self):
"""Return the total number of samples across all sub-datasets."""
return self.total_length
[docs]
class NoClassDataset(Dataset):
"""Flat directory of unlabelled images returned alongside their file paths.
:param data_dir: Directory containing image files.
:param transform: Optional callable applied to each PIL image. If
``None``, images are converted with ``ToTensor``.
:param shuffle: If True, shuffle filename list at construction.
Default ``True``.
:param load_to_memory: If True, decode all images once and hold them
in RAM. Default ``False``.
:param crop_loading_policy: ``declared_uint8_v1`` uses shared crop decoding;
``stored_pil_v1`` preserves untagged historical checkpoints.
"""
def __init__(self, data_dir, transform=None, shuffle=True, load_to_memory=False,
*, crop_loading_policy=DECLARED_UINT8):
"""Enumerate files in ``data_dir`` and optionally preload them."""
self.data_dir = data_dir
self.crop_loading_policy = validate_policy(crop_loading_policy)
self.transform = transform
self.shuffle = shuffle
self.load_to_memory = load_to_memory
self.filenames = [
os.path.join(data_dir, f)
for f in os.listdir(data_dir)
if os.path.isfile(os.path.join(data_dir, f)) and not f.startswith('.')
]
if self.shuffle:
self.shuffle_dataset()
if self.load_to_memory:
self.images = [self.load_image(f) for f in self.filenames]
[docs]
def load_image(self, img_path):
"""Return the image at ``img_path`` decoded as RGB.
:param img_path: Path to the image file.
:returns: PIL ``Image`` in RGB mode.
"""
return read_classification_image(img_path, self.crop_loading_policy)
[docs]
def __len__(self):
"""Return the number of images in the dataset."""
return len(self.filenames)
[docs]
def shuffle_dataset(self):
"""Shuffle the internal filename list in place."""
if self.shuffle:
random.shuffle(self.filenames)
[docs]
def __getitem__(self, index):
"""Return ``(image_tensor, filename)`` for the given index.
:param index: Position within the dataset.
:returns: ``(tensor, path)`` where ``tensor`` is the transformed
image and ``path`` is the source filename.
"""
if self.load_to_memory:
img = self.images[index]
else:
img = self.load_image(self.filenames[index])
if self.transform is not None:
img = self.transform(img)
else:
img = ToTensor()(img)
return img, self.filenames[index]
[docs]
class spacrDataset(Dataset):
"""Image classification dataset that reads class subfolders under ``data_dir``.
:param data_dir: Root directory containing one subdirectory per class.
:param loader_classes: Ordered list of class names — the index in
this list becomes the integer label.
:param transform: Optional callable applied to each PIL image.
:param shuffle: If True, shuffle files+labels at construction.
:param pin_memory: If True, eagerly load every image into RAM via a
multiprocessing pool.
:param specific_files: Optional explicit list of image paths. If
supplied together with ``specific_labels``, directory scanning
is skipped.
:param specific_labels: Labels paired with ``specific_files``.
:param crop_loading_policy: ``declared_uint8_v1`` uses shared crop decoding;
``stored_pil_v1`` preserves untagged historical checkpoints.
:raises ValueError: If no non-hidden image files are found for any
requested class.
"""
def __init__(self, data_dir, loader_classes, transform=None, shuffle=True, pin_memory=False, specific_files=None, specific_labels=None,
*, crop_loading_policy=DECLARED_UINT8):
"""Build the filename/label lists and optionally preload images."""
self.data_dir = data_dir
self.crop_loading_policy = validate_policy(crop_loading_policy)
self.classes = loader_classes
self.transform = transform
self.shuffle = shuffle
self.pin_memory = pin_memory
self.filenames = []
self.labels = []
if specific_files and specific_labels:
self.filenames = specific_files
self.labels = specific_labels
else:
for class_name in self.classes:
class_path = os.path.join(data_dir, class_name)
if not os.path.isdir(class_path):
continue
class_files = [os.path.join(class_path, f) for f in os.listdir(class_path)
if os.path.isfile(os.path.join(class_path, f))
and not f.startswith('.')]
self.filenames.extend(class_files)
self.labels.extend([self.classes.index(class_name)] * len(class_files))
if not self.filenames:
looked = []
for class_name in self.classes:
cp = os.path.join(data_dir, class_name)
if not os.path.isdir(cp):
looked.append(f" {class_name}: NO SUCH FOLDER ({cp})")
else:
n = len([f for f in os.listdir(cp) if not f.startswith('.')])
looked.append(f" {class_name}: {n} file(s) in {cp}")
raise ValueError(
"The training dataset is empty -- no images were found for any "
"class, so there is nothing to train on.\n"
f"Looked under {data_dir} for classes {list(self.classes)}:\n"
+ "\n".join(looked) +
"\n\nThis usually means the dataset-generation step selected no "
"rows. Check that class_metadata values actually occur in the "
"column the Classes editor names, that the annotation column "
"holds the classes in annotated_classes, and that png_type "
"matches the crops that exist.")
if self.shuffle:
self.shuffle_dataset()
if self.pin_memory:
workers = min(len(self.filenames), 32)
with ThreadPoolExecutor(
max_workers=workers,
thread_name_prefix="spacr-image-load",
) as executor:
self.images = list(executor.map(
self.load_image, self.filenames))
else:
self.images = None
[docs]
def load_image(self, img_path):
"""Return the image at ``img_path`` decoded as RGB with EXIF orientation applied.
:param img_path: Path to one crop, taken from ``self.filenames``.
The file is fully decoded and copied before the handle closes, so
the returned image owns its buffer and is safe to hand to a
prefetch thread. Anything PIL can open works: a greyscale crop is
widened to three channels and an RGBA one loses its alpha, which
is what makes every sample the same shape for the transform.
:returns: A ``PIL.Image.Image`` in mode ``RGB``.
"""
return read_classification_image(img_path, self.crop_loading_policy,
legacy_orient=True)
[docs]
def __len__(self):
"""Return the number of samples in the dataset."""
return len(self.filenames)
[docs]
def shuffle_dataset(self):
"""Jointly shuffle ``filenames`` and ``labels`` in place."""
combined = list(zip(self.filenames, self.labels))
random.shuffle(combined)
self.filenames, self.labels = zip(*combined)
[docs]
def get_plate(self, filepath):
"""Return the plate identifier parsed from a filename (leading token before ``_``).
:param filepath: Image path.
:returns: Plate ID string.
"""
filename = os.path.basename(filepath)
return filename.split('_')[0]
[docs]
def __getitem__(self, index):
"""Return ``(image, label, filename)`` for the given index."""
if self.pin_memory:
img = self.images[index]
else:
img = self.load_image(self.filenames[index])
label = self.labels[index]
filename = self.filenames[index]
if self.transform:
img = self.transform(img)
return img, label, filename
[docs]
class spacrDataLoader(_TorchDataLoader):
"""DataLoader that pre-fetches batches into a queue on a background thread.
Wraps ``torch.utils.data.DataLoader`` and runs a daemon thread that
stays one or more batches ahead of consumption to hide I/O latency.
End-of-stream is signalled with a sentinel, so the full batch stream
is always delivered.
:param preload_batches: Number of batches to keep queued ahead.
Default ``1``.
"""
def __init__(self, *args, preload_batches=1, **kwargs):
"""Initialise the underlying DataLoader and the preload queue."""
from .resource_log import _data_loader_arguments
kwargs = _data_loader_arguments(args, kwargs)
super().__init__(*args, **kwargs)
self.preload_batches = preload_batches
self.batch_queue = queue.Queue(maxsize=max(1, preload_batches))
self.thread = None
self.current_batch_index = 0
self._stop_event = False
self._stop_signal = threading.Event()
self._sentinel = object()
self._error = None
self._iteration_active = False
self.pin_memory = kwargs.get('pin_memory', False)
atexit.register(self.cleanup)
def _preload_next_batches(self, q, iterator, stop_signal):
"""Feed every batch of ``iterator`` into ``q``, then a sentinel
marking end-of-stream.
``q`` and ``iterator`` are passed in rather than read off ``self`` so
a producer started by an earlier ``__iter__`` can never write into a
queue created by a later one (that duplicated the whole stream).
"""
try:
for batch in iterator:
if stop_signal.is_set():
break
if self.pin_memory:
batch = self._pin_memory_batch(batch)
while not stop_signal.is_set():
try:
q.put(batch, timeout=0.1)
break
except queue.Full:
continue
except Exception as e:
self._error = e
LOG.exception("spaCR data preloader failed")
finally:
while not stop_signal.is_set():
try:
q.put(self._sentinel, timeout=0.1)
break
except queue.Full:
continue
def _pin_memory_batch(self, batch):
"""Pin a batch's tensors for faster host-to-GPU copies.
Non-tensor members are passed through untouched, so a batch carrying
labels or paths alongside its tensors survives intact.
"""
if isinstance(batch, (list, tuple)):
return [b.pin_memory() if isinstance(b, torch.Tensor) else b for b in batch]
elif isinstance(batch, torch.Tensor):
return batch.pin_memory()
else:
return batch
[docs]
def __iter__(self):
"""Start a fresh pass over the data and return self.
Safe to call more than once (``list(iter(dl))`` calls it twice): any
in-flight producer is stopped first, so the stream is never doubled.
"""
if self._iteration_active:
return self
self.cleanup()
self._stop_event = False
self._stop_signal = threading.Event()
self._error = None
self.current_batch_index = 0
q = queue.Queue(maxsize=max(1, self.preload_batches))
self.batch_queue = q
iterator = iter(super().__iter__())
self._iterator = iterator
self.thread = threading.Thread(
target=self._preload_next_batches,
args=(q, iterator, self._stop_signal),
daemon=True,
name="spacr-data-preloader",
)
self.thread.start()
self._iteration_active = True
return self
[docs]
def __next__(self):
"""Return the next queued batch, or raise ``StopIteration`` at the
sentinel the preloader pushes when the stream is exhausted."""
try:
next_batch = self.batch_queue.get(timeout=60)
except queue.Empty:
self._iteration_active = False
raise StopIteration
if next_batch is self._sentinel:
self._iteration_active = False
if self._error is not None:
err, self._error = self._error, None
raise err
raise StopIteration
self.current_batch_index += 1
return next_batch
[docs]
def cleanup(self):
"""Signal the preloader to stop and join the background thread."""
self._iteration_active = False
self._stop_event = True
stop_signal = getattr(self, '_stop_signal', None)
if stop_signal is not None:
stop_signal.set()
thread = getattr(self, 'thread', None)
if thread is not None and thread.is_alive():
deadline = time.monotonic() + 5
while thread.is_alive() and time.monotonic() < deadline:
try:
while True:
self.batch_queue.get_nowait()
except (queue.Empty, AttributeError):
pass
thread.join(timeout=0.05)
if thread.is_alive():
LOG.error(
"Data preloader did not stop within five seconds; "
"the daemon thread will be abandoned")
[docs]
def __del__(self):
"""Ensure background resources are released on garbage collection."""
self.cleanup()
[docs]
class TarImageDataset(Dataset):
"""Image dataset backed by a tar archive, decoded on demand.
A tar written by :func:`generate_dataset` from on-demand crops carries a
``.spacr_crop_format.json`` member -- the same marker
:mod:`spacr.crops` writes into a crop folder, travelling with the bytes it
describes. It is **not** an image, so it is excluded from the sample list
and surfaced as :attr:`crop_format` instead; an archive without one
reports None, which is every tar written before this existed.
New datasets default to declared channel order and uint8 narrowing through
the shared crop decoder. Inference must pass the model's recorded policy;
untagged older checkpoints use ``stored_pil_v1`` to preserve their pixels.
:param tar_path: Path to the tar archive.
:param transform: Optional callable applied to each PIL image.
:param crop_loading_policy: ``declared_uint8_v1`` (default), or
``stored_pil_v1`` for historical model inputs.
"""
def __init__(self, tar_path, transform=None, *, crop_loading_policy=DECLARED_UINT8):
"""Enumerate archive members without extracting."""
self.tar_path = tar_path
self.transform = transform
self.crop_format = None
self.crop_loading_policy = validate_policy(crop_loading_policy)
self._crop_markers = {}
from . import crops
with tarfile.open(self.tar_path, 'r') as f:
self._archive_names = set(f.getnames())
self.members = []
for m in f.getmembers():
if not m.isfile():
continue
if os.path.basename(m.name) == crops.CROP_FORMAT_SIDECAR:
try:
payload = json.loads(f.extractfile(m).read().decode('utf-8'))
fmt = crops._coerce_format(payload.get('spacr_crop_format'))
if fmt is None:
raise ValueError("unsupported crop format")
directory = os.path.normpath(os.path.dirname(m.name))
if directory in self._crop_markers:
raise ValueError("duplicate crop format marker")
self._crop_markers[directory] = payload
if directory == '.':
self.crop_format = fmt
except (ValueError, TypeError, AttributeError) as exc:
if self.crop_loading_policy == DECLARED_UINT8:
raise ValueError(f"Invalid tar crop marker {m.name}: {exc}") from exc
continue
if m.name.endswith(crops.CROP_MIGRATION_SUFFIX):
continue
self.members.append(m)
[docs]
def __len__(self):
"""Return the number of image members in the archive."""
return len(self.members)
[docs]
def __getitem__(self, idx):
"""Return ``(image, member_name)`` extracted from the tar at ``idx``."""
with tarfile.open(self.tar_path, 'r') as f:
m = self.members[idx]
img_file = f.extractfile(m)
fmt = (self._member_crop_format(m.name)
if self.crop_loading_policy == DECLARED_UINT8 else 1)
img = read_classification_image(BytesIO(img_file.read()),
self.crop_loading_policy, fmt=fmt)
if self.transform:
img = self.transform(img)
return img, m.name
def _member_crop_format(self, name):
"""Resolve a member's nearest folder marker, including migration state."""
from . import crops
directory = os.path.normpath(os.path.dirname(name))
while directory not in self._crop_markers and directory not in ('.', '/'):
directory = os.path.dirname(directory) or '.'
marker = self._crop_markers.get(directory)
if marker is None:
return crops.CROP_FORMAT_LEGACY_BGR
migration = marker.get('migration')
source = crops._coerce_format((migration or marker).get('from')
or marker.get('migrated_from')) or crops.CROP_FORMAT_LEGACY_BGR
basename = os.path.basename(name)
if name + crops.CROP_MIGRATION_SUFFIX in self._archive_names:
return source
if basename in set((migration or marker).get('unconverted') or ()):
return source
if migration:
watermark = migration.get('done_through')
if watermark is None or basename > str(watermark):
return source
return int(marker['spacr_crop_format'])
[docs]
def load_images_from_paths(images_by_key):
"""Load images grouped by key into NumPy arrays.
:param images_by_key: Mapping of key -> list of image paths.
:returns: Mapping of the same keys -> list of ``ndarray`` images.
Paths that fail to load are skipped, recorded on a
:class:`spacr.errors.RunLedger` and reported in a loud summary,
so a short list is never mistaken for a complete one.
"""
images_dict = {}
ledger = RunLedger('load_images_from_paths')
for key, paths in images_by_key.items():
images_dict[key] = []
for path in paths:
with ledger.item(path, stage='load',
echo=f"Error loading image from {path}"):
with Image.open(path) as img:
images_dict[key].append(np.array(img))
ledger.finalize()
return images_dict
_RAW_IMAGE_SUFFIXES = ('.tif', '.tiff', '.png', '.jpg', '.jpeg', '.bmp', '.nd2',
'.czi', '.lif')
def _raw_image_names(folder, img_format=_RAW_IMAGE_SUFFIXES):
"""Name the raw images directly in ``folder``, hidden files excluded.
:param folder: the folder to list. One that is missing or unreadable
holds none.
:param img_format: accepted endings, matched case-sensitively, which is
how :func:`_rename_and_organize_image_files` has always matched them.
:returns: the file names, sorted.
"""
if isinstance(img_format, str):
img_format = [img_format]
try:
names = _listdir_visible(folder)
except OSError:
return []
return sorted(name for name in names
if any(name.endswith(ext) for ext in img_format))
def _stack_field_stems(stack_path):
"""Return the field stems that have a ``.npy`` in ``stack_path``.
:param stack_path: a ``stack/`` folder; a missing one holds no fields.
:returns: set of file names without the ``.npy`` extension.
"""
if not os.path.isdir(stack_path):
return set()
return {os.path.splitext(name)[0] for name in _listdir_visible(stack_path)
if name.endswith('.npy')}
def _rename_and_organize_image_files(src, regex, batch_size=100, metadata_type='', img_format='.tif', timelapse=False, save_original_images=True):
"""
Convert z-stack images to maximum intensity projection (MIP) images and
write the merged multi-channel ``stack/`` arrays **directly** — without
ever creating the intermediate per-channel sub-folders.
Instead of MIP-ing each channel to a ``src/<channel>/`` folder and then
re-reading those folders to merge them (which duplicated the pixel data on
disk), this projects and writes one field at a time. Only filenames are
retained across fields; pixel memory is bounded by one field's channels
plus a decoded plane and the output stack. Each z-plane is folded into
its channel maximum without materializing an entire z-stack. The
merge order and MIP maths are identical to the old folder+\\ ``_merge_file``
path, so the produced stacks are byte-for-byte the same.
Raw images are read from ``src`` and from ``src/orig/``, where an earlier
run set them aside, and a field that already has a ``stack/<fov>.npy`` is
not built again. So a second run on a folder spaCR has already
preprocessed finishes what the first left undone instead of finding
nothing: the fields a killed run never wrote are built, and a
``stack/`` that was emptied is built again from ``orig/``. When
``stack/`` holds fields but none of them is named like a field these
images make, it was written under another naming scheme (an older spaCR,
or channel folders) and nothing is added to it.
Every stack is written atomically, so a killed run leaves no truncated
``.npy`` behind. When no image yields a field at all, nothing is created,
moved or deleted.
Only filenames are kept plate-wide; pixel buffers belong to one field at a
time. Keeping every MIP used plate-sized RAM before the first stack was
written (over 100 GB for a 2000 x 2000 multi-channel plate). One plane is
decoded at a time, also when a regex groups all z slices under one key, so
there is no z-stack-sized temporary, and loop locals are released so the
previous field's last plane is not held while the next is decoded. A failed
read leaves the input folder untouched, including its directory layout:
stack/ is created only when a field is ready to publish.
Args:
src (str): The source directory containing the z-stack images.
regex (str): The regular expression pattern used to match the filenames of the z-stack images.
batch_size (int, optional): Retained for call compatibility; raw ingest
always streams one field at a time regardless of this value.
metadata_type (str, optional): The type of metadata associated with the images. Defaults to ''.
save_original_images (bool, optional): When True (default) the raw input
images are moved aside into ``src/orig/`` for safekeeping. When
False they are deleted after the stack is written, so the pixel data
lives only in ``stack/`` (no duplication); only an image whose field
now has a stack is deleted, and ``orig/`` is never touched.
Defaults to True.
Returns:
int: the number of distinct channels found (0 when nothing was processed).
"""
if isinstance(img_format, str):
img_format = [img_format]
from .utils import _extract_filename_metadata, print_progress
regular_expression = re.compile(regex)
stack_path = os.path.join(src, 'stack')
orig_path = os.path.join(src, 'orig')
files_processed = 0
channels_seen = set()
src_filenames = _raw_image_names(src, img_format)
image_paths_by_key = defaultdict(list)
for folder, names in ((src, src_filenames),
(orig_path, _raw_image_names(orig_path, img_format))):
if folder != src and not names:
continue
print(f'All files: {len(names)} in {folder}')
parsed = _extract_filename_metadata(names, folder, regular_expression, metadata_type)
for key, paths in parsed.items():
if folder != src and 'plateID' not in regular_expression.groupindex:
key = (os.path.basename(src),) + tuple(key[1:])
image_paths_by_key[key].extend(paths)
print(f'All unique FOV: {len(image_paths_by_key)} in {src}')
if not image_paths_by_key:
return 0
stem_of = {key: _escaped_field_stem(key[0], key[1], key[2], key[4])
for key in image_paths_by_key}
wanted = set(stem_of.values())
existing = _stack_field_stems(stack_path)
if existing and not existing & wanted:
print(f'stack/ already holds {len(existing)} field(s), and none is '
f'named like a field these images make (for example '
f'{sorted(wanted)[0]}.npy); it was written under another naming '
f'scheme, so nothing is added to it.')
pending_keys = []
else:
pending_keys = [key for key in image_paths_by_key
if stem_of[key] not in existing]
if existing and pending_keys:
print(f'Resuming: stack/ holds {len(existing & wanted)} of '
f'{len(wanted)} field(s); building the other '
f'{len(wanted - existing)}.')
for key in image_paths_by_key:
if stem_of[key] in existing:
channels_seen.add(key[3])
from .cancellation import checkpoint
field_keys = defaultdict(list)
for key in pending_keys:
field_keys[stem_of[key]].append(key)
sorted_channels = sorted({key[3] for key in image_paths_by_key})
files_to_process = sum(len(image_paths_by_key[key]) for key in pending_keys)
time_ls = []
for stem, keys in field_keys.items():
checkpoint()
start = time.time()
output_filename = stem + '.tif'
new_file = os.path.join(stack_path, stem + '.npy')
if os.path.exists(new_file):
print(f'WARNING: A file with the same name already exists at location {new_file}')
channels_seen.update(key[3] for key in keys)
continue
chan_mips = {}
for key in keys:
channel = key[3]
for path in image_paths_by_key[key]:
checkpoint()
loaded = load_images_from_paths({key: [path]})[key]
for plane in loaded:
previous = chan_mips.get(channel)
if previous is None:
chan_mips[channel] = plane
elif previous.dtype == plane.dtype:
np.maximum(previous, plane, out=previous)
else:
chan_mips[channel] = np.maximum(previous, plane)
channels_seen.add(channel)
loaded.clear()
plane = previous = None
files_processed += 1
planes = []
for channel in sorted_channels:
mip = chan_mips.get(channel)
if mip is None:
print(f"Warning: FOV {output_filename} is missing channel {channel}")
continue
planes.append(np.expand_dims(mip, axis=2))
if planes:
checkpoint()
os.makedirs(stack_path, exist_ok=True)
_save_array_atomic(new_file, np.concatenate(planes, axis=2))
else:
print(f"No valid channels to merge for file {output_filename}")
planes.clear()
chan_mips.clear()
mip = None
time_ls.append(time.time() - start)
print_progress(files_processed, files_to_process, n_jobs=1,
time_ls=time_ls, batch_size=1,
operation_type='Preprocessing filenames')
stacked = _stack_field_stems(stack_path)
if save_original_images:
to_move = []
if wanted & stacked:
to_move = [filename for filename in _listdir_visible(src)
if os.path.splitext(filename)[1] in img_format]
if to_move:
os.makedirs(orig_path, exist_ok=True)
for filename in to_move:
move = os.path.join(orig_path, filename)
if os.path.exists(move):
print(f'WARNING: A file with the same name already exists at location {move}')
else:
shutil.move(os.path.join(src, filename), move)
else:
in_a_stack = {path for key, paths in image_paths_by_key.items()
if stem_of[key] in stacked for path in paths}
for filename in _listdir_visible(src):
if os.path.splitext(filename)[1] in img_format:
path = os.path.join(src, filename)
if path not in in_a_stack:
continue
try:
os.remove(path)
except OSError as e:
print(f"Warning: could not delete original image {filename}: {e}")
files_processed = 0
return len(channels_seen)
def _merge_file(chan_dirs, stack_dir, file_name):
"""
Merge multiple channels into a single stack and save it as a numpy array, using os module for path handling.
Args:
chan_dirs (list): List of directories containing channel images.
stack_dir (str): Directory to save the merged stack.
file_name (str): File name of the channel image.
Returns:
None
"""
file_root, file_ext = os.path.splitext(file_name)
new_file = os.path.join(stack_dir, file_root + '.npy')
if not os.path.exists(new_file):
os.makedirs(stack_dir, exist_ok=True)
channels = []
for i, chan_dir in enumerate(chan_dirs):
img_path = os.path.join(chan_dir, file_name)
img = read_image_rgb(img_path, cv2.IMREAD_UNCHANGED)
if img is None:
print(f"Warning: Failed to read image {img_path}")
continue
chan = np.expand_dims(img, axis=2)
channels.append(chan)
del img
if i % 10 == 0:
gc.collect()
if channels:
stack = np.concatenate(channels, axis=2)
_save_array_atomic(new_file, stack)
else:
print(f"No valid channels to merge for file {file_name}")
def _is_dir_empty(dir_path):
"""
Check if a directory is empty using os module.
"""
return len(os.listdir(dir_path)) == 0
def _generate_time_lists(file_list):
"""
Generate sorted lists of filenames grouped by plate, well, and field.
Args:
file_list (list): A list of filenames.
Returns:
list: A list of sorted file lists, where each file list contains filenames
belonging to the same plate, well, and field, sorted by timepoint.
"""
file_dict = defaultdict(list)
for filename in file_list:
if filename.endswith('.npy'):
parts = filename.split('_')
if len(parts) >= 4:
plate, well, field = parts[:3]
try:
timepoint = int(parts[3].split('.')[0])
except ValueError:
continue
key = (plate, well, field)
file_dict[key].append((timepoint, filename))
else:
continue
sorted_grouped_filenames = [sorted(files, key=lambda x: x[0]) for files in file_dict.values()]
sorted_file_lists = [[filename for _, filename in group] for group in sorted_grouped_filenames]
return sorted_file_lists
def _move_to_chan_folder(src, regex, timelapse=False, metadata_type=''):
"""Sort a flat folder of images into per-channel stacks from their filenames.
Zero padding is undone through ``_int_or_token``, which keeps a token
holding no integer as itself rather than turning it into ``0`` -- a well
called ``0A`` is not well 0.
Every file is parsed inside a run-ledger item, so a name the regex
cannot read is recorded and the rest of the plate still moves; the run
reports what it could not place rather than stopping at the first one.
:param src: the folder to sort.
:param regex: the filename pattern, with named groups for the plate,
well, field, channel and time.
:param timelapse: keep the time component in the output layout.
:param metadata_type: the microscope convention; ``'cq1'`` also converts
the well id, whose scheme differs from the well name it prints.
"""
from .utils import _int_or_token, _convert_cq1_well_id
src_path = src
src = Path(src)
valid_exts = ['.tif', '.png']
ledger = RunLedger('_move_to_chan_folder')
if not (src / 'stack').exists():
for file in src.iterdir():
if file.is_file():
name, ext = file.stem, file.suffix
if ext in valid_exts:
metadata = re.match(regex, file.name)
with ledger.item(
file.name, stage='parse_filename',
echo=(f"Could not extract information from filename "
f"{name}{ext} with {regex}")):
try:
plateID = metadata.group('plateID')
except Exception:
plateID = src.name
wellID = metadata.group('wellID')
fieldID = metadata.group('fieldID')
chanID = metadata.group('chanID')
timeID = metadata.group('timeID')
if wellID[0].isdigit():
wellID = _int_or_token(wellID)
if fieldID[0].isdigit():
fieldID = _int_or_token(fieldID)
if chanID[0].isdigit():
chanID = _int_or_token(chanID)
if timeID[0].isdigit():
timeID = _int_or_token(timeID)
if metadata_type =='cq1':
orig_wellID = wellID
wellID = _convert_cq1_well_id(wellID)
print(f'Converted Well ID: {orig_wellID} to {wellID}')
newname = _escaped_field_stem(
plateID, wellID, fieldID,
timeID if timelapse else '') + ext
newpath = src / chanID
move = newpath / newname
if move.exists():
print(f'WARNING: A file with the same name already exists at location {move}')
else:
newpath.mkdir(exist_ok=True)
shutil.copy(file, move)
valid_exts = ['.tif', '.png']
newpath = os.path.join(src_path, 'orig')
os.makedirs(newpath, exist_ok=True)
for filename in os.listdir(src_path):
if os.path.splitext(filename)[1] in valid_exts:
move = os.path.join(newpath, filename)
if os.path.exists(move):
print(f'WARNING: A file with the same name already exists at location {move}')
else:
shutil.move(os.path.join(src, filename), move)
ledger.finalize()
return
def _channel_folders(src):
"""Name the single-channel folders (``0`` to ``100``, or ``00`` to ``09``) directly in ``src``.
:param src: the plate folder.
:returns: the folder names, sorted as strings.
"""
string_list = [str(i) for i in range(101)]+[f"{i:02d}" for i in range(10)]
try:
names = _listdir_visible(src)
except OSError:
return []
return sorted(d for d in names
if os.path.isdir(os.path.join(src, d)) and d in string_list)
def _merge_channels(src, plot=False):
"""
Merge the channels in the given source directory and save the merged files in a 'stack' directory without using multiprocessing.
Only the fields ``stack/`` lacks are merged, so a run killed while
writing ``stack/`` is finished on the next run, and a stack set aside as
damaged is merged again from the channel folders. When ``stack/`` holds
fields but none is named like a file in the channel folders, it was
written under another naming scheme and nothing is added to it.
:param src: the plate folder holding the channel folders.
:param plot: plot the stacks afterwards.
:returns: the number of channel folders; 0 when there are none.
"""
from .plot import plot_arrays
from .utils import print_progress
stack_dir = os.path.join(src, 'stack')
print(f'generated stack dir at {stack_dir}')
chan_dirs = _channel_folders(src)
num_matching_folders = len(chan_dirs)
print(f'List of folders in src: {chan_dirs}. Single channel folders.')
if not chan_dirs:
print(f'No single-channel folders in {src}; stack/ will be built '
f'directly from the source images.')
return 0
first_dir_path = os.path.join(src, chan_dirs[0])
dir_files = _listdir_visible(first_dir_path)
if not os.path.exists(stack_dir):
os.makedirs(stack_dir, exist_ok=True)
print(f'Generated folder with merged arrays: {stack_dir}')
wanted = {os.path.splitext(name)[0]: name for name in dir_files
if os.path.isfile(os.path.join(first_dir_path, name))}
existing = _stack_field_stems(stack_dir)
if existing and wanted and not existing & set(wanted):
print(f'stack/ already holds {len(existing)} field(s), and none is '
f'named like a file in the channel folders (for example '
f'{sorted(wanted)[0]}.npy); it was written under another '
f'naming scheme, so nothing is added to it.')
pending = []
else:
pending = [name for stem, name in wanted.items()
if stem not in existing]
if existing and pending:
print(f'Resuming: stack/ holds {len(existing & set(wanted))} of '
f'{len(wanted)} field(s); merging the other {len(pending)} '
f'from the channel folders.')
if pending:
time_ls = []
files_to_process = len(pending)
for i, file_name in enumerate(pending):
start_time = time.time()
_merge_file([os.path.join(src, d) for d in chan_dirs], stack_dir, file_name)
stop_time = time.time()
duration = stop_time - start_time
time_ls.append(duration)
files_processed = i + 1
print_progress(files_processed, files_to_process, n_jobs=1, time_ls=time_ls, batch_size=None, operation_type='Merging channels into npy stacks')
if plot:
plot_arrays(os.path.join(src, 'stack'))
return num_matching_folders
def _concatenate_channel(src, channels, randomize=True, timelapse=False, batch_size=100):
"""
Concatenates channel data from multiple files and saves the concatenated data as numpy arrays.
Args:
src (str): The source directory containing the channel data files.
channels (list): The list of channel indices to be concatenated.
randomize (bool, optional): Whether to randomize the order of the files. Defaults to True.
timelapse (bool, optional): Whether the channel data is from a timelapse experiment. Defaults to False.
batch_size (int, optional): The number of files to be processed in each batch. Defaults to 100.
Returns:
str: The directory path where the concatenated channel data is saved.
"""
from .utils import print_progress
channels = [item for item in channels if item is not None]
paths = []
time_ls = []
channel_stack_loc = os.path.join(os.path.dirname(src), 'channel_stack')
os.makedirs(channel_stack_loc, exist_ok=True)
if timelapse:
try:
time_stack_path_lists = _generate_time_lists(os.listdir(src))
for i, time_stack_list in enumerate(time_stack_path_lists):
start = time.time()
stack_region = []
filenames_region = []
for idx, file in enumerate(time_stack_list):
path = os.path.join(src, file)
if idx == 0:
parts = file.split('_')
name = parts[0]+'_'+parts[1]+'_'+parts[2]
array = np.load(path)
array = np.take(array, channels, axis=2)
stack_region.append(array)
filenames_region.append(os.path.basename(path))
stop = time.time()
duration = stop - start
time_ls.append(duration)
files_processed = i+1
files_to_process = len(time_stack_path_lists)
print_progress(files_processed, files_to_process, n_jobs=1, time_ls=time_ls, batch_size=batch_size, operation_type="Concatinating")
stack = np.stack(stack_region)
save_loc = os.path.join(channel_stack_loc, f'{name}.npz')
_savez_atomic(save_loc, data=stack, filenames=filenames_region)
print(save_loc)
del stack
except Exception as e:
print(f"Error processing files, make sure filenames metadata is structured plate_well_field_time.npy")
print(f"Error: {e}")
else:
for file in os.listdir(src):
if file.endswith('.npy'):
path = os.path.join(src, file)
paths.append(path)
if randomize:
random.shuffle(paths)
nr_files = len(paths)
batch_index = 0
stack_ls = []
filenames_batch = []
for i, path in enumerate(paths):
start = time.time()
array = np.load(path)
array = np.take(array, channels, axis=2)
stack_ls.append(array)
filenames_batch.append(os.path.basename(path))
stop = time.time()
duration = stop - start
time_ls.append(duration)
files_processed = i+1
files_to_process = nr_files
print_progress(files_processed, files_to_process, n_jobs=1, time_ls=time_ls, batch_size=batch_size, operation_type="Concatinating")
if (i+1) % batch_size == 0 or i+1 == nr_files:
unique_shapes = {arr.shape[:-1] for arr in stack_ls}
if len(unique_shapes) > 1:
max_dims = np.max(np.array(list(unique_shapes)), axis=0)
print(f'Warning: arrays with multiple shapes found in batch {i+1}. Padding arrays to max X,Y dimentions {max_dims}')
padded_stack_ls = []
for arr in stack_ls:
pad_width = [(0, max_dim - dim) for max_dim, dim in zip(max_dims, arr.shape[:-1])]
pad_width.append((0, 0))
padded_arr = np.pad(arr, pad_width)
padded_stack_ls.append(padded_arr)
stack = np.stack(padded_stack_ls)
else:
stack = np.stack(stack_ls)
save_loc = os.path.join(channel_stack_loc, f'stack_{batch_index}.npz')
_savez_atomic(save_loc, data=stack, filenames=filenames_batch)
batch_index += 1
del stack
stack_ls = []
filenames_batch = []
padded_stack_ls = []
print(f'All files concatenated and saved to:{channel_stack_loc}')
return channel_stack_loc
def _normalize_img_batch(stack, channels, save_dtype, settings):
"""
Normalize the stack of images.
Each channel takes the background floor, signal-to-noise anchor and
background-removal switch of the object whose ``<object>_channel`` names
it: the nucleus, the cell, the pathogen, or any organelle slot the run
enables (``organelle``, ``organelleb``, ...). A slot reads
``<slot>_background``, ``<slot>_signal_to_noise`` and its background
switch, ``remove_background_organelle`` for slot 1 and
``remove_background_organelle_N`` for slot N (the lettered
``remove_background_organelleb`` is still read when a caller passes
settings that were never folded). A channel no object names keeps the generic
``background``, ``Signal_to_noise`` and ``remove_background``. The three
are read one by one, so one a slot does not carry, or carries empty, keeps
the value the channel already had, and when two objects name the same
channel the later one in that order wins each value it carries. A slot
sharing a channel with the nucleus, the cell or the pathogen therefore
normalises it by the slot's floor and anchor.
Args:
stack (numpy.ndarray): The stack of images to normalize.
lower_percentile (int): Lower percentile value for normalization.
save_dtype (numpy.dtype): Data type for saving the normalized stack.
settings (dict): keword arguments
Returns:
numpy.ndarray: The normalized stack.
"""
normalized_stack = np.zeros_like(stack, dtype=np.float32)
return _normalize_img_channels(
normalized_stack, channels, save_dtype, settings,
lambda channel: stack[..., channel])
def _normalize_img_channels(normalized_stack, channels, save_dtype, settings,
load_channel, output_columns=None,
workspace_dir=None):
"""Fill a float32 output from mutable channels in their source order.
:param normalized_stack: zero-filled destination with full or selected channels.
:param channels: source channel indices to normalize.
:param save_dtype: requested output dtype.
:param settings: object-specific background and percentile settings.
:param load_channel: callable returning one mutable source channel.
:param output_columns: optional source-channel to destination-column mapping.
:param workspace_dir: optional private directory for native mapped quantiles.
:returns: normalized output converted to ``save_dtype``.
"""
from .cancellation import checkpoint
from .utils import print_progress
channels = [int(c) for c in channels]
from .organelle_types import _background_switch_key
organelle_slot_channels = [
(role, settings.get(f'{role}_channel'))
for role in enabled_organelle_roles(settings)]
time_ls = []
for i, channel in enumerate(channels):
output_column = channel if output_columns is None else output_columns[channel]
start = time.time()
background = settings.get('background', 100)
signal_threshold = settings.get('Signal_to_noise', 10) * background
remove_background = settings.get('remove_background', False)
if settings.get('nucleus_channel') is not None and channel == settings['nucleus_channel']:
background = settings['nucleus_background']
signal_threshold = settings['nucleus_signal_to_noise']*settings['nucleus_background']
remove_background = settings['remove_background_nucleus']
if settings.get('cell_channel') is not None and channel == settings['cell_channel']:
background = settings['cell_background']
signal_threshold = settings['cell_signal_to_noise']*settings['cell_background']
remove_background = settings['remove_background_cell']
if settings.get('pathogen_channel') is not None and channel == settings['pathogen_channel']:
background = settings['pathogen_background']
signal_threshold = settings['pathogen_signal_to_noise']*settings['pathogen_background']
remove_background = settings['remove_background_pathogen']
for role, role_channel in organelle_slot_channels:
if channel != role_channel:
continue
role_background = settings.get(f'{role}_background')
if role_background is not None:
background = role_background
role_signal_to_noise = settings.get(f'{role}_signal_to_noise')
if role_signal_to_noise is None:
role_signal_to_noise = settings.get('Signal_to_noise', 10)
signal_threshold = role_signal_to_noise * background
role_remove_background = settings.get(
_background_switch_key(role),
settings.get(f'remove_background_{role}'))
if role_remove_background is not None:
remove_background = role_remove_background
single_channel = load_channel(channel)
if workspace_dir is not None and not hasattr(single_channel, '_spacr_native_cleanup'):
_retain_native_memmap(
single_channel, remove_path=os.path.join(
workspace_dir, f'channel-{channel}.npy'))
try:
print(f'Processing channel {channel}: background={background}, signal_threshold={signal_threshold}, remove_background={remove_background}')
quantile_path = None
if workspace_dir is None:
if remove_background:
single_channel[single_channel < background] = 0
non_zero_single_channel = single_channel[single_channel != 0]
else:
non_zero_single_channel, quantile_path = (
_native_nonzero_vector(single_channel, workspace_dir,
channel, background, remove_background))
try:
if not non_zero_single_channel.size:
if workspace_dir is not None:
normalized_stack[..., output_column] = 0
continue
global_lower = np.percentile(
non_zero_single_channel, settings['lower_percentile'],
overwrite_input=True)
global_upper = None
for upper_p in np.linspace(98, 99.5, num=16):
upper_value = np.percentile(
non_zero_single_channel, upper_p, overwrite_input=True)
if upper_value >= signal_threshold:
global_upper = upper_value
break
if global_upper is None:
global_upper = np.percentile(
non_zero_single_channel, 99.5, overwrite_input=True)
finally:
try:
if workspace_dir is not None and quantile_path is not None:
_close_private_memmap(non_zero_single_channel)
finally:
del non_zero_single_channel
if workspace_dir is not None:
checkpoint()
print(f'Channel {channel}: global_lower={global_lower}, global_upper={global_upper}, Signal-to-noise={global_upper / global_lower}')
if workspace_dir is None:
for array_index in range(single_channel.shape[0]):
arr_2d = single_channel[array_index]
for plane_index in np.ndindex(*arr_2d.shape[:-2]):
normalized_stack[(array_index, *plane_index, Ellipsis, output_column)] = (
exposure.rescale_intensity(
arr_2d[plane_index], in_range=(global_lower, global_upper),
out_range=(0, 1)))
else:
for time_index in range(single_channel.shape[0]):
for z_index in range(single_channel.shape[1]):
checkpoint()
plane = single_channel[time_index, z_index]
try:
normalized_stack[time_index, z_index, ..., output_column] = (
exposure.rescale_intensity(
plane, in_range=(global_lower, global_upper),
out_range=(0, 1)))
finally:
del plane
stop = time.time()
duration = stop - start
time_ls.append(duration)
files_processed = i+1
files_to_process = len(channels)
print_progress(files_processed, files_to_process, n_jobs=1, time_ls=time_ls, batch_size=None, operation_type=f"Normalizing")
finally:
if workspace_dir is not None:
try:
_close_private_memmap(single_channel)
finally:
del single_channel
return normalized_stack.astype(save_dtype, copy=False)
@contextmanager
def _native_map_workspace(*, prefix, dir, parent=None):
"""Retain private directories while borrowed native arrays still exist.
:param prefix: private temporary-directory prefix.
:param dir: existing writable parent directory.
:param parent: optional outer workspace retained by the same arrays.
:returns: context yielding the private path and directory-owner tuple.
"""
directory = tempfile.TemporaryDirectory(prefix=prefix, dir=dir)
workspace = {
'name': directory.name,
'owners': (directory, *(parent['owners'] if parent is not None else ())),
}
try:
yield workspace
finally:
workspace['owners'] = ()
directory = None
def _finalize_native_memmap(handle, path, owners):
"""Close before unlink while keeping private parent directories alive.
:param handle: raw mmap handle that does not retain its NumPy owner.
:param path: private disposable file, or None for a staged source.
:param owners: temporary directories retained until the handle closes.
:returns: None.
"""
try:
handle.close()
finally:
if path is not None and os.path.exists(path):
os.unlink(path)
def _retain_native_memmap(mapped, owners=(), *, remove_path=None):
"""Bind private-map cleanup to the last array owner instead of scope exit.
:param mapped: newly opened native normalization map.
:param owners: temporary directories to retain for diagnostic array views.
:param remove_path: optional private disposable file to unlink after closing.
:returns: the unchanged map, with its original dtype, shape and data.
"""
import weakref
mapped._spacr_native_owners = owners
mapped._spacr_native_cleanup = weakref.finalize(
mapped, _finalize_native_memmap, mapped._mmap,
os.fspath(remove_path) if remove_path is not None else None, owners)
return mapped
def _close_private_memmap(mapped):
"""Flush private writes and close maps without registered lifetime ownership.
Registered native maps remain valid until their last borrowed array dies;
their callers drop local references after flushing.
:param mapped: private writable native map to flush and release.
:returns: None.
"""
try:
mapped.flush()
finally:
if not hasattr(mapped, '_spacr_native_cleanup'):
mapped._mmap.close()
def _native_workspace_preflight(workspace, output_shape, selected_count,
source_dtype, filenames):
"""Require space for every simultaneous private map and staged archive."""
import errno
import math
voxels = math.prod(output_shape[:-1])
selected_bytes = voxels * selected_count * np.dtype(np.float32).itemsize
channel_bytes = voxels * np.dtype(source_dtype).itemsize
archive_payload = selected_bytes + np.asarray(filenames).nbytes + 8192
archive_bound = (archive_payload + archive_payload // 4096
+ archive_payload // 16384 + archive_payload // 33554432
+ 26 + 8192)
required = selected_bytes + 2 * channel_bytes + archive_bound
if hasattr(os, 'statvfs'):
stat = os.statvfs(workspace)
available = stat.f_bavail * stat.f_frsize
else:
available = shutil.disk_usage(workspace).free
if available < required:
raise OSError(errno.ENOSPC,
'Native T-by-Z private normalization needs more free disk space',
os.fspath(workspace))
def _reserve_private_memmap(mapped):
"""Allocate a new private map's file before any mapped data writes."""
import errno
from .cancellation import checkpoint
descriptor = os.open(mapped.filename, os.O_RDWR)
try:
size = os.fstat(descriptor).st_size
if hasattr(os, 'posix_fallocate'):
try:
os.posix_fallocate(descriptor, 0, size)
return
except OSError as error:
if error.errno not in (errno.ENOSYS, errno.EOPNOTSUPP):
raise
offset = mapped.offset
block = bytes(1024 * 1024)
while offset < size:
checkpoint()
chunk = block[:min(len(block), size - offset)]
if hasattr(os, 'pwrite'):
written = os.pwrite(descriptor, chunk, offset)
else:
os.lseek(descriptor, offset, os.SEEK_SET)
written = os.write(descriptor, chunk)
if written <= 0:
raise OSError(errno.ENOSPC, 'Native private map reservation failed')
offset += written
os.fsync(descriptor)
finally:
os.close(descriptor)
def _native_nonzero_vector(single_channel, workspace_dir, channel,
background, remove_background):
"""Stage scalar-order nonzero values without a field-sized boolean copy."""
from .cancellation import checkpoint
count = 0
for time_index in range(single_channel.shape[0]):
for z_index in range(single_channel.shape[1]):
checkpoint()
plane = single_channel[time_index, z_index]
if remove_background:
plane[plane < background] = 0
count += np.count_nonzero(plane)
del plane
if not count:
return np.empty(0, dtype=single_channel.dtype), None
path = os.path.join(workspace_dir, f'channel-{channel}-quantiles.bin')
vector = None
try:
vector = _retain_native_memmap(
np.memmap(path, mode='w+', dtype=single_channel.dtype, shape=(count,)),
getattr(single_channel, '_spacr_native_owners', ()), remove_path=path)
_reserve_private_memmap(vector)
offset = 0
for time_index in range(single_channel.shape[0]):
for z_index in range(single_channel.shape[1]):
checkpoint()
plane = single_channel[time_index, z_index]
values = plane[plane != 0]
vector[offset:offset + values.size] = values
offset += values.size
del plane, values
return vector, path
except BaseException:
if vector is not None:
try:
_close_private_memmap(vector)
finally:
vector = None
elif os.path.exists(path):
os.unlink(path)
raise
_PARTIAL_SUFFIX = '.partial'
_DAMAGED_SUFFIX = '.damaged'
def _replace_atomically(output_path, write, prefix='.spacr_tmp_'):
"""Write a file through a hidden sibling and rename it into place.
A process killed while ``write`` runs leaves the sibling behind and the
final name untouched: either absent or still holding the previous
complete file. The sibling is named ``<prefix><random>.partial``, so no
listing that selects ``*.npy`` or ``*.npz`` ever picks it up, and
:func:`_sweep_partial_writes` removes it on a later run.
The sibling is created with :func:`open` in exclusive mode, so it gets
the permissions the process's umask gives any new file, as
:func:`numpy.save` onto the final name did, rather than the owner-only
mode of :func:`tempfile.mkstemp`.
:param output_path: final path.
:param write: callable that receives the open binary handle of the
sibling and writes the complete content to it.
:param prefix: leading part of the sibling's name.
:returns: ``output_path`` after the flushed sibling has replaced it.
"""
output_path = os.fspath(output_path)
directory = os.path.dirname(output_path) or '.'
os.makedirs(directory, exist_ok=True)
while True:
temporary = os.path.join(
directory, f'{prefix}{os.urandom(8).hex()}{_PARTIAL_SUFFIX}')
try:
opened = open(temporary, 'xb')
except FileExistsError:
continue
break
try:
with opened as handle:
write(handle)
handle.flush()
os.fsync(handle.fileno())
os.replace(temporary, output_path)
except BaseException:
try:
os.remove(temporary)
except OSError:
pass
raise
return output_path
def _save_npz_atomic(output_path, **arrays):
"""Write a compressed NumPy archive by atomically replacing its path.
:param output_path: final ``.npz`` path.
:param arrays: named arrays passed to :func:`numpy.savez_compressed`.
:returns: ``output_path`` after the durable replacement.
"""
return _replace_atomically(
output_path, lambda handle: np.savez_compressed(handle, **arrays),
prefix='.spacr_npz_')
def _savez_atomic(output_path, **arrays):
"""Write an uncompressed NumPy archive by atomically replacing its path.
:param output_path: final ``.npz`` path.
:param arrays: named arrays passed to :func:`numpy.savez`.
:returns: ``output_path`` after the durable replacement.
"""
return _replace_atomically(
output_path, lambda handle: np.savez(handle, **arrays),
prefix='.spacr_npz_')
def _normalized_npz_field_ids(src):
"""Return exact field stems carried by V1 normalised mask archives.
:param src: the V1 ``masks/`` directory.
:returns: sorted, de-duplicated field stems as a tuple.
:raises FileNotFoundError: when no normalised archives exist.
:raises ValueError: when an archive has no ``filenames`` manifest.
"""
archives = sorted(
os.path.join(src, name) for name in _listdir_visible(src)
if name.endswith('.npz'))
if not archives:
raise FileNotFoundError(
'preprocess=False with segmentation illumination requires the '
f'normalised V1 .npz inputs in {src}; none were found.')
fields = set()
for path in archives:
with np.load(path, allow_pickle=False) as archive:
if 'filenames' not in archive:
raise ValueError(
'preprocess=False cannot validate segmentation '
f'illumination because {path} has no filenames manifest.')
for name in np.asarray(archive['filenames']).reshape(-1):
fields.add(os.path.splitext(os.path.basename(str(name)))[0])
return tuple(sorted(fields))
def _publish_v1_normalized_archives(staging_dir, output_dir):
"""Replace the complete published V1 NPZ set, with in-process rollback.
:param staging_dir: directory containing only the new, fully written NPZs.
:param output_dir: ``masks/`` directory consumed by the V1 segmenters.
:returns: final archive paths in stable name order.
Existing archives are derived preprocessing artifacts, but they are moved
to a sibling backup until every new archive is in place. A raised replace
therefore restores the prior complete set rather than leaving a mixture
from two runs; a process crash still leaves provenance incomplete, so no
caller can accept a partially published set as corrected.
"""
staged = sorted(
name for name in _listdir_visible(staging_dir) if name.endswith('.npz'))
if not staged:
raise ValueError('cannot publish an empty V1 normalized archive set')
previous = sorted(
name for name in _listdir_visible(output_dir) if name.endswith('.npz'))
backup_dir = tempfile.mkdtemp(
prefix='.spacr_previous_v1_npz_', dir=os.path.dirname(output_dir))
moved_previous = []
moved_staged = []
try:
for name in previous:
os.replace(os.path.join(output_dir, name),
os.path.join(backup_dir, name))
moved_previous.append(name)
for name in staged:
os.replace(os.path.join(staging_dir, name),
os.path.join(output_dir, name))
moved_staged.append(name)
except BaseException:
for name in reversed(moved_staged):
published = os.path.join(output_dir, name)
if os.path.exists(published):
os.replace(published, os.path.join(staging_dir, name))
for name in reversed(moved_previous):
saved = os.path.join(backup_dir, name)
if os.path.exists(saved):
os.replace(saved, os.path.join(output_dir, name))
shutil.rmtree(backup_dir, ignore_errors=True)
raise
shutil.rmtree(backup_dir)
shutil.rmtree(staging_dir)
return tuple(os.path.join(output_dir, name) for name in staged)
def _invalidate_v1_object_masks(output_dir):
"""Remove object masks made from the superseded normalised NPZ set.
:param output_dir: V1 ``masks/`` directory containing derived mask stacks.
:returns: removed mask-stack directory names in stable order.
A fresh illumination application changes the pixels presented to every
segmenter. Keeping an existing ``*_mask_stack`` would let both the outer
completeness check and per-batch resume check reuse masks drawn from the
old, potentially uncorrected pixels.
"""
removed = []
for name in sorted(os.listdir(output_dir)):
path = os.path.join(output_dir, name)
if name.endswith('_mask_stack') and os.path.isdir(path):
shutil.rmtree(path)
removed.append(name)
return tuple(removed)
def _invalidate_v1_segmentation_outputs(src):
"""Remove V1 outputs derived from a superseded segmentation-input set.
:param src: experiment root containing ``masks/`` and optional ``merged/``.
:returns: names of removed mask-stack directories and whether ``merged/``
was removed.
This runs only from the full Mask pipeline immediately before it redraws
masks. Removing the complete ``merged/`` directory is intentional: a
shorter rerun must not leave an old field that is absent from the new raw
stack but would otherwise still be measured later.
"""
masks_dir = os.path.join(src, 'masks')
removed_masks = (
_invalidate_v1_object_masks(masks_dir)
if os.path.isdir(masks_dir) else ())
merged_dir = os.path.join(src, 'merged')
removed_merged = os.path.isdir(merged_dir)
if removed_merged:
shutil.rmtree(merged_dir)
return removed_masks, removed_merged
def _correct_v1_segmentation_batch(
stack, filenames, channels, settings, illumination_session, psf_session=None):
"""Correct selected V1 channels on a private batch copy.
:returns: ``(working_stack, field_ids)``; without a session the original
stack and an empty tuple are returned unchanged.
"""
if illumination_session is None and psf_session is None:
return stack, ()
from .measure_hooks import PreprocessingContext
working = np.array(stack, copy=True, dtype=(
np.float32 if psf_session is not None and psf_session.processes
else None))
field_ids = []
for index, filename in enumerate(filenames):
field_id = os.path.splitext(os.path.basename(str(filename)))[0]
context = PreprocessingContext(
file_name=os.path.basename(str(filename)),
channels=list(channels),
settings=settings,
)
field = (psf_session.unmix(stack[index]) if psf_session is not None
else stack[index])
selected = field[..., list(channels)]
corrected = (illumination_session.correct(field_id, selected, context)
if illumination_session is not None else selected)
if psf_session is not None:
corrected = psf_session.correct(corrected)
working[index][..., list(channels)] = corrected
field_ids.append(field_id)
return working, tuple(field_ids)
def _concatenate_and_normalize_impl(
src, channels, save_dtype=np.float32, settings=None,
illumination_session=None, archive_output_fldr=None,
only_fields=None, first_batch_index=0, psf_session=None):
"""Concatenate per-file channel arrays and normalise them into a single stack.
:param src: Directory containing per-FOV ``.npy`` channel arrays.
:param channels: Channel indices to keep in the output stack.
:param save_dtype: NumPy dtype for the saved normalised arrays.
Default ``np.float32``.
:param settings: Preprocessing settings dict. **Required** — it must
contain the background, signal-to-noise, randomize, timelapse,
batch_size and plotting keys used elsewhere in preprocessing. The
``None`` in the signature is kept only so the argument can still be
passed positionally; omitting it is an error.
:param illumination_session: optional segmentation-only illumination
session. It corrects private copies of the selected channels before
normalisation and records completion only after each NPZ is durable.
:param psf_session: optional PSF session captured for this run. Unmixes
each whole raw field first when unmixing is on, then applies the PSF,
or the whole enhancement chain when one is on, after illumination on
each field before padding or normalization;
preserves floating point intensities and records archive identities.
:param only_fields: when given, the field stems to normalise; every other
``.npy`` in ``src`` is left out. Used, without a timelapse, to rebuild
only the fields a damaged or missing archive held.
:param first_batch_index: number of the first ``stack_<n>_norm.npz``
written, so archives added next to a previous run's do not replace
them.
:returns: Path to the directory where normalised arrays were saved.
:raises ValueError: if ``settings`` is not supplied.
"""
if settings is None:
raise ValueError(
"concatenate_and_normalize requires a settings dict (it reads "
"'timelapse', 'randomize', 'batch_size', 'lower_percentile' and the "
"per-channel background / Signal_to_noise keys); pass the dict "
"returned by settings.set_default_settings_preprocess_img_data.")
from .utils import print_progress
from .plot import plot_arrays
channels = [int(c) for c in channels if c is not None]
"""
Concatenates and normalizes channel data from multiple files and saves the normalized data.
Args:
src (str): The source directory containing the channel data files.
channels (list): The list of channel indices to be concatenated and normalized.
randomize (bool, optional): Whether to randomize the order of the files. Defaults to True.
timelapse (bool, optional): Whether the channel data is from a timelapse experiment. Defaults to False.
batch_size (int, optional): The number of files to be processed in each batch. Defaults to 100.
backgrounds (list, optional): Background values for each channel. Defaults to [100, 100, 100].
remove_backgrounds (list, optional): Whether to remove background values for each channel. Defaults to [False, False, False].
lower_percentile (int, optional): Lower percentile value for normalization. Defaults to 2.
save_dtype (numpy.dtype, optional): Data type for saving the normalized stack. Defaults to np.float32.
signal_to_noise (list, optional): Signal-to-noise ratio thresholds for each channel. Defaults to [5, 5, 5].
signal_thresholds (list, optional): Signal thresholds for each channel. Defaults to [1000, 1000, 1000].
Returns:
str: The directory path where the concatenated and normalized channel data is saved.
"""
channels = [item for item in channels if item is not None]
print(f"Generating concatenated and normalized channel data for channels: {channels}")
paths = []
time_ls = []
output_fldr = os.path.join(os.path.dirname(src), 'masks')
os.makedirs(output_fldr, exist_ok=True)
archive_output_fldr = archive_output_fldr or output_fldr
ledger = RunLedger('concatenate_and_normalize')
intended_fields = []
if settings['timelapse']:
try:
source_npy_names = sorted(
name for name in _listdir_visible(src) if name.endswith('.npy'))
time_stack_path_lists = _generate_time_lists(source_npy_names)
grouped_names = sorted(
filename for group in time_stack_path_lists
for filename in group)
if ((illumination_session is not None or psf_session is not None) and
grouped_names != source_npy_names):
missing = sorted(set(source_npy_names) - set(grouped_names))
raise ValueError(
'illumination correction could not group every source '
f'NPY as plate_well_field_time: {missing}')
intended_fields = [
os.path.splitext(os.path.basename(str(filename)))[0]
for group in time_stack_path_lists for filename in group
]
for i, time_stack_list in enumerate(time_stack_path_lists):
start = time.time()
stack_region = []
filenames_region = []
for idx, file in enumerate(time_stack_list):
path = os.path.join(src, file)
if idx == 0:
parts = file.split('_')
name = parts[0] + '_' + parts[1] + '_' + parts[2]
array = np.load(path)
stack_region.append(array)
filenames_region.append(os.path.basename(path))
stop = time.time()
duration = stop - start
time_ls.append(duration)
files_processed = i+1
files_to_process = len(time_stack_path_lists)
print_progress(files_processed, files_to_process, n_jobs=1, time_ls=time_ls, batch_size=None, operation_type="Concatinating")
stack = np.stack(stack_region)
stack, _field_ids = _correct_v1_segmentation_batch(
stack, filenames_region, channels, settings,
illumination_session, psf_session)
normalized_stack = _normalize_img_batch(stack=stack,
channels=channels,
save_dtype=save_dtype,
settings=settings)
normalized_stack = normalized_stack[..., channels]
save_loc = os.path.join(
archive_output_fldr, f'{name}_norm_timelapse.npz')
arrays = dict(data=normalized_stack,
filenames=filenames_region)
_save_npz_atomic(save_loc, **arrays)
if i == 0 and settings.get('plot'):
plot_arrays(save_loc, settings['figuresize'], settings['cmap'], nr=settings['nr'], normalize=False)
print(save_loc)
del stack, normalized_stack
except Exception as e:
print(f"Error processing files, make sure filenames metadata is structured plate_well_field_time.npy")
print(f"Error: {e}")
if illumination_session is not None or psf_session is not None:
raise
else:
for file in sorted(_listdir_visible(src)):
if file.endswith('.npy'):
if (only_fields is not None and
os.path.splitext(file)[0] not in only_fields):
continue
path = os.path.join(src, file)
paths.append(path)
if settings['randomize']:
random.shuffle(paths)
intended_fields = [
os.path.splitext(os.path.basename(path))[0] for path in paths
]
nr_files = len(paths)
batch_index = first_batch_index
stack_ls = []
filenames_batch = []
time_ls = []
files_processed = 0
for i, path in enumerate(paths):
start = time.time()
with ledger.item(path, stage='load_npy',
echo=f"Error loading file {path}"):
array = np.load(path)
stack_ls.append(array)
filenames_batch.append(os.path.basename(path))
stop = time.time()
duration = stop - start
time_ls.append(duration)
files_processed += 1
files_to_process = nr_files
print_progress(files_processed, files_to_process, n_jobs=1, time_ls=time_ls, batch_size=None, operation_type="Concatinating")
if stack_ls and ((i + 1) % settings['batch_size'] == 0 or i + 1 == nr_files):
if psf_session is not None:
stack_ls = [
_correct_v1_segmentation_batch(
array[None], [filename], channels, settings,
illumination_session, psf_session)[0][0]
for array, filename in zip(stack_ls, filenames_batch)]
unique_shapes = {arr.shape[:-1] for arr in stack_ls}
if len(unique_shapes) > 1:
max_dims = np.max(np.array(list(unique_shapes)), axis=0)
print(f'Warning: arrays with multiple shapes found in batch {i + 1}. Padding arrays to max X,Y dimensions {max_dims}')
padded_stack_ls = []
for arr in stack_ls:
pad_width = [(0, max_dim - dim) for max_dim, dim in zip(max_dims, arr.shape[:-1])]
pad_width.append((0, 0))
padded_arr = np.pad(arr, pad_width)
padded_stack_ls.append(padded_arr)
stack = np.stack(padded_stack_ls)
else:
stack = np.stack(stack_ls)
if psf_session is None:
stack, _field_ids = _correct_v1_segmentation_batch(
stack, filenames_batch, channels, settings,
illumination_session)
normalized_stack = _normalize_img_batch(stack=stack,
channels=channels,
save_dtype=save_dtype,
settings=settings)
normalized_stack = normalized_stack[..., channels]
save_loc = os.path.join(
archive_output_fldr, f'stack_{batch_index}_norm.npz')
arrays = dict(data=normalized_stack,
filenames=filenames_batch)
_save_npz_atomic(save_loc, **arrays)
if batch_index == 0 and settings.get('plot'):
print(f"plotting: {save_loc}")
plot_arrays(save_loc, settings['figuresize'], settings['cmap'], nr=settings['nr'], normalize=False)
batch_index += 1
del stack, normalized_stack
stack_ls = []
filenames_batch = []
padded_stack_ls = []
if illumination_session is not None or psf_session is not None:
staged_fields = _normalized_npz_field_ids(archive_output_fldr)
if set(staged_fields) != set(intended_fields):
missing = sorted(set(intended_fields) - set(staged_fields))
extra = sorted(set(staged_fields) - set(intended_fields))
raise RuntimeError(
'incomplete illumination fields before V1 publication: '
f'missing={missing}, unexpected={extra}')
if psf_session is not None:
from .cancellation import checkpoint
checkpoint()
_publish_v1_normalized_archives(
archive_output_fldr, output_fldr)
_invalidate_v1_segmentation_outputs(os.path.dirname(src))
settings['resume'] = False
for session in (illumination_session, psf_session):
if session is not None:
for field_id in staged_fields:
session.mark_completed(field_id)
session.finish(intended_fields)
print(f'All files concatenated and normalized. Saved to: {output_fldr}')
ledger.finalize()
return output_fldr
[docs]
def concatenate_and_normalize(
src, channels, save_dtype=np.float32, settings=None,
illumination_session=None, psf_session=None):
"""Concatenate, optionally correct, and normalise V1 field arrays.
:param src: directory containing per-field ``.npy`` channel arrays.
:param channels: channel indices retained in the output archives.
:param save_dtype: NumPy dtype for normalised arrays. Defaults to float32.
:param settings: required preprocessing settings mapping.
:param illumination_session: optional segmentation-only correction
session. Corrected archives are staged privately and published as one
complete set; the staging directory is removed on success or failure.
:param psf_session: optional PSF session from ``spacr.psf_pipeline``.
Uses the same complete-set publication; cancellation leaves its
provenance incomplete and prevents reuse of partially processed data.
:returns: the ``masks/`` directory containing normalised NPZ archives.
"""
if illumination_session is None and psf_session is None:
return _concatenate_and_normalize_impl(
src, channels, save_dtype=save_dtype, settings=settings)
output_fldr = os.path.join(os.path.dirname(src), 'masks')
os.makedirs(output_fldr, exist_ok=True)
with tempfile.TemporaryDirectory(
prefix='.spacr_v1_npz_', dir=os.path.dirname(output_fldr)
) as staging_dir:
return _concatenate_and_normalize_impl(
src, channels, save_dtype=save_dtype, settings=settings,
illumination_session=illumination_session,
archive_output_fldr=staging_dir, psf_session=psf_session)
def _get_lists_for_normalization(settings):
"""
Get lists for normalization based on the provided settings.
Args:
settings (dict): A dictionary containing the settings for normalization.
Returns:
tuple: A tuple containing three lists - backgrounds, signal_to_noise, and signal_thresholds.
"""
backgrounds = []
signal_to_noise = []
signal_thresholds = []
remove_background = []
for ch in [settings['nucleus_channel'], settings['cell_channel'], settings['pathogen_channel']]:
if not ch is None:
if ch == settings['nucleus_channel']:
backgrounds.append(settings['nucleus_background'])
signal_to_noise.append(settings['nucleus_signal_to_noise'])
signal_thresholds.append(settings['nucleus_signal_to_noise']*settings['nucleus_background'])
remove_background.append(settings['remove_background_nucleus'])
elif ch == settings['cell_channel']:
backgrounds.append(settings['cell_background'])
signal_to_noise.append(settings['cell_signal_to_noise'])
signal_thresholds.append(settings['cell_signal_to_noise']*settings['cell_background'])
remove_background.append(settings['remove_background_cell'])
elif ch == settings['pathogen_channel']:
backgrounds.append(settings['pathogen_background'])
signal_to_noise.append(settings['pathogen_signal_to_noise'])
signal_thresholds.append(settings['pathogen_signal_to_noise']*settings['pathogen_background'])
remove_background.append(settings['remove_background_pathogen'])
return backgrounds, signal_to_noise, signal_thresholds, remove_background
def _normalize_stack(src, backgrounds=None, remove_backgrounds=None, lower_percentile=2, save_dtype=np.float32, signal_to_noise=None, signal_thresholds=None):
"""
Normalize the stack of images.
Args:
src (str): The source directory containing the stack of images.
backgrounds (list, optional): Background values for each channel. Defaults to [100, 100, 100].
remove_background (list, optional): Whether to remove background values for each channel. Defaults to [False, False, False].
lower_percentile (int, optional): Lower percentile value for normalization. Defaults to 2.
save_dtype (numpy.dtype, optional): Data type for saving the normalized stack. Defaults to np.float32.
signal_to_noise (list, optional): Signal-to-noise ratio thresholds for each channel. Defaults to [5, 5, 5].
signal_thresholds (list, optional): Signal thresholds for each channel. Defaults to [1000, 1000, 1000].
Returns:
None
"""
if backgrounds is None:
backgrounds = [100, 100, 100]
if remove_backgrounds is None:
remove_backgrounds = [False, False, False]
if signal_to_noise is None:
signal_to_noise = [5, 5, 5]
if signal_thresholds is None:
signal_thresholds = [1000, 1000, 1000]
paths = [os.path.join(src, file) for file in os.listdir(src) if file.endswith('.npz')]
output_fldr = os.path.join(os.path.dirname(src), 'masks')
os.makedirs(output_fldr, exist_ok=True)
time_ls = []
for file_index, path in enumerate(paths):
with np.load(path) as data:
stack = data['data']
filenames = data['filenames']
normalized_stack = np.zeros_like(stack, dtype=np.float32)
file = os.path.basename(path)
name, _ = os.path.splitext(file)
for chan_index, channel in enumerate(range(stack.shape[-1])):
n_channels = stack.shape[-1]
setting_lists = {
"backgrounds": backgrounds,
"remove_backgrounds": remove_backgrounds,
"signal_to_noise": signal_to_noise,
"signal_thresholds": signal_thresholds,
}
short = {
name: len(values) for name, values in setting_lists.items()
if len(values) < n_channels
}
if short:
details = ", ".join(
f"{name}={length}" for name, length in short.items())
raise ValueError(
f"Normalization stack has {n_channels} channels but "
f"per-channel settings are incomplete ({details}).")
single_channel = stack[:, :, :, channel]
background = backgrounds[chan_index]
signal_threshold = signal_thresholds[chan_index]
remove_background = remove_backgrounds[chan_index]
signal_2_noise = signal_to_noise[chan_index]
print(f'chan_index:{chan_index} background:{background} signal_threshold:{signal_threshold} remove_background:{remove_background} signal_2_noise:{signal_2_noise}')
if remove_background:
single_channel[single_channel < background] = 0
non_zero_single_channel = single_channel[single_channel != 0]
upper_p = 98.0
if non_zero_single_channel.size:
for upper_p in np.linspace(98, 100, num=100).tolist():
global_upper = np.percentile(
non_zero_single_channel, upper_p)
if global_upper >= signal_threshold:
break
arr_2d_normalized = np.zeros_like(single_channel, dtype=single_channel.dtype)
signal_to_noise_ratio_ls = []
time_ls = []
lower = upper = 0.0
for array_index in range(single_channel.shape[0]):
start = time.time()
arr_2d = single_channel[array_index, :, :]
non_zero_arr_2d = arr_2d[arr_2d != 0]
if non_zero_arr_2d.size > 0:
lower, upper = np.percentile(non_zero_arr_2d, (lower_percentile, upper_p))
signal_to_noise_ratio = upper / lower
else:
lower, upper = 0.0, 0.0
signal_to_noise_ratio = 0
signal_to_noise_ratio_ls.append(signal_to_noise_ratio)
average_stnr = np.mean(signal_to_noise_ratio_ls) if len(signal_to_noise_ratio_ls) > 0 else 0
if signal_to_noise_ratio > signal_2_noise:
arr_2d_rescaled = exposure.rescale_intensity(arr_2d, in_range=(lower, upper), out_range=(0, 1))
arr_2d_normalized[array_index, :, :] = arr_2d_rescaled
else:
arr_2d_normalized[array_index, :, :] = arr_2d
stop = time.time()
duration = (stop - start) * single_channel.shape[0]
time_ls.append(duration)
average_time = np.mean(time_ls) if len(time_ls) > 0 else 0
print(f'channels:{chan_index}/{stack.shape[-1] - 1}, arrays:{array_index + 1}/{single_channel.shape[0]}, Signal:{upper:.1f}, noise:{lower:.1f}, Signal-to-noise:{average_stnr:.1f}, Time/channel:{average_time:.2f}sec')
normalized_stack[:, :, :, channel] = arr_2d_normalized
save_loc = os.path.join(output_fldr, f'{name}_norm_stack.npz')
_savez_atomic(save_loc, data=normalized_stack.astype(save_dtype), filenames=filenames)
del normalized_stack, single_channel, arr_2d_normalized, stack, filenames
gc.collect()
return print(f'Saved stacks: {output_fldr}')
def _normalize_timelapse(src, lower_percentile=2, save_dtype=np.float32):
"""
Normalize the timelapse data by rescaling the intensity values based on percentiles.
Args:
src (str): The source directory containing the timelapse data files.
lower_percentile (int, optional): The lower percentile used to calculate the intensity range. Defaults to 1.
save_dtype (numpy.dtype, optional): The data type to save the normalized stack. Defaults to np.float32.
"""
paths = [os.path.join(src, file) for file in os.listdir(src) if file.endswith('.npz')]
output_fldr = os.path.join(os.path.dirname(src), 'masks')
os.makedirs(output_fldr, exist_ok=True)
for file_index, path in enumerate(paths):
with np.load(path) as data:
stack = data['data']
filenames = data['filenames']
normalized_stack = np.zeros_like(stack, dtype=save_dtype)
file = os.path.basename(path)
name, _ = os.path.splitext(file)
for chan_index in range(stack.shape[-1]):
single_channel = stack[:, :, :, chan_index]
for array_index in range(single_channel.shape[0]):
arr_2d = single_channel[array_index]
non_zero = arr_2d[arr_2d != 0]
if non_zero.size:
q_low, q_high = np.percentile(
non_zero, (lower_percentile, 98))
if q_high > q_low:
normalized_stack[array_index, :, :, chan_index] = (
exposure.rescale_intensity(
arr_2d, in_range=(q_low, q_high),
out_range='dtype'))
else:
normalized_stack[array_index, :, :, chan_index] = arr_2d
else:
normalized_stack[array_index, :, :, chan_index] = arr_2d
print(f'channels:{chan_index+1}/{stack.shape[-1]}, arrays:{array_index+1}/{single_channel.shape[0]}', end='\r')
save_loc = os.path.join(output_fldr, f'{name}_norm_timelapse.npz')
_savez_atomic(save_loc, data=normalized_stack, filenames=filenames)
del normalized_stack, stack, filenames
gc.collect()
print(f'\nSaved normalized stacks: {output_fldr}')
def _create_movies_from_npy_per_channel(src, fps=10):
"""
Create movies from numpy files per channel.
Args:
src (str): The source directory containing the numpy files.
fps (int, optional): Frames per second for the output movies. Defaults to 10.
"""
from .timelapse import _npz_to_movie
master_path = os.path.dirname(src)
save_path = os.path.join(master_path,'movies')
os.makedirs(save_path, exist_ok=True)
files = [f for f in _listdir_visible(src) if f.endswith('.npy')]
organized_files = {}
for f in files:
match = re.match(r'(\w+)_(\w+)_(\w+)_(\d+)\.npy', f)
if match:
plate, well, field, time = match.groups()
key = (plate, well, field)
if key not in organized_files:
organized_files[key] = []
organized_files[key].append((int(time), os.path.join(src, f)))
for key, file_list in organized_files.items():
plate, well, field = key
file_list.sort(key=lambda x: x[0])
_times, paths = zip(*file_list)
arrays = np.stack(tuple(map(np.load, paths)), axis=0)
filenames = list(map(os.path.basename, paths))
for channel in range(arrays.shape[-1]):
channel_arrays = arrays[..., channel]
channel_data_flat = channel_arrays.reshape(-1)
p1, p99 = np.percentile(channel_data_flat, [1, 99])
normalized_channel_arrays = [(np.clip((arr - p1) / (p99 - p1), 0, 1) * 255).astype(np.uint8) for arr in channel_arrays]
normalized_channel_arrays_3d = [arr[..., np.newaxis] for arr in normalized_channel_arrays]
channel_save_path = os.path.join(save_path, f'{plate}_{well}_{field}_channel_{channel}.mp4')
_npz_to_movie(normalized_channel_arrays_3d, filenames, channel_save_path, fps)
[docs]
def delete_empty_subdirectories(folder_path):
"""Recursively delete every empty subdirectory under ``folder_path``.
:param folder_path: Root directory to scan.
:returns: None
"""
for dirpath, dirnames, filenames in os.walk(folder_path, topdown=False):
for dirname in dirnames:
full_dir_path = os.path.join(dirpath, dirname)
try:
os.rmdir(full_dir_path)
print(f"Deleted empty directory: {full_dir_path}")
except OSError:
continue
[docs]
def select_fields(names, fields):
"""Keep only the ``names`` whose field is in ``fields``.
This filter allows mask generation to be rerun for selected fields without
processing every field on the plate.
:param names: stack file names, as written by
`_rename_and_organize_image_files`.
:param fields: what to keep. ``None`` or empty keeps everything, which
is the default and the behaviour every existing run has. A list, or
a comma-separated string, of field ids in any spelling the rest of
spaCR accepts -- ``'f3'``, ``3``, ``'F003'`` -- or a glob such as
``'f1*'`` matched against the field id.
:returns: the kept names, in the order given.
"""
import fnmatch
from . import schema
if fields is None or (hasattr(fields, '__len__') and not len(fields)):
return list(names)
if isinstance(fields, str):
wanted = [part.strip() for part in fields.split(',') if part.strip()]
elif isinstance(fields, (list, tuple, set)):
wanted = [str(part).strip() for part in fields if str(part).strip()]
else:
wanted = [str(fields).strip()]
if not wanted:
return list(names)
def field_of(token):
"""One field's canonical token, so 'f3', '3' and 'F003' are one field.
NORMALISED THROUGH `schema`, which is the same rule the file names
themselves were written by -- matching the raw text instead would
make a selection depend on which of three spellings the user typed.
Anything schema cannot read is lower-cased and passed through, so a
token from a convention it has not met still selects itself.
"""
try:
index = schema.field_index(token)
except Exception: # noqa: BLE001
index = None
return f'f{index}' if index is not None else str(token).strip().lower()
patterns = [field_of(w) if not any(c in str(w) for c in '*?[')
else str(w).strip().lower() for w in wanted]
kept = []
for name in names:
try:
this = schema.parse_field_stem(name).fieldID
except Exception: # noqa: BLE001
continue
this = str(this).lower()
if any(fnmatch.fnmatch(this, pattern) for pattern in patterns):
kept.append(name)
return kept
def _name_list(names, limit=8):
"""Join names for a message, saying how many were left out.
:param names: the names, in the order to show them.
:param limit: how many to show.
:returns: the joined names, ending in ``and <n> more`` when cut short.
"""
names = list(names)
shown = ', '.join(names[:limit])
if len(names) > limit:
shown += f' and {len(names) - limit} more'
return shown
def _sweep_partial_writes(folder):
"""Remove the siblings of atomic writes that were killed before their rename.
:param folder: a ``stack/`` or ``masks/`` folder; a missing one is skipped.
:returns: the names removed, sorted.
"""
if not os.path.isdir(folder):
return []
removed = []
for name in sorted(os.listdir(folder)):
if not (name.startswith(('.spacr_tmp_', '.spacr_npz_'))
and name.endswith(_PARTIAL_SUFFIX)):
continue
try:
os.remove(os.path.join(folder, name))
except OSError:
continue
removed.append(name)
if removed:
print(f'Removed {len(removed)} unfinished write(s) a killed run left '
f'in {folder}: {_name_list(removed)}')
return removed
def _set_aside(path):
"""Rename a damaged file so no listing of ``.npy`` or ``.npz`` files picks it up.
:param path: the damaged file.
:returns: its new path, ``<path>.damaged`` or, when that name is taken,
``<path>.damaged.<n>``.
"""
target = path + _DAMAGED_SUFFIX
counter = 1
while os.path.exists(target):
target = f'{path}{_DAMAGED_SUFFIX}.{counter}'
counter += 1
os.replace(path, target)
return target
def _npy_is_whole(path):
"""Decide, without reading its pixels, whether a ``.npy`` was written to the end.
The header is parsed and the file's length compared with the length its
declared shape needs (:func:`spacr.resume.validate_merged_field`). A dtype
that comparison cannot size is memory-mapped instead, which fails in the
same way on a short file.
:param path: the ``.npy`` file.
:returns: ``(ok, reason)``; ``reason`` is ``'done'`` when ``ok``, otherwise
``'empty'``, ``'truncated'`` or ``'unreadable'``.
"""
from .resume import REASON_DONE, REASON_UNREADABLE, validate_merged_field
ok, reason = validate_merged_field(path)
if ok or reason != REASON_UNREADABLE:
return ok, reason
try:
mapped = np.load(path, mmap_mode='r', allow_pickle=False)
except Exception:
return False, reason
del mapped
return True, REASON_DONE
def _inspect_normalized_archive(path, *, field_axis=0):
"""Decide whether a normalised ``.npz`` archive is whole, without inflating its pixels.
``numpy.savez_compressed`` writes the zip directory last, so an archive
cut short by a killed run has none and does not open. Beyond that, every
member's recorded extent has to fit inside the file. ``data.npy`` needs
a readable numeric-array header and enough declared bytes for its shape
and dtype; object-valued pixels are refused. The small ``filenames.npy``
array is read and must contain one filename per batch field. Legacy
object-valued filenames are not unpickled and cannot establish field
coverage. Pixel arrays are not materialized or fully CRC-scanned.
:param path: the ``.npz`` archive.
:param field_axis: axis of ``data`` named by the filenames vector; normally
zero, or the declared time axis for a time-stack archive.
:returns: ``(ok, reason, fields, planes)``. ``reason`` is ``'done'`` when
``ok``. ``fields`` is the tuple of field stems the archive lists, or
``None`` when it is damaged or lists them as an object array, which is
not read without unpickling. ``planes`` is the length of the last axis
of ``data``, or ``None`` when it is damaged.
"""
import math
import zipfile
import zlib
from numpy.lib import format as npy_format
try:
size = os.path.getsize(path)
except OSError as exc:
return False, f'unreadable ({exc})', None, None
if size == 0:
return False, 'empty: the write never started', None, None
try:
with zipfile.ZipFile(path) as archive:
members = {info.filename: info for info in archive.infolist()}
for required in ('data.npy', 'filenames.npy'):
if required not in members:
return False, f'holds no {required}', None, None
for info in members.values():
if info.header_offset + info.compress_size > size:
return (False, 'truncated: a member runs past the end of '
'the file', None, None)
with archive.open('data.npy') as member:
if npy_format.read_magic(member) == (1, 0):
shape, _, dtype = npy_format.read_array_header_1_0(member)
else:
shape, _, dtype = npy_format.read_array_header_2_0(member)
if dtype.hasobject:
return False, 'data.npy requires unpickling', None, None
expected = member.tell() + math.prod(shape) * dtype.itemsize
if members['data.npy'].file_size < expected:
return False, 'truncated: data.npy pixels are incomplete', None, None
planes = int(shape[-1]) if shape else None
with archive.open('filenames.npy') as member:
try:
names = npy_format.read_array(member, allow_pickle=False)
except ValueError as exc:
if 'allow_pickle' not in str(exc):
raise
return True, 'done', None, planes
if (not shape or not 0 <= field_axis < len(shape)
or names.ndim != 1 or names.size != shape[field_axis]):
return False, 'filenames count does not match data.npy fields', None, None
except zipfile.BadZipFile as exc:
return False, f'not a complete zip archive ({exc})', None, None
except (OSError, ValueError, EOFError, KeyError, zlib.error) as exc:
return (False, f'unreadable ({type(exc).__name__}: {exc})', None,
None)
fields = tuple(os.path.splitext(os.path.basename(str(name)))[0]
for name in np.asarray(names).reshape(-1))
return True, 'done', fields, planes
def _mask_batch_manifest(src, *, field_axis=0):
"""Inventory immutable batch identities before assigning mask workers.
Inspect headers and filename vectors without inflating image arrays, then
hash each archive in bounded chunks. Reject ambiguous output ownership,
unsafe filenames, unreadable archives and inputs changed during inspection.
No inputs or outputs are modified and no model is loaded.
:param src: directory of prepared NPZ batches.
:param field_axis: data axis identified by the filenames vector.
:returns: JSON-compatible records in archive-name order, carrying absolute
paths, SHA256 digests, byte sizes, field filenames and channel counts.
:raises ValueError: if a batch cannot safely belong to one worker.
:raises FileNotFoundError: if there are no prepared batches.
"""
import hashlib
root = Path(src).resolve()
paths = sorted(root / name for name in _listdir_visible(root)
if name.endswith('.npz'))
if not paths:
raise FileNotFoundError(f'No prepared NPZ mask batches in {root}')
owners = {}
records = []
for path in paths:
before = path.stat()
ok, reason, fields, planes = _inspect_normalized_archive(
path, field_axis=field_axis)
if not ok or fields is None:
raise ValueError(f'Cannot dispatch {path.name}: {reason if not ok else "unreadable field identities"}')
with np.load(path, allow_pickle=False) as archive:
names = archive['filenames']
if names.dtype.kind != 'U':
raise ValueError(f'{path.name}: field filenames must be Unicode strings')
names = names.tolist()
if not names:
raise ValueError(f'{path.name}: batch contains no fields')
for name in names:
if (not name or name.startswith('.') or '/' in name or '\\' in name
or ':' in name or not name.endswith('.npy')
or '\x00' in name):
raise ValueError(f'{path.name}: unsafe field filename {name!r}')
identity = os.path.normcase(name)
if identity in owners:
raise ValueError(f'Field {name!r} belongs to both {owners[identity]} and {path.name}')
owners[identity] = path.name
digest = hashlib.sha256()
with path.open('rb') as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b''):
digest.update(chunk)
after = path.stat()
if any(getattr(before, attr) != getattr(after, attr)
for attr in ('st_dev', 'st_ino', 'st_size', 'st_mtime_ns')):
raise ValueError(f'Batch changed while building its manifest: {path}')
records.append({'path': str(path), 'sha256': digest.hexdigest(),
'bytes': after.st_size, 'fields': names, 'planes': planes})
return records
def _set_aside_damaged_stacks(stack_path):
"""Check every ``stack/*.npy`` before it is reused, and set the damaged ones aside.
:param stack_path: the ``stack/`` folder; a missing one is skipped.
:returns: ``(name, reason)`` for each file renamed to ``<name>.damaged``.
"""
if not os.path.isdir(stack_path):
return []
_sweep_partial_writes(stack_path)
names = sorted(name for name in _listdir_visible(stack_path)
if name.endswith('.npy'))
damaged = []
for name in names:
path = os.path.join(stack_path, name)
ok, reason = _npy_is_whole(path)
if not ok:
_set_aside(path)
damaged.append((name, reason))
if damaged:
print(f'Checked {len(names)} field stack(s) in {stack_path}: '
f'{len(damaged)} damaged, set aside as <name>.damaged: '
f'{_name_list(f"{name} ({reason})" for name, reason in damaged)}')
return damaged
def _set_aside_names(folder, extension):
"""Name the files in ``folder`` that were set aside as damaged, by the name they had.
:param folder: a ``stack/`` or ``masks/`` folder; a missing one holds none.
:param extension: ``'.npy'`` or ``'.npz'``.
:returns: the original names (``<name>.damaged`` and
``<name>.damaged.<n>`` both give ``<name>``), sorted and de-duplicated.
"""
pattern = re.compile(
r'(.+' + re.escape(extension) + r')' + re.escape(_DAMAGED_SUFFIX)
+ r'(\.\d+)?$')
try:
names = _listdir_visible(folder)
except OSError:
return []
return sorted({match.group(1) for match in map(pattern.match, names)
if match})
def _report_unrebuilt_stacks(stack_path, src):
"""Report the field stacks that were set aside as damaged and never built again.
Such a field has a ``stack/<name>.damaged`` and no ``stack/<name>``, so it
gets no normalised archive, no masks and no ``merged/`` array. Each one is
recorded as a failure, so the run ends on ``RUN INCOMPLETE`` instead of
finishing with the field missing, and it is reported again on every run
until its raw images are back or the ``.damaged`` file is deleted.
:param stack_path: the ``stack/`` folder.
:param src: the plate folder, named in the message.
:returns: the names of the stacks that are still missing.
"""
missing = [name for name in _set_aside_names(stack_path, '.npy')
if not os.path.exists(os.path.join(stack_path, name))]
if not missing:
return []
print(f'{len(missing)} damaged field stack(s) could not be built again, '
f'because no raw image left in {src}, its orig/ or its channel '
f'folders builds a field of that name: {_name_list(missing)}. Those '
f'fields get no masks and no merged/ array. Put their raw images '
f'back and run again to build them, or delete the <name>.damaged '
f'file(s) in {stack_path} to go on without them.')
ledger = RunLedger('field_stacks')
for name in missing:
ledger.record_failure(
os.path.join(stack_path, name), stage='rebuild_damaged_stack',
exc=(f'set aside as {name}{_DAMAGED_SUFFIX} because it was '
f'damaged, and not built again: no raw image left builds '
f'this field'))
ledger.finalize()
return missing
def _check_normalized_archives(masks_path):
"""Check every ``masks/*.npz`` before it is reused, and set the damaged ones aside.
:param masks_path: the ``masks/`` folder.
:returns: dict with ``archives`` (the names checked), ``damaged``
(``(name, reason)`` for each renamed to ``<name>.damaged``),
``earlier`` (archives an earlier run set aside), ``unlisted`` (whole
archives that list their fields as an object array), ``covered``
(the field stems the whole archives list) and ``planes`` (the
channel counts of the whole archives).
"""
_sweep_partial_writes(masks_path)
earlier = _set_aside_names(masks_path, '.npz')
archives = sorted(name for name in _listdir_visible(masks_path)
if name.endswith('.npz'))
damaged, unlisted, covered, plane_counts = [], [], set(), set()
for name in archives:
ok, reason, fields, planes = _inspect_normalized_archive(
os.path.join(masks_path, name))
if not ok:
_set_aside(os.path.join(masks_path, name))
damaged.append((name, reason))
continue
plane_counts.add(planes)
if fields is None:
unlisted.append(name)
else:
covered.update(fields)
if archives:
print(f'Checked {len(archives)} normalised archive(s) in {masks_path}: '
f'{len(archives) - len(damaged)} whole, {len(damaged)} damaged.')
if damaged:
print(f'Set aside as <name>.damaged: '
f'{_name_list(f"{name} ({reason})" for name, reason in damaged)}')
if earlier:
print(f'{masks_path} also holds {len(earlier)} archive(s) an earlier '
f'run set aside as damaged: {_name_list(earlier)}.')
return {'archives': archives, 'damaged': damaged, 'earlier': earlier,
'unlisted': unlisted, 'covered': covered, 'planes': plane_counts}
def _check_archives_without_preprocessing(src):
"""Check ``masks/*.npz`` before a run with ``preprocess`` off segments them.
With ``preprocess`` off nothing normalises the plate again, so an
archive a killed run cut short is not rebuilt. It is set aside as
``<name>.damaged`` like any other, and the run stops with an error that
names it and says what to do, rather than with the
:class:`zipfile.BadZipFile` the segmenter would raise on it. Fields of
``stack/`` that no whole archive lists are reported, not normalised.
Earlier quarantines continue to raise on subsequent runs until a valid
same-name replacement exists, or readable archives cover every existing
stack field. Legacy archives with unreadable object-valued filenames
cannot establish that coverage. Quarantined evidence is retained.
:param src: the plate folder holding ``masks/``.
:returns: the names of the whole archives.
:raises FileNotFoundError: when an archive is damaged.
"""
masks_path = os.path.join(src, 'masks')
stack_path = os.path.join(src, 'stack')
checked = _check_normalized_archives(masks_path)
stack_fields = _stack_field_stems(stack_path)
earlier = set(checked['earlier']) - set(checked['archives'])
if (stack_fields and not checked['unlisted']
and stack_fields <= checked['covered']):
earlier.clear()
damaged = sorted({name for name, _ in checked['damaged']} | earlier)
if damaged:
raw = (_raw_image_names(src) or
_raw_image_names(os.path.join(src, 'orig')) or
_channel_folders(src))
if stack_fields:
way_out = (f' Turn preprocess on and run again: the fields they '
f'held are normalised again from {stack_path}, into '
f'new archives.')
elif raw:
way_out = (' Turn preprocess on and run again: the fields they '
'held are built again from the raw images and '
'normalised into new archives.')
else:
way_out = (f' Neither stack/ nor the raw images are left to build '
f'them again from. Point src at a copy of the plate\'s '
f'raw images, or move the <name>.damaged file(s) out of '
f'{masks_path} and run again to segment only the '
f'fields the whole archives hold.')
raise FileNotFoundError(
f'{len(damaged)} normalised archive(s) in {masks_path} were '
f'damaged by an earlier run '
f'({_name_list(damaged)}) and have been set '
f'aside as <name>.damaged. preprocess is off, so they are not '
f'built again.{way_out}')
if not checked['unlisted']:
missing = stack_fields - checked['covered']
if missing and checked['archives']:
print(f'{len(missing)} field(s) in stack/ are in no archive in '
f'{masks_path}: {_name_list(sorted(missing))}. preprocess '
f'is off, so they are not normalised and get no masks; '
f'turn preprocess on to add them.')
return checked['archives']
def _next_archive_index(masks_path):
"""Return the first ``n`` that no ``stack_<n>_norm.npz`` in ``masks_path`` uses.
:param masks_path: the ``masks/`` folder. Archives already set aside as
damaged count as using their number.
:returns: one more than the highest number in use, or 0.
"""
used = [-1]
for name in _listdir_visible(masks_path):
match = re.match(r'stack_(\d+)_norm\.npz', name)
if match:
used.append(int(match.group(1)))
return max(used) + 1
def _rebuild_stacks_from_raw(settings, src):
"""Build the field stacks ``stack/`` lacks, from raw images ``orig/`` still holds.
:param settings: the preprocessing settings. ``metadata_type``,
``custom_regex``, ``batch_size``, ``timelapse`` and
``save_original_images`` are read; nothing is written to them.
:param src: the plate folder.
:returns: the number of stacks built; 0 when there are no raw images.
"""
from .utils import _get_regex
raw = (_raw_image_names(src) or
_raw_image_names(os.path.join(src, 'orig')))
if not raw:
return 0
image_format = Counter(
name.rsplit('.', 1)[-1].lower() for name in raw).most_common(1)[0][0]
metadata_type = settings.get('metadata_type', 'cellvoyager')
regex = _get_regex(metadata_type, image_format,
settings.get('custom_regex'))
stack_path = os.path.join(src, 'stack')
before = len(_stack_field_stems(stack_path))
_rename_and_organize_image_files(
src, regex, int(settings.get('batch_size') or 50), metadata_type,
list(_RAW_IMAGE_SUFFIXES), timelapse=settings.get('timelapse', False),
save_original_images=settings.get('save_original_images', True))
return len(_stack_field_stems(stack_path)) - before
def _resume_normalized_archives(settings, src, mask_channels):
"""Check the archives an earlier run left in ``masks/`` before they are reused.
A plate folder that already holds ``masks/`` skips preprocessing. Before
it does, each ``stack/*.npy`` (:func:`_npy_is_whole`) and each
``masks/*.npz`` (:func:`_inspect_normalized_archive`) is checked, and a
damaged file is renamed to ``<name>.damaged`` and reported by name. A
field stack that is missing is built again from the raw images in the
plate folder or ``orig/``, or from its channel folders, when they are
there; one that cannot be is recorded as a failure
(:func:`_report_unrebuilt_stacks`). The fields of ``stack/`` that no
whole archive lists -- those of a damaged archive, and those a killed run
never reached -- are normalised again, into new archives numbered after
the existing ones.
:param settings: the preprocessing settings; not modified.
:param src: the plate folder holding ``masks/`` and ``stack/``.
:param mask_channels: the channel indices the archives keep.
:returns: True when ``masks/`` holds a whole archive for every field in
``stack/`` and preprocessing can be skipped. False when fields are
missing from an illumination-corrected or timelapse set, which has
to be rebuilt whole from ``stack/``.
:raises FileNotFoundError: when ``masks/`` holds an archive set aside as
damaged, by this run or an earlier one, and ``stack/`` holds no field
to rebuild it from.
"""
stack_path = os.path.join(src, 'stack')
masks_path = os.path.join(src, 'masks')
from zipfile import BadZipFile
from .psf_pipeline import (validate_psf_resume, _record_path,
processing_requested)
psf_tracked = (processing_requested(settings) or
_record_path(src).exists())
if psf_tracked:
try:
validate_psf_resume(settings, src, mask_channels,
expected_fields=_normalized_npz_field_ids(masks_path))
except (ValueError, OSError, EOFError, BadZipFile):
return False
_set_aside_damaged_stacks(stack_path)
try:
_rebuild_stacks_from_raw(settings, src)
if _channel_folders(src):
_merge_channels(src, plot=False)
except Exception as exc:
print(f'Could not build missing field stacks from the raw images: '
f'{type(exc).__name__}: {exc}')
_report_unrebuilt_stacks(stack_path, src)
checked = _check_normalized_archives(masks_path)
damaged = checked['damaged']
set_aside = sorted({name for name, _ in damaged} | set(checked['earlier']))
unlisted, covered = checked['unlisted'], checked['covered']
plane_counts = checked['planes']
stack_fields = _stack_field_stems(stack_path)
if set_aside and not stack_fields:
raise FileNotFoundError(
f'{len(set_aside)} normalised archive(s) in {masks_path} were '
f'damaged by an earlier run ({_name_list(set_aside)}) '
f'and are set aside as <name>.damaged, and neither {stack_path} '
f'nor the raw images in {src} or its orig/ are left to build '
f'their fields again from. Point src at a copy of the plate\'s '
f'raw images to preprocess it again, or move the <name>.damaged '
f'file(s) out of {masks_path} and run again to segment only the '
f'fields the whole archives hold.')
if unlisted:
print(f'{len(unlisted)} archive(s) list their fields as an object '
f'array, which is not read without unpickling '
f'({_name_list(unlisted)}), so which fields are missing is not '
f'known and none is rebuilt.')
return True
missing = stack_fields - covered
if not missing:
return True
if (settings.get('illumination_correction', False) or psf_tracked or
settings.get('timelapse', False)):
print(f'{len(missing)} field(s) in stack/ are in no whole archive; '
f'this archive set is rebuilt whole from stack/.')
return False
if plane_counts and plane_counts != {len(mask_channels)}:
print(f'{len(missing)} field(s) in stack/ are in no whole archive, '
f'but the archives in masks/ hold '
f'{_name_list(str(n) for n in sorted(plane_counts, key=str))} '
f'channel(s) and these settings select {len(mask_channels)}, '
f'so none is added to them. Move masks/ out of {src} and run '
f'again to normalise the whole plate with these channels.')
return True
from .settings import set_default_settings_preprocess_img_data
norm_settings = set_default_settings_preprocess_img_data(dict(settings))
first = _next_archive_index(masks_path)
before = set(_listdir_visible(masks_path))
print(f'{len(missing)} field(s) in stack/ are in no whole archive; '
f'normalising them again, into new archives from '
f'stack_{first}_norm.npz on.')
_concatenate_and_normalize_impl(
stack_path, mask_channels, save_dtype=np.float32,
settings=norm_settings, only_fields=missing,
first_batch_index=first)
written = sorted(name for name in set(_listdir_visible(masks_path)) - before
if name.endswith('.npz'))
print(f'Wrote {len(written)} archive(s) for those fields: '
f'{_name_list(written)}')
return True
def _sample_stacks_for_test_mode(source, test_folder, test_images,
random_test=True):
"""Give test mode fields to work on when a plate's raw images are gone.
Test mode copies raw images, and a plate an earlier run preprocessed may
hold none: ``save_original_images`` off deletes them once ``stack/`` is
written. Its ``stack/`` still holds a stack per field, and a sample of
those is copied into ``test_folder/stack/`` instead.
:param source: the plate folder the user chose.
:param test_folder: the ``test/`` folder test mode works in.
:param test_images: how many fields to copy.
:param random_test: pick them at random, with the seed test mode uses,
rather than taking the first in name order.
:returns: the names copied; empty when ``source/stack/`` holds no whole
field stack.
"""
stack_path = os.path.join(source, 'stack')
names = sorted(name for name in _stack_field_stems(stack_path)
if _npy_is_whole(os.path.join(stack_path, name + '.npy'))[0])
if not names:
return []
if random_test:
random.Random(42).shuffle(names)
chosen = names[:max(int(test_images or 1), 1)]
destination = os.path.join(test_folder, 'stack')
for name in chosen:
with open(os.path.join(stack_path, name + '.npy'), 'rb') as original:
_replace_atomically(
os.path.join(destination, name + '.npy'),
lambda handle: shutil.copyfileobj(original, handle))
print(f'Test mode: {source} holds no raw images, so {len(chosen)} of the '
f'{len(names)} field stack(s) in its stack/ were copied into '
f'{destination}.')
return [name + '.npy' for name in chosen]
def _describe_processed_folder(folder):
"""Say what an earlier spaCR run left in ``folder``.
:param folder: a plate folder.
:returns: ``(summary, advice)``. ``summary`` lists each output folder
found there with what it holds, and is ``''`` when there is no
``orig/``, ``stack/``, ``channel_stack/``, ``masks/`` or ``merged/``.
``advice`` says what that leaves a new run to start from, or is
``''``.
"""
def count(name, suffixes):
"""Count the files in ``folder/name`` ending in ``suffixes``."""
try:
entries = _listdir_visible(os.path.join(folder, name))
except OSError:
return None
return sum(1 for entry in entries if entry.endswith(suffixes))
held = {
'orig': (count('orig', _RAW_IMAGE_SUFFIXES), 'raw image(s)'),
'stack': (count('stack', ('.npy',)), 'field stack(s)'),
'channel_stack': (count('channel_stack', ('.npz',)), 'archive(s)'),
'masks': (count('masks', ('.npz',)), 'normalised archive(s)'),
'merged': (count('merged', ('.npy',)), 'merged field(s)'),
}
parts = [f'{name}/ holds {number} {noun}'
for name, (number, noun) in held.items() if number is not None]
if not parts:
return '', ''
if os.path.isfile(os.path.join(folder, 'measurements', 'measurements.db')):
parts.append('measurements/ holds measurements.db')
advice = ''
if not held['orig'][0] and not held['stack'][0]:
if held['merged'][0]:
advice = (' Its raw images and field stacks were removed when that '
'run finished (keep_original_images and keep_intermediate '
'were off), so nothing is left here to preprocess; the '
'results are in merged/ and measurements/. To segment it '
'again, point src at a copy of the raw images.')
else:
advice = (' Neither orig/ nor stack/ holds anything a new run can '
'start from. Point src at the folder that holds the raw '
'images.')
return '; '.join(parts), advice
def _no_stacks_error(src, requested_src, regex, metadata_type):
"""Build the error for a run that produced no field stack, naming the likely cause.
:param src: the folder preprocessing ran in; ``test/`` in test mode.
:param requested_src: the folder the user chose.
:param regex: the filename pattern the images were matched against.
:param metadata_type: the naming convention that pattern came from.
:returns: a :class:`FileNotFoundError` to raise.
"""
entries = []
try:
entries = sorted(_listdir_visible(src))
except OSError:
pass
images = [name for name in entries
if name.lower().endswith(_RAW_IMAGE_SUFFIXES)]
subject = requested_src or src
hint = ''
if subject != src:
hint = (f' Test mode copies a sample of the raw images in {subject} '
f'(or in its orig/) into {src}, and found none it could use.')
try:
entries = sorted(_listdir_visible(subject))
except OSError:
entries = []
subdirs = [name for name in entries
if os.path.isdir(os.path.join(subject, name))][:6]
subject_images = [name for name in entries
if name.lower().endswith(_RAW_IMAGE_SUFFIXES)]
summary, advice = _describe_processed_folder(subject)
if summary:
hint += (f' {subject} is a folder spaCR has already processed, not a '
f'folder of plates: {summary}.{advice}')
if subject_images or _raw_image_names(os.path.join(subject, 'orig')):
hint += (f' None of the image files in {subject} or its orig/ could '
f'be read as a field: spaCR reads files ending in '
f'{", ".join(_RAW_IMAGE_SUFFIXES)} whose names match the '
f'pattern for metadata_type={metadata_type!r}: {regex}. A '
f'file whose name matches and that still gave no field '
f'could not be opened; the log above names each one.')
elif not summary and subdirs:
hint += (f" It holds no images but does hold sub-folders "
f"({', '.join(subdirs)}) — if those are plates, point "
f"src at one of them rather than at their parent.")
return FileNotFoundError(
f"No image stacks were produced from {src}. spaCR found "
f"{len(images)} image file(s) directly in that folder.{hint}")
def _volume_file_hash(path):
"""Hash one immutable source or derived volume in bounded chunks."""
import hashlib
before = os.stat(path)
digest = hashlib.sha256()
with open(path, 'rb') as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b''):
digest.update(chunk)
after = os.stat(path)
if any(getattr(before, key) != getattr(after, key)
for key in ('st_dev', 'st_ino', 'st_size', 'st_mtime_ns')):
raise ValueError(f'Volume changed while being read: {path}')
return digest.hexdigest()
def _preprocess_volume_tiffs(settings):
"""Preserve explicitly labelled ZYX TIFFs through raw Mask preprocessing.
Each file is one complete field/channel volume. Filename metadata owns
field/channel identity, TIFF axes own spatial identity. Raw files stay
untouched; an exact source/output receipt guards reuse of existing stacks.
Unsupported layouts are refused before writing any stack.
Exclusive directory creation protects even against a late empty-folder
collision. The completion receipt is published only after every immutable
stack exists, so interrupted folders cannot be reused.
"""
from .settings import set_default_settings_preprocess_img_data
from .utils import _get_regex, _extract_filename_metadata
from .psf_pipeline import processing_requested
from .cancellation import checkpoint
settings = set_default_settings_preprocess_img_data(settings)
src = os.fspath(settings['src'])
if settings.get('timelapse') or settings.get('t_stack'):
raise ValueError('Raw volumetric TIFF ingest does not combine time and Z; use prepared volumes')
if settings.get('z_axis', 0) not in (None, 0):
raise ValueError('Raw volumetric TIFF ingest writes canonical ZYXC: z_axis must be 0')
if settings.get('test_mode'):
raise ValueError('Raw volumetric TIFF ingest needs a source folder of complete volumes; turn test_mode off')
if settings.get('illumination_correction') or processing_requested(settings):
raise ValueError('Raw volumetric TIFF ingest does not yet support illumination or PSF preprocessing')
names = _raw_image_names(src)
if not names or any(Path(name).suffix.lower() not in ('.tif', '.tiff') for name in names):
raise ValueError('Raw volumetric ingest requires metadata-labelled ZYX TIFF files in src')
pattern = re.compile(_get_regex(settings['metadata_type'], 'tif', settings['custom_regex']))
parsed = _extract_filename_metadata(names, src, pattern, settings['metadata_type'])
if sum(map(len, parsed.values())) != len(names):
raise ValueError('Every volumetric TIFF must match the filename metadata pattern')
fields = defaultdict(dict)
inputs = []
for key, paths in sorted(parsed.items(), key=lambda item: str(item[0])):
stem = _escaped_field_stem(key[0], key[1], key[2], key[4])
channel = key[3]
if len(paths) != 1 or channel in fields[stem]:
raise ValueError(f'Volume {stem} channel {channel} has multiple files; provide one complete ZYX TIFF per channel')
path = paths[0]
checkpoint()
digest = _volume_file_hash(path)
with tifffile.TiffFile(path) as image:
if len(image.series) != 1:
raise ValueError(f'Volume {path} contains multiple TIFF series; select one explicitly')
series = image.series[0]
if series.axes != 'ZYX' or len(series.shape) != 3 or min(series.shape) < 2:
raise ValueError(f'Volume {path} must explicitly declare ZYX axes with at least two Z planes; got {series.axes} {series.shape}')
if series.dtype.kind not in 'uif':
raise ValueError(f'Volume {path} must contain real numeric intensities')
shape, dtype = tuple(series.shape), str(series.dtype)
record = dict(name=os.path.basename(path), field=stem, channel=channel,
sha256=digest, shape=list(shape), dtype=dtype, axes='ZYX')
fields[stem][channel] = record
inputs.append(record)
channels = sorted(
{channel for field in fields.values() for channel in field},
key=lambda channel: (not channel.isdecimal(),
int(channel) if channel.isdecimal() else channel))
for stem, field in fields.items():
if set(field) != set(channels):
raise ValueError(f'Volume {stem} is missing channel companions')
if len({tuple(row['shape']) for row in field.values()}) != 1:
raise ValueError(f'Volume {stem} channels have different ZYX shapes')
channel_keys = ['nucleus_channel', 'cell_channel', 'pathogen_channel',
*(f'{role}_channel' for role in ORGANELLE_ROLES)]
selected = list(dict.fromkeys(settings[key] for key in channel_keys
if settings.get(key) is not None))
if not selected or any(type(value) is not int or not 0 <= value < len(channels) for value in selected):
raise ValueError('Volumetric object channels must index the sorted source-channel list')
recipe_keys = {'background', 'Signal_to_noise', 'remove_background',
'lower_percentile', 'normalize', 'segmentation_backend',
'magnification', 'anisotropy', 'seg_qc', 'adjust_cells'}
recipe_prefixes = ('nucleus_', 'cell_', 'pathogen_', 'organelle',
'remove_background_', 'z_', 'voxel_size_')
recipe = json.loads(json.dumps({
key: value for key, value in settings.items()
if key in recipe_keys or key.startswith(recipe_prefixes)}, default=str))
receipt_path = os.path.join(src, 'stack', '.spacr_volume_ingest.json')
stack_path = os.path.dirname(receipt_path)
existing = _stack_field_stems(stack_path)
if os.path.isfile(receipt_path):
with open(receipt_path, encoding='utf8') as handle:
content = handle.read(16 * 1024 * 1024 + 1)
if len(content) > 16 * 1024 * 1024:
raise ValueError('Volumetric source receipt exceeds 16 MiB')
receipt = json.loads(content)
if not isinstance(receipt, dict) or receipt.get('version') != 1:
raise ValueError('Volumetric source receipt has an unsupported format')
if (receipt.get('inputs') != inputs or receipt.get('channels') != channels
or receipt.get('recipe') != recipe):
raise ValueError('Volumetric source identity or processing settings changed; use a fresh output folder')
if existing != set(fields):
raise ValueError('Volumetric stack inventory differs from its receipt; use a fresh output folder')
if set(receipt.get('stacks', {})) != {stem + '.npy' for stem in fields}:
raise ValueError('Volumetric stack receipt is incomplete')
for name, digest in receipt['stacks'].items():
if _volume_file_hash(os.path.join(stack_path, name)) != digest:
raise ValueError(f'Volumetric stack changed: {name}')
elif existing or any(os.path.exists(os.path.join(src, name)) for name in ('masks', 'merged')):
raise ValueError('Existing outputs have no volumetric source receipt; use a fresh output folder')
else:
staging = tempfile.mkdtemp(prefix='.spacr-volume-', dir=src)
try:
stacks = {}
for stem, field in fields.items():
checkpoint()
volumes = []
for channel in channels:
row = field[channel]
path = os.path.join(src, row['name'])
with tifffile.TiffFile(path) as image:
volume = image.series[0].asarray()
if (list(volume.shape) != row['shape'] or str(volume.dtype) != row['dtype']
or not np.isfinite(volume).all()
or _volume_file_hash(path) != row['sha256']):
raise ValueError(f'Volume changed or contains nonfinite intensities: {path}')
volumes.append(volume)
output = os.path.join(staging, stem + '.npy')
_save_array_atomic(output, np.stack(volumes, axis=-1))
stacks[stem + '.npy'] = _volume_file_hash(output)
del volumes, volume
for row in inputs:
if _volume_file_hash(os.path.join(src, row['name'])) != row['sha256']:
raise ValueError('Volumetric input changed before stack publication')
with open(os.path.join(staging, '.spacr_volume_ingest.json'), 'w', encoding='utf8') as handle:
json.dump(dict(version=1, axes='ZYXC', inputs=inputs,
channels=channels, recipe=recipe, stacks=stacks), handle, indent=2)
handle.write('\n')
os.mkdir(stack_path)
published = []
try:
for name in [*stacks, '.spacr_volume_ingest.json']:
destination = os.path.join(stack_path, name)
os.link(os.path.join(staging, name), destination)
published.append(destination)
except BaseException:
for destination in reversed(published):
os.unlink(destination)
try:
os.rmdir(stack_path)
except OSError:
pass
raise
finally:
if os.path.isdir(staging):
shutil.rmtree(staging)
print('Volumetric TIFF ingest preserves raw files and ZYX planes; normalizing one field at a time.')
normalization = dict(settings, batch_size=1, randomize=False, plot=False)
concatenate_and_normalize(stack_path, selected, np.float32, normalization)
for key in channel_keys:
if settings.get(key) is not None:
settings[f'cellpose_{key}'] = selected.index(settings[key])
settings['channels'] = list(range(len(channels)))
settings['z_axis'] = 0
return settings, src
def _preprocess_mapped_volume_series(settings):
"""Build one native TZYXC Mask archive per fixed-map field series.
The Convert map names every acquired YX plane. This route preserves those
files and assembles ZYXC stacks for each timepoint before normalizing the
complete TZYXC field. It deliberately performs independent 3-D Mask
segmentation only; timelapse tracking and Measure are separate contracts.
:param settings: raw Mask settings with explicit TZYX axes and native Z.
:returns: normalized settings and the unchanged source directory.
:raises ValueError: for unsupported recipes, incomplete maps, changed
sources, ambiguous axes or incompatible planes.
"""
import csv
import hashlib
import io as standard_io
from . import core
from .cancellation import checkpoint
from .psf_pipeline import processing_requested
from .settings import set_default_settings_preprocess_generate_masks
from .zstack import plan_4d_from_settings
settings = set_default_settings_preprocess_generate_masks(settings)
src = os.fspath(settings['src'])
if (settings.get('timelapse') or settings.get('microscope_feedback')
or settings.get('watch_pipeline') not in (None, '', 'mask')
or settings.get('apply_model_to_dataset')
or settings.get('generate_training_dataset')
or settings.get('real_object_classifier')):
raise ValueError('Native T-by-Z batch ingest supports Mask only; '
'Measure, tracking, Classify and feedback need a '
'separate 3-D-over-time contract.')
if (settings.get('test_mode') or settings.get('illumination_correction')
or processing_requested(settings)):
raise ValueError('Native T-by-Z batch ingest does not support test-mode '
'sampling, illumination or PSF preprocessing.')
if (settings.get('plot') or settings.get('adjust_cells')
or settings.get('segmentation_backend', 'cellpose') != 'cellpose'):
raise ValueError('Native T-by-Z Mask supports Cellpose without 2-D '
'plots or cell-adjustment operations.')
if (settings.get('metadata_type') != 'cellvoyager'
or settings.get('custom_regex') not in (None, '', 'None')):
raise ValueError('Native T-by-Z batch ingest requires a fixed Convert '
'map and its CellVoyager target names.')
plan = plan_4d_from_settings(settings)
if (plan is None or plan.t_axis != 0 or plan.z_axis != 1
or plan.z_mode != 'volumetric' or plan.frame_interval_s is None):
raise ValueError('Native T-by-Z batch ingest requires explicit TZYX '
'axes, frame interval and volumetric Z spacing.')
map_settings = dict(settings, timelapse=True)
manifest, map_sha256 = core._watch_map_manifest(src, map_settings)
if manifest is None:
raise ValueError('Native T-by-Z batch ingest requires a complete fixed '
'conversion_map.csv before processing.')
data = core._watch_map_bytes(src)
if hashlib.sha256(data).hexdigest() != map_sha256:
raise ValueError('Native T-by-Z conversion map changed during preflight.')
names = {name for members in manifest.values() for name in members}
if set(_raw_image_names(src)) != names:
raise ValueError('Native T-by-Z source images differ from the exact '
'conversion-map targets.')
fields = defaultdict(lambda: defaultdict(lambda: defaultdict(dict)))
for row in csv.DictReader(standard_io.StringIO(data.decode('utf-8-sig'))):
field = (row['plate'], row['well'], int(row['field']))
fields[field][int(row['t'])][int(row['channel'])][int(row['z'])] = row['target']
if any(len(times) < 2 or any(len(planes) < 2
for channels in times.values() for planes in channels.values())
for times in fields.values()):
raise ValueError('Native T-by-Z batch ingest requires at least two '
'timepoints and two Z planes per mapped channel.')
channel_ids = sorted({channel for times in fields.values()
for channels in times.values() for channel in channels})
channel_keys = ['nucleus_channel', 'cell_channel', 'pathogen_channel',
*(f'{role}_channel' for role in ORGANELLE_ROLES)]
selected = list(dict.fromkeys(settings[key] for key in channel_keys
if settings.get(key) is not None))
if (not selected or any(type(value) is not int or value < 0
or value >= len(channel_ids) for value in selected)):
raise ValueError('Native T-by-Z object channels must index the fixed '
'mapped channel list.')
recipe = json.loads(json.dumps({key: value for key, value in settings.items()
if key != 'src'}, default=str, sort_keys=True))
inputs = {}
for name in sorted(names):
path = os.path.join(src, name)
identity = core._watch_file_identity(path)
if identity is None:
raise ValueError(f'Native T-by-Z input is missing or unsafe: {name}')
inputs[name] = {'identity': identity,
'sha256': core._watch_artifact_sha256(path)}
if core._watch_file_identity(path) != identity:
raise ValueError(f'Native T-by-Z input changed during preflight: {name}')
stack_path = os.path.join(src, 'stack')
masks_path = os.path.join(src, 'masks')
receipt_path = os.path.join(stack_path, '.spacr_volume_series_ingest.json')
expected_stack_names = {
_escaped_field_stem(*field, time) + '.npy'
for field, times in fields.items() for time in times}
expected_archive_names = {
_escaped_field_stem(*field, '') + 'norm_timelapse.npz'
for field in fields}
if os.path.isfile(receipt_path):
if (os.path.islink(stack_path) or os.path.islink(masks_path)
or os.path.islink(receipt_path)):
raise ValueError('Native T-by-Z outputs or receipt cannot be links.')
with open(receipt_path, encoding='utf8') as handle:
content = handle.read(16 * 1024 * 1024 + 1)
if len(content) > 16 * 1024 * 1024:
raise ValueError('Native T-by-Z source receipt exceeds 16 MiB.')
receipt = json.loads(content)
if (receipt.get('version') != 1 or receipt.get('axes') != 'TZYXC'
or receipt.get('map_sha256') != map_sha256
or receipt.get('inputs') != inputs or receipt.get('recipe') != recipe
or set(receipt.get('stacks', {})) != expected_stack_names
or set(receipt.get('archives', {})) != expected_archive_names):
raise ValueError('Native T-by-Z source, map or processing recipe '
'differs from its completed receipt; use a fresh folder.')
if ({name for name in os.listdir(stack_path) if name.endswith('.npy')}
!= expected_stack_names or
{name for name in os.listdir(masks_path)
if name.endswith('_norm_timelapse.npz')}
!= expected_archive_names):
raise ValueError('Native T-by-Z output inventory differs from '
'its completed receipt; use a fresh folder.')
for folder, key in ((stack_path, 'stacks'), (masks_path, 'archives')):
for name, digest in receipt[key].items():
if core._watch_file_identity(os.path.join(folder, name)) is None or (
_volume_file_hash(os.path.join(folder, name)) != digest):
raise ValueError(f'Native T-by-Z output changed: {name}')
elif any(os.path.lexists(path) for path in (stack_path, masks_path,
os.path.join(src, 'merged'))):
raise ValueError('Existing outputs have no native T-by-Z source receipt; '
'use a fresh output folder.')
else:
with _native_map_workspace(prefix='.spacr-volume-series-', dir=src) as stage_owner:
stage = stage_owner['name']
stage_stack = os.path.join(stage, 'stack')
stage_masks = os.path.join(stage, 'masks')
os.mkdir(stage_stack)
os.mkdir(stage_masks)
stacks, archives = {}, {}
for field, times in sorted(fields.items()):
filenames, shape, dtype, frame_shape = [], None, None, None
for time, channels in sorted(times.items()):
volumes = []
for channel in channel_ids:
planes = []
for z, name in sorted(channels[channel].items()):
checkpoint()
path = os.path.join(src, name)
flags = (os.O_RDONLY | getattr(os, 'O_NOFOLLOW', 0)
| getattr(os, 'O_NONBLOCK', 0))
with os.fdopen(os.open(path, flags), 'rb') as handle:
info = os.fstat(handle.fileno())
opened = [info.st_dev, info.st_ino, info.st_size,
info.st_mtime_ns, info.st_ctime_ns]
if opened != inputs[name]['identity']:
raise ValueError(f'Native T-by-Z input changed: {name}')
with tifffile.TiffFile(handle, name=path) as image:
if len(image.series) != 1 or image.series[0].axes != 'YX':
raise ValueError(f'Native T-by-Z input must be one YX plane: {name}')
plane = image.series[0].asarray()
info = os.fstat(handle.fileno())
if [info.st_dev, info.st_ino, info.st_size,
info.st_mtime_ns, info.st_ctime_ns] != opened:
raise ValueError(f'Native T-by-Z input changed: {name}')
if (plane.ndim != 2 or plane.dtype.kind not in 'uif'
or not np.isfinite(plane).all()):
raise ValueError(f'Native T-by-Z input has invalid intensities: {name}')
if shape is None:
shape, dtype = plane.shape, plane.dtype.str
if plane.shape != shape or plane.dtype.str != dtype:
raise ValueError('Native T-by-Z planes differ in shape or dtype.')
if (core._watch_file_identity(path) != inputs[name]['identity']
or core._watch_artifact_sha256(path)
!= inputs[name]['sha256']):
raise ValueError(f'Native T-by-Z input changed: {name}')
planes.append(plane)
volumes.append(np.stack(planes))
del planes, plane
frame = np.stack(volumes, axis=-1)
del volumes
if frame_shape is None:
frame_shape = frame.shape
filename = _escaped_field_stem(*field, time) + '.npy'
output = os.path.join(stage_stack, filename)
_save_array_atomic(output, frame)
stacks[filename] = _volume_file_hash(output)
del frame
filenames.append(filename)
staged_paths = [os.path.join(stage_stack, name) for name in filenames]
for name, path in zip(filenames, staged_paths):
checkpoint()
if os.path.islink(path) or _volume_file_hash(path) != stacks[name]:
raise ValueError(f'Native T-by-Z staged stack changed: {name}')
output_shape = (len(filenames), *frame_shape)
with _native_map_workspace(
prefix='.spacr-normalize-', dir=stage,
parent=stage_owner) as workspace_owner:
workspace = workspace_owner['name']
_native_workspace_preflight(
workspace, output_shape, len(selected), dtype, filenames)
normalized = _retain_native_memmap(
np.lib.format.open_memmap(
os.path.join(workspace, 'selected.npy'), mode='w+',
dtype=np.float32,
shape=(*output_shape[:-1], len(selected))),
workspace_owner['owners'],
remove_path=os.path.join(workspace, 'selected.npy'))
try:
_reserve_private_memmap(normalized)
def load_channel(channel):
"""Read one private channel from verified staged stacks."""
channel_path = os.path.join(
workspace, f'channel-{channel}.npy')
values = _retain_native_memmap(
np.lib.format.open_memmap(
channel_path, mode='w+', dtype=np.dtype(dtype),
shape=output_shape[:-1]),
workspace_owner['owners'], remove_path=channel_path)
try:
_reserve_private_memmap(values)
for index, path in enumerate(staged_paths):
checkpoint()
mapped = _retain_native_memmap(
np.load(path, mmap_mode='r', allow_pickle=False),
stage_owner['owners'])
try:
if (mapped.shape != frame_shape
or mapped.dtype.str != dtype):
raise ValueError(
'Native T-by-Z staged stack shape or dtype changed.')
for z_index in range(frame_shape[0]):
checkpoint()
values[index, z_index] = (
mapped[z_index, ..., channel])
finally:
del mapped
return values
except BaseException:
try:
_close_private_memmap(values)
finally:
values = None
raise
normalized = _normalize_img_channels(
normalized, selected, np.float32, settings,
load_channel,
output_columns={channel: index
for index, channel in enumerate(selected)},
workspace_dir=workspace)
for name, path in zip(filenames, staged_paths):
checkpoint()
if os.path.islink(path) or _volume_file_hash(path) != stacks[name]:
raise ValueError(f'Native T-by-Z staged stack changed: {name}')
archive_name = _escaped_field_stem(*field, '') + 'norm_timelapse.npz'
archive = os.path.join(stage_masks, archive_name)
_save_npz_atomic(archive, data=normalized,
filenames=filenames)
archives[archive_name] = _volume_file_hash(archive)
finally:
try:
_close_private_memmap(normalized)
finally:
del normalized
if (hashlib.sha256(core._watch_map_bytes(src)).hexdigest() != map_sha256
or any(core._watch_file_identity(os.path.join(src, name))
!= row['identity'] or core._watch_artifact_sha256(os.path.join(src, name))
!= row['sha256'] for name, row in inputs.items())):
raise ValueError('Native T-by-Z inputs or map changed before publication.')
with open(os.path.join(stage_stack, '.spacr_volume_series_ingest.json'),
'w', encoding='utf8') as handle:
json.dump(dict(version=1, axes='TZYXC', map_sha256=map_sha256,
inputs=inputs, recipe=recipe, stacks=stacks,
archives=archives), handle, indent=2)
handle.write('\n')
published = []
try:
os.mkdir(stack_path)
published.append(stack_path)
os.mkdir(masks_path)
published.append(masks_path)
for folder, stage_folder, names_to_link in (
(stack_path, stage_stack, [*stacks, '.spacr_volume_series_ingest.json']),
(masks_path, stage_masks, archives)):
for name in names_to_link:
destination = os.path.join(folder, name)
os.link(os.path.join(stage_folder, name), destination)
published.append(destination)
except BaseException:
for path in reversed(published):
if os.path.isdir(path):
os.rmdir(path)
else:
os.unlink(path)
raise
for key in channel_keys:
if settings.get(key) is not None:
settings[f'cellpose_{key}'] = selected.index(settings[key])
settings['channels'] = list(range(len(channel_ids)))
return settings, src
[docs]
def preprocess_img_data(settings):
"""Convert raw microscopy images into normalized, channel-merged ``.npy`` stacks ready for mask generation.
Usually invoked internally by
:func:`spacr.core.preprocess_generate_masks`, but callable directly
when you only want the preprocessing half. By default it converts z-stacks to MIPs,
renames files into the Yokogawa/spacr layout, merges per-channel
folders into stacked ``.npy`` arrays with optional background
subtraction and percentile normalization, and (in ``test_mode``)
emits example plots.
With ``z_stack=True`` and ``z_segmentation_mode='volumetric'``, a
separate raw TIFF route preserves Z. Each field/channel must be one
complete TIFF series explicitly labelled ``ZYX``. Filename metadata
identifies the field and channel; ambiguous axes, duplicate channels
and missing companions are refused before any stack is written. This
route retains the original files even if ``save_original_images`` is
false, publishes canonical ``ZYXC`` stacks with a completion record,
and normalises one field at a time. Reuse verifies source and stack
hashes and relevant processing settings. An old projected stack cannot
be reused as a volume. Individual slice-file layouts, time-series,
test-mode sampling and illumination or PSF preprocessing are not
supported by this raw volumetric route.
A separate Mask-only route accepts a fixed Convert map with every
explicitly labelled YX plane in a dense channel-by-Z-by-time grid when
``z_stack`` and ``t_stack`` are both on. It builds native TZYXC archives
for the existing 4-D segmenter; physical Z and frame spacing and TZYX
axis order must be declared. The planar sources remain unchanged.
Timelapse tracking, Measure and Classify are not enabled by this route.
Running it again on a plate folder it has already processed resumes
rather than starting over. Raw images an earlier run moved into
``orig/`` are read from there, and only the fields ``stack/`` lacks are
built. Every ``stack/*.npy`` and ``masks/*.npz`` an earlier run left is
checked before it is reused: a file cut short is renamed to
``<name>.damaged``, named in the log, and built again, a field stack
from the raw images or channel folders and an archive from ``stack/``.
A field stack with nothing left to build it from is recorded as a
failure, so the run ends incomplete instead of quietly short of that
field, and an archive with nothing left to build it from stops the run
with an error that names it. When no field can be built at all, the
error says what the folder does hold. In ``test_mode`` a plate whose
raw images are gone is sampled from its ``stack/`` instead.
:param settings: Preprocessing settings dict, canonicalized via
:func:`spacr.settings.set_default_settings_preprocess_img_data`.
Key entries:
- ``src`` — folder of raw images (``.tif/.nd2/.czi/.lif`` etc.).
- ``metadata_type`` — ``'cellvoyager'`` / ``'auto'``; drives
filename regex.
- ``custom_regex`` — override the built-in regex.
- ``cell_channel``, ``nucleus_channel``, ``pathogen_channel``,
``organelle_channel``, ``channels`` — channel selection.
- ``z_stack`` and ``z_segmentation_mode`` select the explicit
volumetric TIFF route described above. Other raw layouts retain
the existing per-field/channel projection behavior.
- ``remove_background_cell`` / ``_nucleus`` / ``_pathogen`` /
``_organelle`` and each object's ``*_background`` and
``*_signal_to_noise`` values. Every organelle slot the run
enables uses its own pair, e.g. ``organelleb_background``.
- ``normalize``, ``lower_percentile``, ``save_dtype``.
- ``batch_size``, ``randomize``, ``test_mode``, ``test_images``,
``plot``, ``cmap``, ``figuresize``.
:returns: Tuple ``(settings, src)`` — ``settings`` with defaults
applied and ``src`` pointing at the folder containing the
generated ``stack/`` / ``channel_stack/`` outputs (the
downstream mask stage reads from here).
Example:
.. code-block:: python
from spacr.io import preprocess_img_data
settings = {
'src': '/data/plate01',
'metadata_type': 'cellvoyager',
'cell_channel': 0, 'nucleus_channel': 1, 'pathogen_channel': 2,
'channels': [0, 1, 2, 3], 'normalize': True,
}
settings, src = preprocess_img_data(settings)
See Also:
:func:`spacr.core.preprocess_generate_masks` — full pipeline
wrapper that calls this then generates masks.
"""
if (settings.get('z_stack') and settings.get('t_stack')
and settings.get('z_segmentation_mode') == 'volumetric'):
return _preprocess_mapped_volume_series(settings)
if settings.get('z_stack') and settings.get('z_segmentation_mode') == 'volumetric':
return _preprocess_volume_tiffs(settings)
src = settings['src']
requested_src = src
if len(_listdir_visible(src)) < 100:
delete_empty_subdirectories(src)
from .object_roles import ORGANELLE_ROLES
mask_channel_keys = (
'nucleus_channel', 'cell_channel', 'pathogen_channel',
*(f'{role}_channel' for role in ORGANELLE_ROLES),
)
mask_channels_raw = [settings.get(key) for key in mask_channel_keys]
seen = {}
mask_channels = []
for ch in mask_channels_raw:
if ch is None:
continue
try:
ch = int(ch)
except (TypeError, ValueError):
continue
if ch not in seen:
seen[ch] = len(mask_channels)
mask_channels.append(ch)
files = _listdir_visible(src)
valid_ext = ['tif', 'tiff', 'png', 'jpg', 'jpeg', 'bmp', 'nd2', 'czi', 'lif']
extensions = [file.split('.')[-1].lower() for file in files]
valid_extensions = [ext for ext in extensions if ext in valid_ext]
img_format = None
if valid_extensions:
extension_counts = Counter(valid_extensions)
most_common_extension = Counter(valid_extensions).most_common(1)[0][0]
img_format = most_common_extension
print(f"Found {extension_counts[most_common_extension]} {most_common_extension} files")
else:
print(f"Could not find any {valid_ext} files in {src}")
print(f"{files} in {src}")
print(f"Please check the folder and try again")
if os.path.exists(os.path.join(src,'stack')):
print('Found existing stack folder.')
if os.path.exists(os.path.join(src,'channel_stack')):
print('Found existing channel_stack folder.')
if (os.path.exists(os.path.join(src, 'masks')) and
settings.get('test_mode', False)):
print('Found existing masks folder; test mode works on a sample '
'in test/ and does not reuse it.')
elif (os.path.exists(os.path.join(src, 'masks')) and
_resume_normalized_archives(settings, src, mask_channels)):
print('Found existing masks folder. Skipping preprocessing')
if (settings.get('illumination_correction', False) and
settings.get('masks', True)):
from .illumination import (
load_segmentation_illumination_resume,
)
mask_src = os.path.join(src, 'masks')
load_segmentation_illumination_resume(
settings,
provenance_path=os.path.join(
src, 'illumination',
'segmentation_application.json'),
pipeline_style='v1',
expected_fields=_normalized_npz_field_ids(mask_src),
verbose=settings.get('verbose', True),
)
return settings, src
from .settings import set_default_settings_preprocess_img_data
from .utils import _get_regex, _run_test_mode
from .plot import plot_arrays
settings = set_default_settings_preprocess_img_data(settings)
regex_format = img_format
if regex_format is None:
set_aside = _raw_image_names(os.path.join(src, 'orig'))
if set_aside:
regex_format = Counter(
name.rsplit('.', 1)[-1].lower() for name in set_aside
).most_common(1)[0][0]
regex = _get_regex(settings['metadata_type'], regex_format, settings['custom_regex'])
if settings['test_mode']:
print(f"Running spacr in test mode")
settings['plot'] = True
if os.path.exists(os.path.join(src,'test')):
try:
os.rmdir(os.path.join(src, 'test'))
print(f"Deleted test directory: {os.path.join(src, 'test')}")
except OSError as e:
print(f"Error deleting test directory: {e}")
print(f"Delete manually before running test mode")
pass
src = _run_test_mode(settings['src'], regex, settings['timelapse'], settings['test_images'], settings['random_test'])
settings['src'] = src
if (not settings['timelapse'] and not _raw_image_names(src) and
not _stack_field_stems(os.path.join(src, 'stack'))):
_sample_stacks_for_test_mode(
requested_src, src, settings['test_images'],
settings['random_test'])
stack_path = os.path.join(src, 'stack')
_set_aside_damaged_stacks(stack_path)
if img_format == None:
if not os.path.exists(stack_path) or _channel_folders(src):
_merge_channels(src, plot=False)
resuming = os.path.exists(stack_path)
raw_waiting = bool(_raw_image_names(src) or
_raw_image_names(os.path.join(src, 'orig')))
if not resuming or raw_waiting:
stacks_before = len(_stack_field_stems(stack_path))
try:
img_format = ['.tif', '.tiff', '.png', '.jpg', '.jpeg', '.bmp', '.nd2', '.czi', '.lif']
nr_channel_folders = _rename_and_organize_image_files(
src, regex, settings['batch_size'], settings['metadata_type'], img_format,
timelapse=settings['timelapse'],
save_original_images=settings.get('save_original_images', True))
all_imgs = len([f for f in _listdir_visible(stack_path) if f.endswith('.npy')]) if os.path.isdir(stack_path) else 0
if resuming and all_imgs == stacks_before:
print(f'Nothing was added to stack/; resuming from the '
f'{all_imgs} field stack(s) it holds.')
else:
batch_size = int(settings.get('batch_size') or 0)
full_batches = all_imgs // batch_size if batch_size else 0
last_batch_size = all_imgs % batch_size if batch_size else 0
if last_batch_size == 1:
if full_batches == 0:
print(f"Warning: Only one batch of size 1 detected (all images: {all_imgs}). Adjust the batch size.")
else:
print(f"all images: {all_imgs}, full batch: {full_batches}, last batch: {last_batch_size}")
print("Warning: Last batch of size 1 detected. Adjust the batch size.")
if nr_channel_folders and len(settings['channels']) != nr_channel_folders:
print(f"Number of channels does not match number of channel folders. channels: {settings['channels']} channel folders: {nr_channel_folders}")
new_channels = list(range(nr_channel_folders))
print(f"Changing channels from {settings['channels']} to {new_channels}")
settings['channels'] = new_channels
if settings['timelapse']:
_create_movies_from_npy_per_channel(stack_path, fps=settings['fps'])
if settings['plot']:
print(f"plotting {settings['nr']} images from {src}/stack")
plot_arrays(stack_path, settings['figuresize'], settings['cmap'], nr=settings['nr'], normalize=settings['normalize'])
except Exception as e:
print(f"Error: {e}")
_report_unrebuilt_stacks(stack_path, src)
stacked = ([f for f in _listdir_visible(stack_path) if f.endswith('.npy')]
if os.path.isdir(stack_path) else [])
stacked = select_fields(stacked, settings.get('fields'))
if not stacked:
raise _no_stacks_error(src, requested_src, regex,
settings['metadata_type'])
illumination_session = None
if (settings.get('illumination_correction', False) and
settings.get('masks', True)):
from .illumination import prepare_segmentation_illumination
illumination_session = prepare_segmentation_illumination(
settings,
src=stack_path,
channels=mask_channels,
pipeline_style='v1',
)
from .psf_pipeline import _prepare_segmentation_psf
psf_session = _prepare_segmentation_psf(settings, src, mask_channels,
stack_dir=stack_path)
concatenate_and_normalize(src=stack_path,
channels=mask_channels,
save_dtype=np.float32,
settings=settings,
illumination_session=illumination_session,
psf_session=psf_session)
for key in mask_channel_keys:
ch = settings.get(key)
if ch is None:
continue
try:
ch = int(ch)
except (TypeError, ValueError):
continue
settings[f"cellpose_{key}"] = seen[ch]
return settings, src
def _check_masks(batch, batch_filenames, output_folder, resume=False):
"""
Check the masks in a batch and filter out the ones that already exist in the output folder.
Args:
batch (list): List of masks.
batch_filenames (list): List of filenames corresponding to the masks.
output_folder (str): Path to the output folder.
resume (bool): Accepted for the callers that pass it. Existing
``.npy`` files are validated before they are skipped whether or
not it is set, and a damaged one is returned for processing.
Returns:
tuple: A tuple containing the filtered batch (numpy array) and the filtered filenames (list).
"""
from .resume import validate_merged_field
def needs_processing(filename):
"""Report whether a field still has to be generated.
Args:
filename (str): Name relative to the enclosing
``output_folder``, not a full path — it is joined onto that
folder here. An existing file is validated by its header and
length, so an empty or truncated ``.npy`` left behind by a
killed run counts as missing, is named in the log, and is
generated again.
"""
path = os.path.join(output_folder, filename)
if not os.path.isfile(path):
return True
ok, reason = validate_merged_field(path)
if not ok:
print(f"{path} is damaged ({reason}); generating it again.")
return not ok
existing_files_mask = [
needs_processing(filename) for filename in batch_filenames]
filtered_batch = [b for b, exists in zip(batch, existing_files_mask) if exists]
filtered_filenames = [f for f, exists in zip(batch_filenames, existing_files_mask) if exists]
return np.array(filtered_batch), filtered_filenames
def _get_avg_object_size(masks):
"""
Calculate:
- average number of objects per image
- average object size over all objects
Parameters:
masks (list): A list of 2D or 3D masks with labeled objects.
Returns:
tuple:
avg_num_objects_per_image (float)
avg_object_size (float)
"""
per_image_counts = []
all_areas = []
for idx, mask in enumerate(masks):
if mask.ndim in [2, 3] and np.any(mask):
props = measure.regionprops(mask)
areas = [prop.area for prop in props]
per_image_counts.append(len(areas))
all_areas.extend(areas)
else:
per_image_counts.append(0)
if not np.any(mask):
print(f"Warning: Mask {idx} is empty.")
else:
print(f"Warning: Mask {idx} has invalid dimension: {mask.ndim}")
if per_image_counts:
avg_num_objects_per_image = sum(per_image_counts) / len(per_image_counts)
else:
avg_num_objects_per_image = 0
if all_areas:
avg_object_size = sum(all_areas) / len(all_areas)
else:
avg_object_size = 0
return avg_num_objects_per_image, avg_object_size
def _save_figure(fig, src, text, dpi=None, i=1, all_folders=1):
"""
Save a figure to a specified location.
Parameters:
fig (matplotlib.figure.Figure): The figure to be saved.
src (str): The source file path.
text (str): The text to be included in the figure name.
dpi (int, optional): Resolution. ``None`` (the default) follows the
user's figure-resolution preference -- see spacr.plot.save_figure.
"""
from .utils import print_progress
save_folder = os.path.dirname(src)
obj_type = os.path.basename(src)
name = os.path.basename(save_folder)
save_folder = os.path.join(save_folder, 'figure')
os.makedirs(save_folder, exist_ok=True)
fig_name = f'{obj_type}_{name}_{text}.pdf'
save_location = os.path.join(save_folder, fig_name)
from .plot import save_figure
save_location = save_figure(fig, save_location, dpi=dpi,
bbox_inches='tight')
files_processed = i
files_to_process = all_folders
print_progress(files_processed, files_to_process, n_jobs=1, time_ls=None, batch_size=None, operation_type="Saving Figures")
print(f'Saved single cell figure: {os.path.basename(save_location)}')
plt.close(fig)
del fig
gc.collect()
[docs]
class TimelapseKeyMismatch(ValueError):
"""One side of the ``png_list`` join carries a timepoint and the other does not.
A timepoint column is written only by a timelapse run — by
:func:`spacr.utils.filepaths_to_database` onto ``png_list`` and by
:func:`spacr.utils._merge_and_save_to_database` onto every object table —
so a database where one of the two has it and the other does not was
written by two runs that disagreed about whether the experiment was a
timelapse. Joining them without the timepoint silently multiplies every
object by the number of frames, which is precisely the failure this
exception exists to stop.
"""
[docs]
class JoinFanOut(ValueError):
"""A left join returned more rows than the frame it started from.
The object tables carry one row per object per field per timepoint and
``png_list`` carries one crop per the same key, so the join is one-to-one
(with unmatched rows permitted by the left join) and the row count cannot
grow. If it did, a join key is duplicated and every downstream measurement
is multiplied.
"""
[docs]
class MergeCardinalityError(JoinFanOut):
"""A database merge violated its declared relationship.
This is a :class:`JoinFanOut` for backwards compatibility with callers
that already handle duplicate crop rows, but it is used for every
explicitly validated object-table relationship.
"""
def _merge_key_details(frame, *, columns=None, use_index=False):
"""Return a readable key description and duplicate examples for one side."""
if use_index:
duplicated = frame.index.duplicated(keep=False)
examples = frame.index[duplicated].unique().tolist()[:3]
return "index", examples
key_columns = [columns] if isinstance(columns, str) else list(columns or [])
if not key_columns:
return "unspecified keys", []
duplicated = frame.duplicated(subset=key_columns, keep=False)
examples = (
frame.loc[duplicated, key_columns]
.drop_duplicates()
.head(3)
.itertuples(index=False, name=None)
)
return repr(key_columns), list(examples)
def _merge_with_cardinality(
left, right, *, validate, left_name, right_name, **merge_kwargs):
"""Merge frames under an explicit pandas cardinality contract.
pandas correctly rejects invalid relationships through ``validate=`` but
its error does not name the database tables or show the offending keys.
This wrapper retains pandas' enforcement and turns cardinality violations
into an actionable spaCR error. Other merge errors (for example, colliding
suffixes) remain ordinary :class:`pandas.errors.MergeError` instances.
"""
try:
return left.merge(right, validate=validate, **merge_kwargs)
except pd.errors.MergeError as exc:
on = merge_kwargs.get("on")
left_on = merge_kwargs.get("left_on", on)
right_on = merge_kwargs.get("right_on", on)
left_key, left_examples = _merge_key_details(
left,
columns=left_on,
use_index=bool(merge_kwargs.get("left_index")),
)
right_key, right_examples = _merge_key_details(
right,
columns=right_on,
use_index=bool(merge_kwargs.get("right_index")),
)
invalid = []
if validate in {"one_to_one", "1:1", "one_to_many", "1:m"}:
if left_examples:
invalid.append(
f"{left_name} has duplicated {left_key}: {left_examples}")
if validate in {"one_to_one", "1:1", "many_to_one", "m:1"}:
if right_examples:
invalid.append(
f"{right_name} has duplicated {right_key}: "
f"{right_examples}")
if not invalid:
raise
raise MergeCardinalityError(
f"Cannot merge {left_name} with {right_name}: the declared "
f"validate={validate!r} relationship was violated; "
+ "; ".join(invalid)
+ ". De-duplicate or repair the named source table before "
"continuing; otherwise rows and measurements would be "
"silently multiplied."
) from exc
[docs]
class CropModeMismatch(ValueError):
"""``png_list`` holds no crops of the object this join is anchored on.
``measure_crop`` writes one object-id column per ``crop_mode`` --
``cell_id`` for ``crop_mode=['cell']``, ``nucleus_id`` for ``['nucleus']``
and so on; the mapping is :data:`spacr.utils.PNG_OBJECT_ID_COLUMNS`.
:func:`_read_and_join_tables` anchors on the ``cell`` table, so it needs
``cell_id``. A database measured only with nucleus crops does not have that
column, and used to fail with ``KeyError: "['cell_id'] not in index"`` --
which names neither the table, nor the column's absence, nor the setting
that caused it.
The cell a nucleus crop belongs to *is* in the crop's file name --
:func:`spacr.utils._generate_names` writes
``<field>_<cell>_<nucleus>.png`` -- but it is not stored:
:func:`spacr.utils.filepaths_to_database` keeps only the last token.
Recovering it would mean a second file-name parser beside
:mod:`spacr.schema`, so this is a refusal rather than a guess.
"""
def _report_fan_out(left, merged, join_cols, left_name='cell',
right_name='png_list'):
"""Raise :class:`JoinFanOut` if ``merged`` grew, naming the offending keys.
The invariant checked is ``len(merged) == len(left)``, not
``len(merged) <= len(png_list)``. For a LEFT join those are different
statements and only the first one is true of a healthy database: crops are
routinely a strict subset of the measured objects — ``save_png`` can be off
for some fields, a crop can fail to write, and ``png_list`` is appended per
field so an interrupted run leaves fewer crops than cells. In all of those
cases ``len(merged) == len(cell) > len(png_list)`` and nothing is wrong.
What cannot happen is the join *growing* the left frame, and that is the
exact signature of the timelapse bug (12 cell rows in, 36 out).
"""
if len(merged) <= len(left):
return
raise JoinFanOut(
f"Joining {left_name} to {right_name} on {list(join_cols)} turned "
f"{len(left)} {left_name} rows into {len(merged)}: {right_name} holds "
f"more than one row per {list(join_cols)}, so every measurement in the "
f"result is duplicated. This usually means the crop step ran twice and "
f"appended a second set of rows to {right_name}; de-duplicate that "
f"table before reading."
)
def _read_and_join_tables(db_path, table_names=None,
duplicate_column_policy='warn',
keep_uninfected=True,
collapse_duplicate_identity=True,
require_crops=True):
"""
Reads and joins tables from a SQLite database.
**A column that arrives from two tables is compared, not duplicated.**
Cell and cytoplasm are both keyed to the same object, so a join hands
back ``plateID`` and ``plateID_cytoplasm``, ``prcf`` and
``prcf_cytoplasm``, and so on. Where the pair agrees, one copy is kept.
Where it disagrees the two tables describe different objects under the
same identity, which no analysis should quietly average over:
``duplicate_column_policy='warn'`` prints and keeps the cell table's
value, ``'raise'`` stops the run. ``collapse_duplicate_identity=False``
turns the whole comparison off and keeps both copies -- for a caller
that has its OWN reason to want the suffixed columns, which
:mod:`spacr.anndata_export` does: its ``drop_redundant_identity=False``
is documented to keep them, and a collapse here would have made that
option a no-op.
``require_crops=False`` is for a caller whose analysis does not need a
picture: a UMAP of MEASUREMENTS is still valid for an object whose crop
never wrote, and dropping it silently shrinks the embedding. The default
keeps the inner join, because the callers that do need a crop -- the
classifier, the annotator, the image grids -- are the majority and
carrying an unusable row into them is what that join prevents.
**Which cells survive a join is a decision, not a default.** A cell with
no nucleus is debris, and a cell with no crop cannot be classified, so
both of those joins are inner. A cell with no PATHOGEN is an uninfected
cell -- usually the control population -- so that join keeps it, and
``keep_uninfected=False`` is how an analysis restricts itself to infected
cells deliberately rather than by accident. See
:func:`spacr.object_roles.join_how`.
``png_list`` is joined to the object tables on plate / row / column / field
**and on the timepoint when both sides carry one**. Without the timepoint
the join is many-to-many on a timelapse database: every frame's crop
matches every frame's object row, so N objects x T frames came back as
N x T x T rows with the wrong PNG attached to most of them.
Either spelling of the timepoint is accepted on read (``timeID`` is
canonical, ``time_id`` is what ``png_list`` was written with before the two
were unified). :func:`spacr.utils.rename_columns_in_db` runs above and
repairs an old database in place, but a database opened read-only, or one
carrying both spellings, still reads correctly here.
**The object id is migrated, not assumed.** ``png_list.cell_id`` is text
(``'o5'``) and every object table's key is an integer, so the two are
reconciled through :func:`spacr.utils.object_label_from_png_id` -- one
implementation, shared with anything else that has to cross that boundary.
The migration this replaces was ``.str[1:].astype(int)``, which died on
four values the real writers produce every day: ``'omulti'`` and
``'onone'`` (a crop overlapping several cells, or none), ``'error'`` (an
unparseable crop name) and ``NULL`` (any row belonging to a *different*
crop mode, in a database measured with more than one). Those rows are now
dropped from the ``png_list`` side and counted out loud: the object keeps
its measurements and simply has no crop path, which is the same state as a
crop that was never written.
Args:
db_path (str): The path to the SQLite database file.
table_names (list, optional): The names of the tables to read and join. Defaults to ['cell', 'cytoplasm', 'nucleus', 'pathogen', 'png_list'].
Returns:
pandas.DataFrame: The joined DataFrame containing the data from the specified tables, or None if an error occurs.
Raises:
TimelapseKeyMismatch: when exactly one of ``png_list`` and ``cell``
carries a timepoint column.
MergeCardinalityError: when either table repeats an object key.
CropModeMismatch: when ``png_list`` carries no ``cell_id`` column, i.e.
it holds crops of some other object.
"""
if table_names is None:
table_names = ['cell', 'cytoplasm', 'nucleus', 'pathogen', 'png_list']
from .database_schema import ensure_database_schema
from .utils import (PNG_CROP_MODE_BY_ID_COLUMN, PNG_OBJECT_ID_COLUMNS,
TIME_COLUMN_ALIASES, object_label_from_png_id,
_time_column)
ensure_database_schema(db_path)
from .database_concurrency import connect as _connect_database
conn = _connect_database(db_path)
dataframes = {}
for table_name in table_names:
try:
dataframes[table_name] = pd.read_sql(f"SELECT * FROM {table_name}", conn)
except (sqlite3.OperationalError, pd.io.sql.DatabaseError) as e:
print(f"Table {table_name} not found in the database.")
print(e)
conn.close()
if 'png_list' in dataframes:
png_raw = dataframes['png_list']
id_column = PNG_OBJECT_ID_COLUMNS['cell']
if id_column not in png_raw.columns:
present = [c for c in png_raw.columns
if c in PNG_CROP_MODE_BY_ID_COLUMN]
modes = sorted(PNG_CROP_MODE_BY_ID_COLUMN[c] for c in present)
raise CropModeMismatch(
f"png_list in {db_path} has no {id_column!r} column, so its "
f"crops cannot be attached to the cell table this join is "
f"anchored on. It holds "
+ (f"{', '.join(modes)} crops ({', '.join(present)})"
if present else "no object-id column at all")
+ f". Re-run the Measure module with 'cell' in crop_mode to "
f"write cell crops alongside the ones already there."
)
png_cols = [id_column, 'png_path', 'plateID', 'rowID', 'columnID',
'fieldID']
png_time = _time_column(png_raw.columns)
if png_time is not None:
png_cols = png_cols + [png_time]
png_list_df = png_raw[png_cols].copy()
labels = object_label_from_png_id(png_list_df[id_column])
usable = labels.notna()
if not usable.all():
raw = png_list_df[id_column]
other_mode = int((raw.isna() & ~usable).sum())
unreadable = raw[~usable & raw.notna()]
if other_mode:
print(f"png_list: {other_mode} of {len(png_list_df)} rows have "
f"no {id_column} and belong to another crop mode; they "
f"are not cell crops and take no part in this join.")
if len(unreadable):
sample = ', '.join(repr(v) for v in unreadable.unique()[:4])
print(f"png_list: {len(unreadable)} row(s) carry a "
f"{id_column} that is not an object number ({sample}"
f"{' ...' if unreadable.nunique() > 4 else ''}); those "
f"crops cannot be matched to an object and are skipped. "
f"'omulti'/'onone' mean the crop overlapped several "
f"cells or none, 'error' means its file name could not "
f"be parsed.")
png_list_df = png_list_df.loc[usable].copy()
labels = labels.loc[usable]
png_list_df[id_column] = labels.astype('int64')
png_list_df.rename(columns={id_column: 'object_label'}, inplace=True)
if 'cell' in dataframes:
join_cols = ['object_label', 'plateID', 'rowID', 'columnID','fieldID']
cell_time = _time_column(dataframes['cell'].columns)
if png_time is not None and cell_time is not None:
if png_time != cell_time:
png_list_df = png_list_df.rename(columns={png_time: cell_time})
join_cols = join_cols + [cell_time]
elif png_time is not None or cell_time is not None:
raise TimelapseKeyMismatch(
f"png_list and cell disagree about the timepoint: png_list "
f"has {png_time!r} and cell has {cell_time!r} (of "
f"{list(TIME_COLUMN_ALIASES)}). One of the two was written "
f"by a non-timelapse run, so there is no timepoint to join "
f"on and joining without it would match every frame's crop "
f"to every frame's object. Re-run the missing step with the "
f"same 'timelapse' setting."
)
_before_png = len(dataframes['cell'])
merged = _merge_with_cardinality(
dataframes['cell'],
png_list_df,
on=join_cols,
how=(join_how('png_list', keep_uninfected=keep_uninfected)
if require_crops else 'left'),
validate='one_to_one',
left_name='cell',
right_name='png_list',
)
_lost_png = _before_png - len(merged)
if _lost_png > 0:
print(f"png_list: {_lost_png} of {_before_png} measured "
f"cell(s) have no crop that can be matched to them and "
f"are not in the joined table. They were measured; "
f"they simply cannot be shown or classified.")
dataframes['cell'] = merged
else:
print("Cell table not found in database tables.")
return png_list_df
for entity in CHILD_ROLES:
if entity in dataframes:
if 'cell_id' not in dataframes[entity].columns:
print(f"{entity} was measured without a cell mask, so its rows "
f"carry no cell_id and cannot be rolled up onto the cell "
f"table; {entity} features are left out of the join. "
f"Re-run Measure with cell_mask_dim set to link them.")
del dataframes[entity]
continue
numeric_cols = dataframes[entity].select_dtypes(include=[np.number]).columns.tolist()
non_numeric_cols = dataframes[entity].select_dtypes(exclude=[np.number]).columns.tolist()
agg_dict = {col: 'mean' for col in numeric_cols}
agg_dict.update({col: 'first' for col in non_numeric_cols if col not in ['cell_id', 'prcf']})
grouping_cols = ['cell_id', 'prcf']
agg_df = dataframes[entity].groupby(grouping_cols).agg(agg_dict)
agg_df['count_' + entity] = dataframes[entity].groupby(grouping_cols).size()
dataframes[entity] = agg_df
joined_df = None
if 'cell' in dataframes:
joined_df = dataframes['cell']
if 'cytoplasm' in dataframes:
joined_df = _merge_with_cardinality(
joined_df,
dataframes['cytoplasm'],
on=['object_label', 'prcf'],
how=join_how('cytoplasm', keep_uninfected=keep_uninfected),
suffixes=('', '_cytoplasm'),
validate='one_to_one',
left_name='cell',
right_name='cytoplasm',
)
if collapse_duplicate_identity:
joined_df = reconcile_duplicates(
joined_df, '_cytoplasm', left_name='cell',
right_name='cytoplasm', on_conflict=duplicate_column_policy)
for entity in CHILD_ROLES:
if entity in dataframes:
joined_df = _merge_with_cardinality(
joined_df,
dataframes[entity],
left_on=['object_label', 'prcf'],
right_index=True,
how=join_how(entity, keep_uninfected=keep_uninfected),
suffixes=('', f'_{entity}'),
validate='one_to_one',
left_name='cell',
right_name=f'aggregated {entity}',
)
if collapse_duplicate_identity:
joined_df = reconcile_duplicates(
joined_df, f'_{entity}', left_name='cell',
right_name=entity, on_conflict=duplicate_column_policy)
from . import schema as _schema
joined_df = _schema.normalise_plate_columns(joined_df)
return joined_df
#: Table holding the settings of the run that wrote the database **last**.
#: Two columns, ``setting_key`` / ``setting_value``, one row per setting.
SETTINGS_TABLE = 'settings'
#: Append-only companion to :data:`SETTINGS_TABLE`: every stage that has ever
#: written settings into this database, oldest first.
SETTINGS_HISTORY_TABLE = 'settings_history'
#: Columns of :data:`SETTINGS_HISTORY_TABLE`. ``setting_key`` /
#: ``setting_value`` come last and are spelled identically to the ``settings``
#: table, so ``SELECT setting_key, setting_value FROM settings_history`` reads
#: exactly like the table it archives.
SETTINGS_HISTORY_COLUMNS = ('run_id', 'stage', 'stamped_utc', 'setting_key',
'setting_value')
def _settings_history_rows(conn):
"""Read :data:`SETTINGS_HISTORY_TABLE`, or ``[]`` when there is none."""
try:
return conn.execute(
f'SELECT {", ".join(SETTINGS_HISTORY_COLUMNS)} '
f'FROM "{SETTINGS_HISTORY_TABLE}" ORDER BY rowid').fetchall()
except sqlite3.Error:
return []
[docs]
def read_settings_history(db_path):
"""Every settings snapshot ever written into ``db_path``, oldest first.
:param db_path: path to a ``measurements.db``.
:returns: list of ``{'run_id', 'stage', 'stamped_utc', 'settings'}``, one
entry per recorded run, oldest first. A database that predates the
history table returns ``[]``.
Example:
.. code-block:: python
from spacr.io import read_settings_history
for run in read_settings_history('.../measurements/measurements.db'):
print(run['stamped_utc'], run['stage'],
run['settings'].get('crop_mode'))
"""
if not os.path.isfile(str(db_path)):
return []
conn = sqlite3.connect(str(db_path), timeout=5)
try:
rows = _settings_history_rows(conn)
finally:
conn.close()
runs = []
index = {}
for run_id, stage, stamped, key, value in rows:
marker = (run_id, stage, stamped)
if marker not in index:
index[marker] = {'run_id': run_id, 'stage': stage,
'stamped_utc': stamped, 'settings': {}}
runs.append(index[marker])
index[marker]['settings'][key] = value
return runs
def _save_settings_to_db(settings, stage=None):
"""Record this run's settings in the database it is about to write.
Two tables, because two different questions are being asked:
* ``settings`` — what the **most recent** run was configured with.
Replaced, which is what :func:`spacr.resume.read_recorded_settings`
reads and compares against before a resume; it has to be exactly one
run's settings or that comparison means nothing.
* ``settings_history`` — **every** run that has ever written settings
here, appended, each tagged with a ``run_id``, a stage name and a UTC
timestamp.
The second is the repair. ``settings`` alone is replace-only, so a database
written by more than one stage — or measured twice, once for cell crops and
once for pathogen crops, both appending to the same ``png_list`` — kept the
last stage's settings only, and every row the earlier ones wrote was left
with no record of how it was produced. Worse, this call happens *before*
any field is measured, so a run that recorded its settings and then died
replaced the settings of the run that actually produced the rows on disk.
Measured on two saves with different ``crop_mode``/``channels``: 1 of 2
stages recoverable before, 2 of 2 after.
A ``settings`` table written before this history existed is copied into the
history the first time this runs, so a database already on disk keeps its
one snapshot instead of losing it to the next run.
:param settings: settings dict; must contain ``src``.
:param stage: name of the pipeline stage, e.g. ``'measure_crop'``. When
None it is taken from ``settings['stage']`` or ``settings['module']``
if either is set, and recorded as ``'unknown'`` otherwise — ``run_id``
and ``stamped_utc`` still keep the runs apart.
:returns: None.
"""
import uuid
from .errors import _utcnow
if stage is None:
for key in ('stage', 'module'):
candidate = settings.get(key)
if isinstance(candidate, str) and candidate.strip():
stage = candidate.strip()
break
stage = stage or 'unknown'
run_id = uuid.uuid4().hex
stamped = _utcnow()
settings_df = pd.DataFrame(list(settings.items()), columns=['setting_key', 'setting_value'])
settings_df['setting_value'] = settings_df['setting_value'].apply(str)
src = os.path.dirname(settings['src'])
directory = f'{src}/measurements'
os.makedirs(directory, exist_ok=True)
conn = sqlite3.connect(f'{directory}/measurements.db', timeout=5)
try:
from .database_schema import migrate_connection
migrate_connection(
conn,
path=os.path.abspath(f'{directory}/measurements.db'),
)
conn.execute(
f'CREATE TABLE IF NOT EXISTS "{SETTINGS_HISTORY_TABLE}" ('
'run_id TEXT, stage TEXT, stamped_utc TEXT, '
'setting_key TEXT, setting_value TEXT)')
insert = (f'INSERT INTO "{SETTINGS_HISTORY_TABLE}" '
f'({", ".join(SETTINGS_HISTORY_COLUMNS)}) VALUES (?,?,?,?,?)')
if not _settings_history_rows(conn):
try:
previous = conn.execute(
f'SELECT setting_key, setting_value '
f'FROM "{SETTINGS_TABLE}"').fetchall()
except sqlite3.Error:
previous = []
if previous:
conn.executemany(insert, [('', 'before-history', '', key, value)
for key, value in previous])
conn.executemany(insert, [(run_id, stage, stamped, key, value)
for key, value
in zip(settings_df['setting_key'],
settings_df['setting_value'])])
settings_df.to_sql(SETTINGS_TABLE, conn, if_exists='replace', index=False)
conn.commit()
finally:
conn.close()
MASK_MOVIE_DPI = 100
MASK_MOVIE_MIN_PX = 320
MASK_MOVIE_MAX_PX = 1024
def _mask_movie_frame_geometry(masks, *, dpi=MASK_MOVIE_DPI,
min_px=MASK_MOVIE_MIN_PX,
max_px=MASK_MOVIE_MAX_PX):
"""Size one movie frame, and its lettering, from the masks it will show.
The frame keeps the mask's own aspect ratio and its own resolution, with
the long side held inside ``[min_px, max_px]``: a small mask is scaled up
far enough for the label numbers to be legible, and a whole-slide field is
scaled down instead of being written at a resolution nobody can play.
The lettering has to be derived here rather than kept at the old constants.
24 pt at dpi 80 on a 4000 px canvas is 0.7 % of the frame height; the same
24 pt on a 512 px frame would be half the picture, and the old ratio on a
512 px frame would be three pixels. Text is therefore sized as a fraction
of the frame, with a floor so it never disappears entirely.
:param masks: the frames the movie will show. Ragged input is allowed and
the largest frame decides, so a mask that grew mid-series is not
cropped.
:param dpi: dots per inch handed to the writer. Only the product
``figsize * dpi`` reaches the file, but the two are kept separate
because point sizes are relative to inches.
:returns: ``dict`` with ``figsize``, ``dpi``, ``frame_px`` (width, height),
``label_pt``, ``title_pt``, ``caption_pt`` and ``band`` -- the fraction
of the figure reserved above and below the image for the two captions.
:raises ValueError: when there is no mask to measure, which would otherwise
surface as an unreadable empty GIF.
"""
shapes = [np.asarray(mask).shape[:2] for mask in masks]
shapes = [(int(shape[0]), int(shape[1])) for shape in shapes
if len(shape) == 2 and shape[0] > 0 and shape[1] > 0]
if not shapes:
raise ValueError(
'cannot size a mask movie: no frame has a non-empty shape. '
'An empty animation writes a GIF no player can open.')
height = max(h for h, _ in shapes)
width = max(w for _, w in shapes)
long_side = max(height, width)
scale = 1.0
if long_side < min_px:
scale = min_px / long_side
elif long_side > max_px:
scale = max_px / long_side
frame_h = max(1, int(round(height * scale)))
frame_w = max(1, int(round(width * scale)))
def _points(fraction):
"""A font size in POINTS for a fraction of the frame's short side.
Matplotlib sizes text in points and this geometry is in pixels, so
the conversion has to use the dpi the writer will actually use --
computing it against a default dpi puts the labels at the wrong size
in the file while looking right on screen. Floored at 5 pt, below
which a label is ink rather than text.
"""
return max(5.0, round(min(frame_h, frame_w) * fraction * 72.0 / dpi, 1))
return {
'figsize': (frame_w / dpi, frame_h / dpi),
'dpi': dpi,
'frame_px': (frame_w, frame_h),
'label_pt': _points(1 / 28),
'title_pt': _points(1 / 22),
'caption_pt': _points(1 / 26),
'band': 0.06,
}
def _save_mask_timelapse_as_gif(masks, tracks_df, path, cmap, norm, filenames):
"""
Save a timelapse animation of masks as a GIF.
The frame is sized from the mask by :func:`_mask_movie_frame_geometry`
rather than at a fixed 50 x 50 inches, and a band is reserved at the top
and bottom of the figure for the two captions. The frame counter used to
be drawn by ``ax.set_title`` into a figure whose axes had been given every
last inch by ``subplots_adjust(top=1)``, so it was clipped away on every
frame: measured on a rendered GIF, zero lit pixels in the top 5 % of the
canvas against 222 in the bottom 5 % where the filename sits.
Parameters:
- masks (list): List of mask frames.
- tracks_df (pandas.DataFrame): DataFrame containing track information.
- path (str): Path to save the GIF file.
- cmap (str or matplotlib.colors.Colormap): Colormap for displaying the masks.
- norm (matplotlib.colors.Normalize): Normalization for the colormap.
- filenames (list): List of filenames corresponding to each mask frame.
Returns:
None
"""
geometry = _mask_movie_frame_geometry(masks)
band = geometry['band']
with figure_style(theme_target()):
fig, ax = plt.subplots(figsize=geometry['figsize'], facecolor='black')
from .figures.bundle import _register_figure_data
_register_figure_data(fig, None, kind="mask", title="Mask timelapse")
ax.set_facecolor('black')
ax.axis('off')
plt.subplots_adjust(left=0, right=1, top=1 - band, bottom=band,
wspace=0, hspace=0)
filename_text_obj = None
frame_text_obj = None
def _update(frame):
"""
Update the frame of the animation.
Parameters:
- frame (int): The frame number to update.
Returns:
None
"""
nonlocal filename_text_obj, frame_text_obj
if filename_text_obj is not None:
filename_text_obj.remove()
if frame_text_obj is not None:
frame_text_obj.remove()
ax.clear()
ax.axis('off')
current_mask = masks[frame]
ax.imshow(current_mask, cmap=cmap, norm=norm)
frame_text_obj = fig.text(0.5, 1 - band / 2, f'Frame: {frame}',
ha='center', va='center',
fontsize=geometry['title_pt'], color='white')
filename_text = filenames[frame]
filename_text_obj = fig.text(0.5, band / 2, filename_text, ha='center', va='center', fontsize=geometry['caption_pt'], color='white')
for label_value in np.unique(current_mask):
if label_value == 0: continue
y, x = np.mean(np.where(current_mask == label_value), axis=1)
ax.text(x, y, str(label_value), color='white', fontsize=geometry['label_pt'], ha='center', va='center')
if tracks_df is not None:
for track in tracks_df['track_id'].unique():
_track = tracks_df[tracks_df['track_id'] == track]
ax.plot(_track['x'], _track['y'], '-w', linewidth=1)
anim = FuncAnimation(fig, _update, frames=len(masks), blit=False)
anim.save(
path,
writer='pillow',
fps=2,
dpi=geometry['dpi'],
savefig_kwargs={'facecolor': 'black', 'transparent': False},
)
plt.close(fig)
print(f'Saved timelapse to {path}')
def _save_object_counts_to_database(arrays, object_type, file_names, db_path, added_string):
"""
Save the counts of unique objects in masks to a SQLite database.
Args:
arrays (List[np.ndarray]): List of masks.
object_type (str): Type of object.
file_names (List[str]): List of file names corresponding to the masks.
db_path (str): Path to the SQLite database.
added_string (str): Additional string to append to the count type.
Returns:
None
"""
def _count_objects(mask):
"""Count unique objects in a mask, assuming 0 is the background."""
unique, counts = np.unique(mask, return_counts=True)
if unique[0] == 0:
return len(unique) - 1
return len(unique)
records = []
for mask, file_name in zip(arrays, file_names):
object_count = _count_objects(mask)
count_type = f"{object_type}{added_string}"
records.append((file_name, count_type, object_count))
from .database_concurrency import connect as _connect_database
conn = _connect_database(db_path)
try:
from .database_schema import migrate_connection
migrate_connection(conn, path=os.path.abspath(db_path))
cursor = conn.cursor()
cursor.execute('''
CREATE TABLE IF NOT EXISTS object_counts (
file_name TEXT,
count_type TEXT,
object_count INTEGER,
PRIMARY KEY (file_name, count_type)
)
''')
cursor.executemany('''
INSERT INTO object_counts (file_name, count_type, object_count)
VALUES (?, ?, ?)
ON CONFLICT(file_name, count_type) DO UPDATE SET
object_count = excluded.object_count
''', records)
conn.commit()
finally:
conn.close()
def _create_database(db_path):
"""Create a SQLite database at the current spaCR schema version.
Existing databases are migrated in place. A database from a newer spaCR
release raises :class:`spacr.database_schema.DatabaseSchemaTooNewError`
rather than being silently opened with older code.
"""
try:
from .database_concurrency import connect as _connect_database
conn = _connect_database(db_path)
except sqlite3.Error as error:
print(error)
return
try:
from .database_schema import migrate_connection
migrate_connection(conn, path=os.path.abspath(db_path))
finally:
conn.close()
[docs]
def save_object_mask(output_folder, filename, mask, compression='lzw'):
"""Save an integer label mask as a lossless, compressed TIFF.
Masks are saved as TIFF (not .npy) so they're readable by ImageJ/other
tools, with lossless compression (default LZW). Object labels are NEVER
altered — the array is written verbatim as uint16, exactly as recorded in
the measurements database.
:param output_folder: destination folder (e.g. ``masks/cell_mask_stack``).
:param filename: reference filename (the stack basename; extension ignored).
:param mask: 2-D integer label array.
:param compression: lossless codec — ``'lzw'`` | ``'zlib'`` | ``'none'``.
:returns: the path written.
"""
base = os.path.splitext(os.path.basename(filename))[0]
out_path = os.path.join(output_folder, base + '.tif')
comp = None if str(compression).lower() in ('none', '', 'no', 'false') else str(compression).lower()
write_tiff(
out_path,
np.asarray(mask).astype(np.uint16),
compression=comp,
)
return out_path
def _mask_variant_path(folder, ref_filename):
"""Return the path to ``ref_filename``'s array in ``folder``, preferring a
compressed ``.tif`` mask, then legacy ``.npy``, then the exact name."""
base = os.path.splitext(ref_filename)[0]
for cand in (base + '.tif', base + '.tiff', base + '.npy',
ref_filename):
p = os.path.join(folder, cand)
if os.path.isfile(p):
return p
return None
def _listdir_visible(folder):
"""List ``folder`` like :func:`os.listdir`, leaving out every name that starts with a dot.
The folders the Mask pipeline writes and re-reads (``stack/``,
``masks/``, ``masks/<object>_mask_stack/``, ``merged/``, ``test/``)
can hold two kinds of dot-file that end in ``.npy`` or ``.npz`` and are
not arrays: the AppleDouble ``._<name>`` sidecar macOS writes beside a
file on a volume that cannot store extended attributes natively (exFAT,
FAT, many SMB shares), and the ``.spacr_tmp_*.npy`` / ``.spacr_npz_*.npz``
temporaries that spaCR versions before the ``.partial`` naming of
:func:`_replace_atomically` left behind when a run was killed mid-write.
:func:`numpy.load` reads a sidecar as a pickle and refuses it, and reads
such a temporary as a truncated array. spaCR never names a field with a
leading dot.
:param folder: the directory to list.
:returns: the entry names, in :func:`os.listdir` order, without the
dot-files.
:raises OSError: whatever :func:`os.listdir` raises for ``folder``.
"""
return [name for name in os.listdir(folder) if not name.startswith('.')]
def _save_array_atomic(output_path, array):
"""Write ``array`` to ``output_path`` as ``.npy`` atomically.
``np.save(path, arr)`` writes straight onto the destination, so a
process killed part-way through — full disk, OOM killer, Ctrl-C —
leaves a *short* file at the final name. It still starts with the
``.npy`` magic and still parses as a header, so anything that decides
"this field is done because the file is there" will happily accept it
and measure whatever bytes happened to land. That is the failure mode
that turns a resume into silently corrupt output.
Writing to a sibling temporary file and then ``os.replace``-ing it
into position makes the destination atomic within the filesystem: it
is either the previous content or the complete new array, never a
prefix of it. The temp file is removed if anything goes wrong, and one
left by a killed process ends in ``.partial``, not ``.npy``, so no
listing of fields mistakes it for one (see :func:`_replace_atomically`).
:param output_path: final ``.npy`` path.
:param array: array to write.
:returns: ``output_path``.
"""
return _replace_atomically(
output_path, lambda handle: np.save(handle, array),
prefix='.spacr_tmp_')
def _load_array_any(path):
"""Load a ``.tif``/``.tiff`` (via tifffile) or ``.npy`` array.
Never unpickles. Every mask spaCR writes is a plain ``uint16`` array, so
a ``.npy`` holding pickled objects is refused with numpy's
:class:`ValueError` rather than run.
"""
if path.endswith(('.tif', '.tiff')):
import tifffile
return tifffile.imread(path)
return np.load(path, allow_pickle=False)
def _load_and_concatenate_arrays(
src, channels, cell_chann_dim, nucleus_chann_dim,
pathogen_chann_dim, organelle_chann_dim, resume=False,
organelle_chann_dims=None, mask_folders=None):
"""
Load and concatenate arrays from multiple folders.
Every merged stack is written **atomically** — to a temporary file in
the destination folder, then ``os.replace``\\ d into place. The plain
``np.save`` this used to do wrote straight onto the destination, so a
run killed mid-write (full disk, OOM, Ctrl-C) left a short file that
still looked like a valid ``.npy`` to anything that only checked
whether it existed. ``os.replace`` is atomic within a filesystem, so
``merged/<field>.npy`` is now either absent or complete.
Args:
src (str): The source directory containing the arrays.
channels (list): List of channel indices to select from the arrays.
cell_chann_dim (int): Dimension of the cell channel.
nucleus_chann_dim (int): Dimension of the nucleus channel.
pathogen_chann_dim (int): Dimension of the pathogen channel.
organelle_chann_dim (int or None): Dimension of the organelle channel. If None, organelle masks are included only if the folder exists.
resume (bool): Opt-in checkpointing. When True, fields whose merged
stack is already present **and verified complete** are skipped, so
a run that died at field 900 of 1000 does not redo the first 900.
Verification is deliberately not ``os.path.exists``: files left
behind by older, non-atomic versions of this function can be
truncated, and those are re-merged rather than trusted. Default
False, which redoes every field exactly as before.
mask_folders (dict or None): optional role-to-folder overrides for
finalized masks, for example ``{'cell': adjusted_cell_folder}``.
Other roles retain their ordinary masks folders. Unknown roles
and missing override directories are refused before any output.
Returns:
None
"""
from .utils import print_progress
from .resume import completed_fields_in_merged, format_resume, plan_resume
overrides = dict(mask_folders or {})
if set(overrides) - {'cell', 'nucleus', 'pathogen', *ORGANELLE_ROLES}:
raise ValueError('Unknown object role in mask folder overrides')
for role, folder in overrides.items():
if not os.path.isdir(folder):
raise ValueError(f'Mask folder override for {role} is not a directory: {folder}')
folder_paths = [os.path.join(src+'/stack')]
mask_roles = []
try:
_mask_stacks = set(_listdir_visible(os.path.join(src, 'masks')))
except OSError:
_mask_stacks = set()
def add_mask_folder(role, enabled):
"""Queue one object's mask stack, if this run has that object.
EITHER the caller named a channel dimension for it OR the folder is
on disk: a run that segmented an object always has the folder, and a
run being re-read from settings may name the object before the
folder is written. Requiring both would drop a mask stack that is
sitting right there.
:param role: the object, e.g. ``'cell'``.
:param enabled: that object's channel dimension, or None.
"""
folder = overrides.get(role, os.path.join(src, 'masks', f'{role}_mask_stack'))
if enabled is not None or role in overrides or f'{role}_mask_stack' in _mask_stacks:
folder_paths.append(folder)
mask_roles.append(role)
add_mask_folder('cell', cell_chann_dim)
add_mask_folder('nucleus', nucleus_chann_dim)
add_mask_folder('pathogen', pathogen_chann_dim)
add_mask_folder('organelle', organelle_chann_dim)
extra_dims = dict(organelle_chann_dims or {})
for role in ORGANELLE_ROLES[1:]:
add_mask_folder(role, extra_dims.get(role))
output_folder = src+'/merged'
reference_folder = folder_paths[0]
os.makedirs(output_folder, exist_ok=True)
count=0
reference_files = _listdir_visible(reference_folder)
from .image_quality import excluded_fields
rejected_quality = excluded_fields(src)
reference_files = [name for name in reference_files if name not in rejected_quality]
all_imgs = len(reference_files)
time_ls = []
layout_written = False
reference_npy = next(
(name for name in reference_files if name.endswith('.npy')), None)
if reference_npy is not None:
if channels is None:
reference_array = np.load(
os.path.join(reference_folder, reference_npy), mmap_mode='r')
n_intensity_for_layout = int(reference_array.shape[-1])
else:
n_intensity_for_layout = len(channels)
intended_layout = {
'version': 1,
'intensity_channels': (
list(channels) if channels is not None
else list(range(n_intensity_for_layout))),
'mask_plane_order': list(mask_roles),
'mask_dims': {
role: n_intensity_for_layout + index
for index, role in enumerate(mask_roles)
},
}
manifest_path = os.path.join(output_folder, MERGED_LAYOUT_SIDECAR)
if resume and os.path.isfile(manifest_path):
try:
with open(manifest_path, 'r', encoding='utf-8') as handle:
existing_layout = json.load(handle)
except (OSError, json.JSONDecodeError) as exc:
raise ValueError(
f'Cannot resume: merged plane layout is unreadable: '
f'{manifest_path}: {exc}') from exc
if existing_layout != intended_layout:
raise ValueError(
'Cannot resume merged arrays with a different plane '
f'layout. Existing: {existing_layout!r}; requested: '
f'{intended_layout!r}. Start a non-resume merge into a '
'clean destination so every field uses one layout.')
already_done = set()
if resume:
candidates = [os.path.splitext(f)[0] for f in reference_files
if f.endswith('.npy')]
rejected = {}
already_done = completed_fields_in_merged(
output_folder, reasons=rejected, fields=candidates)
print(format_resume(plan_resume(candidates, already_done,
reasons=rejected, enabled=True,
src=output_folder)))
for idx, filename in enumerate(reference_files):
start = time.time()
stack_ls = []
if filename.endswith('.npy') and os.path.splitext(filename)[0] not in already_done:
count += 1
exists_in_all_folders = all(
_mask_variant_path(folder, filename) is not None
for folder in folder_paths)
if exists_in_all_folders:
ref_array_path = os.path.join(reference_folder, filename)
concatenated_array = np.load(ref_array_path)
if channels is not None:
concatenated_array = np.take(concatenated_array, channels, axis=-1)
if not layout_written:
n_intensity = int(concatenated_array.shape[-1])
layout = {
'version': 1,
'intensity_channels': (
list(channels) if channels is not None
else list(range(n_intensity))),
'mask_plane_order': list(mask_roles),
'mask_dims': {
role: n_intensity + index
for index, role in enumerate(mask_roles)
},
}
fd, temporary = tempfile.mkstemp(
prefix='.spacr_plane_layout_', suffix='.json',
dir=output_folder)
try:
with os.fdopen(fd, 'w', encoding='utf-8') as handle:
json.dump(layout, handle, indent=2, sort_keys=True)
handle.write('\n')
handle.flush()
os.fsync(handle.fileno())
os.replace(
temporary,
os.path.join(output_folder,
MERGED_LAYOUT_SIDECAR))
except BaseException:
try:
os.remove(temporary)
except OSError:
pass
raise
layout_written = True
stack_ls.append(concatenated_array)
for folder in folder_paths[1:]:
array_path = _mask_variant_path(folder, filename)
array = _load_array_any(array_path)
if array.ndim in (2, concatenated_array.ndim - 1):
array = np.expand_dims(array, axis=-1)
stack_ls.append(array)
if len(stack_ls) > 0:
stack_ls = [np.expand_dims(arr, axis=-1) if arr.ndim == 2 else arr for arr in stack_ls]
unique_shapes = {arr.shape[:-1] for arr in stack_ls}
if len(unique_shapes) > 1:
max_tuple_length = max(len(shape) for shape in unique_shapes)
padded_shapes = [shape + (0,) * (max_tuple_length - len(shape)) for shape in unique_shapes]
max_dims = np.max(np.array(padded_shapes), axis=0)
print(f'Warning: arrays with multiple shapes found. Padding arrays to max X,Y dimentions {max_dims}')
padded_stack_ls = []
for arr in stack_ls:
pad_width = [(0, max_dim - dim) for max_dim, dim in zip(max_dims, arr.shape[:-1])]
pad_width.append((0, 0))
padded_arr = np.pad(arr, pad_width)
padded_stack_ls.append(padded_arr)
stack = np.concatenate(padded_stack_ls, axis=-1)
else:
stack = np.concatenate(stack_ls, axis=-1)
if stack.shape[-1] > concatenated_array.shape[-1]:
output_path = os.path.join(output_folder, filename)
_save_array_atomic(output_path, stack)
stop = time.time()
duration = stop - start
time_ls.append(duration)
files_processed = idx+1
files_to_process = all_imgs
print_progress(files_processed, files_to_process, n_jobs=1, time_ls=time_ls, batch_size=None, operation_type="Merging Arrays")
return
def _results_to_csv(src, df, df_well):
"""
Save the given dataframes as CSV files in the specified directory.
Args:
src (str): The directory path where the CSV files will be saved.
df (pandas.DataFrame): The dataframe containing cell data.
df_well (pandas.DataFrame): The dataframe containing well data.
Returns:
tuple: A tuple containing the cell dataframe and well dataframe.
"""
cells = df
wells = df_well
results_loc = src+'/results'
wells_loc = results_loc+'/wells.csv'
cells_loc = results_loc+'/cells.csv'
os.makedirs(results_loc, exist_ok=True)
wells.to_csv(wells_loc, index=True, header=True)
cells.to_csv(cells_loc, index=True, header=True)
return cells, wells
[docs]
def read_plot_model_stats(train_file_path, val_file_path ,save=False):
"""Plot training vs. validation curves from a saved model's per-epoch CSVs.
:param train_file_path: Path to the training stats CSV.
:param val_file_path: Path to the validation stats CSV.
:param save: If True, write the figures next to the training CSV instead
of showing them. The file format follows the user's figure preference
via :func:`spacr.plot.save_figure`, which also corrects the extension.
Default ``False``.
:returns: None
"""
def _plot_and_save(train_df, val_df, column='accuracy', save=False, path=None, dpi=None):
"""Draw one training curve -- train against validation -- and write it.
One function per COLUMN rather than per figure because the caller
asks for accuracy, loss and the rest by name, and every one of them
is the same plot of the same two frames.
:param train_df: per-epoch training statistics.
:param val_df: the same for validation.
:param column: which statistic to draw.
:param save: write a PDF beside the model rather than only showing it.
:param path: the folder to write into.
:param dpi: resolution for the written file.
"""
pdf_path = os.path.join(path, f'{column}.pdf')
with figure_style(theme_target()):
fig, axes = plt.subplots(1, 2, figsize=(20, 10), sharey=True)
from .figures.bundle import _register_figure_data
_register_figure_data(fig, lambda: pd.concat([train_df.assign(split="train"), val_df.assign(split="val")], ignore_index=True), x="epoch", y=column, hue="split", kind="line")
sns.lineplot(ax=axes[0], x='epoch', y=column, data=train_df, marker='o', color='red')
sns.lineplot(ax=axes[1], x='epoch', y=column, data=val_df, marker='o', color='blue')
axes[0].set_title(f'Train {column} vs. Epoch', fontsize=20)
axes[0].set_xlabel('Epoch', fontsize=16)
axes[0].set_ylabel(column, fontsize=16)
axes[0].tick_params(axis='both', which='major', labelsize=12)
axes[1].set_title(f'Validation {column} vs. Epoch', fontsize=20)
axes[1].set_xlabel('Epoch', fontsize=16)
axes[1].tick_params(axis='both', which='major', labelsize=12)
plt.tight_layout()
if save:
from .plot import save_figure
pdf_path = save_figure(plt.gcf(), pdf_path, dpi=dpi)
else:
plt.show()
train_df = pd.read_csv(train_file_path, index_col=0)
val_df = pd.read_csv(val_file_path, index_col=0)
fldr_1 = os.path.dirname(train_file_path)
with plt.rc_context():
if save:
sns.set(style="whitegrid")
_plot_and_save(train_df, val_df, column='accuracy', save=save, path=fldr_1)
_plot_and_save(train_df, val_df, column='neg_accuracy', save=save, path=fldr_1)
_plot_and_save(train_df, val_df, column='pos_accuracy', save=save, path=fldr_1)
_plot_and_save(train_df, val_df, column='loss', save=save, path=fldr_1)
_plot_and_save(train_df, val_df, column='prauc', save=save, path=fldr_1)
_plot_and_save(train_df, val_df, column='optimal_threshold', save=save, path=fldr_1)
def _save_model(model, model_type, results_dict, dst, epoch, epochs,
intermedeate_save=None,
channels=None,
val_dict=None, optimizer=None, scheduler=None,
best_metric=None, is_best=None,
epochs_without_improvement=0, preprocessing=None,
classes=None, force_last=False):
"""
Save the model based on certain conditions during training.
Args:
model (torch.nn.Module): The trained model to be saved.
model_type (str): The type of the model.
results_df (pandas.DataFrame): The dataframe containing the validation results.
dst (str): The destination directory to save the model.
epoch (int): The current epoch number.
epochs (int): The total number of epochs.
intermedeate_save (list, optional): List of accuracy thresholds to trigger intermediate model saves.
Defaults to [0.99, 0.98, 0.95, 0.94].
channels (list, optional): List of channels used. Defaults to ['r', 'g', 'b'].
"""
from .torch_artifacts import save_model_artifact
if channels is None:
channels = ['r', 'g', 'b']
channels_str = ''.join(channels)
check_dict = val_dict if val_dict is not None else results_dict
acc = float(check_dict.get('accuracy', float('nan')))
common = {
'optimizer': optimizer,
'scheduler': scheduler,
'epoch': epoch,
'metrics': check_dict,
'best_metric': best_metric,
'epochs_without_improvement': epochs_without_improvement,
'preprocessing': preprocessing,
'classes': classes,
'channels': channels,
}
saved_best = None
if is_best:
saved_best = os.path.join(
dst, f'{model_type}_best_channels_{channels_str}.pth')
save_model_artifact(
model, saved_best, artifact_role='best', **common)
print(f"Saved new best model at epoch {epoch} "
f"(validation accuracy {acc:.4f}): {saved_best}")
if intermedeate_save is True or intermedeate_save is None:
thresholds = [0.99, 0.98, 0.95, 0.94]
elif intermedeate_save is False:
thresholds = []
else:
thresholds = sorted(
{float(value) for value in intermedeate_save}, reverse=True)
saved_archive = None
if is_best is not False and np.isfinite(acc):
crossed = next((value for value in thresholds if acc >= value), None)
if crossed is not None:
saved_archive = os.path.join(
dst,
f'{model_type}_epoch_{epoch}_acc_{acc * 100:.4f}_'
f'channels_{channels_str}.pth')
save_model_artifact(
model, saved_archive, artifact_role='milestone', **common)
saved_last = None
if force_last or epoch % 100 == 0 or epoch == epochs:
saved_last = os.path.join(
dst, f'{model_type}_last_channels_{channels_str}.pth')
save_model_artifact(
model, saved_last, artifact_role='last', **common)
return saved_best or saved_last or saved_archive
def _save_progress(dst, train_df, validation_df):
"""
Save the progress of the classification model.
Parameters:
dst (str): The destination directory to save the progress.
train_df (pandas.DataFrame): The DataFrame containing training stats.
validation_df (pandas.DataFrame): The DataFrame containing validation stats (if available).
Returns:
None
"""
def _save_df_to_csv(file_path, df):
"""
Save the given DataFrame to the specified CSV file, either creating a new file or appending to an existing one.
Parameters:
file_path (str): The file path where the CSV will be saved.
df (pandas.DataFrame): The DataFrame to save.
"""
if not os.path.exists(file_path):
with open(file_path, 'w') as f:
df.to_csv(f, index=True, header=True)
f.flush()
else:
with open(file_path, 'a') as f:
df.to_csv(f, index=True, header=False)
f.flush()
os.makedirs(dst, exist_ok=True)
results_path_train = os.path.join(dst, 'train.csv')
results_path_validation = os.path.join(dst, 'validation.csv')
_save_df_to_csv(results_path_train, train_df)
if validation_df is not None:
_save_df_to_csv(results_path_validation, validation_df)
read_plot_model_stats(results_path_train, results_path_validation, save=True)
return
def _copy_missclassified(df):
"""Copy every misclassified crop into a ``missclassified`` folder beside its source.
Split into ``pc`` and ``nc`` by what the original path contains, so the
two failure directions can be looked at separately -- a model that only
ever errs one way is a different problem from one that errs both.
:param df: predictions carrying ``true_label``, ``predicted_label`` and
``filename``.
"""
misclassified = df[df['true_label'] != df['predicted_label']]
for _, row in misclassified.iterrows():
original_path = row['filename']
filename = os.path.basename(original_path)
dest_folder = os.path.dirname(os.path.dirname(original_path))
if "pc" in original_path:
new_path = os.path.join(dest_folder, "missclassified/pc", filename)
else:
new_path = os.path.join(dest_folder, "missclassified/nc", filename)
os.makedirs(os.path.dirname(new_path), exist_ok=True)
shutil.copy(original_path, new_path)
print(f"Copied {len(misclassified)} misclassified images.")
return
def _read_db(db_loc, tables):
"""Read tables out of a measurements database as data frames.
A ``~`` path is expanded HERE, once, for every reader. A ``src``
beginning with ``~`` produced a literal ``~/...`` path that the schema
migration resolved against the WORKING DIRECTORY and then refused. It is
fixed here rather than in the migration, whose own docstring states
non-expansion as a deliberate contract.
:param db_loc: the database path.
:param tables: the table names to read.
:returns: one data frame per requested table, in the order asked.
"""
import gc
import os
import pathlib
import sqlite3
import pandas as pd
from .database_schema import ensure_database_schema
from .tabular import _backend_of, read_database
backend = _backend_of(db_loc)
if backend != 'postgres' and isinstance(db_loc, (str, os.PathLike)):
db_loc = os.path.expanduser(os.path.expandvars(os.fspath(db_loc)))
from .utils import correct_metadata
def _quote_identifier(name):
"""Safely quote SQLite identifiers (e.g., table names)."""
if not isinstance(name, str) or not name:
raise ValueError(f"Invalid table name: {name!r}")
return '"' + name.replace('"', '""') + '"'
tables = [tables] if isinstance(tables, str) else list(tables)
for table in tables:
_quote_identifier(table)
if backend != 'sqlite':
return [correct_metadata(frame) for frame in read_database(
db_loc, tables, canonicalise=False, report=None,
migrate=False, read_only=True)]
directory = os.path.dirname(os.path.abspath(db_loc)) or "."
writable = (os.access(db_loc, os.W_OK) and os.access(directory, os.W_OK))
if writable:
ensure_database_schema(db_loc)
dfs = []
chunksize = 100_000
if writable:
connect_to = db_loc
connect_kwargs = {}
else:
connect_to = f"file:{pathlib.Path(db_loc).as_uri()[7:]}?mode=ro"
connect_kwargs = {"uri": True}
with sqlite3.connect(connect_to, timeout=30, **connect_kwargs) as conn:
existing_tables = {
row[0]
for row in conn.execute(
"SELECT name FROM sqlite_master WHERE type='table'"
).fetchall()
}
print('existing_tables:', existing_tables)
print('tables:', tables)
for table in tables:
if table not in existing_tables:
raise ValueError(f"Table not found in database: {table}")
quoted_table = _quote_identifier(table)
query = f"SELECT * FROM {quoted_table}"
chunks = []
for chunk in pd.read_sql_query(query, conn, chunksize=chunksize):
chunks.append(chunk)
if len(chunks) == 0:
df = pd.read_sql_query(f"SELECT * FROM {quoted_table} LIMIT 0", conn)
elif len(chunks) == 1:
df = chunks[0]
else:
df = pd.concat(chunks, ignore_index=True)
del chunks
gc.collect()
df = correct_metadata(df)
dfs.append(df)
del df
gc.collect()
return dfs
def _read_and_merge_data(
locs, tables, verbose=False, nuclei_limit=10, pathogen_limit=10,
change_plate=False, acquisition_conflict="raise",
keep_uninfected=True):
"""Read object tables and merge their measurements by parent object.
Shared acquisition-stamp values are coalesced when one table is missing a
value. Conflicting non-null values are rejected by default because features
measured with different dimensionality, units, or voxel calibration cannot
safely share one row. A caller that has independently established which
source is authoritative may explicitly pass ``"prefer_left"`` or
``"prefer_right"``.
:param locs: measurement database paths.
:param tables: tables to read and merge.
:param verbose: print table and merge sizes.
:param nuclei_limit: maximum nuclei per parent cell, or ``None``.
:param pathogen_limit: maximum pathogens per parent cell, or ``None``.
:param change_plate: replace plate IDs according to database order.
:param acquisition_conflict: ``"raise"`` (default), ``"prefer_left"``,
or ``"prefer_right"``.
:raises AcquisitionMetadataConflictError: when two non-null acquisition
stamps differ under the default policy.
:raises ValueError: when ``acquisition_conflict`` is not a known policy.
:returns: merged feature frame and the ungrouped object-table frames.
"""
from pandas.api.types import is_object_dtype
from .utils import MEASUREMENT_STAMP_COLUMNS, _split_data
conflict_policies = {"raise", "prefer_left", "prefer_right"}
if acquisition_conflict not in conflict_policies:
raise ValueError(
"acquisition_conflict must be one of "
f"{sorted(conflict_policies)}, got {acquisition_conflict!r}."
)
pathogen_counts = None
metadata_key = 'object_label'
shared_metadata_columns = set(MEASUREMENT_STAMP_COLUMNS)
def _merge_grouped(left, right, right_name="grouped object data"):
"""Merge grouped tables while keeping only one copy of shared acquisition metadata.
THE JOIN TYPE NOW COMES FROM THE REGISTRY. It used to be inner
unconditionally -- pandas' default, since no ``how=`` was passed --
and this docstring said the choice was "deliberately left alone"
because the decision had not been made. It has since: `object_roles.
join_how` records it and `_read_and_join_tables` already reads it, so
the two readers of the same tables were disagreeing about which
objects exist.
nucleus INNER a cell with no nucleus is debris
png_list INNER a cell with no crop cannot be classified
cytoplasm LEFT one row per cell; it makes no difference
pathogen LEFT an UNINFECTED cell is usually the control
organelle LEFT same reasoning
Inner for pathogen was the consequential one: it silently conditioned
every result on infection, deleting the control population from the
denominator without a word.
A ``right_name`` the registry does not know keeps the historical
inner join rather than being guessed at -- the metadata and stamp
merges go through here too, and they are not object tables.
What is NOT defensible is doing it in silence, which is what this
used to do. The discontinuity is brutal: on a 100-cell plate where
NO crop carries a usable object id, png_list drops out before the
join and all 100 cells survive; where exactly ONE does, the merge
keeps that one and deletes the other 99. Nothing printed either way,
and every shipped caller passes verbose=False.
So the shortfall is reported, named by table. This is the mirror of
`_report_fan_out` for the shrinking direction.
"""
if left.empty:
return right.copy()
if right.empty:
return left.copy()
shared = [col for col in shared_metadata_columns if col in left.columns and col in right.columns]
for col in shared:
common_idx = left.index.intersection(right.index)
if len(common_idx):
a = left.loc[common_idx, col]
b = right.loc[common_idx, col]
mismatch = a.notna() & b.notna() & a.ne(b)
if mismatch.any():
conflict_index = common_idx[mismatch.to_numpy()]
examples = ", ".join(
f"{index!r}: {a.loc[index]!r} != {b.loc[index]!r}"
for index in conflict_index[:3]
)
count = int(mismatch.sum())
if acquisition_conflict == "raise":
raise AcquisitionMetadataConflictError(
f"{count} conflicting value(s) for acquisition "
f"metadata column {col!r}. Examples: {examples}. "
"These tables were measured under incompatible "
"acquisition settings. Repair or remeasure them, "
"or explicitly pass acquisition_conflict="
"'prefer_left'/'prefer_right' only when that "
"choice is scientifically justified."
)
if acquisition_conflict == "prefer_right":
left.loc[conflict_index, col] = b.loc[conflict_index]
print(
f"Resolved {count} conflicting value(s) for "
f"acquisition metadata column {col!r} with "
f"acquisition_conflict={acquisition_conflict!r}."
)
missing_left = a.isna() & b.notna()
if missing_left.any():
fill_index = common_idx[missing_left.to_numpy()]
left.loc[fill_index, col] = right.loc[fill_index, col]
right = right.drop(columns=shared)
before = len(left)
from .object_roles import JOIN_HOW
how = (join_how(right_name, keep_uninfected=keep_uninfected)
if str(right_name).strip().lower() in JOIN_HOW else "inner")
result = _merge_with_cardinality(
left,
right,
left_index=True,
right_index=True,
how=how,
validate="one_to_one",
left_name="grouped object data",
right_name=right_name,
)
lost = before - len(result)
if lost > 0:
print(
f"{lost} of {before} objects have no row in {right_name} and "
f"were removed from the merged data. If {right_name} is "
f"expected to cover every object, that is a gap in the "
f"database rather than a filter."
)
return result
def _split_object_data(frame, group_by, object_type):
"""Group object data while retaining its complete provenance stamp."""
numeric, non_numeric = _split_data(frame, group_by, object_type)
stamp_columns = [
column for column in MEASUREMENT_STAMP_COLUMNS
if column in frame.columns
]
if stamp_columns:
grouped_stamp = (
frame.set_index(group_by)[stamp_columns]
.groupby(level=0, sort=False)
.first()
)
numeric = _merge_grouped(numeric, grouped_stamp)
return numeric, non_numeric
data_dict = {table: [] for table in tables}
for idx, loc in enumerate(locs):
db_dfs = _read_db(loc, tables)
if change_plate:
for df in db_dfs:
df['plateID'] = f'plate{idx+1}'
df['prc'] = df['plateID'].astype(str) + '_' + df['rowID'].astype(str) + '_' + df['columnID'].astype(str)
for table, df in zip(tables, db_dfs):
data_dict[table].append(df)
for table, dfs in data_dict.items():
if dfs:
frame = pd.concat(dfs, axis=0)
if 'prcf' in frame.columns and is_object_dtype(frame['prcf']):
frame['prcf'] = frame['prcf'].astype(str)
data_dict[table] = frame
if verbose:
print(f"{table}: {len(data_dict[table])}")
merged_df = pd.DataFrame()
if 'cell' in data_dict:
cells = data_dict['cell'].copy()
cells = cells.assign(object_label=lambda x: 'o' + x['object_label'].astype(int).astype(str))
cells = cells.assign(prcfo=lambda x: x['prcf'] + '_' + x['object_label'])
cells_g_df, metadata = _split_object_data(
cells, 'prcfo', 'object_label')
merged_df = cells_g_df.copy()
if verbose:
print(f'cells: {len(cells)}, cells grouped: {len(cells_g_df)}')
if 'cytoplasm' in data_dict:
cytoplasms = data_dict['cytoplasm'].copy()
cytoplasms = cytoplasms.assign(object_label=lambda x: 'o' + x['object_label'].astype(int).astype(str))
cytoplasms = cytoplasms.assign(prcfo=lambda x: x['prcf'] + '_' + x['object_label'])
if 'cell' not in data_dict:
merged_df, metadata = _split_object_data(
cytoplasms, 'prcfo', 'object_label')
if verbose:
print(f'cytoplasms: {len(cytoplasms)}, cytoplasms grouped: {len(merged_df)}')
else:
cytoplasms_g_df, _ = _split_object_data(
cytoplasms, 'prcfo', 'object_label')
merged_df = _merge_grouped(merged_df, cytoplasms_g_df, 'cytoplasm')
if verbose:
print(f'cytoplasms: {len(cytoplasms)}, cytoplasms grouped: {len(cytoplasms_g_df)}')
if 'nucleus' in data_dict:
nucleus = data_dict['nucleus'].copy()
nucleus = nucleus.dropna(subset=['cell_id'])
nucleus = nucleus.assign(object_label=lambda x: 'o' + x['object_label'].astype(int).astype(str))
nucleus = nucleus.assign(cell_id=lambda x: 'o' + x['cell_id'].astype(int).astype(str))
nucleus = nucleus.assign(prcfo=lambda x: x['prcf'] + '_' + x['cell_id'])
nucleus['nucleus_prcfo_count'] = nucleus.groupby('prcfo')['prcfo'].transform('count')
if nuclei_limit is not None:
if nuclei_limit is True:
nucleus = nucleus[nucleus['nucleus_prcfo_count'] == 1]
elif isinstance(nuclei_limit, (float, int)):
nucleus = nucleus[nucleus['nucleus_prcfo_count'] <= int(nuclei_limit)]
if all(key not in data_dict for key in ['cell', 'cytoplasm']):
merged_df, metadata = _split_object_data(
nucleus, 'prcfo', 'cell_id')
metadata_key = 'cell_id'
if verbose:
print(f'nucleus: {len(nucleus)}, nucleus grouped: {len(merged_df)}')
else:
nucleus_g_df, _ = _split_object_data(
nucleus, 'prcfo', 'cell_id')
merged_df = _merge_grouped(merged_df, nucleus_g_df, 'nucleus')
if verbose:
print(f'nucleus: {len(nucleus)}, nucleus grouped: {len(nucleus_g_df)}')
if 'pathogen' in data_dict:
pathogens = data_dict['pathogen'].copy()
pathogens = pathogens.dropna(subset=['cell_id'])
pathogens = pathogens.assign(object_label=lambda x: 'o' + x['object_label'].astype(int).astype(str))
pathogens = pathogens.assign(cell_id=lambda x: 'o' + x['cell_id'].astype(int).astype(str))
pathogens = pathogens.assign(prcfo=lambda x: x['prcf'] + '_' + x['cell_id'])
pathogens['pathogen_prcfo_count'] = pathogens.groupby('prcfo')['prcfo'].transform('count')
if pathogen_limit is not None:
if pathogen_limit is True:
pathogens = pathogens[pathogens['pathogen_prcfo_count'] <= 1]
elif isinstance(pathogen_limit, (float, int)):
pathogens = pathogens[pathogens['pathogen_prcfo_count'] <= int(pathogen_limit)]
if all(key not in data_dict for key in ['cell', 'cytoplasm', 'nucleus']):
merged_df, metadata = _split_object_data(
pathogens, 'prcfo', 'cell_id')
metadata_key = 'cell_id'
if verbose:
print(f'pathogens: {len(pathogens)}, pathogens grouped: {len(merged_df)}')
else:
pathogens_g_df, _ = _split_object_data(
pathogens, 'prcfo', 'cell_id')
merged_df = _merge_grouped(merged_df, pathogens_g_df, 'pathogen')
if verbose:
print(f'pathogens: {len(pathogens)}, pathogens grouped: {len(pathogens_g_df)}')
pathogen_counts = pathogens.groupby('prcfo')['prcfo'].size().rename('pathogen_prcfo_count')
for organelle_role in ORGANELLE_ROLES:
if organelle_role not in data_dict:
continue
organelles = data_dict[organelle_role].copy()
organelles = organelles.dropna(subset=['cell_id'])
organelles = organelles.assign(object_label=lambda x: 'o' + x['object_label'].astype(int).astype(str))
organelles = organelles.assign(cell_id=lambda x: 'o' + x['cell_id'].astype(int).astype(str))
organelles = organelles.assign(prcfo=lambda x: x['prcf'] + '_' + x['cell_id'])
count_column = f'{organelle_role}_prcfo_count'
organelles[count_column] = organelles.groupby('prcfo')['prcfo'].transform('count')
earlier = ['cell', 'cytoplasm', 'nucleus', 'pathogen'] + list(
ORGANELLE_ROLES[:ORGANELLE_ROLES.index(organelle_role)])
if all(key not in data_dict for key in earlier):
merged_df, metadata = _split_object_data(
organelles, 'prcfo', 'cell_id')
metadata_key = 'cell_id'
if verbose:
print(f'{organelle_role}: {len(organelles)}, '
f'{organelle_role} grouped: {len(merged_df)}')
else:
organelles_g_df, _ = _split_object_data(
organelles, 'prcfo', 'cell_id')
merged_df = _merge_grouped(
merged_df, organelles_g_df, organelle_role)
if verbose:
print(f'{organelle_role}: {len(organelles)}, '
f'{organelle_role} grouped: {len(organelles_g_df)}')
if 'png_list' in data_dict:
from .utils import PNG_CROP_MODE_BY_ID_COLUMN, PNG_OBJECT_ID_COLUMNS, object_label_from_png_id
png_list = data_dict['png_list'].copy()
id_column = PNG_OBJECT_ID_COLUMNS['cell']
if id_column not in png_list.columns:
present = [c for c in png_list.columns if c in PNG_CROP_MODE_BY_ID_COLUMN]
modes = sorted(PNG_CROP_MODE_BY_ID_COLUMN[c] for c in present)
raise CropModeMismatch(
f"png_list has no {id_column!r} column, so its crops cannot be keyed onto the objects being merged. It holds "
+ (f"{', '.join(modes)} crops ({', '.join(present)})" if present else "no object-id column at all")
+ ". Re-run the Measure module with 'cell' in crop_mode."
)
keep = object_label_from_png_id(png_list[id_column]).notna()
if not keep.all():
print(
f"png_list: {int((~keep).sum())} of {len(png_list)} rows are not usable cell crops "
f"(another crop mode, or an object id that is not a number); they take no part in the merge."
)
png_list = png_list.loc[keep].copy()
png_list_g_df_numeric, png_list_g_df_non_numeric = _split_data(png_list, 'prcfo', id_column)
png_list_g_df_non_numeric.drop(
columns=['plateID', 'rowID', 'columnID', 'fieldID', 'file_name', 'cell_id', 'prcf'],
inplace=True,
errors='ignore',
)
if verbose:
print(f'png_list: {len(png_list)}, png_list grouped: {len(png_list_g_df_numeric)}')
print(f"Added png_list columns: {png_list_g_df_numeric.columns}, {png_list_g_df_non_numeric.columns}")
merged_df = _merge_grouped(merged_df, png_list_g_df_numeric, 'png_list')
merged_df = _merge_grouped(merged_df, png_list_g_df_non_numeric, 'png_list')
metadata = metadata.assign(prc=lambda x: x['plateID'] + '_' + x['rowID'] + '_' + x['columnID'])
metadata = metadata.assign(
prcfo=lambda x: x['prcf'] + '_' + x[metadata_key])
cells_well = metadata.groupby('prc')['prcfo'].nunique().reset_index(name='cells_per_well')
metadata = _merge_with_cardinality(
metadata,
cells_well,
on='prc',
validate='many_to_one',
left_name='object metadata',
right_name='well counts',
)
metadata.set_index('prcfo', inplace=True)
merged_df = _merge_grouped(metadata, merged_df)
merged_df.drop(columns=['label_list_morphology', 'label_list_intensity'], errors='ignore', inplace=True)
if pathogen_counts is not None:
merged_df['pathogen_prcfo_count'] = merged_df.index.to_series().map(pathogen_counts).fillna(0).astype('Int64')
if verbose:
print(f'Generated dataframe with: {len(merged_df.columns)} columns and {len(merged_df)} rows')
object_table_order = [
'cell', 'cytoplasm', 'nucleus', 'pathogen', *ORGANELLE_ROLES]
obj_df_ls = [data_dict[table] for table in object_table_order
if table in data_dict]
return merged_df, obj_df_ls
def _read_mask(mask_path):
"""Read a mask file as 16-bit labels.
:param mask_path: the mask file.
:returns: the labels as ``uint16`` -- converted rather than cast, so an
8-bit mask keeps its object ids instead of having them rescaled.
"""
mask = imageio2.imread(mask_path)
if mask.dtype != np.uint16:
mask = img_as_uint(mask)
return mask
[docs]
def convert_numpy_to_tiff(folder_path, limit=None):
"""Convert every ``.npy`` array in ``folder_path`` to a TIFF under ``folder_path/tiff``.
:param folder_path: Folder containing the ``.npy`` files.
:param limit: If set, stop after processing this many files.
:returns: None
"""
tiff_subdir = os.path.join(folder_path, 'tiff')
os.makedirs(tiff_subdir, exist_ok=True)
files = os.listdir(folder_path)
for i, filename in enumerate(files):
if limit is not None and i >= limit:
break
if not filename.endswith('.npy'):
continue
file_path = os.path.join(folder_path, filename)
numpy_array = np.load(file_path)
tiff_filename = os.path.splitext(filename)[0] + '.tif'
tiff_file_path = os.path.join(tiff_subdir, tiff_filename)
write_tiff(tiff_file_path, numpy_array)
print(f"Converted {filename} to {tiff_filename} and saved in 'tiff' subdirectory.")
return
[docs]
def generate_cellpose_train_test(src, test_split=0.1):
"""Split image/mask pairs in ``src`` into ``train`` and ``test`` sibling folders.
Only images that have a corresponding mask in ``src/masks`` are
considered.
:param src: Folder containing images and a ``masks`` subfolder.
:param test_split: Fraction of pairs to route into the test set.
Default ``0.1``.
:returns: None
"""
mask_src = os.path.join(src, 'masks')
img_paths = glob.glob(os.path.join(src, '*.tif'))
img_filenames = [os.path.basename(file) for file in img_paths]
img_filenames = [file for file in img_filenames if os.path.exists(os.path.join(mask_src, file))]
print(f'Found {len(img_filenames)} images with masks')
random.shuffle(img_filenames)
split_index = int(len(img_filenames) * test_split)
train_files = img_filenames[split_index:]
test_files = img_filenames[:split_index]
list_of_lists = [test_files, train_files]
print(f'Split dataset into Train {len(train_files)} and Test {len(test_files)} files')
train_dir = os.path.join(os.path.dirname(src), 'train')
train_dir_masks = os.path.join(train_dir, 'masks')
test_dir = os.path.join(os.path.dirname(src), 'test')
test_dir_masks = os.path.join(test_dir, 'masks')
os.makedirs(train_dir, exist_ok=True)
os.makedirs(train_dir_masks, exist_ok=True)
os.makedirs(test_dir, exist_ok=True)
os.makedirs(test_dir_masks, exist_ok=True)
for i, ls in enumerate(list_of_lists):
if i == 0:
dst = test_dir
dst_mask = test_dir_masks
_type = 'Test'
else:
dst = train_dir
dst_mask = train_dir_masks
_type = 'Train'
for idx, filename in enumerate(ls):
img_path = os.path.join(src, filename)
mask_path = os.path.join(mask_src, filename)
new_img_path = os.path.join(dst, filename)
new_mask_path = os.path.join(dst_mask, filename)
shutil.copy(img_path, new_img_path)
shutil.copy(mask_path, new_mask_path)
print(f'Copied {idx+1}/{len(ls)} images to {_type} set')
#: How a mate is spelled in a FASTQ filename, mapped to the key spaCR uses.
#:
#: ``R1``/``R2`` is the Illumina convention. ``1``/``2`` is what ENA and the
#: SRA publish -- every file downloaded from those archives is
#: ``<run>_1.fastq.gz`` and ``<run>_2.fastq.gz`` -- and not recognising it surfaced as ``KeyError: 'R1'`` after a successful download of
#: the project's own reads.
_MATE_SPELLINGS = {
"r1": "R1", "1": "R1", "read1": "R1", "fwd": "R1",
"r2": "R2", "2": "R2", "read2": "R2", "rev": "R2",
}
[docs]
def parse_gz_files(folder_path):
"""Group ``.fastq.gz`` files in ``folder_path`` by sample name and read direction.
Accepts both naming conventions in the wild: ``<sample>_R1_...`` from an
Illumina run, and ``<run>_1.fastq.gz`` from ENA or the SRA. See
:data:`_MATE_SPELLINGS`.
A file whose mate cannot be identified contributes NOTHING rather than an
empty entry. The previous version created ``{sample: {}}`` for it, which
turned an unrecognised filename into a ``KeyError: 'R1'`` several frames
later in :func:`spacr.sequencing.generate_barecode_mapping` -- a crash that
named neither the file nor the problem.
:param folder_path: Directory containing gzipped FASTQ files.
:returns: Mapping ``{sample_name: {"R1": path, "R2": path}}``. Samples may
have only one of the two.
"""
files = os.listdir(folder_path)
gz_files = [f for f in files if f.endswith('.fastq.gz')]
samples_dict = {}
for gz_file in gz_files:
stem = gz_file[:-len('.fastq.gz')]
parts = stem.split('_')
if len(parts) < 2:
samples_dict.setdefault(stem, {})['R1'] = os.path.join(
folder_path, gz_file)
continue
sample_name = '_'.join(parts[:-1])
mate = _MATE_SPELLINGS.get(parts[-1].strip().lower())
if mate is None:
for position, token in enumerate(parts):
candidate = _MATE_SPELLINGS.get(token.strip().lower())
if candidate is not None and position > 0:
mate = candidate
sample_name = '_'.join(parts[:position])
break
if mate is None:
LOG.warning(
"%s: cannot tell which mate this is, so it is skipped. "
"Expected a name like <sample>_R1.fastq.gz or "
"<run>_1.fastq.gz.", gz_file)
continue
samples_dict.setdefault(sample_name, {})[mate] = os.path.join(
folder_path, gz_file)
return samples_dict
#: Object types that can be cut on demand.
#: Crop order, which differs from the hook order on purpose. Membership is
#: checked against spacr.object_roles; the order stays here.
CROP_OBJECT_TYPES = (
'cell', 'nucleus', 'pathogen', 'cytoplasm', *ORGANELLE_ROLES)
#: ``png_list`` column holding the object id (``'o<N>'``) for each crop mode,
#: as written by :func:`spacr.utils.filepaths_to_database`.
#: Column name used to carry a per-row crop handle through the frames in this
#: module without colliding with a measurement column.
CROP_REF_COLUMN = '_spacr_crop_ref'
[docs]
def crop_object_type(png_type, default='cell'):
"""Return the object type named by a ``png_type`` / ``file_metadata`` string.
``'cell_png'`` -> ``'cell'``, ``'…/nucleus_png/…'`` -> ``'nucleus'``. A
string that names no object type (a plate prefix, say) falls back to
``default``, because that filter is about *which rows*, not *which mask*.
:param png_type: the setting value, or any path/substring containing it.
:param default: object type to assume when nothing is named.
:returns: one of :data:`CROP_OBJECT_TYPES`.
"""
text = str(png_type or '').lower()
for obj in CROP_OBJECT_TYPES:
if f'{obj}_png' in text:
return obj
return default
#: ``measure_crop`` settings that shape a crop, and that a later run may
#: legitimately override when cutting one on demand.
CROP_SHAPE_KEYS = ('png_dims', 'png_size', 'normalize', 'normalize_by',
'crop_mode', 'use_bounding_box', 'dialate_pngs',
'dialate_png_ratios', 'cell_mask_dim', 'nucleus_mask_dim',
'pathogen_mask_dim', 'organelle_mask_dim',
*(f'{role}_mask_dim' for role in ORGANELLE_ROLES[1:]))
def _crop_shape_overrides(settings):
"""Return the crop-shaping settings that may override the saved snapshot.
Only the keys that mean the same thing to a crop as they do to
``measure_crop``. ``normalize`` is the trap and the reason this function
exists: ``measure_crop`` writes a ``[p1, p2]`` percentile pair, and
``train_test_model`` / ``deep_spacr`` write a **bool** meaning "normalise
the tensor". Forwarding the bool would replace the ``[1, 99]`` stretch the
PNG folder was written with by a full 0-100 one and change every pixel,
silently, on the one path whose entire purpose is to be pixel-identical
to that folder.
"""
out = {}
for key in CROP_SHAPE_KEYS:
if key not in settings:
continue
value = settings[key]
if key == 'normalize' and not (
value is False
or (isinstance(value, (list, tuple)) and len(value) == 2)):
continue
out[key] = value
return out
#: The two vocabularies for the same idea, and the map between them.
#:
#: `spacr.settings` writes the USER-FACING words -- 'pre_generated' means the
#: PNGs already cut to disk, 'on_demand' means cut them now from the merged
#: stacks. `crops.resolve_crop_source` speaks in terms of the SOURCE it will
#: build: 'png' or 'merged'. Nothing translated between them, so
#: `crop_source='on_demand'` -- which the Classify screen validates and
#: accepts -- reached `resolve_crop_source`, raised CropError, was swallowed,
#: and the run trained on pre-cut PNGs instead. The user's explicit choice
#: was ignored in silence.
#:
#: `load_images` and `stream_images` are the CURRENT spelling -- the one the
#: Classify panel writes -- and they were missing here, so the defect the
#: paragraph above describes came straight back under new names. Worse than
#: before: `settings._canonical_image_source` rewrites EVERY choice into this
#: pair, so `crop_source='merged'` became `'stream_images'` and no value the
#: panel could produce survived the lookup. Streaming was unreachable from
#: the GUI; every run trained on pre-cut PNGs and said nothing.
#: ONE TABLE, TWO READERS. The spellings themselves live in
#: `spacr.crop_source`, which is where the training path resolves them, so a
#: settings file cannot be accepted by one reader and refused by the other --
#: which is exactly the failure both paragraphs above describe, twice. Two
#: entries differ, and each because the question differs:
#:
#: * 'generate' names an ACTION to training (write a crop set, then load it)
#: and a SOURCE here (once written, the images are PNGs);
#: * 'auto' means LOAD IMAGES to training, which has to pick a mode, and
#: stays 'auto' here, because `crops.resolve_crop_source` answers "what is
#: available in this project" and is the thing that computes it.
CROP_SOURCE_ALIASES = {
**_crop_source.CROP_SOURCE_ALIASES,
'generate': 'png',
'auto': 'auto',
}
def _canonical_crop_source(choice):
"""The word `crops.resolve_crop_source` understands.
An unrecognised value is passed through UNCHANGED rather than coerced to
'auto', so `resolve_crop_source` raises on it and names it. Quietly
substituting a default is how the original defect behaved.
"""
if not choice:
return 'auto'
return CROP_SOURCE_ALIASES.get(str(choice).strip().lower(), str(choice))
[docs]
def open_crop_source(settings, src=None, object_type=None, verbose=True):
"""Return the :class:`spacr.crops.CropSource` a run should read crops from.
Thin, non-raising wrapper over :func:`spacr.crops.resolve_crop_source`:
it reads ``settings['crop_source']`` (``'auto'`` | ``'png'`` |
``'merged'``), prints which source was chosen and why, and returns None
when neither is available -- so a caller can fall back to whatever it did
before instead of failing on a project that predates ``merged/``.
The run's own crop-shaping settings are forwarded (see
:func:`_crop_shape_overrides`), which is what makes "cut fresh at the
current crop settings" true rather than a slogan:
``resolve_crop_source`` starts from the ``measure_crop`` snapshot in
``measurements.db`` and lets those override it, so a run that asks for
96 px crops gets 96 px crops out of ``merged/`` even though the folder on
disk holds 48 px ones.
:param settings: settings dict (or a source path) holding ``crop_source``.
:param src: the experiment root; defaults to ``settings['src']`` (its
first entry when that is a list).
:param object_type: default object type for a merged source.
:param verbose: print the chosen source.
:returns: a :class:`spacr.crops.CropSource`, or None.
"""
from . import crops
if isinstance(settings, dict):
request = _crop_shape_overrides(settings)
choice = _canonical_crop_source(settings.get('crop_source'))
if src is None:
src = settings.get('src')
else:
request = {}
choice = 'auto'
if src is None:
src = settings
if isinstance(src, (list, tuple)):
src = src[0] if len(src) else None
if not src:
return None
request['src'] = src
request['crop_source'] = choice
try:
source = crops.resolve_crop_source(request, object_type=object_type)
except crops.CropError as exc:
print(f"crop_source={choice!r}: {exc}")
return None
if verbose:
print(f"Crop source: {source.describe()}")
return source
[docs]
class LazyCropPNG:
"""A PNG-shaped byte stream that is only produced when something opens it.
``spacr.utils.plot_umap_images`` and ``spacr.utils.plot_clusters_grid``
reach for every thumbnail with ``PIL.Image.open(image_paths[i])``.
``Image.open`` accepts a path *or* any seekable binary stream, so an
instance of this class can sit in that list exactly where a path string
used to and the plotting code needs no change -- which is the point: the
PNG list and the on-demand list have to be interchangeable, or the two
sources are not really alternatives.
Nothing is read until something opens the object, so building one per row
of a large screen costs a small dict and only the handful of thumbnails
actually drawn ever touch ``merged/``.
The bytes are always a **current-format (RGB) crop PNG**: the array comes
from ``CropSource.get``, which is ``crops.png_view`` for the merged source
and ``crops.read_crop_png`` for the PNG one, and both of those return the
corrected order. A legacy folder is therefore corrected on the way through
here, the same way the Annotate screen corrects it.
:param source: the :class:`spacr.crops.CropSource` to cut/read with.
:param row: the row mapping identifying the object.
:param name: the crop's file name, for messages and tar members.
"""
__slots__ = ('source', 'row', 'name', '_buf')
def __init__(self, source, row, name=''):
"""Record the source and row; produce nothing yet."""
self.source = source
self.row = row
self.name = name
self._buf = None
[docs]
def array(self):
"""Return the crop as an ``(H, W, 3)`` uint8 RGB array."""
return self.source.get(self.row)
[docs]
def png_bytes(self):
"""Return the crop encoded as a current-format (RGB) PNG."""
return self._stream().getvalue()
def _raw_bytes(self):
"""Return the bytes already on disk, for a PNG source. None otherwise."""
resolve = getattr(self.source, 'resolve', None)
if resolve is None:
return None
try:
with open(resolve(self.row), 'rb') as handle:
return handle.read()
except Exception:
return None
def _stream(self):
"""Materialise (once) and return the BytesIO holding the PNG."""
if self._buf is None:
buf = BytesIO()
try:
Image.fromarray(self.array()).save(buf, format='PNG')
except Exception:
raw = self._raw_bytes()
if raw is None:
raise
buf = BytesIO(raw)
buf.seek(0)
self._buf = buf
return self._buf
[docs]
def read(self, size=-1):
"""Read up to ``size`` bytes of the PNG.
The first read is what actually cuts and encodes the crop; every later
one is served from the same buffer and continues where the previous
read stopped, so re-reading the image needs an explicit ``seek(0)``.
:param size: Byte count. The default ``-1`` (or any negative value)
reads to the end of the PNG, which is what PIL does when it is
handed this object as a file.
:returns: The bytes read, empty once the buffer is exhausted.
"""
return self._stream().read(size)
[docs]
def seek(self, offset, whence=0):
"""Seek within the PNG.
Seeking materialises the crop if it has not been produced yet, so the
cost of the first ``seek`` is the cost of encoding the whole image,
not of moving a cursor.
:param offset: How far to move, in bytes, and it must be an integer —
a float raises ``TypeError``. Seeking past the end of the encoded
PNG is allowed rather than an error: the position lands beyond the
buffer and the next ``read`` returns empty.
:param whence: ``0`` from the start (the default), ``1`` from the
current position, ``2`` from the end. Semantics are the
underlying ``BytesIO`` ones, so a negative ``offset`` is a
``ValueError`` under ``0`` but legal under ``1`` and ``2``.
:returns: The new absolute position.
"""
return self._stream().seek(offset, whence)
[docs]
def tell(self):
"""Return the current offset."""
return self._stream().tell()
[docs]
def readable(self):
"""Return True -- the stream is readable."""
return True
[docs]
def seekable(self):
"""Return True -- the stream is seekable."""
return True
[docs]
def writable(self):
"""Return False -- the stream is read-only."""
return False
[docs]
def close(self):
"""Drop the materialised bytes; a later read produces them again."""
self._buf = None
@property
[docs]
def closed(self):
"""Return False -- this object is never permanently closed."""
return False
[docs]
def __enter__(self):
"""Return self, so the object can be used as a context manager."""
return self
[docs]
def __exit__(self, *exc):
"""Release the materialised bytes."""
self.close()
return False
[docs]
def __repr__(self):
"""Return a short description naming the crop."""
kind = getattr(self.source, 'kind', '?')
return f"<LazyCropPNG {self.name or '?'} from {kind}>"
[docs]
def crop_png_name(file_name, object_type, object_label, cell_id=None):
"""Return the file name :func:`spacr.utils._generate_names` gives this crop.
The name matters downstream: :func:`_png_group_id` parses the plate / well
/ field out of it for group-aware cross-validation, and
:func:`spacr.utils.process_vision_results` parses the ``prcfo`` that
:func:`spacr.deep_spacr.merge_predictions_into_db` merges on. A crop cut
on demand has to carry the same name as the one the PNG folder would have
held or those two stop lining up.
:param file_name: the merged array's stem (``plate1_A01_1``).
:param object_type: which crop mode this is.
:param object_label: the object's integer label.
:param cell_id: the parent cell label, for nucleus/pathogen crops.
:returns: the crop's file name, ending in ``.png``.
"""
stem = os.path.splitext(os.path.basename(str(file_name)))[0]
label = int(object_label)
if object_type in ('nucleus', 'pathogen'):
parent = _object_id_int(cell_id)
parent_str = 'none' if not parent else str(parent)
return f"{stem}_{parent_str}_{label}.png"
return f"{stem}_{label}.png"
[docs]
def crop_rows_from_object_table(db_path, object_type='cell', verbose=True):
"""Return one crop row per object, straight off the measurement table.
This is the path for a project that never wrote a PNG folder at all, so
there is no ``png_list`` to start from: ``object_label``, ``path_name``
and the well keys are already on every measurement row.
:param db_path: the ``measurements.db``.
:param object_type: which object table to read.
:param verbose: print what was found.
:returns: a DataFrame with ``path_name``, ``object_label``, the well keys,
``png_name`` and ``png_path`` (the path the crop *would* have had).
"""
if not os.path.isfile(db_path):
return pd.DataFrame()
select = ('object_label, plateID, rowID, columnID, fieldID, prcf, '
'file_name, path_name')
from .database_concurrency import connect as _connect_database
conn = _connect_database(db_path)
try:
try:
df = pd.read_sql(f'SELECT {select} FROM "{object_type}"', conn)
except Exception:
try:
df = pd.read_sql(
f'SELECT object_label, plateID, rowID, columnID, fieldID, '
f'file_name, path_name FROM "{object_type}"', conn)
except Exception:
if verbose:
print(f"crop_rows_from_object_table: no '{object_type}' "
f"table in {db_path}")
return pd.DataFrame()
parents = {}
if object_type in ('nucleus', 'pathogen'):
try:
link = pd.read_sql(
f'SELECT object_label, prcf, cell_id FROM "{object_type}"',
conn)
parents = {(r.prcf, r.object_label): r.cell_id
for r in link.itertuples()}
except Exception:
parents = {}
finally:
conn.close()
if df.empty:
return df
df['object_type'] = object_type
df['png_name'] = [
crop_png_name(row.file_name, object_type, row.object_label,
parents.get((getattr(row, 'prcf', None), row.object_label)))
for row in df.itertuples()
]
df['png_path'] = [
os.path.join(str(plate) + '_' + str(well), f'{object_type}_png', name)
for plate, well, name in zip(df['plateID'], df['rowID'], df['png_name'])
]
if verbose:
print(f"crop_rows_from_object_table: {len(df)} '{object_type}' objects "
f"in {db_path}")
return df
[docs]
def crop_refs_for_rows(source, df, object_type='cell', name_column=None):
"""Return one :class:`LazyCropPNG` per row of ``df``.
:param source: the crop source to cut/read with.
:param df: rows carrying whatever the source needs (``png_path`` for the
PNG source, ``path_name`` + ``object_label`` for the merged one).
:param object_type: object type stamped onto each row.
:param name_column: column holding the crop's file name; defaults to
the basename of ``png_path``.
:returns: list of :class:`LazyCropPNG`.
"""
n = len(df)
def _col(name):
"""One column as a plain list, or a column of None when it is absent.
PLAIN LISTS RATHER THAN `itertuples`, for two measured reasons: the
joined UMAP frame carries a couple of hundred columns, so building a
namedtuple per row of it costs more than reading the crops does, and
itertuples silently RENAMES any column whose name is not a valid
identifier -- which is how a lookup starts missing a column that is
plainly there.
:param name: the column, or a falsy value for "this frame has none".
"""
if name and name in df.columns:
return df[name].tolist()
return [None] * n
def _missing(value):
"""Whether a cell carries no answer.
NaN AS WELL AS None, because a column read out of pandas holds NaN
where a row had nothing and `None is not float('nan')`. Testing only
for None lets a NaN through as if it were a value, and it then
reaches a path name or an object label.
"""
return value is None or (isinstance(value, float) and np.isnan(value))
png_paths = _col('png_path')
path_names = _col('path_name')
labels = _col('object_label')
names = _col(name_column)
refs = []
for i in range(n):
entry = {'object_type': object_type}
if not _missing(png_paths[i]):
entry['png_path'] = png_paths[i]
if not _missing(path_names[i]):
entry['path_name'] = path_names[i]
if not _missing(labels[i]):
entry['object_label'] = int(labels[i])
if not _missing(names[i]):
name = str(names[i])
elif not _missing(png_paths[i]):
name = os.path.basename(str(png_paths[i]))
else:
name = ''
refs.append(LazyCropPNG(source, entry, name=name))
return refs
[docs]
def mark_crop_output_folder(folder, fmt=None, source_folder=None,
db_path=None, **extra):
"""Stamp a folder spaCR has just filled with crop PNGs.
Called *before* the folder is filled, exactly as
:func:`spacr.crops.stamp_crop_folder` is on the measure path, so an
interrupted run leaves a marked folder holding fewer crops rather than an
unmarked folder of corrected ones -- the one state that is silently
misread.
``fmt=None`` inherits the format from ``source_folder``. Byte-for-byte
copies retain their source format: formats 1 and 3 use declared order,
while format 2 needs channel reversal when read by a declared-order model.
:param folder: the folder about to be filled.
:param fmt: the format to record; None inherits from ``source_folder``.
:param source_folder: the folder the crops are being copied from.
:param db_path: ``measurements.db`` consulted when ``source_folder``
carries no sidecar.
:param extra: extra keys recorded in the sidecar.
:returns: the sidecar path, or None when it could not be written.
"""
from . import crops
if fmt is None:
if source_folder:
try:
fmt = crops.crop_folder_format(source_folder, db_path=db_path)
except Exception:
fmt = crops.CROP_FORMAT_LEGACY_BGR
else:
fmt = crops.CROP_FORMAT_CURRENT
try:
return crops.write_crop_folder_marker(folder, fmt=int(fmt), **extra)
except Exception as exc:
print(f"Warning: could not stamp the crop format on {folder}: {exc}")
return None
[docs]
def generate_dataset(settings=None):
"""Pack per-object PNGs referenced by one or more ``measurements.db`` files into a single tar for inference or upload.
Selects PNG paths (via the ``png_list`` table plus optional
``file_metadata`` filter) from each source's measurements database,
optionally random-subsamples, then bundles the images in parallel
into a dated tar under the first source's ``datasets/`` folder.
Use this to produce the ``tar_path`` consumed by
:func:`spacr.deep_spacr.deep_spacr` / ``apply_model_to_tar``.
``crop_source`` chooses where the images come from. ``'png'`` (and
``'auto'`` wherever a crop folder exists) is the behaviour above,
unchanged for uniform source formats: files are byte-copied into the tar
with their format marker. Mixed formats are decoded into declared uint8
copies in the archive only. ``'merged'`` (and
``'auto'`` on a project with no crop folder) cuts every crop out of
``merged/*.npy`` through :mod:`spacr.crops` instead, so the tar can be
built with no PNG folder on disk at all, and is rebuilt at the *current*
crop settings rather than whatever the folder was generated with. The
members are named exactly as the PNG folder would have named them, so
everything that parses a crop file name downstream -- fold grouping,
``prcfo``, the prediction merge -- keeps working either way.
:param settings: Settings dict, canonicalized via
:func:`spacr.settings.set_generate_dataset_defaults`. Key
entries:
- ``src`` (str or list of str) — folder(s) containing
``measurements/measurements.db`` and the PNG crops.
- ``file_metadata`` — filter/join key applied against
``png_list``.
- ``sample`` — ``int`` or ``[int]`` cap on selected PNGs
(random subsample); omit for all.
- ``experiment`` — string suffix used in the tar filename.
- ``crop_source`` — ``'auto'`` | ``'png'`` | ``'merged'``.
:returns: Absolute path to the created ``…/datasets/<date>_<
experiment>.tar``.
:raises RuntimeError: if ``src`` is not a string / list of strings,
no images are selected, no image could be written, or the
destination folder cannot be resolved.
The tar-writing pool is closed and joined, not left to the ``with``
block's ``terminate()``: that sends SIGTERM to idle workers, and a
worker whose SIGTERM handler needs a lock the interrupted code holds
(coverage's ``sigterm = True`` data save does) never exits, so the
shutdown waited on it for ever. Workers that finished their tasks are
let go by a normal exit instead.
Example:
.. code-block:: python
from spacr.io import generate_dataset
tar_path = generate_dataset({
'src': ['/data/plate01', '/data/plate02'],
'experiment': 'screen_v1',
'sample': 100000,
})
See Also:
:func:`training_dataset_from_annotation` — build a labeled
``train/`` / ``test/`` tree instead of a flat tar.
:func:`spacr.deep_spacr.deep_spacr` — consumes the tar via
``apply_model_to_tar``.
"""
if settings is None:
settings = {}
import os, tarfile, shutil, random, datetime
from multiprocessing import Value, Lock, cpu_count
from .resource_log import _parallel_pool as Pool
from .utils import (
initiate_counter, add_images_to_tar, save_settings,
generate_path_list_from_db, correct_paths
)
from .settings import set_generate_dataset_defaults
settings = set_generate_dataset_defaults(settings)
save_settings(settings, 'generate_dataset', show=True)
if isinstance(settings['src'], str):
settings['src'] = [settings['src']]
object_type = crop_object_type(
settings.get('file_metadata') or settings.get('path_string')
or settings.get('png_type'))
if isinstance(settings['src'], list):
all_paths = []
n_on_demand = 0
dst = None
for i, src in enumerate(settings['src']):
db_path = os.path.join(src, 'measurements', 'measurements.db')
if i == 0:
dst = os.path.join(src, 'datasets')
source = open_crop_source(settings, src, object_type=object_type)
if source is not None and getattr(source, 'kind', 'png') == 'merged':
refs = _dataset_crop_refs(db_path, source, settings, object_type)
n_on_demand += len(refs)
all_paths.extend(refs)
continue
paths = generate_path_list_from_db(db_path, file_metadata=settings['file_metadata'])
if not paths:
print(f"No png_list rows selected from {db_path}.")
continue
paths = correct_paths(paths, src)
all_paths.extend(paths)
if isinstance(settings['sample'], int) and settings['sample']:
k = min(int(settings['sample']), len(all_paths))
selected_paths = random.sample(all_paths, k) if k else []
print(f"Random selection of {len(selected_paths)} paths")
elif isinstance(settings['sample'], list) and settings['sample']:
k = min(int(settings['sample'][0]), len(all_paths))
selected_paths = random.sample(all_paths, k) if k else []
print(f"Random selection of {len(selected_paths)} paths")
else:
selected_paths = list(all_paths)
random.shuffle(selected_paths)
print(f"All paths: {len(selected_paths)} paths")
else:
raise RuntimeError("settings['src'] must be a string or list of strings.")
total_images = len(selected_paths)
print(f"Found {total_images} images")
if total_images == 0:
raise RuntimeError("No images selected; nothing to tar.")
os.makedirs(dst, exist_ok=True)
date_name = datetime.date.today().strftime('%y%m%d')
if len(settings['src']) > 1:
date_name = f"{date_name}_combined"
tar_name = f"{date_name}_{settings['experiment']}.tar"
tar_name = os.path.join(dst, tar_name)
if os.path.exists(tar_name):
number = random.randint(1, 100)
tar_name_2 = f"{date_name}_{settings['experiment']}_{settings['file_metadata']}_{number}.tar"
print(f"Warning: {os.path.basename(tar_name)} exists, saving as {os.path.basename(tar_name_2)} ")
tar_name = os.path.join(dst, tar_name_2)
source_format = _crop_format_of_items(selected_paths)
if n_on_demand or source_format is None:
written, skipped = _write_crop_tar(selected_paths, tar_name, settings)
if written == 0:
raise RuntimeError(
f"No image could be written to {tar_name}: all "
f"{total_images} selected crops failed. Check that "
f"merged/*.npy is where measurements.db says it is.")
if skipped:
print(f"Warning: {skipped} of {total_images} crops could not be "
f"produced and are NOT in the tar.")
print(f"\nSaved {written} images to {tar_name}")
return tar_name
temp_dir = os.path.join(dst, "temp_tars")
os.makedirs(temp_dir, exist_ok=True)
num_procs = max(1, min(max(2, cpu_count() - 2), total_images))
from .resource_log import _array_file_nbytes, _guard_workers
num_procs = _guard_workers('dataset', num_procs, _array_file_nbytes(
selected_paths[0]) if selected_paths else 0, settings=settings)
chunk_size = total_images // num_procs
remainder = total_images % num_procs
paths_chunks = []
start = 0
for i in range(num_procs):
end = start + chunk_size + (1 if i < remainder else 0)
paths_chunks.append(selected_paths[start:end])
start = end
temp_tar_files = [os.path.join(temp_dir, f"temp_{i}.tar") for i in range(num_procs)]
print(f"Generating temporary tar files in {dst}")
counter = Value('i', 0)
lock = Lock()
with Pool(processes=num_procs, initializer=initiate_counter, initargs=(counter, lock)) as pool:
pool.starmap(
add_images_to_tar,
[(paths_chunks[i], temp_tar_files[i], total_images) for i in range(num_procs)]
)
pool.close()
pool.join()
print(f"Merging temporary files")
written = 0
with tarfile.open(tar_name, 'w') as final_tar:
from .crops import CROP_FORMAT_SIDECAR
marker = json.dumps({'spacr_crop_format': source_format}).encode('utf-8')
info = tarfile.TarInfo(CROP_FORMAT_SIDECAR)
info.size = len(marker)
final_tar.addfile(info, BytesIO(marker))
for temp_tar_path in temp_tar_files:
with tarfile.open(temp_tar_path, 'r') as temp_tar:
for member in temp_tar.getmembers():
if member.isfile():
file_obj = temp_tar.extractfile(member)
final_tar.addfile(member, file_obj)
written += 1
os.remove(temp_tar_path)
shutil.rmtree(temp_dir)
if written == 0:
raise RuntimeError(
f"No image could be written to {tar_name}: none of the "
f"{total_images} selected PNG paths exist on disk. The crop "
f"folder has been deleted or moved -- set crop_source='merged' "
f"to cut the crops out of merged/*.npy instead.")
if written < total_images:
print(f"Warning: {total_images - written} of {total_images} selected "
f"PNGs were missing and are NOT in the tar.")
print(f"\nSaved {written} images to {tar_name}")
return tar_name
def _dataset_crop_refs(db_path, source, settings, object_type, verbose=True):
"""Return the on-demand crops one source folder contributes to a dataset tar.
Prefers ``png_list`` when there is one, so the tar holds exactly the crops
the PNG path would have held, under exactly the same names and the same
``file_metadata`` filter. Falls back to the object measurement table for a
project that never wrote a PNG folder at all.
:param db_path: the source's ``measurements/measurements.db``.
:param source: the merged :class:`spacr.crops.CropSource`.
:param settings: the ``generate_dataset`` settings.
:param object_type: which crop mode to cut.
:param verbose: print what was selected.
:returns: list of :class:`LazyCropPNG`.
"""
file_metadata = settings.get('file_metadata')
png_df = None
if os.path.isfile(db_path):
from .database_concurrency import connect as _connect_database
conn = _connect_database(db_path)
try:
png_df = pd.read_sql('SELECT * FROM png_list', conn)
except Exception:
png_df = None
finally:
conn.close()
def _filter(frame, column):
"""Keep the rows whose ``column`` contains any of the wanted terms.
SUBSTRING AND NOT REGEX (`regex=False`): the terms come from a user
naming plates or wells, and a stray `(` or `+` in one of them would
otherwise raise out of pandas rather than simply matching nothing.
Any term matching is enough -- several terms are alternatives, which
is what a user listing them means.
:param frame: the rows to filter.
:param column: the column to search; an absent one filters nothing,
because a frame that never had it cannot contradict the request.
"""
if not file_metadata or column not in frame.columns:
return frame
terms = file_metadata if isinstance(file_metadata, (list, tuple)) else [file_metadata]
text = frame[column].astype(str)
mask = np.zeros(len(frame), dtype=bool)
for term in terms:
mask |= text.str.contains(str(term), regex=False, na=False).to_numpy()
return frame[mask]
if png_df is not None and len(png_df):
png_df = _filter(png_df, 'png_path')
rows = crop_rows_from_png_list(db_path, png_df, object_type,
verbose=verbose)
return crop_refs_for_rows(source, rows, object_type)
rows = crop_rows_from_object_table(db_path, object_type, verbose=verbose)
if len(rows):
rows = _filter(rows, 'png_path')
return crop_refs_for_rows(source, rows, object_type,
name_column='png_name')
def _write_crop_tar(items, tar_name, settings=None):
"""Write ``items`` into ``tar_name``, cutting on-demand crops as it goes.
``items`` may mix PNG paths and :class:`LazyCropPNG` handles. Uniform
formats retain their original bytes; mixed formats become declared uint8
copies. Source files are never modified.
The archive also carries a ``.spacr_crop_format.json`` member, the same
marker :mod:`spacr.crops` writes into a crop folder, so "which channel
order is this tar in?" is answerable from the tar alone.
:class:`TarImageDataset` skips it and reports it as ``crop_format``.
:param items: paths and/or :class:`LazyCropPNG` handles.
:param tar_name: destination archive.
:param settings: optional settings, recorded in the marker.
:returns: ``(written, skipped)`` counts.
"""
from . import crops
from .utils import print_progress
total = len(items)
written = 0
skipped = 0
used = set()
fmt = _crop_format_of_items(items)
canonicalize = fmt is None
if canonicalize:
fmt = crops.CROP_FORMAT_CURRENT
print("Mixed crop formats: writing declared-order uint8 copies into the tar.")
with tarfile.open(tar_name, 'w') as tar:
marker = json.dumps({
'spacr_crop_format': fmt,
'note': ('Uniform source formats preserve stored pixels. Mixed '
'formats are decoded to declared uint8 copies.'),
'png_dims': list((settings or {}).get('png_dims') or []),
}, indent=2, sort_keys=True).encode('utf-8')
info = tarfile.TarInfo(crops.CROP_FORMAT_SIDECAR)
info.size = len(marker)
tar.addfile(info, BytesIO(marker))
for i, item in enumerate(items):
try:
if isinstance(item, LazyCropPNG):
payload = item.png_bytes()
name = item.name or f"crop_{i}.png"
else:
name = os.path.basename(str(item))
if canonicalize:
buf = BytesIO()
Image.fromarray(crops.read_crop_png(str(item))).save(buf, format='PNG')
payload = buf.getvalue()
else:
with open(str(item), 'rb') as handle:
payload = handle.read()
except Exception as exc:
skipped += 1
if skipped <= 5:
print(f"Could not read crop {item!r}: {exc}")
continue
if name in used:
stem, ext = os.path.splitext(name)
name = f"{stem}__{i}{ext}"
used.add(name)
info = tarfile.TarInfo(name)
info.size = len(payload)
tar.addfile(info, BytesIO(payload))
written += 1
if written % 100 == 0 or written == total:
print_progress(written, total, n_jobs=1, time_ls=None,
batch_size=None,
operation_type="generating .tar dataset")
return written, skipped
#: Accepted values for the ``class_balance`` setting.
CLASS_BALANCE_MODES = ('none', 'weighted_sampler', 'sqrt_weighted_sampler', 'weighted_loss')
#: Accepted values for the ``cv_group_by`` setting.
CV_GROUP_LEVELS = ('cell', 'field', 'well', 'plate')
#: max/min class-count ratio at or above which the data is called skewed.
IMBALANCE_RATIO_WARN = 1.5
#: max/min class-count ratio at or above which the skew is called severe.
IMBALANCE_RATIO_SEVERE = 10.0
def _png_group_id(path, level):
"""Return the plate / well / field group id encoded in a spacr crop filename.
spacr object crops are named ``<plate>_<well>_<field>_..._<object>.png``
(see ``spacr.utils._generate_names``), so the grouping key is a prefix of
the underscore-separated basename.
:param path: image path or bare filename.
:param level: ``'plate'``, ``'well'`` or ``'field'``.
:returns: group id string, or None when the name has too few parts to
carry the requested level.
:raises ValueError: if ``level`` is not a supported grouping level.
"""
if level not in ('plate', 'well', 'field'):
raise ValueError(
f"group level {level!r} is not one of {('plate', 'well', 'field')}")
stem = os.path.splitext(os.path.basename(str(path)))[0]
parts = stem.split('_')
n_needed = {'plate': 1, 'well': 2, 'field': 3}[level]
if len(parts) < n_needed or any(p == '' for p in parts[:n_needed]):
return None
return '_'.join(parts[:n_needed])
[docs]
def dataset_labels(dataset):
"""Return the integer class label of every sample in ``dataset``.
Handles the three shapes that flow through the training path: a
``spacrDataset`` (labels are already a list), a ``torch.utils.data.Subset``
of one (produced by ``random_split`` and by the fold splitter), and the
plain list of ``(image, label, filename)`` tuples that ``augment_dataset``
returns. Only the last shape has to be walked, and it holds tensors in
memory already, so nothing here decodes an image.
:param dataset: dataset, Subset, or sequence of ``(img, label, name)``.
:returns: list of int labels, positionally aligned with the dataset.
"""
if isinstance(dataset, Subset):
parent = dataset_labels(dataset.dataset)
return [parent[i] for i in dataset.indices]
labels = getattr(dataset, 'labels', None)
if labels is not None:
return [int(v) for v in labels]
return [int(item[1]) for item in dataset]
[docs]
def dataset_filenames(dataset):
"""Return the source filename of every sample in ``dataset``.
Mirrors :func:`dataset_labels` so group ids can be derived without
touching pixels.
:param dataset: dataset, Subset, or sequence of ``(img, label, name)``.
:returns: list of filename strings, positionally aligned with the dataset.
"""
if isinstance(dataset, Subset):
parent = dataset_filenames(dataset.dataset)
return [parent[i] for i in dataset.indices]
names = getattr(dataset, 'filenames', None)
if names is not None:
return [str(v) for v in names]
return [str(item[2]) for item in dataset]
[docs]
def summarize_class_imbalance(labels, classes=None):
"""Measure the class skew of a label vector.
:param labels: iterable of integer class labels.
:param classes: ordered class names; index i names label i. Defaults to
``['class_0', ...]`` sized to the largest label seen.
:returns: dict with ``counts``, ``fractions``, ``imbalance_ratio``
(majority/minority, ``inf`` when a class is empty), ``minority``,
``majority``, ``empty_classes``, ``skewed`` and ``severe``.
"""
labels = [int(v) for v in labels]
if classes is None:
n_classes = (max(labels) + 1) if labels else 0
classes = [f'class_{i}' for i in range(n_classes)]
classes = list(classes)
counts = [0] * len(classes)
unknown = 0
for v in labels:
if 0 <= v < len(counts):
counts[v] += 1
else:
unknown += 1
total = sum(counts)
fractions = [(c / total) if total else 0.0 for c in counts]
hi = max(counts) if counts else 0
lo = min(counts) if counts else 0
if lo > 0:
ratio = hi / lo
elif hi > 0:
ratio = float('inf')
else:
ratio = 1.0
empty = [classes[i] for i, c in enumerate(counts) if c == 0]
return {
'classes': classes,
'counts': counts,
'fractions': fractions,
'n': total,
'unknown_labels': unknown,
'imbalance_ratio': ratio,
'majority': classes[counts.index(hi)] if counts else None,
'minority': classes[counts.index(lo)] if counts else None,
'minority_fraction': (lo / total) if total else 0.0,
'empty_classes': empty,
'skewed': bool(counts) and ratio >= IMBALANCE_RATIO_WARN,
'severe': bool(counts) and ratio >= IMBALANCE_RATIO_SEVERE,
}
[docs]
def class_sampling_weights(counts, mode):
"""Per-class sampling weight for a ``WeightedRandomSampler``.
``'weighted_sampler'`` uses ``1/n_c``, which makes every class equally
likely to be drawn. ``'sqrt_weighted_sampler'`` uses ``1/sqrt(n_c)``, a
partial correction that moves the realised frequencies toward balance
without oversampling a tiny class so hard that the model memorises its
handful of crops.
:param counts: per-class sample counts.
:param mode: ``'weighted_sampler'`` or ``'sqrt_weighted_sampler'``.
:returns: list of per-class weights, scaled so they sum to 1.
:raises ValueError: if ``mode`` does not describe a sampler.
"""
if mode == 'weighted_sampler':
power = 1.0
elif mode == 'sqrt_weighted_sampler':
power = 0.5
else:
raise ValueError(
f"class_balance mode {mode!r} does not build a sampler; "
f"expected 'weighted_sampler' or 'sqrt_weighted_sampler'")
raw = [(1.0 / (float(c) ** power)) if c > 0 else 0.0 for c in counts]
total = sum(raw)
return [w / total for w in raw] if total > 0 else raw
[docs]
def expected_sampled_fractions(counts, mode):
"""Class frequencies the loader is expected to realise under ``mode``.
This is what makes the effect visible before a single epoch runs: the
report prints the observed fractions next to these.
:param counts: per-class sample counts.
:param mode: any value of ``CLASS_BALANCE_MODES``.
:returns: list of expected per-class draw probabilities.
"""
total = sum(counts)
if mode not in ('weighted_sampler', 'sqrt_weighted_sampler'):
return [(c / total) if total else 0.0 for c in counts]
per_class = class_sampling_weights(counts, mode)
mass = [n * w for n, w in zip(counts, per_class)]
s = sum(mass)
return [m / s for m in mass] if s > 0 else mass
[docs]
def make_class_balance_sampler(labels, mode, num_samples=None, generator=None):
"""Build the ``WeightedRandomSampler`` for a class-balance mode.
:param labels: integer labels of the split being sampled.
:param mode: any value of ``CLASS_BALANCE_MODES``; the non-sampler modes
return ``(None, None)``.
:param num_samples: draws per epoch. Defaults to ``len(labels)`` so the
epoch keeps its usual length.
:param generator: optional ``torch.Generator`` for reproducible draws.
:returns: ``(sampler, per_sample_weights)``, or ``(None, None)``.
:raises ValueError: if ``mode`` is not a recognised class-balance mode.
"""
if mode not in CLASS_BALANCE_MODES:
raise ValueError(
f"class_balance {mode!r} is not one of {CLASS_BALANCE_MODES}")
if mode in ('none', 'weighted_loss'):
return None, None
labels = [int(v) for v in labels]
if not labels:
return None, None
n_classes = max(labels) + 1
counts = [0] * n_classes
for v in labels:
counts[v] += 1
per_class = class_sampling_weights(counts, mode)
weights = torch.as_tensor([per_class[v] for v in labels], dtype=torch.double)
sampler = WeightedRandomSampler(
weights=weights,
num_samples=int(num_samples) if num_samples is not None else len(labels),
replacement=True,
generator=generator,
)
return sampler, weights
[docs]
def report_class_balance(labels, classes=None, class_balance='none',
split_name='train', verbose=True):
"""Measure class skew, decide what was done about it, and say so out loud.
A silent auto-fix is worse than none: the printed report always names the
per-class counts, the imbalance ratio and the concrete action taken, and
when ``class_balance='none'`` on skewed data it names the modes that would
have helped instead of quietly doing nothing.
:param labels: integer labels of the split.
:param classes: ordered class names.
:param class_balance: requested mode, one of ``CLASS_BALANCE_MODES``.
:param split_name: split being described (``'train'``, ``'validation'``, ``'test'``).
:param verbose: print the report. The dict is returned either way.
:returns: the summary dict, extended with ``mode``, ``action``,
``recommendation`` and ``report``.
:raises ValueError: if ``class_balance`` is not a recognised mode.
"""
if class_balance not in CLASS_BALANCE_MODES:
raise ValueError(
f"class_balance {class_balance!r} is not one of {CLASS_BALANCE_MODES}")
summary = summarize_class_imbalance(labels, classes=classes)
summary['mode'] = class_balance
summary['split'] = split_name
recommendation = ''
if split_name != 'train':
summary['action'] = (f"none - {split_name} data is never resampled or "
f"reweighted, so its metrics keep the real class prior")
elif class_balance == 'weighted_sampler':
summary['action'] = ("WeightedRandomSampler on the train loader only "
"(per-class weight 1/n, draws ~uniform across classes)")
elif class_balance == 'sqrt_weighted_sampler':
summary['action'] = ("WeightedRandomSampler on the train loader only "
"(per-class weight 1/sqrt(n), partial correction)")
elif class_balance == 'weighted_loss':
summary['action'] = ("loss reweighting - loss_type switched to "
"'ce_weighted' (inverse-frequency class weights); "
"sampling is unchanged")
elif summary['severe']:
summary['action'] = 'none (no rebalancing applied)'
recommendation = (
"severe skew - set class_balance='weighted_sampler' to draw classes "
"uniformly, or 'weighted_loss' to reweight cross-entropy instead; "
"'sqrt_weighted_sampler' is the safer choice when the minority class "
"is small enough to be memorised")
elif summary['skewed']:
summary['action'] = 'none (no rebalancing applied)'
recommendation = (
"the data is skewed - consider class_balance='sqrt_weighted_sampler' "
"or 'weighted_loss'; accuracy will flatter the majority class as it is")
else:
summary['action'] = 'none needed (classes are within 1.5x of each other)'
if summary['empty_classes'] and not recommendation:
recommendation = (f"classes {summary['empty_classes']} have no {split_name} "
f"samples and cannot be learned or scored")
summary['recommendation'] = recommendation
summary['report'] = format_class_balance_report(summary, class_balance, split_name)
if verbose:
print(summary['report'])
return summary
[docs]
def make_cv_folds(labels, n_splits, groups=None, seed=0):
"""Split indices into ``n_splits`` class-stratified, optionally grouped folds.
Every index lands in exactly one validation fold, so the k folds partition
the dataset. With ``groups`` supplied, a whole group is assigned to a
single fold — crops from the same well never straddle the train/val line —
and groups are placed greedily into whichever fold currently leaves the
per-class proportions most even, which is how stratification survives
grouping.
:param labels: integer labels, one per sample.
:param n_splits: number of folds, must be >= 2.
:param groups: optional group id per sample (same length as ``labels``).
:param seed: seed for the shuffle. Both branches use it: ungrouped, it
shuffles each class and picks the starting fold; grouped, it orders
groups of equal size and breaks ties between folds the greedy pass
rates equally, so re-running with a different seed gives a different
— and equally stratified — partition. Where only one partition is
feasible (as many groups as folds, say) no seed can change it.
:returns: list of ``(train_idx, val_idx)`` numpy integer arrays.
:raises ValueError: if ``n_splits`` < 2, if ``groups`` is the wrong
length, or if there are fewer samples/groups than folds.
"""
labels = np.asarray([int(v) for v in labels])
n = len(labels)
k = int(n_splits)
if k < 2:
raise ValueError(f"n_splits must be >= 2 for k-fold, got {n_splits!r}")
if n < k:
raise ValueError(f"cannot build {k} folds from {n} samples")
if groups is not None and len(groups) != n:
raise ValueError(
f"groups has {len(groups)} entries but there are {n} samples")
rng = np.random.default_rng(seed)
n_classes = int(labels.max()) + 1 if n else 0
fold_of = np.empty(n, dtype=int)
if groups is None:
for c in range(n_classes):
idx = np.flatnonzero(labels == c)
if idx.size == 0:
continue
rng.shuffle(idx)
offset = int(rng.integers(k))
fold_of[idx] = (np.arange(idx.size) + offset) % k
else:
groups = np.asarray([str(g) for g in groups])
uniq = np.unique(groups)
if uniq.size < k:
raise ValueError(
f"cannot build {k} group-aware folds from {uniq.size} distinct "
f"group(s); lower cross_validation_folds or group at a finer "
f"level (e.g. cv_group_by='field')")
hist = {g: np.zeros(n_classes, dtype=float) for g in uniq}
members = {g: np.flatnonzero(groups == g) for g in uniq}
for g in uniq:
for c in labels[members[g]]:
hist[g][c] += 1.0
order = sorted(rng.permutation(uniq), key=lambda g: -hist[g].sum())
class_totals = np.bincount(labels, minlength=n_classes).astype(float)
class_totals[class_totals == 0] = 1.0
fold_class = np.zeros((k, n_classes), dtype=float)
fold_size = np.zeros(k, dtype=float)
for g in order:
fold_rank = rng.permutation(k)
best_f, best_cost = None, None
for f in range(k):
fold_class[f] += hist[g]
cost = round(
float(np.mean(np.std(fold_class / class_totals, axis=0))), 12)
fold_class[f] -= hist[g]
key = (cost, fold_size[f], int(fold_rank[f]))
if best_cost is None or key < best_cost:
best_cost, best_f = key, f
fold_class[best_f] += hist[g]
fold_size[best_f] += hist[g].sum()
fold_of[members[g]] = best_f
all_idx = np.arange(n)
folds = []
for f in range(k):
val_idx = all_idx[fold_of == f]
train_idx = all_idx[fold_of != f]
folds.append((train_idx, val_idx))
return folds
[docs]
def make_validation_holdout(labels, validation_fraction, groups, seed=0):
"""Choose one stratified group fold closest to a requested holdout size.
The ordinary Classify validation split used to call ``random_split`` and
could put crops from one well on both sides even though grouped CV did not.
This helper uses the same group-stratified partitioner as CV and selects
the candidate fold closest to the requested size and class distribution.
:param labels: One integer class label per sample, in dataset order; the
returned indices point back into that same order.
:param validation_fraction: Target share of samples to hold out, strictly
between 0 and 1; any number outside that range, ``nan`` and ``inf``
included, is a ``ValueError``. It is only a target. The holdout is one
whole fold of a split into ``max(2, round(1 / fraction))`` folds,
itself capped at the number of distinct groups, so the realised share
is quantised to whole groups and can miss in either direction, by a
lot. Over eight equal groups, 0.05 holds out 0.125 (no finer split is
available) and 0.7 holds out 0.5 (the two-fold floor); with one
dominant group — 70 of 100 samples across four groups — those same two
requests instead hold out 0.10 and 0.70. A miss that big is no longer
silent: whenever the realised share is further than
``max(0.01, 0.1 * validation_fraction)`` from the requested one, a
``UserWarning`` names both numbers and why they differ.
:param groups: Group id per sample — well, field, or whatever
``cv_group_by`` names — and required, not optional, because the point
of this function is that a group never straddles the split. Needs the
same length as ``labels`` and at least two distinct values.
:param seed: Seed passed to :func:`make_cv_folds` and used to break ties
between equally suitable folds. Different seeds can produce different
holdouts when several partitions satisfy the constraints. The seed has
no effect when the groups permit only one partition.
:returns: One ``(train_idx, val_idx)`` pair of numpy integer arrays.
:raises ValueError: if ``validation_fraction`` is outside ``(0, 1)``, if
``groups`` is missing or the wrong length, or if fewer than two
distinct groups are present.
:warns UserWarning: if whole-group quantisation makes the realised holdout
share miss ``validation_fraction`` by more than the tolerance above.
"""
labels = np.asarray([int(value) for value in labels])
fraction = float(validation_fraction)
if not 0.0 < fraction < 1.0:
raise ValueError("validation_fraction must be strictly between 0 and 1")
if groups is None or len(groups) != len(labels):
raise ValueError("group-aware validation requires one group per sample")
distinct = len(set(str(group) for group in groups))
if distinct < 2:
raise ValueError(
"A leakage-safe validation split needs at least two distinct groups; "
"choose a finer cv_group_by level or add another group."
)
requested_folds = max(2, int(round(1.0 / fraction)))
n_splits = min(requested_folds, distinct)
candidates = make_cv_folds(
labels, n_splits, groups=groups, seed=seed,
)
target_size = fraction * len(labels)
total_distribution = np.bincount(
labels, minlength=(int(labels.max()) + 1 if len(labels) else 0)
).astype(float)
total_distribution /= max(total_distribution.sum(), 1.0)
def score(candidate):
"""Rank one candidate fold; lower is a better holdout.
:param candidate: A ``(train_idx, val_idx)`` pair as produced by
:func:`make_cv_folds`; only the validation half is looked at.
:returns: ``(cost, n_validation)`` where cost adds the size error
(as a fraction of the dataset) to the mean absolute per-class
deviation from the whole dataset's distribution. The trailing
count is a tie-break, so equally good folds resolve to the
smaller holdout.
"""
_train, validation = candidate
distribution = np.bincount(
labels[validation], minlength=len(total_distribution)
).astype(float)
distribution /= max(distribution.sum(), 1.0)
size_cost = abs(len(validation) - target_size) / max(len(labels), 1)
class_cost = float(np.mean(np.abs(
distribution - total_distribution
))) if len(total_distribution) else 0.0
return size_cost + class_cost, len(validation)
tie_break = np.random.default_rng(seed).permutation(len(candidates))
train_idx, val_idx = min(
enumerate(candidates),
key=lambda item: score(item[1]) + (int(tie_break[item[0]]),),
)[1]
realised = len(val_idx) / max(len(labels), 1)
tolerance = max(0.01, 0.1 * fraction)
if abs(realised - fraction) > tolerance:
if requested_folds > distinct:
reason = (
f"only {distinct} distinct group(s) are available, so the split "
f"could not go finer than {n_splits} folds"
)
elif fraction > 0.5:
reason = (
"a holdout larger than half the data would need fewer than two "
"folds, so the split floors at 2 and holds out about half"
)
else:
reason = (
"the holdout is one whole fold and groups are never split, so "
"the share is quantised to whole groups"
)
warnings.warn(
f"validation_fraction={fraction:.4g} was requested but "
f"{len(val_idx)} of {len(labels)} samples "
f"({realised:.4g}) were held out: {reason}. Group at a finer "
f"cv_group_by level for more groups to choose from, or ask for "
f"{realised:.4g} so the setting matches what you get.",
UserWarning,
stacklevel=2,
)
return train_idx, val_idx
[docs]
def summarize_cv_folds(labels, folds, classes=None, groups=None):
"""Tabulate fold sizes and per-class validation counts.
:param labels: integer labels, one per sample.
:param folds: list of ``(train_idx, val_idx)`` from :func:`make_cv_folds`.
:param classes: ordered class names.
:param groups: optional group id per sample; adds a distinct-group column.
:returns: DataFrame with one row per fold.
"""
labels = np.asarray([int(v) for v in labels])
if classes is None:
classes = [f'class_{i}' for i in range(int(labels.max()) + 1 if labels.size else 0)]
rows = []
for i, (train_idx, val_idx) in enumerate(folds, start=1):
y_val = labels[val_idx]
row = {'fold': i, 'n_train': len(train_idx), 'n_val': len(val_idx)}
missing = []
for c, name in enumerate(classes):
cnt = int(np.sum(y_val == c))
row[f'val_{name}'] = cnt
if cnt == 0:
missing.append(name)
if groups is not None:
g = np.asarray([str(x) for x in groups])
row['val_groups'] = int(np.unique(g[val_idx]).size)
row['val_classes_missing'] = ','.join(missing)
rows.append(row)
return pd.DataFrame(rows)
[docs]
def report_cv_folds(labels, folds, classes=None, groups=None, group_by='none',
verbose=True):
"""Print the fold table and every warning the split earned.
Two failure modes are called out rather than allowed to surface later as
mysterious metrics: a class too rare to reach every fold's validation set
(its recall is undefined there), and ungrouped folds on object crops
(which leak well identity between train and validation).
:param labels: integer labels, one per sample.
:param folds: list of ``(train_idx, val_idx)``.
:param classes: ordered class names.
:param groups: optional group id per sample.
:param group_by: the grouping level that produced ``groups``.
:param verbose: print the table and warnings.
:returns: ``(fold_table, warnings)``.
"""
table = summarize_cv_folds(labels, folds, classes=classes, groups=groups)
warnings_out = []
missing = table[table['val_classes_missing'] != '']
for _, row in missing.iterrows():
warnings_out.append(
f"fold {int(row['fold'])}: no validation samples for class(es) "
f"{row['val_classes_missing']} - per-class scores for those classes "
f"are undefined in this fold and are dropped from the fold spread")
if (table['n_val'] == 0).any():
bad = table.loc[table['n_val'] == 0, 'fold'].tolist()
warnings_out.append(f"fold(s) {bad} have an empty validation set")
from .classifier_evaluation import normalize_split_level
if groups is None or normalize_split_level(group_by) == 'cell':
warnings_out.append(
"folds are NOT group-aware: crops from the same well or field can "
"land on both sides of a fold, which leaks and inflates every "
"score - set cv_group_by to 'field', 'well' or 'plate' for object crops")
if verbose:
print(f"--- Cross-validation folds (k={len(folds)}, "
f"grouping={group_by}) ---")
print(table.to_string(index=False))
for w in warnings_out:
print(f" WARNING: {w}")
return table, warnings_out
def _resolve_channel_indices(channels, verbose=False):
"""Map ``['r','g','b']``-style channel names to tensor channel indices."""
if channels is None:
channels = ['r', 'g', 'b']
chans = []
if 'r' in channels: chans.append(1)
if 'g' in channels: chans.append(2)
if 'b' in channels: chans.append(3)
if verbose:
print(f'Training a network on channels: {chans}')
print(f'Channel 1: Red, Channel 2: Green, Channel 3: Blue')
return chans
def _classification_data_dir(src, mode, classes):
"""Return ``src/<mode>`` after checking it and every class subfolder exists.
:param src: dataset root holding ``train/`` and ``test/``.
:param mode: ``'train'`` or ``'test'``.
:param classes: ordered class-folder names.
:returns: validated path to the split folder.
:raises FileNotFoundError: if the split folder or any class folder is absent.
"""
data_dir = os.path.join(src, mode)
if not os.path.isdir(data_dir):
raise FileNotFoundError(
f"No '{mode}/' folder found at: {data_dir}\n"
f"The classifier trains on a '{src}/train' and '{src}/test' split of\n"
f"annotated crops, which doesn't exist yet. Generate it first with the\n"
f"'Generate Training Data' step (spacr.io.generate_training_dataset),\n"
f"then point the Classify (CV) 'src' at the folder that contains\n"
f"train/ and test/ (each with class subfolders, e.g. 1/ and 2/)."
)
missing = [c for c in classes if not os.path.isdir(os.path.join(data_dir, c))]
if missing:
available = sorted([d for d in os.listdir(data_dir)
if os.path.isdir(os.path.join(data_dir, d))])
raise FileNotFoundError(
f"Class folders missing in {data_dir}:\n"
f" Missing: {missing}\n"
f" Available: {available}"
)
return data_dir
def _classification_transform(image_size, channel_idx, normalize):
"""Compose the resize / channel-select / normalise transform for crops."""
from .utils import SelectChannels
n_ch = len(channel_idx)
norm_transforms = (
[transforms.Normalize(mean=(0.5,) * n_ch, std=(0.5,) * n_ch)]
if normalize else []
)
return transforms.Compose([
transforms.ToTensor(),
transforms.CenterCrop(size=(image_size, image_size)),
SelectChannels(channel_idx),
*norm_transforms,
])
def _cv_group_ids(filenames, group_by, verbose=True):
"""Derive per-sample group ids at the requested plate/well/field level.
Every grouped filename must carry a verifiable identity. Anonymous names
are refused rather than converted to singleton pseudo-groups, which would
make an ordinary random split look leakage-safe.
:param filenames: image paths.
:param group_by: one of ``CV_GROUP_LEVELS``.
:param verbose: print the grouping summary.
:returns: ``(group_ids, 0)`` — ``(None, 0)`` for ``cell``/legacy ``none``.
:raises ValueError: if ``group_by`` is not a supported level.
"""
from .classifier_evaluation import normalize_split_level, split_group_values
level = normalize_split_level(group_by)
if level == 'cell':
return None, 0
_level, values = split_group_values(
group_by=level, paths=filenames, table='classification dataset')
ids = values.tolist()
if verbose:
print(f"Grouping folds by {level}: {len(set(ids))} distinct "
f"{level}(s) across {len(ids)} crops")
return ids, 0
[docs]
def generate_cv_loaders(src, n_splits, mode='train', image_size=224, batch_size=32,
classes=None, n_jobs=None, pin_memory=False, normalize=False,
channels=None, augment=False, verbose=False,
group_by='well', class_balance='none', seed=0,
crop_loading_policy=DECLARED_UINT8):
"""Build one ``(train_loader, val_loader)`` pair per cross-validation fold.
The dataset under ``src/<mode>`` is read once and then re-split k ways, so
every crop is used for validation exactly once. Folds are class-stratified
and, by default, grouped by well so that crops from the same well stay on
one side of the split. Class balancing is applied to the fold's train
loader only.
:param src: dataset root containing ``train``/``test`` subfolders.
:param n_splits: number of folds, must be >= 2.
:param mode: which split to fold — normally ``'train'``.
:param image_size: square resize target in pixels.
:param batch_size: loader batch size.
:param classes: ordered class names matching the subfolder names.
:param n_jobs: DataLoader worker count.
:param pin_memory: if True, pin batches to page-locked memory.
:param normalize: if True, apply per-channel normalisation.
:param channels: subset of RGB channels to keep.
:param augment: if True, 8-fold augment each fold's train split.
:param verbose: log configuration to stdout.
:param group_by: fold grouping level, one of ``CV_GROUP_LEVELS``.
:param class_balance: one of ``CLASS_BALANCE_MODES``, train loaders only.
:param seed: seed for the deterministic fold assignment.
:param crop_loading_policy: crop decoding policy recorded on each loader;
defaults to declared channel order and high-byte uint8 narrowing.
:returns: ``(fold_loaders, info)`` where ``fold_loaders`` is a list of
``(train_loader, val_loader)`` and ``info`` holds ``fold_table``,
``warnings``, ``imbalance`` and ``groups``.
:raises ValueError: if ``n_splits`` < 2 or a setting value is unknown.
"""
from .utils import augment_dataset
if int(n_splits) < 2:
raise ValueError(
f"cross_validation_folds must be >= 2 to build folds, got {n_splits!r}; "
f"0 or 1 means the single train/validation split")
if classes is None:
classes = ['nc', 'pc']
channel_idx = _resolve_channel_indices(channels, verbose=verbose)
data_dir = _classification_data_dir(src, mode, classes)
transform = _classification_transform(image_size, channel_idx, normalize)
data = spacrDataset(data_dir, classes, transform=transform,
shuffle=True, pin_memory=pin_memory,
crop_loading_policy=crop_loading_policy)
labels = dataset_labels(data)
filenames = dataset_filenames(data)
groups, _ = _cv_group_ids(filenames, group_by, verbose=True)
folds = make_cv_folds(labels, int(n_splits), groups=groups, seed=seed)
fold_table, fold_warnings = report_cv_folds(
labels, folds, classes=classes, groups=groups, group_by=group_by,
verbose=True)
imbalance = report_class_balance(labels, classes=classes,
class_balance=class_balance,
split_name='train', verbose=True)
num_workers = max(0, int(n_jobs)) if n_jobs is not None else 0
if num_workers > 0:
from .resource_log import _guard_workers, _loader_unit_bytes
num_workers = _guard_workers('classify', num_workers, _loader_unit_bytes(
batch_size, image_size, channels))
use_persistent = num_workers > 0
fold_loaders = []
for train_idx, val_idx in folds:
train_dataset = Subset(data, list(train_idx))
val_dataset = Subset(data, list(val_idx))
if augment:
train_dataset = augment_dataset(
train_dataset, is_grayscale=(len(channel_idx) == 1))
sampler, _ = make_class_balance_sampler(
dataset_labels(train_dataset), class_balance)
train_loader = DataLoader(
train_dataset, batch_size=batch_size,
shuffle=(sampler is None), sampler=sampler,
num_workers=num_workers, pin_memory=pin_memory,
persistent_workers=use_persistent)
val_loader = DataLoader(
val_dataset, batch_size=batch_size, shuffle=False,
num_workers=num_workers, pin_memory=pin_memory,
persistent_workers=use_persistent)
train_loader.crop_loading_policy = data.crop_loading_policy
val_loader.crop_loading_policy = data.crop_loading_policy
fold_loaders.append((train_loader, val_loader))
info = {
'fold_table': fold_table,
'warnings': fold_warnings,
'imbalance': imbalance,
'groups': groups,
'group_by': group_by,
'labels': labels,
'folds': folds,
'classes': list(classes),
'dataset': data,
}
return fold_loaders, info
[docs]
def generate_loaders(src, mode='train', image_size=224, batch_size=32,
classes=None, n_jobs=None, validation_split=0.0,
pin_memory=False, normalize=False, channels=None,
augment=False, verbose=False, class_balance='none',
seed=42, group_by='none', crop_loading_policy=DECLARED_UINT8):
"""Build ``spacrDataLoader`` objects for training, validation, or testing.
Reads class subfolders under ``src/<mode>``, applies the requested
transforms (channel selection, optional normalisation, optional
augmentation) and returns loaders sized to ``batch_size``.
:param src: Root folder containing ``train``/``test`` subfolders.
:param mode: Which split to load — ``'train'`` or ``'test'``.
:param image_size: Square resize target in pixels. Default ``224``.
:param batch_size: Loader batch size. Default ``32``.
:param classes: Ordered class names. Default ``['nc', 'pc']``.
:param n_jobs: DataLoader worker count. ``None`` (the default) means 0
workers, i.e. batches are read in the calling process.
:param validation_split: Fraction of the train split to hold out.
:param pin_memory: If True, pin batches to page-locked memory.
:param normalize: If True, apply per-channel normalisation.
:param channels: Subset of RGB channels to keep, e.g. ``['r', 'g']``.
:param augment: If True, apply the training augmentation pipeline.
:param verbose: If True, log configuration to stdout.
:param class_balance: One of ``CLASS_BALANCE_MODES``. ``'none'`` (default)
leaves sampling untouched; the sampler modes attach a
``WeightedRandomSampler`` to the TRAIN loader only. The skew is
reported either way.
:param seed: Reproducible train/validation split and loader order.
:param group_by: ``field``, ``well`` or ``plate`` keeps that acquisition
identity entirely on one side of the ordinary validation holdout.
``none`` retains the legacy per-object random split.
:param crop_loading_policy: crop decoding policy recorded on each loader;
defaults to declared channel order and high-byte uint8 narrowing.
:returns: For ``mode='train'``, a tuple of loaders and a plot handle;
for ``mode='test'``, the test loader (plus optional metadata).
:raises ValueError: if ``class_balance`` is not a recognised mode.
"""
if classes is None:
classes = ['nc', 'pc']
from .utils import augment_dataset
if class_balance not in CLASS_BALANCE_MODES:
raise ValueError(
f"class_balance {class_balance!r} is not one of {CLASS_BALANCE_MODES}")
channels = _resolve_channel_indices(channels, verbose=verbose)
if mode == 'train':
print('Loading Train and validation datasets')
elif mode == 'test':
validation_split = 0.0
print('Loading test dataset')
else:
print(f'mode:{mode} is not valid, use mode = train or test')
return
data_dir = _classification_data_dir(src, mode, classes)
transform = _classification_transform(image_size, channels, normalize)
data = spacrDataset(data_dir, classes, transform=transform,
shuffle=True, pin_memory=pin_memory,
crop_loading_policy=crop_loading_policy)
num_workers = max(0, int(n_jobs)) if n_jobs is not None else 0
if num_workers > 0:
from .resource_log import _guard_workers, _loader_unit_bytes
num_workers = _guard_workers('classify', num_workers, _loader_unit_bytes(
batch_size, image_size, channels))
use_persistent = num_workers > 0
if validation_split > 0 and mode == 'train':
from .classifier_evaluation import (grouped_split,
normalize_split_level)
level = normalize_split_level(group_by)
filenames = dataset_filenames(data)
groups, _unparsed = _cv_group_ids(filenames, level, verbose=True)
if groups is None:
groups = np.arange(len(data), dtype=object)
train_idx, val_idx, split_report = grouped_split(
groups, dataset_labels(data), validation_split, seed=seed,
group_by=level)
print(split_report.summary())
train_dataset = Subset(data, list(train_idx))
val_dataset = Subset(data, list(val_idx))
train_size, val_size = len(train_idx), len(val_idx)
if not augment:
print(f'Train data:{train_size}, Validation data:{val_size}')
generator = torch.Generator().manual_seed(int(seed))
if augment:
print(f'Data before augmentation: Train: {len(train_dataset)}, Validation:{len(val_dataset)}')
train_dataset = augment_dataset(train_dataset, is_grayscale=(len(channels) == 1))
print(f'Data after augmentation: Train: {len(train_dataset)}')
report_class_balance(dataset_labels(train_dataset), classes=classes,
class_balance=class_balance, split_name='train')
report_class_balance(dataset_labels(val_dataset), classes=classes,
class_balance='none', split_name='validation')
sampler, _ = make_class_balance_sampler(
dataset_labels(train_dataset), class_balance)
print(f'Generating Dataloader with {num_workers} workers')
train_loaders = DataLoader(train_dataset, batch_size=batch_size,
shuffle=(sampler is None),
sampler=sampler,
generator=generator,
num_workers=num_workers,
pin_memory=pin_memory,
persistent_workers=use_persistent)
val_loaders = DataLoader(val_dataset, batch_size=batch_size,
shuffle=False,
num_workers=num_workers,
pin_memory=pin_memory,
persistent_workers=use_persistent)
train_fig = None
train_loaders.crop_loading_policy = data.crop_loading_policy
val_loaders.crop_loading_policy = data.crop_loading_policy
return train_loaders, val_loaders, train_fig
else:
split_name = 'train' if mode == 'train' else 'test'
effective_balance = class_balance if split_name == 'train' else 'none'
report_class_balance(dataset_labels(data), classes=classes,
class_balance=effective_balance,
split_name=split_name)
sampler = None
if split_name == 'train':
sampler, _ = make_class_balance_sampler(dataset_labels(data),
class_balance)
train_loaders = DataLoader(data, batch_size=batch_size,
shuffle=(sampler is None),
sampler=sampler,
num_workers=num_workers,
pin_memory=pin_memory,
persistent_workers=use_persistent)
val_loaders = []
train_fig = None
train_loaders.crop_loading_policy = data.crop_loading_policy
return train_loaders, val_loaders, train_fig
[docs]
def generate_training_dataset(settings):
"""
Build a balanced training/testing dataset from one of:
- metadata rules (exact matches or compound 'where' rules)
- annotation columns (each <col>_<value> is a standalone class)
- measurement rules (numeric ranges/bins; supports multiple conditions per class)
New behavior (annotation mode):
- If a column has only one annotated value (e.g., only '1's), we add a
'<column>_random' class using unannotated rows for that column (same size as positives).
- Optional: persist that random selection into DB as a new INT column named '<column>_random' with 1's.
``crop_source`` chooses where the pixels come from. ``'png'`` (and
``'auto'`` wherever a crop folder exists) copies the pre-generated PNGs,
unchanged. ``'merged'`` (and ``'auto'`` with no crop folder) cuts each
selected crop out of ``merged/*.npy`` through :mod:`spacr.crops` instead:
the labels still come from ``png_list``, but the pixels are cut fresh at
the current crop settings, so the training set costs no standing disk and
cannot be built out of a folder that has gone stale. A project with no
``png_list`` at all falls back to the object measurement table, which
still carries everything the metadata rules select on.
:param settings: Settings dict, first completed by
:func:`spacr.settings.set_generate_training_dataset_defaults`. Keys
read here:
- ``src``: plate root, or a list of roots merged into one dataset.
- ``dataset_mode``: ``'metadata'``, ``'annotation'`` or ``'measurement'``.
- ``test_split``: fraction of each class routed to ``test/``.
- ``cv_group_by``: acquisition identity kept intact across the
train/test boundary. Default ``'well'``.
- ``path_string`` (legacy alias ``png_type``): substring a crop path
must contain. Default ``'cell_png'``.
- ``crop_source``: ``'auto'``, ``'png'`` or ``'merged'``, as above.
- ``balance_to_smallest``: downsample every class to the smallest one.
Default True.
- ``random_seed``: seeds balancing and random-class sampling.
Default 42.
- ``tables``: object tables the project has. Only clears
``nuclei_limit`` and ``pathogen_limit`` when the matching table is
absent; it does not select where the crop list comes from.
- metadata mode: ``metadata_rules``, or ``class_metadata`` values
matched against the column the Classes editor names.
- annotation mode: ``annotation_columns`` (legacy
``annotation_column``), optional ``annotation_values`` filter, and
``write_random_annotation_column``.
The resulting ``class_folder_names`` and ``nr_classes`` are written
back into the dict for downstream training. A pre-split list-shaped
``classes`` entry is retired only after those folders are written;
dict-shaped class definitions remain untouched.
:returns: ``(train_class_dir, test_class_dir)`` — the ``train/`` and
``test/`` roots written under ``<src>/datasets/training``
(``training_all`` when several sources are combined), suffixed to stay
unique.
:raises ValueError: if ``dataset_mode`` is unrecognised, a rule names a
missing column or unsupported operator, or a requested class selected
no crops.
"""
import os, random, operator, sqlite3
import numpy as np
from .utils import save_settings
from .settings import set_generate_training_dataset_defaults
settings = set_generate_training_dataset_defaults(settings)
balance_to_smallest = bool(settings.get('balance_to_smallest', True))
png_type = settings.get('path_string') or settings.get('png_type', 'cell_png')
tables = settings.get('tables') or ['cell', 'nucleus', 'pathogen', 'cytoplasm']
write_rand_col = bool(settings.get('write_random_annotation_column', False))
rng = random.Random(int(settings.get('random_seed', 42)))
if 'nucleus' not in tables:
settings['nuclei_limit'] = False
if 'pathogen' not in tables:
settings['pathogen_limit'] = 0
save_settings(settings, 'cv_dataset', show=True)
if isinstance(settings['src'], str):
settings['src'] = [settings['src']]
def _ensure_unique_dir(dst_base):
"""``dst_base``, or the first ``dst_base_N`` that does not exist yet.
A TRAINING SET IS NEVER WRITTEN OVER ONE THAT IS ALREADY THERE. The
folder is the record of what a model was trained on, so reusing the
name would leave a model whose training data cannot be reconstructed.
:param dst_base: the folder that was asked for.
:returns: a folder path nothing occupies.
"""
dst = dst_base
if os.path.exists(dst):
base = dst
j = 1
while os.path.exists(f"{base}_{j}"):
j += 1
dst = f"{base}_{j}"
print(f'Creating new directory for training: {dst}')
return dst
def _load_png_table(db_path, object_type='cell'):
"""The per-object crop table, or the measurements standing in for it.
`png_list` ALONE, deliberately: joining it against the measurement
tables would drop every object those tables do not also carry, and a
training set is allowed to be a subset of what was measured.
NO `png_list` MEANS NO PNG FOLDER WAS EVER WRITTEN, which is not the
same as no data. The objects are still in the measurement table with
the same well metadata the class rules select on, so that is read
instead -- otherwise a project holding everything it needs reports
"0 classes".
:param db_path: the measurements database.
:param object_type: which object's table to fall back to.
"""
try:
[png_df] = _read_db(db_loc=db_path, tables=['png_list'])
png_df = png_df.copy()
except Exception:
png_df = pd.DataFrame()
if len(png_df):
return png_df
print(f"No 'png_list' rows in {db_path}; falling back to the "
f"'{object_type}' measurement table for the crop list.")
return crop_rows_from_object_table(db_path, object_type)
def _class_items(frame):
"""Return the per-row crop entries a class list is built from.
On-demand handles when the merged source is in play, PNG paths
otherwise -- generate_dataset_from_lists takes either.
"""
if CROP_REF_COLUMN in frame.columns:
return [ref for ref in frame[CROP_REF_COLUMN].tolist()
if ref is not None]
return frame['png_path'].dropna().tolist()
def _fix_path_under_src(src_root, p):
"""Make sure png_path lives under the current src root (portable absolute fix).
THE RULE ITSELF LIVES IN `spacr.portable_paths` and is shared with the
montage, which needs exactly this and used to get none of it -- the
rule was a nested local here, reachable only from this generator, so a
screen that had moved computer showed the montage 60,816 dead paths
while this function resolved every one.
The `/data/` rebuild is now only applied when it lands on a file that
EXISTS. Rewriting to somewhere equally absent is strictly worse than
leaving the recorded path alone: the copy below then fails naming a
folder the user never chose.
"""
from .portable_paths import reroot_crop_path
if not isinstance(p, str) or p.strip() == "":
return None
rerooted = reroot_crop_path(p, src_root)
if rerooted != p:
return rerooted
if not os.path.isabs(p):
return os.path.join(src_root, p.lstrip('/'))
return p
def _apply_where(df, where):
"""where: list of {'column','op','value'} AND-combined."""
if not where:
return df
OPS = {
'==': operator.eq, '!=': operator.ne,
'<': operator.lt, '<=': operator.le,
'>': operator.gt, '>=': operator.ge,
'in': lambda a,b: a.isin(b) if hasattr(a, 'isin') else False,
'notin': lambda a,b: ~a.isin(b) if hasattr(a, 'isin') else False,
}
mask = np.ones(len(df), dtype=bool)
for cond in where:
if not isinstance(cond, dict):
raise ValueError(
f"Dataset selection conditions must be mappings; "
f"received {cond!r}.")
col, op = cond.get('column'), cond.get('op')
val = cond.get('value', None)
if col not in df.columns:
raise ValueError(
f"Dataset selection rule references unknown column "
f"{col!r}. Available columns: "
f"{sorted(map(str, df.columns))}.")
if op not in OPS:
raise ValueError(
f"Dataset selection rule uses unsupported operator "
f"{op!r}. Choose from {sorted(OPS)}.")
series = df[col]
if op in ('in','notin'):
vals = val if isinstance(val, (list,tuple,set)) else [val]
mask &= OPS[op](series, vals)
else:
mask &= OPS[op](series, val)
return df[mask]
def _balance_lists(list_of_lists):
"""Cut every class down to the smallest one, when that was asked for.
A CLASSIFIER TRAINED ON 9,000 negatives and 300 positives learns to
say "negative", so balancing is the ordinary case rather than an
exotic one -- but it THROWS AWAY DATA, which is why it is a setting
and not a default of this function.
Sampled rather than truncated: the first N crops of a class share a
plate, a well and often a field, so taking them in order would trade
a class imbalance for a batch imbalance.
The no-classes gate immediately before the call rejects an empty
list, so only populated collections reach this.
:param list_of_lists: one list of crop paths per class.
"""
if not balance_to_smallest:
return list_of_lists
sizes = [len(x) for x in list_of_lists]
size = min(sizes) if sizes else 0
print(f"Class sizes: {sizes} -> balancing to {size}")
out = []
for paths in list_of_lists:
if len(paths) > size:
out.append(rng.sample(paths, size))
else:
out.append(paths)
return out
def _annotation_classes_from_columns(png_df, ann_cols, ann_vals_filter=None, db_path=None):
"""
Build classes per (column,value). If a column only has one annotated value in {1,2},
also create '<column>_random' from unannotated rows (same count as positives).
Optionally persist '<column>_random' as a new INT column with 1's for sampled rows.
Returns (names, lists) aligned.
"""
names, lists = [], []
df = png_df.copy()
keep_cols = ['png_path'] + (
[CROP_REF_COLUMN] if CROP_REF_COLUMN in df.columns else []
) + [c for c in ann_cols if c in df.columns]
df = df[keep_cols]
for col in ann_cols:
if col not in df.columns:
print(f"Warning: annotation column '{col}' not in png_list; skipping.")
continue
col_series = df[col].dropna()
try:
vals = sorted(set(col_series.astype(int).tolist()))
except Exception:
vals = sorted(set(col_series.tolist()))
if ann_vals_filter and col in ann_vals_filter:
allow = set(ann_vals_filter[col])
vals = [v for v in vals if v in allow]
distinct_vals = []
for v in vals:
cls_name = f"{col}_{v}"
sel = _class_items(df[df[col] == v])
distinct_vals.append((v, sel))
names.append(cls_name)
lists.append(sel)
if len(distinct_vals) == 1:
v, pos_paths = distinct_vals[0]
pos_n = len(pos_paths)
unann_paths = _class_items(df[df[col].isna()])
if not unann_paths:
print(f"Column '{col}': no unannotated rows available for <{col}_random>; skipping random class.")
continue
if pos_n == 0:
print(f"Column '{col}': only one value present but it has 0 rows; skipping random class.")
continue
if len(unann_paths) >= pos_n:
rand_paths = rng.sample(unann_paths, pos_n)
else:
rand_paths = unann_paths
names.append(f"{col}_random")
lists.append(rand_paths)
if write_rand_col and db_path:
rand_col = f"{col}_random"
qcol = rand_col.replace('"', '""')
with sqlite3.connect(db_path, timeout=30) as conn:
cur = conn.cursor()
cur.execute('PRAGMA table_info("png_list")')
existing = {r[1] for r in cur.fetchall()}
if rand_col not in existing:
cur.execute(f'ALTER TABLE "png_list" ADD COLUMN "{qcol}" INTEGER')
conn.commit()
for p in rand_paths:
png_path = (p.row.get('png_path')
if isinstance(p, LazyCropPNG) else p)
if not png_path:
continue
cur.execute(
f'UPDATE "png_list" SET "{qcol}" = 1 WHERE png_path = ?',
(png_path,)
)
conn.commit()
return names, lists
class_path_list = None
class_names = None
first_src = settings['src'][0]
dst_final = _ensure_unique_dir(os.path.join(
first_src, 'datasets',
'training_all' if len(settings['src']) > 1 else 'training'))
crop_db_path = None
selection_context = []
for i, src in enumerate(settings['src']):
db_path = os.path.join(src, 'measurements', 'measurements.db')
object_type = crop_object_type(png_type)
png_df = _load_png_table(db_path, object_type)
fixed_paths = [ _fix_path_under_src(src, p) for p in png_df['png_path'] ]
png_df['png_path'] = fixed_paths
if png_type:
png_df = png_df[png_df['png_path'].astype(str).str.contains(png_type, na=False)]
source = open_crop_source(settings, src, object_type=object_type)
if source is not None and getattr(source, 'kind', 'png') == 'merged':
rows = crop_rows_from_png_list(db_path, png_df, object_type)
refs = crop_refs_for_rows(source, rows, object_type)
rows = rows.copy()
rows[CROP_REF_COLUMN] = refs
png_df = rows
crop_db_path = db_path if os.path.isfile(db_path) else None
from .training_basis import resolve_basis
mode = resolve_basis(settings)
this_names, this_lists = [], []
if mode == 'metadata':
rules = settings.get('metadata_rules')
if rules:
if all('name' in r for r in rules):
for r in rules:
where = r.get('where')
col, op, val = r.get('column'), r.get('op'), r.get('value')
if where is None and col is not None and op is not None:
where = [{'column': col, 'op': op, 'value': val}]
df_sel = _apply_where(png_df, where)
this_names.append(r['name'])
this_lists.append(_class_items(df_sel))
else:
for r in rules:
col, op, val = r['column'], r['op'], r['value']
df_sel = _apply_where(png_df, [{'column': col, 'op': op, 'value': val}])
name = r.get('name', f"{col}{op}{val}")
this_names.append(name)
this_lists.append(_class_items(df_sel))
else:
class_meta = settings.get('class_metadata') or []
if isinstance(class_meta, str):
import ast as _ast
try:
parsed = _ast.literal_eval(class_meta.strip())
except (ValueError, SyntaxError):
parsed = [p.strip() for p in class_meta.split(',') if p.strip()]
class_meta = (
parsed if isinstance(parsed, (list, tuple))
else [parsed]
)
if not class_meta:
raise ValueError(
"metadata dataset mode requires at least one "
"class_metadata value.")
meta_col = _class_column(settings)
if meta_col not in png_df.columns:
raise ValueError(
f"metadata mode: column '{meta_col}' is not in png_list, "
f"so no class can be selected. Present columns: "
f"{sorted(map(str, png_df.columns))}. Set the Classes "
f"editor's column to one of those (usually 'columnID' "
f"or 'rowID'), or switch 'dataset_mode' to "
f"'annotation'."
)
meta_values = png_df[meta_col].astype(str)
selection_context.append(
f"{src}: available {meta_col} values are "
f"{sorted(meta_values.dropna().unique().tolist())}"
)
for cm in class_meta:
if isinstance(cm, (list, tuple, set)):
wanted = [str(v) for v in cm]
else:
wanted = [str(cm)]
name = wanted[0] if len(wanted) == 1 else '_'.join(wanted)
sel = png_df[meta_values.isin(wanted)]
this_names.append(name)
this_lists.append(_class_items(sel))
else:
ann_cols = settings.get('annotation_columns')
if not ann_cols:
ann_cols = [settings.get('annotation_column')]
ann_cols = [c for c in (ann_cols or []) if c]
if not ann_cols:
raise ValueError(
"annotation dataset mode requires at least one "
"annotation_columns entry (or annotation_column).")
ann_vals = settings.get('annotation_values')
this_names, this_lists = _annotation_classes_from_columns(
png_df, ann_cols, ann_vals_filter=ann_vals, db_path=db_path
)
if class_path_list is None:
class_path_list = [[] for _ in range(len(this_lists))]
class_names = this_names[:]
if this_names != class_names:
print("Warning: class name/order mismatch across sources; aligning by index. "
"Make sure your rules are identical for all 'src' roots.")
for idx in range(min(len(class_path_list), len(this_lists))):
class_path_list[idx].extend(this_lists[idx])
if not class_path_list or sum(len(x) for x in class_path_list) == 0:
details = "\n".join(f" {line}" for line in selection_context)
raise ValueError(
"Training-dataset generation selected no crops for any class. "
"Check class_metadata against the column the Classes editor "
"names, or choose an annotation value that occurs in the "
"database."
+ (f"\n{details}" if details else "")
)
empty_classes = [
name for name, items in zip(class_names or [], class_path_list)
if not items
]
if empty_classes:
counts = ", ".join(
f"{name}={len(items)}"
for name, items in zip(class_names or [], class_path_list)
)
details = "\n".join(f" {line}" for line in selection_context)
raise ValueError(
"Training-dataset generation cannot balance or train because "
f"these classes selected no crops: {empty_classes}. "
f"Selected counts: {counts}. Change class_metadata/rules to "
"values that exist, or add annotations for the missing class."
+ (f"\n{details}" if details else "")
)
class_path_list = _balance_lists(class_path_list)
from .io import generate_dataset_from_lists
final_names = class_names or [f"class_{i}" for i in range(len(class_path_list))]
print(f"class_path_list: {len(class_path_list)} classes")
train_class_dir, test_class_dir = generate_dataset_from_lists(
dst_final,
class_data=class_path_list,
classes=final_names,
test_split=settings['test_split'],
db_path=crop_db_path,
random_seed=settings.get('random_seed', 42),
group_by=settings.get('cv_group_by', 'well'),
)
from .classify_classes import _record_generated_folder_names
_record_generated_folder_names(settings, final_names)
settings['nr_classes'] = len(final_names)
try:
save_settings(settings, 'cv_dataset', show=False)
except Exception as exc:
LOG.warning("the cv_dataset settings snapshot was not written (%s); "
"the dataset in %s cannot be reproduced from disk.",
exc, train_class_dir)
return train_class_dir, test_class_dir
[docs]
def training_dataset_from_annotation(db_path, dst, annotation_column='test', annotated_classes=(1, 2)):
"""Group per-object PNG paths by manual annotation values so they can be turned into a CNN training set.
Reads the ``png_list`` table of a spacr ``measurements.db``, buckets
PNG paths by the value found in ``annotation_column`` (typically
filled by the spacr annotation GUI), and, when only one class has
been annotated, samples an equal-sized "other" class from
unannotated rows. The returned list-of-lists is consumed by
:func:`generate_dataset_from_lists` to lay out
``train/<class>/*.png`` / ``test/<class>/*.png``.
:param db_path: SQLite ``measurements.db`` containing a
``png_list`` table with ``png_path`` plus ``annotation_column``.
:param dst: Output root (currently unused; kept for API symmetry
with sister builders).
:param annotation_column: Column in ``png_list`` holding class
labels. Default ``'test'``.
:param annotated_classes: Class values to pull from
``annotation_column``. When length is 1, an equal-sized "other"
class is sampled from rows whose annotation != that value.
:returns: List of lists — one list of PNG paths per output class,
in the same order as ``annotated_classes``.
Example:
.. code-block:: python
from spacr.io import training_dataset_from_annotation, generate_dataset_from_lists
class_data = training_dataset_from_annotation(
'/data/plate01/measurements/measurements.db',
dst='/data/plate01/dataset',
annotation_column='test', annotated_classes=(1, 2),
)
generate_dataset_from_lists('/data/plate01/dataset', class_data, classes=['neg','pos'])
See Also:
:func:`training_dataset_from_annotation_metadata` — same, but
first restricts rows by plate row/column metadata.
:func:`generate_dataset_from_lists` — turns the returned lists
into a ``train/`` / ``test/`` folder tree.
"""
all_paths = []
print(f'Reading DataBase: {db_path}')
with sqlite3.connect(db_path, timeout=30) as conn:
cursor = conn.cursor()
query = f"SELECT png_path, {annotation_column} FROM png_list"
cursor.execute(query)
while True:
rows = cursor.fetchmany(1000)
if not rows:
break
for row in rows:
all_paths.append(row)
print('Total paths retrieved:', len(all_paths))
class_paths = []
for class_ in annotated_classes:
class_paths_temp = [path for path, annotation in all_paths if annotation == class_]
class_paths.append(class_paths_temp)
print(f'Found {len(class_paths_temp)} images in class {class_}')
if len(annotated_classes) == 1:
target_class = annotated_classes[0]
count_target_class = len(class_paths[0])
print(f'Annotated class: {target_class} with {count_target_class} images')
alt_class_paths = [path for path, annotation in all_paths if annotation != target_class]
print('Alternative paths available:', len(alt_class_paths))
balanced_count = min(count_target_class, len(alt_class_paths))
print(f'Sampling {balanced_count} images for each class')
sampled_target_class_paths = random.sample(class_paths[0], balanced_count)
sampled_alt_class_paths = random.sample(alt_class_paths, balanced_count)
class_paths[0] = sampled_target_class_paths
class_paths.append(sampled_alt_class_paths)
print(f'Generated a list of lists from annotation of {len(class_paths)} classes')
for i, ls in enumerate(class_paths):
print(f'Class {i}: {len(ls)} images')
return class_paths
def _class_column(settings) -> str:
"""The png_list column the classes are defined by.
ONE PLACE, and `classes` is it: every class in the Classes editor already
carries the column its value came from, so asking for the column a second
time under its own setting was asking the user to restate something
spaCR knows -- and giving them a way to say it differently.
An older settings file that still names `metadata_type_by` is honoured
first, so a CSV written before the removal runs unchanged.
:param settings: the run settings.
:returns: the column name, defaulting to 'columnID'.
"""
from collections.abc import Mapping
legacy = str(settings.get('metadata_type_by') or '').strip()
if legacy:
return legacy
classes = settings.get('classes')
if isinstance(classes, Mapping):
for rule in classes.values():
if isinstance(rule, Mapping):
column = str(rule.get('column') or '').strip()
if column:
return column
return 'columnID'
def _crop_format_of_items(items, db_path=None):
"""Return the crop format the items share, or None when they disagree.
Resolve each PNG independently so interrupted migrations are respected.
Uniform copies retain that format; mixed sources need normalization in the
destination. Crops cut on demand are always current.
"""
from . import crops
formats = set()
for item in items:
if isinstance(item, LazyCropPNG):
formats.add(crops.CROP_FORMAT_CURRENT)
else:
formats.add(crops.crop_format_for_png(str(item), db_path=db_path))
if len(formats) == 1:
return formats.pop()
return None
def _write_class_item(item, dst_dir, *, canonicalize=False, db_path=None):
"""Put one crop into ``dst_dir``: copy a path, cut a :class:`LazyCropPNG`."""
if isinstance(item, LazyCropPNG):
out = os.path.join(dst_dir, item.name or 'crop.png')
with open(out, 'wb') as handle:
handle.write(item.png_bytes())
return out
out = os.path.join(dst_dir, os.path.basename(str(item)))
if canonicalize:
from .crops import read_crop_png
Image.fromarray(read_crop_png(str(item), db_path=db_path)).save(out)
else:
shutil.copy(str(item), out)
return out
[docs]
def generate_dataset_from_lists(dst, class_data, classes, test_split=0.1,
db_path=None, random_seed=42,
group_by='well'):
"""Put the crops listed per class into ``dst/train/<class>`` and ``dst/test/<class>``.
An entry may be a **path**, which is copied byte for byte exactly as
before, or a :class:`LazyCropPNG`, which is cut out of ``merged/*.npy``
through :mod:`spacr.crops` and written as a current-format (RGB) crop PNG.
The two are interchangeable, so a training set can be built with no crop
folder on disk at all.
Each destination class folder is stamped before it is filled. Uniform
source formats keep their original bytes and marker. Mixed source formats
are decoded into declared uint8 crops in the destination only, so no
generated folder silently loses its channel-order record.
:param dst: Output root; ``train`` and ``test`` subfolders are created.
:param class_data: Sequence of per-class lists of paths and/or
:class:`LazyCropPNG` handles.
:param classes: Class names paired positionally with ``class_data``.
:param test_split: Fraction of each class routed to ``test/``.
Default ``0.1``.
:param db_path: optional ``measurements.db`` consulted for the crop format
of a source folder that carries no sidecar.
:param random_seed: Reproducible global train/test split seed.
:param group_by: acquisition identity kept intact across the permanent
train/test boundary. Default ``well``. ``cell`` is the leakiest
per-object choice; legacy ``none`` aliases it.
:returns: ``(train_dir, test_dir)`` tuple of the top-level split paths.
:raises ValueError: if ``len(class_data) != len(classes)``.
"""
from .utils import print_progress
if len(class_data) != len(classes):
raise ValueError("class_data and classes must have the same length.")
total_files = sum(len(data) for data in class_data)
processed_files = 0
time_ls = []
failed = 0
every_item = [item for data in class_data for item in data]
fmt = _crop_format_of_items(every_item, db_path=db_path)
canonicalize = bool(every_item and fmt is None)
if canonicalize:
from .crops import CROP_FORMAT_CURRENT
fmt = CROP_FORMAT_CURRENT
print(f"This dataset mixes crops of more than one format; writing "
f"declared-order uint8 copies into {dst}. Source images are unchanged.")
if every_item:
os.makedirs(dst, exist_ok=True)
from .crops import write_crop_folder_marker
write_crop_folder_marker(dst, fmt=fmt, classes=list(map(str, classes)),
split='train/test')
from .classifier_evaluation import grouped_split, split_group_values
grouped_splits = None
split_report = None
if class_data:
flat_items = []
flat_labels = []
for class_index, data in enumerate(class_data):
for item in data:
flat_items.append(item)
flat_labels.append(class_index)
if flat_items:
names = [
item.name if isinstance(item, LazyCropPNG) else str(item)
for item in flat_items
]
level, flat_groups = split_group_values(
group_by=group_by, paths=names,
table='generated training dataset')
train_indices, test_indices, split_report = grouped_split(
flat_groups, flat_labels, test_split, seed=random_seed,
group_by=level,
)
train_index_set = set(map(int, train_indices))
test_index_set = set(map(int, test_indices))
grouped_splits = {
class_index: ([], []) for class_index in range(len(classes))
}
for index, (item, class_index) in enumerate(
zip(flat_items, flat_labels)
):
destination = (
grouped_splits[class_index][0]
if index in train_index_set
else grouped_splits[class_index][1]
)
destination.append(item)
print(split_report.summary())
os.makedirs(dst, exist_ok=True)
with open(os.path.join(dst, '.spacr_split.json'), 'w') as handle:
json.dump(split_report.to_dict(), handle, indent=2,
sort_keys=True)
for class_index, (cls, data) in enumerate(zip(classes, class_data)):
train_class_dir = os.path.join(dst, f'train/{cls}')
test_class_dir = os.path.join(dst, f'test/{cls}')
os.makedirs(train_class_dir, exist_ok=True)
os.makedirs(test_class_dir, exist_ok=True)
print('data',len(data), test_split)
if not data:
print(f"Class {cls!r} selected no crops; its folders are empty.")
continue
write_crop_folder_marker(train_class_dir, fmt=fmt)
write_crop_folder_marker(test_class_dir, fmt=fmt)
train_data, test_data = grouped_splits[class_index]
for item in train_data:
start = time.time()
try:
_write_class_item(item, train_class_dir, canonicalize=canonicalize,
db_path=db_path)
except Exception as exc:
failed += 1
if failed <= 5:
print(f"Could not add {item!r} to {train_class_dir}: {exc}")
duration = time.time() - start
time_ls.append(duration)
print_progress(processed_files, total_files, n_jobs=1, time_ls=None, batch_size=None, operation_type="Copying files for Train dataset")
processed_files += 1
for item in test_data:
start = time.time()
try:
_write_class_item(item, test_class_dir, canonicalize=canonicalize,
db_path=db_path)
except Exception as exc:
failed += 1
if failed <= 5:
print(f"Could not add {item!r} to {test_class_dir}: {exc}")
duration = time.time() - start
time_ls.append(duration)
print_progress(processed_files, total_files, n_jobs=1, time_ls=None, batch_size=None, operation_type="Copying files for Test dataset")
processed_files += 1
empty = []
for cls in classes:
train_class_dir = os.path.join(dst, f'train/{cls}')
test_class_dir = os.path.join(dst, f'test/{cls}')
n_train = len([f for f in os.listdir(train_class_dir) if not f.startswith('.')])
n_test = len([f for f in os.listdir(test_class_dir) if not f.startswith('.')])
print(f'Train class {cls}: {n_train}, Test class {cls}: {n_test}')
if n_train == 0:
empty.append(cls)
if failed:
print(f"Warning: {failed} of {total_files} crops could not be written "
f"into {dst}.")
if failed == total_files:
raise RuntimeError(
f"No crop could be written into {dst}: all {total_files} "
f"selected crops failed. If the PNG crop folder has been "
f"deleted or moved, set crop_source='merged' to cut the crops "
f"out of merged/*.npy instead.")
if empty:
print(f"Warning: class(es) {', '.join(map(str, empty))} have no "
f"training images; the model cannot learn them.")
return os.path.join(dst, 'train'), os.path.join(dst, 'test')
def _next_synthetic_yokogawa_well(used_wells, n_wells=384):
"""Return the next free ``plate<N>_<well>`` id, and claim it.
Fills one plate before starting the next, and **never returns an id
that is already in** ``used_wells``. The version this replaces fell out
of its ``for`` loop and returned ``f"plate{plate}_A01"`` unconditionally,
so the 386th caller got ``plate2_A01`` a second time and its TIFF
overwrote the 385th's — 386 inputs, 385 outputs, nothing said.
:param used_wells: set of ids already handed out; mutated in place.
:param n_wells: plate format to fill, a key of
:data:`spacr.schema.PLATE_FORMATS`.
:returns: the claimed ``plate<N>_<well>`` id.
"""
sequence = _cv.well_sequence(n_wells)
plate = 1
while True:
for well in sequence:
name = f"plate{plate}_{well}"
if name not in used_wells:
used_wells.add(name)
return name
plate += 1
[docs]
def convert_separate_files_to_yokogawa(folder, regex):
"""Rename per-slice TIFFs in ``folder`` into the Yokogawa CV filename convention.
Files are grouped by ``(plateID, wellID, fieldID, timeID, chanID)``
parsed from the regex. Groups with multiple Z-slices are max-
projected before saving, and the mapping is logged to
``rename_log.csv``.
Well naming, in full:
1. A ``wellID`` that **is** a well address keeps it —
:func:`spacr.convert.normalise_well` reads ``a1``, ``A-01``, ``Q01``
(row 17) and ``AA13`` (row 27) alike. Every one of those used to be
thrown away and replaced with the next free synthetic id, so a real
1536-plate came out relabelled ``A01, A02, …`` with only
``rename_log.csv`` to say what had happened.
2. Anything else — ``1``, ``well_left``, a positional number — is
handed a synthetic id, in ``_natural_key`` order so the same folder
always converts the same way. It used to follow ``os.listdir`` order,
which is the filesystem's business and not reproducible.
3. Each distinct source ``plateID`` gets its own ``plate<N>`` token, so
well ``A01`` of two source plates stays two wells.
:param folder: Folder containing the source TIFFs.
:param regex: Pattern with named groups ``wellID`` (required) plus
optional ``plateID``, ``fieldID``, ``timeID``, ``chanID``,
``sliceID``.
:returns: None
:raises ValueError: when a file the regex matches carries a ``fieldID``,
``timeID``, ``chanID`` or ``sliceID`` that is not a whole number
(before anything is written), or cannot be read and converted. The
message names the file, and in the second case says how many
converted files were written before the conversion stopped;
``rename_log.csv`` is written only when every file converted.
"""
pattern = re.compile(regex, re.I)
files_by_region = {}
rename_log = []
csv_path = os.path.join(folder, "rename_log.csv")
used_wells = set()
region_to_well = {}
for file in sorted(_listdir_visible(folder)):
match = pattern.match(file)
if not match:
print(f"Skipping {file}: does not match regex.")
continue
meta = match.groupdict()
if 'wellID' not in meta or meta['wellID'] is None:
print(f"Skipping {file}: missing mandatory wellID.")
continue
wellID = meta['wellID']
plateID = meta.get('plateID', '1') or '1'
fieldID = meta.get('fieldID', '1') or '1'
try:
int(fieldID)
timeID = int(meta.get('timeID', 1) or 1)
chanID = int(meta.get('chanID', 1) or 1)
sliceID = meta.get('sliceID')
sliceID = int(sliceID) if sliceID is not None else None
except ValueError as exc:
raise ValueError(
f"{file} matched the regex, but its fieldID, timeID, chanID "
f"or sliceID is not a whole number ({exc}). Nothing was "
f"converted.") from exc
region_key = (plateID, wellID, fieldID, timeID, chanID)
files_by_region.setdefault(region_key, []).append((file, sliceID))
source_wells = sorted({region[:2] for region in files_by_region},
key=lambda pair: (_cv.natural_key(pair[0]),
_cv.natural_key(pair[1])))
plate_tokens = {plate_key: f'plate{index}' for index, plate_key in enumerate(
sorted({plate_key for plate_key, _ in source_wells},
key=_cv.natural_key), start=1)}
canonical_wells = {}
for plate_key, well_key in source_wells:
canonical = _cv.normalise_well(well_key)
if canonical is not None:
canonical_wells[(plate_key, well_key)] = canonical
n_wells = _cv.plate_format_for_names(0, sorted(set(canonical_wells.values())))
for key, canonical in canonical_wells.items():
name = f'{plate_tokens[key[0]]}_{canonical}'
if name in used_wells:
continue
region_to_well[key] = name
used_wells.add(name)
for key in source_wells:
if key in region_to_well:
continue
region_to_well[key] = _next_synthetic_yokogawa_well(used_wells, n_wells)
print(f"Well {key[1]!r} is not a plate address; converted as "
f"{region_to_well[key]} (see {os.path.basename(csv_path)}).")
for region, file_list in files_by_region.items():
assigned_well = region_to_well[region[:2]]
plateID, wellID, fieldID, timeID, chanID = region
slice_ids = [sid for _, sid in file_list if sid is not None]
unique_slices = set(slice_ids)
original_files = ";".join(f[0] for f in file_list)
new_filename = f"{assigned_well}_T{timeID:04d}F{int(fieldID):03d}L01C{chanID:02d}.tif"
try:
images = []
for filename, _ in sorted(file_list, key=lambda x: x[1] or 1):
img = tifffile.imread(os.path.join(folder, filename))
images.append(img)
if len(unique_slices) > 1:
img_to_save = np.max(np.stack(images), axis=0)
else:
img_to_save = images[0]
dtype = img_to_save.dtype
new_filepath = os.path.join(folder, new_filename)
write_tiff(new_filepath, img_to_save.astype(dtype))
except Exception as exc:
raise ValueError(
f"{original_files} matched the regex but could not be "
f"converted into {new_filename}: {type(exc).__name__}: {exc}. "
f"{len(rename_log)} of {len(files_by_region)} converted "
f"file(s) had been written to {folder} before it stopped, "
f"and {os.path.basename(csv_path)} was not written.") from exc
rename_log.append({"Original File(s)": original_files, "Renamed TIFF": new_filename})
pd.DataFrame(rename_log).to_csv(csv_path, index=False)
print(f"Processing complete. Files saved in {folder} and rename log saved as {csv_path}.")
[docs]
def convert_to_yokogawa(folder):
"""Convert every image in ``folder`` to Yokogawa-style naming with a MIP.
ND2, CZI, LIF and plain TIFF/PNG/JPEG inputs are detected by
extension, max-projected over Z, and written out as
``plate<N>_<well>_T####F###L01C##.tif``. A ``rename_log.csv``
records the original-to-new mapping.
A file that cannot be read is skipped so the rest of the folder
still converts — but the skip is recorded on a
:class:`spacr.errors.RunLedger`, printed as a loud summary at the
end, and **stamped into a sibling** ``rename_log.run_status.json``.
That sidecar is what lets a later reader (or
:func:`spacr.errors.run_is_complete`) tell that the converted
folder is missing inputs, instead of quietly analysing a subset.
:param folder: Directory of raw images, converted in place.
:returns: the :class:`spacr.errors.RunLedger` for the conversion.
:raises ValueError: If the folder already contains Yokogawa-named
converted images or a previous ``rename_log.csv``. The check runs
before writing any image or log, including after the converted
images have been moved into ``orig/``. Read an already
converted folder with ``metadata_type='cellvoyager'``, or retry raw
conversion in a separate folder containing only the original inputs.
"""
files = sorted(_listdir_visible(folder))
converted_name = re.compile(
r"plate\d+_[A-Z]+\d+_T\d+F\d+L\d+(?:A\d+)?(?:Z\d+)?C\d+\.tiff?",
re.IGNORECASE,
)
for file in files:
if converted_name.fullmatch(file) or file == "rename_log.csv":
existing = ("a conversion log" if file == "rename_log.csv"
else "converted images")
raise ValueError(
f"{folder} already contains {existing}, including {file}. "
"Automatic conversion would risk overwriting images or changing "
"well assignments. Use metadata_type='cellvoyager' to read the "
"converted images, or convert the original inputs in a separate "
"folder. No images or rename log were changed."
)
def _get_next_well(used_wells):
"""Return the next free well, filling one plate before the next.
The well ids come from :func:`spacr.convert.well_sequence`, which
builds them out of :data:`spacr.schema.PLATE_FORMATS` — one
definition instead of the three copies of ``"ABCDEFGHIJKLMNOP"`` and
``range(1, 25)`` this module used to carry.
The plate format stays 384 here: unlike
:func:`convert_separate_files_to_yokogawa`, the inputs carry no well
names at all, so nothing in them can ask for a bigger plate and the
addresses are synthetic either way.
"""
return _next_synthetic_yokogawa_well(used_wells, 384)
rename_log = []
csv_path = os.path.join(folder, "rename_log.csv")
used_wells = set()
ledger = RunLedger('convert_to_yokogawa')
for file in files:
path = os.path.join(folder, file)
ext = file.lower().split('.')[-1]
well = _get_next_well(used_wells)
if ext == 'nd2':
with ledger.item(file, stage='nd2',
echo=f"Error processing ND2 file {file}"):
nd2 = ND2Reader(path)
metadata = nd2.metadata
timepoints = list(range(len(metadata.get("frames", [0])))) or [0]
fields = list(range(len(metadata.get("fields_of_view", [0])))) or [0]
z_levels = list(metadata.get("z_levels", range(1))) if metadata.get("z_levels") else [0]
channels = metadata.get("channels", [])
for t_idx in timepoints:
for f_idx in fields:
for c_idx, channel in enumerate(channels):
try:
mip_image = np.maximum.reduce([
nd2.get_frame_2D(t=t_idx, v=f_idx, z=z_idx, c=c_idx)
for z_idx in z_levels
], axis=0)
dtype = mip_image.dtype
filename = f"{well}_T{t_idx+1:04d}F{f_idx+1:03d}L01C{c_idx+1:02d}.tif"
filepath = os.path.join(folder, filename)
write_tiff(filepath, mip_image.astype(dtype))
rename_log.append({"Original File": file,
"Renamed TIFF": filename,
"ext": ext,
"time": t_idx,
"field": f_idx,
"channel": channel,
"z": z_levels})
except IndexError as frame_err:
ledger.record_failure(
f"{file}:T{t_idx}F{f_idx}C{c_idx}",
stage='nd2_frame', exc=frame_err)
print(f"Warning: ND2 file {file} has an incomplete data structure. Skipping.")
elif ext == 'czi':
with ledger.item(file, stage='czi',
echo=f"Error processing CZI file {file}"):
czi_reader = pyczi or _load_pylibczi()
with czi_reader.open_czi(path) as czidoc:
bbox = czidoc.total_bounding_box
_, tlen = bbox.get('T', (0,1))
_, clen = bbox.get('C', (0,1))
_, zlen = bbox.get('Z', (0,1))
scenes_bb = czidoc.scenes_bounding_rectangle
scenes = sorted(scenes_bb.keys()) if scenes_bb else [None]
folder = os.path.dirname(path)
for scene in scenes:
scene_well = _get_next_well(used_wells)
F_idx = scene + 1 if scene is not None else 1
A_idx = scene + 1 if scene is not None else 1
for t in range(tlen):
for c in range(clen):
for z in range(zlen):
arr = czidoc.read(
plane={'T': t, 'C': c, 'Z': z},
scene=scene
)
plane = np.squeeze(arr)
fn = (
f"{scene_well}_"
f"T{t+1:04d}"
f"F{F_idx:03d}"
f"L01"
f"A{A_idx:02d}"
f"Z{z+1:02d}"
f"C{c+1:02d}.tif"
)
outpath = os.path.join(folder, fn)
write_tiff(
outpath,
plane.astype(plane.dtype),
compression='zlib'
)
rename_log.append({
"Original File": file,
"Renamed TIFF": fn,
"ext": ext,
"scene": scene,
"time": t,
"slice": z,
"field": F_idx,
"channel": c,
"well": scene_well
})
elif ext == 'lif':
with ledger.item(file, stage='lif',
echo=f"Error processing LIF file {file}"):
lif_file = readlif.reader.LifFile(path)
for image_idx, image in enumerate(lif_file.get_iter_image()):
timepoints = range(getattr(image.dims, 't', 1))
z_levels = range(getattr(image.dims, 'z', 1))
channels = range(getattr(image, 'channels', 1) or 1)
for t_idx in timepoints:
for c_idx in channels:
z_stack = []
for z_idx in z_levels:
try:
frame = image.get_frame(z=z_idx, t=t_idx, c=c_idx)
z_stack.append(frame)
except IndexError as frame_err:
ledger.record_failure(
f"{file}:T{t_idx}Z{z_idx}C{c_idx}",
stage='lif_frame', exc=frame_err)
print(f"Missing frame: T{t_idx}, Z{z_idx}, C{c_idx} in {file}, skipping frame.")
if z_stack:
mip_image = np.max(np.stack(z_stack), axis=0)
dtype = mip_image.dtype
filename = f"{well}_T{t_idx+1:04d}F{image_idx+1:03d}L01C{c_idx+1:02d}.tif"
filepath = os.path.join(folder, filename)
write_tiff(filepath, mip_image.astype(dtype))
rename_log.append({"Original File": file, "Renamed TIFF": filename})
elif ext in ['tif', 'tiff', 'png', 'jpg', 'jpeg', 'bmp'] and not file.startswith("plate"):
with ledger.item(file, stage='tiff',
echo=f"Error processing standard image file {file}"):
with tifffile.TiffFile(path) as tif:
images = tif.asarray()
ndim = images.ndim
t_dim = c_dim = 1
if ndim == 2:
mip_image = images
filename = f"{well}_T0001F001L01C01.tif"
write_tiff(os.path.join(folder, filename), mip_image)
rename_log.append({"Original File": file, "Renamed TIFF": filename})
continue
elif ndim == 3:
if images.shape[0] <= 4:
c_dim = images.shape[0]
for c in range(c_dim):
mip_image = images[c, :, :]
filename = f"{well}_T0001F001L01C{c+1:02d}.tif"
write_tiff(
os.path.join(folder, filename), mip_image)
rename_log.append({"Original File": file, "Renamed TIFF": filename})
else:
mip_image = np.max(images, axis=0)
filename = f"{well}_T0001F001L01C01.tif"
write_tiff(
os.path.join(folder, filename), mip_image)
rename_log.append({"Original File": file, "Renamed TIFF": filename})
elif ndim == 4:
try:
axes = (tif.series[0].axes or '').upper()
except Exception:
axes = ''
t_axis = 0
if axes[:2] == 'ZT':
t_axis = 1
elif axes[:2] != 'TZ':
print(f"WARNING: {file} is 4-D but declares axes "
f"{axes or '(none)'}; reading it as "
f"(T, Z, Y, X). If it is really (Z, T, Y, X), "
f"every timepoint written below is a z-plane "
f"and every projection is over time.")
t_dim = images.shape[t_axis]
for t in range(t_dim):
plane_stack = images[t] if t_axis == 0 else images[:, t]
mip_image = np.max(plane_stack, axis=0)
filename = f"{well}_T{t+1:04d}F001L01C01.tif"
write_tiff(
os.path.join(folder, filename), mip_image)
rename_log.append({"Original File": file, "Renamed TIFF": filename})
else:
raise ValueError(f"Unsupported TIFF dimensions: {images.shape}")
pd.DataFrame(rename_log).to_csv(csv_path, index=False)
print(f"Processing complete. Files saved in {folder} and rename log saved as {csv_path}.")
ledger.finalize(artifact=csv_path)
return ledger
[docs]
def apply_augmentation(image, method):
"""Return ``image`` transformed by a named geometric augmentation.
:param image: NumPy image array.
:param method: One of ``'rotate90'``, ``'rotate180'``, ``'rotate270'``,
``'flip_h'``, ``'flip_v'``; any other value returns the input
unchanged.
:returns: Augmented image array.
"""
if method == 'rotate90':
return cv2.rotate(image, cv2.ROTATE_90_CLOCKWISE)
elif method == 'rotate180':
return cv2.rotate(image, cv2.ROTATE_180)
elif method == 'rotate270':
return cv2.rotate(image, cv2.ROTATE_90_COUNTERCLOCKWISE)
elif method == 'flip_h':
return cv2.flip(image, 1)
elif method == 'flip_v':
return cv2.flip(image, 0)
return image
[docs]
def process_instruction(entry):
"""Copy one image/mask pair described by ``entry``, applying an optional augmentation.
:param entry: Dict with keys ``src_img``, ``src_msk``, ``dst_img``,
``dst_msk`` and ``augment`` (augmentation name or falsy).
:returns: ``1`` on success — used for progress counting.
"""
img = tifffile.imread(entry["src_img"])
msk = tifffile.imread(entry["src_msk"])
if entry["augment"]:
img = apply_augmentation(img, entry["augment"])
msk = apply_augmentation(msk, entry["augment"])
write_tiff(entry["dst_img"], img)
write_tiff(entry["dst_msk"], msk)
return 1
[docs]
def prepare_cellpose_dataset(input_root, augment_data=False, train_fraction=0.8, n_jobs=None):
"""Aggregate image/mask pairs from sibling dataset folders into a Cellpose training split.
Discovers ``<input_root>/*/masks`` layouts, balances datasets to a
common size (with augmentations if requested) and copies the
selected pairs into ``<input_root>/cellpose_dataset/train`` and
``.../test``.
:param input_root: Directory containing one subfolder per dataset.
:param augment_data: If True, expand under-sized datasets by
applying geometric augmentations. Default ``False``.
:param train_fraction: Fraction of pairs routed to the train split.
Default ``0.8``.
:param n_jobs: Worker count for parallel copies. Default: CPU count.
:returns: None
:raises ValueError: if no valid ``<subdir>/masks`` datasets are found.
"""
from .utils import print_progress
time_ls = []
input_root = os.path.abspath(input_root)
output_root = os.path.join(input_root, "cellpose_dataset")
def get_augmentations():
"""Return the list of augmentation names used to expand datasets."""
return ['rotate90', 'rotate180', 'rotate270', 'flip_h', 'flip_v']
def find_image_mask_pairs(dataset_path):
"""Return ``(image_path, mask_path)`` pairs found under ``dataset_path``.
:param dataset_path: One dataset folder holding its images at the top
level and a ``masks`` subfolder in which each mask carries exactly
the same file name as its image. Only ``.tif``/``.tiff`` images
are considered, matching is by name rather than by order, and an
image whose mask is absent is dropped without a message — so a
short pair count here means missing or renamed masks.
"""
mask_dir = os.path.join(dataset_path, "masks")
pairs = []
for fname in os.listdir(dataset_path):
if fname.lower().endswith((".tif", ".tiff")):
img_path = os.path.join(dataset_path, fname)
msk_path = os.path.join(mask_dir, fname)
if os.path.isfile(msk_path):
pairs.append((img_path, msk_path))
return pairs
def prepare_output_folders(base):
"""Create ``train/{images,masks}`` and ``test/{images,masks}`` under ``base``.
:param base: Output root, which the caller sets to
``<input_root>/cellpose_dataset``. Creation is ``exist_ok``, so
rerunning does not fail, but nothing is emptied first and the
copies are renumbered from ``00000``, so a smaller second run
leaves the tail of a larger earlier one mixed into the split.
"""
for subset in ["train", "test"]:
os.makedirs(os.path.join(base, subset, "images"), exist_ok=True)
os.makedirs(os.path.join(base, subset, "masks"), exist_ok=True)
print("Scanning datasets...")
datasets = []
for subdir in os.listdir(input_root):
dataset_path = os.path.join(input_root, subdir)
if os.path.isdir(dataset_path) and os.path.isdir(os.path.join(dataset_path, "masks")):
pairs = find_image_mask_pairs(dataset_path)
if pairs:
datasets.append(pairs)
print(f" Found {len(pairs)} images in {dataset_path}")
if not datasets:
raise ValueError("No valid datasets with images and masks found.")
prepare_output_folders(output_root)
min_size = min(len(pairs) for pairs in datasets)
target_size = min_size if not augment_data else max(len(pairs) for pairs in datasets)
print("\nPreparing instruction list...")
instructions = []
global_index = 0
for pairs in datasets:
dataset_len = len(pairs)
sampled_pairs = []
if dataset_len >= target_size:
sampled_pairs = random.sample(pairs, target_size)
else:
sampled_pairs = pairs.copy()
needed = target_size - dataset_len
aug_methods = get_augmentations()
combos = [(img_path, msk_path, aug)
for aug in aug_methods
for (img_path, msk_path) in pairs]
pool = []
while len(pool) < needed:
round_ = combos[:]
random.shuffle(round_)
pool.extend(round_)
sampled_pairs.extend(pool[:needed])
augmented_sampled = [
(tup[0], tup[1], None) if len(tup) == 2 else tup
for tup in sampled_pairs
]
random.shuffle(augmented_sampled)
split_idx = int(train_fraction * len(augmented_sampled))
split_sets = {
"train": augmented_sampled[:split_idx],
"test": augmented_sampled[split_idx:]
}
for subset, items in split_sets.items():
for img_path, msk_path, aug in items:
dst_img = os.path.join(output_root, subset, "images", f"{global_index:05d}.tif")
dst_msk = os.path.join(output_root, subset, "masks", f"{global_index:05d}.tif")
instructions.append({
"src_img": img_path,
"src_msk": msk_path,
"dst_img": dst_img,
"dst_msk": dst_msk,
"augment": aug
})
global_index += 1
print(f"Total files to process: {len(instructions)}")
print("Processing images with multiprocessing...")
if n_jobs is None:
n_jobs = max(1, cpu_count() - 1)
else:
n_jobs = int(n_jobs)
if instructions:
from .resource_log import _array_file_nbytes, _guard_workers
n_jobs = _guard_workers('cellpose_dataset', n_jobs,
_array_file_nbytes(instructions[0]["src_img"]))
with Pool(n_jobs) as pool:
for i, _ in enumerate(pool.imap_unordered(process_instruction, instructions), 1):
print_progress(i, len(instructions), n_jobs=n_jobs, time_ls=time_ls, batch_size=None, operation_type="cellpose dataset")
print(f"Done. Dataset saved to: {output_root}")