Source code for spacr.normalization

"""Control crop storage precision and model-input normalization.

``.npy`` crops keep whatever
the merged stack held, PNG crops narrow once through
:func:`spacr.crops.narrow_to_uint8` -- the high byte of a uint16, a linear
rescale rather than a clip -- and measurements are taken from the
full-precision array. So the structure the request asks for is the structure:
original precision wherever it can be, one declared narrowing at the only
boundary that requires one.

Storage dtype and model input scaling are separate choices.
``transforms.ToTensor()`` divides by 255 and
hands the model a float in [0, 1] whatever the file held, so "scale to 0-1"
is already what happens at the point it matters. The dtype on disk decides
file size and what other tools can open the crop -- a storage decision.

The subsequent normalization setting changes the model input. spaCR has
historically normalized with

    mean = std = (0.5, 0.5, 0.5)

which maps [0, 1] onto [-1, 1]. Every ImageNet-pretrained torchvision model
was fitted on

    mean = (0.485, 0.456, 0.406)   std = (0.229, 0.224, 0.225)

so a finetune starts by handing pretrained weights inputs distributed
differently from the ones they learned. A long finetune adapts; a short one,
or a frozen backbone, pays for it. That was a literal in two places and is
now a choice.

The default remains ``symmetric`` for compatibility. This is what spaCR has
always done, and switching it silently would move every existing model's
scores with nothing in the artifact to say why.
"""
from __future__ import annotations

import logging
from itertools import islice
from typing import Any, Dict, Optional, Sequence, Tuple

import numpy as np

LOG = logging.getLogger("spacr.normalization")

__all__ = [
    "CROP_DTYPES",
    "IMAGENET_MEAN",
    "IMAGENET_STD",
    "NORMALIZATIONS",
    "CLIP_MEAN",
    "CLIP_STD",
    "INCEPTION_MEAN",
    "INCEPTION_STD",
    "apply_crop_dtype",
    "dataset_statistics",
    "normalization_stats",
    "describe_normalization",
]

#: What a crop file may hold. ``original`` is what the pipeline produced --
#: uint16 for a 16-bit camera -- and is the default because it is the only
#: one that loses nothing.
CROP_DTYPES: Tuple[str, ...] = ("original", "uint8", "uint16")

#: The statistics every ImageNet-pretrained torchvision model was fitted on.
IMAGENET_MEAN: Tuple[float, float, float] = (0.485, 0.456, 0.406)
IMAGENET_STD: Tuple[float, float, float] = (0.229, 0.224, 0.225)

#: OpenAI CLIP's, for a CLIP or OpenCLIP backbone. Close to ImageNet's and
#: NOT the same; a CLIP model fed ImageNet statistics is being handed inputs
#: half a standard deviation off on the blue channel.
CLIP_MEAN: Tuple[float, float, float] = (0.48145466, 0.4578275, 0.40821073)
CLIP_STD: Tuple[float, float, float] = (0.26862954, 0.26130258, 0.27577711)

#: Inception / TF-slim's, and what spaCR's historic 0.5/0.5 actually is.
#: Named so a user recognises it rather than having to recognise the numbers.
INCEPTION_MEAN: Tuple[float, float, float] = (0.5, 0.5, 0.5)
INCEPTION_STD: Tuple[float, float, float] = (0.5, 0.5, 0.5)

#: How the loader normalises after ``ToTensor()`` has already produced [0, 1].
#:
#: ``symmetric``  what spaCR has always done: [0, 1] -> [-1, 1]. The same
#:                thing Inception and TF-slim call their preprocessing.
#: ``imagenet``   what a torchvision pretrained backbone expects.
#: ``clip``       what a CLIP / OpenCLIP backbone expects. Close to
#:                ImageNet's and not the same.
#: ``dataset``    the mean and standard deviation of THIS dataset, per
#:                channel. See :func:`dataset_statistics` -- for fluorescence
#:                this is the one with an argument behind it.
#: ``custom``     numbers the user supplies, for a backbone whose
#:                preprocessing is none of the above.
#: ``none``       leave it in [0, 1]. For training from scratch, where there
#:                are no pretrained statistics to match and centring is the
#:                optimiser's problem rather than the data's.
NORMALIZATIONS: Tuple[str, ...] = (
    "symmetric", "imagenet", "clip", "dataset", "custom", "none")


[docs] def normalization_stats(mode: Any, *, mean: Optional[Sequence[float]] = None, std: Optional[Sequence[float]] = None, channels: int = 3 ) -> Optional[Tuple[Tuple[float, ...], Tuple[float, ...]]]: """``(mean, std)`` for ``mode``, or None when nothing should be applied. :param mode: one of :data:`NORMALIZATIONS`. Anything unrecognised falls back to ``symmetric`` with a log line -- a typo must not silently train a model on statistics nobody chose, but it must not stop a run either, and ``symmetric`` is what the run would have used before this setting existed. :param mean: for ``custom`` and ``dataset``, the per-channel means. One value is broadcast to every channel, which is what a single-stain dataset wants. :param std: as ``mean``. A zero is replaced by 1.0 rather than dividing by it -- a channel with no variance is a constant channel, and dividing it by its own zero spread produces inf and then a loss of nan, several minutes into training, with nothing saying why. :param channels: how many planes the model will see. Only used to broadcast a single supplied value. """ name = str(mode or "symmetric").strip().lower() if name not in NORMALIZATIONS: LOG.info("input_statistics %r is not one of %s; using symmetric, " "which is what spaCR did before this setting existed", mode, list(NORMALIZATIONS)) name = "symmetric" if name == "none": return None if name == "imagenet": return (IMAGENET_MEAN, IMAGENET_STD) if name == "clip": return (CLIP_MEAN, CLIP_STD) if name in ("custom", "dataset"): if mean is None or std is None: raise ValueError( f"input_statistics={name!r} needs both mean and std. " f"For 'dataset', compute them with " f"spacr.normalization.dataset_statistics; for 'custom', " f"supply the numbers your backbone was trained with.") return (_broadcast(mean, channels), _clean_std(std, channels)) return (INCEPTION_MEAN, INCEPTION_STD)
def _broadcast(values: Sequence[float], channels: int) -> Tuple[float, ...]: """Convert normalization values to floats and expand a singleton. :param values: Per-channel values, or one value shared by all channels. :param channels: Requested channel count for singleton expansion. :returns: Float tuple, with one input repeated to at least one channel; multi-value inputs retain their original length. """ out = [float(v) for v in values] if len(out) == 1: return tuple(out * max(1, int(channels))) return tuple(out) #: Below this, a channel is constant. NOT ``== 0``: the sum-of-squares #: identity in :func:`dataset_statistics` leaves a constant channel at about #: 7e-09 rather than at zero through floating-point cancellation, and #: dividing by 7e-09 multiplies that channel by 1.3e8 -- the same disaster as #: dividing by zero, arrived at by a route an equality check does not catch. #: On [0, 1] data nothing real has a spread this small. CONSTANT_CHANNEL_STD = 1e-6 def _clean_std(values: Sequence[float], channels: int) -> Tuple[float, ...]: """Per-channel spreads, with a constant channel normalised by one. A constant channel has no spread, and dividing by it produces inf or an enormous number, then a nan loss several minutes into training with nothing saying why. A spread of 1 leaves that channel alone, which is the only sensible thing to do with a channel that carries no variation. """ out = [] for value in _broadcast(values, channels): number = abs(float(value)) if not np.isfinite(number) or number < CONSTANT_CHANNEL_STD: LOG.info("a channel has no spread (%.3g); normalising it by 1.0 " "rather than multiplying it by %.3g", number, 1.0 / number if number else float("inf")) number = 1.0 out.append(number) return tuple(out)
[docs] def dataset_statistics(loader: Any, *, max_batches: Optional[int] = None ) -> Tuple[Tuple[float, ...], Tuple[float, ...]]: """Return the dataset's per-channel mean and standard deviation. Dataset-specific statistics are appropriate when fluorescence channels differ substantially from the natural-image distribution represented by ImageNet statistics. Fluorescence channels may contain independently exposed stains and a large background fraction. Computed in one streaming pass with the sum-of-squares identity, so a dataset that does not fit in memory still yields exact statistics rather than statistics of whatever fitted. :param loader: anything iterable yielding ``(images, ...)`` batches, or bare image tensors, shaped ``(N, C, H, W)`` on [0, 1]. :param max_batches: stop after this many. The mean of a plate converges in far fewer batches than the plate has, and a full pass over a million crops to compute two numbers per channel is a cost with no matching gain. ``None`` reads everything. :returns: ``(mean, std)``, per channel. :raises ValueError: the loader yielded nothing, so the answer would be the statistics of an empty set rather than of this dataset. """ total = None total_sq = None count = 0 source = loader if max_batches is None else islice(loader, max_batches) for batch in source: images = batch[0] if isinstance(batch, (tuple, list)) else batch values = np.asarray(getattr(images, "numpy", lambda: images)(), dtype=np.float64) if values.ndim != 4: continue flat = values.transpose(1, 0, 2, 3).reshape(values.shape[1], -1) if total is None: total = flat.sum(axis=1) total_sq = (flat ** 2).sum(axis=1) else: total += flat.sum(axis=1) total_sq += (flat ** 2).sum(axis=1) count += flat.shape[1] if not count or total is None: raise ValueError( "the loader yielded no images, so these would be the statistics " "of an empty set rather than of this dataset") mean = total / count variance = np.maximum(total_sq / count - mean ** 2, 0.0) return (tuple(float(v) for v in mean), _clean_std(tuple(float(v) for v in np.sqrt(variance)), len(mean)))
[docs] def describe_normalization(mode: Any, **kwargs) -> str: """One line for the log, so a model card records what it was trained on. :param mode: normalization preset or explicit normalization mode. """ try: stats = normalization_stats(mode, **kwargs) except ValueError as exc: return str(exc) if stats is None: return "inputs left in [0, 1]; no mean/std normalisation" mean, std = stats known = {IMAGENET_MEAN: "ImageNet", CLIP_MEAN: "CLIP", INCEPTION_MEAN: "Inception/TF-slim"} name = known.get(tuple(mean)) if name == "Inception/TF-slim": return (f"inputs normalised with mean={mean}, std={std} " f"({name}), mapping [0, 1] to [-1, 1]") if name: return (f"inputs normalised with the {name} statistics " f"mean={mean}, std={std} -- what that pretrained backbone " f"expects") return (f"inputs normalised with mean={mean}, std={std}, measured from " f"this dataset or supplied by hand")
[docs] def apply_crop_dtype(array: np.ndarray, dtype: Any = "original") -> np.ndarray: """Return ``array`` in the dtype a crop file should hold. :param array: the crop as the pipeline produced it. :param dtype: one of :data:`CROP_DTYPES`. :returns: the array, unchanged for ``original``. Conversion to ``uint8`` delegates to :func:`spacr.crops.narrow_to_uint8`, which applies the project's declared 16-to-8-bit linear mapping. Conversion to ``uint16`` is a cast rather than an intensity stretch; values from an 8-bit input remain unchanged. """ name = str(dtype or "original").strip().lower() if name not in CROP_DTYPES: LOG.info("crop_dtype %r is not one of %s; keeping the original", dtype, list(CROP_DTYPES)) return array if name == "original": return array array = np.asarray(array) if name == "uint8": if array.dtype == np.uint8: return array from .crops import narrow_to_uint8 return narrow_to_uint8(array) if array.dtype == np.uint16: return array if np.issubdtype(array.dtype, np.floating): from .crops import narrow_to_uint8 return narrow_to_uint8(array).astype(np.uint16) return array.astype(np.uint16)