"""Say what each object in a mask is, after any segmenter has made it.
Segmentation answers "where is there an object"; this answers
"what is it". The two are kept apart on purpose: the same classification
must be available whether spaCR segmented the field a moment ago or is
handed a mask Cellpose wrote last week.
TWO WAYS IN:
* :func:`segment_and_classify` takes a segmentation backend and its
options, segments, and classifies what it made.
* :func:`classify_objects` takes an image and a mask that already exists
and classifies it as it stands. Nothing is re-segmented.
WHAT A HEAD IS. A head is a trained classifier: it is handed a batch of
crops and returns a class and a probability for each. It is a protocol
rather than a base class, so an adapter around a torch module or a
scikit-learn estimator can supply it. Two applications are planned -- parasites per
vacuole, and what each object is -- and the PV one is three heads in
practice, trained on Hoechst, the parasite stain and CellMask separately,
so a plate counts from whichever stain it has.
Geometry reports border contact and suggests pairs that may have been split
by a segmenter. A shared boundary is a heuristic, not proof that two labels
belong to one biological object. Border contact and split candidates are
reported separately; automatic merging excludes either object touching the
field edge. Merge accuracy still requires independently labelled examples.
IDS ARE NEVER RENUMBERED. Everything here keys on the label ids the mask
arrived with, and every mask it returns keeps them. `canonical_labels` in
`mask_engine` says what renumbering costs: the measurements, the crops and
the tracks are all keyed on those ids.
"""
from __future__ import annotations
import logging
import math
import operator
import os
from typing import Any, Dict, Iterable, List, Mapping, Optional, Sequence, Tuple
import numpy as np
from .crop_source import crop_object
LOG = logging.getLogger(__name__)
#: The classes the object head answers with. The two "cut" classes are
#: separate from their whole counterparts because a cut object is not a
#: measurement anybody should use, and separate from each other because
#: only one of them can be repaired here.
OBJECT_CLASSES: Tuple[str, ...] = (
"cell", "nucleus", "cut cell", "cut nucleus", "artifact",
)
#: How the parasite count is reported once a head has predicted it. The
#: head predicts a COUNT; these are the bins the count is shown in. A
#: vacuole of three stays a three in the model and is binned only for the
#: report.
PV_BINS: Tuple[int, ...] = (1, 2, 4, 8, 16)
[docs]
def bin_parasite_count(count: float) -> str:
"""Put a predicted parasite count in the bin a report shows it in.
THE BIN IS A DISPLAY CHOICE, NOT A TRAINED ONE. The head predicts how
many parasites it sees; this is only how that number is shown. That
matters because the rule below can change -- or be replaced by a different one
per report -- without retraining anything.
TIES GO DOWN, and 12 is the case: it is four from 8 and four from 16.
A vacuole of 12 has finished eight and is partway to sixteen, so it is
reported as the round it has completed rather than the one it has
not.
:param count: finite, nonnegative count, including fractional predictions.
:returns: one of ``"1"``, ``"2"``, ``"4"``, ``"8"``, ``"16"`` or
``">16"``.
:raises ValueError: the count is negative or not finite.
"""
value = float(count)
if not math.isfinite(value) or value < 0:
raise ValueError("count must be finite and nonnegative")
if value > PV_BINS[-1]:
return f">{PV_BINS[-1]}"
nearest = min(PV_BINS, key=lambda edge: (abs(edge - value), edge))
return str(nearest)
[docs]
class ParasiteCountHead:
"""Adapt a count regressor to the object-head prediction interface.
:param regressor: model with ``predict(crops)`` returning one finite,
nonnegative count per crop, shaped ``(N,)`` or ``(N, 1)``. The
caller selects this model's stain through ``classify_objects``'s
``channels`` argument; separate models can use the same labels.
Counts retain their regression values. Display classes use
:func:`bin_parasite_count`; a regressor supplies no class probability.
"""
def __init__(self, regressor: Any):
"""Retain the stain-specific count regressor without loading or fitting it."""
self.regressor = regressor
[docs]
def predict(self, crops: Sequence[np.ndarray]) -> List[Dict[str, Any]]:
"""Return unrounded counts and display classes for a crop batch.
:param crops: channel-selected crops in object order.
:returns: one count, display class and absent probability per crop.
:raises ValueError: predictions have the wrong shape or invalid counts.
"""
values = np.asarray(self.regressor.predict(crops), dtype=float)
if values.shape == (len(crops), 1):
values = values[:, 0]
if values.shape != (len(crops),):
raise ValueError("count predictions must have one value per crop")
return [{"count": float(value), "class": bin_parasite_count(value),
"probability": None} for value in values]
def _label_mask(mask: np.ndarray) -> np.ndarray:
"""Validate a nonempty 2-D image of nonnegative integer labels."""
field = np.asarray(mask)
if field.ndim != 2 or not all(field.shape):
raise ValueError("mask must be a nonempty 2-D label image")
if field.dtype.kind not in "iu" or np.any(field < 0):
raise ValueError("mask labels must be nonnegative integers")
return field
def _labels_of(mask: np.ndarray) -> np.ndarray:
"""Every object id in ``mask``, background excluded."""
ids = np.unique(np.asarray(mask))
return ids[ids != 0]
[docs]
def touches_border(mask: np.ndarray, label: int) -> bool:
"""Whether object ``label`` runs into the edge of the field.
:param mask: the label image.
:param label: the object.
:returns: True when any of its pixels is in the first or last row or
column, which is what makes it a cut object nothing here can mend.
"""
field = np.asarray(mask)
if field.size == 0:
return False
edges = (field[0, :], field[-1, :], field[:, 0], field[:, -1])
return any(bool(np.any(edge == label)) for edge in edges)
def _perimeter_pixels(field: np.ndarray, label: int) -> int:
"""How many of ``label``'s pixels sit against something that is not it."""
here = field == label
if not here.any():
return 0
padded = np.pad(here, 1, constant_values=False)
neighbours = (padded[:-2, 1:-1] & padded[2:, 1:-1]
& padded[1:-1, :-2] & padded[1:-1, 2:])
return int(np.count_nonzero(here & ~neighbours))
def _shared_border(field: np.ndarray, first: int, second: int) -> int:
"""How many pixels of ``first`` are orthogonally against ``second``."""
here = field == first
there = field == second
if not here.any() or not there.any():
return 0
padded = np.pad(there, 1, constant_values=False)
against = (padded[:-2, 1:-1] | padded[2:, 1:-1]
| padded[1:-1, :-2] | padded[1:-1, 2:])
return int(np.count_nonzero(here & against))
[docs]
def split_candidates(mask: np.ndarray,
share: float = 0.25) -> List[Tuple[int, int]]:
"""Pairs of labels that look like one object cut in two.
The rule is a geometry heuristic: two labels are a
candidate when they touch, and the boundary they share is at least
``share`` of the smaller one's perimeter. A cell lying against another
cell shares a little of its edge; a cell the segmenter cut down the
middle shares most of it.
This finds candidates. It does not merge them -- :func:`merge_halves`
does, and only when asked, because a wrong merge destroys two objects
to make one.
:param mask: the label image.
:param share: fraction of the smaller perimeter that must be shared,
between zero and one. Shared pixels are counted on the lower-id
label, with orthogonal adjacency. The default 0.25 has synthetic
checks but no measured biological merge precision.
:returns: pairs ``(low id, high id)``, each pair once.
"""
field = _label_mask(mask)
if not math.isfinite(share) or not 0 <= share <= 1:
raise ValueError("share must be finite and between zero and one")
padded = np.pad(field, 1)
boundary = ((field != padded[:-2, 1:-1])
| (field != padded[2:, 1:-1])
| (field != padded[1:-1, :-2])
| (field != padded[1:-1, 2:])) & (field != 0)
ids, counts = np.unique(field[boundary], return_counts=True)
perimeters = dict(zip(map(int, ids), map(int, counts)))
contacts = []
for first, second, axis in ((field[:-1], field[1:], 0),
(field[:, :-1], field[:, 1:], 1)):
rows, cols = np.nonzero((first != second) & (first != 0) & (second != 0))
if not rows.size:
continue
left, right = first[rows, cols], second[rows, cols]
first_pixel = (rows.astype(np.uint64) * field.shape[1]
+ cols.astype(np.uint64))
second_pixel = first_pixel + (field.shape[1] if axis == 0 else 1)
contacts.append(np.column_stack((
np.minimum(left, right).astype(np.uint64),
np.maximum(left, right).astype(np.uint64),
np.where(left < right, first_pixel, second_pixel))))
if not contacts:
return []
unique_pixels = np.unique(np.concatenate(contacts), axis=0)
pairs, shared_counts = np.unique(unique_pixels[:, :2], axis=0,
return_counts=True)
return [(int(first), int(second))
for (first, second), shared in zip(pairs, shared_counts)
if shared / min(perimeters[int(first)], perimeters[int(second)])
>= share]
[docs]
def merge_halves(mask: np.ndarray,
pairs: Iterable[Tuple[int, int]]) -> np.ndarray:
"""Join each pair into one object, keeping the lower id.
The lower id survives so that a merge is stable however the pairs are
ordered, and so the object keeps the id that measurements and tracks
already know it by wherever that is the lower one.
:param mask: the label image.
:param pairs: pairs of ids, as :func:`split_candidates` returns them.
:returns: a NEW mask; the one passed in is not touched.
:raises ValueError: a pair names background, an absent id, or a noninteger.
"""
field = _label_mask(mask).copy()
survivor = {int(label): int(label) for label in _labels_of(field)}
def root(label):
"""Resolve a merge group and compress the path to its lowest id."""
while survivor[label] != label:
survivor[label] = survivor[survivor[label]]
label = survivor[label]
return label
for first, second in pairs:
try:
first, second = operator.index(first), operator.index(second)
except TypeError as exc:
raise ValueError("merge pairs must contain integer label ids") from exc
if first not in survivor or second not in survivor:
raise ValueError("merge pairs must name existing nonzero labels")
low, high = sorted((root(first), root(second)))
survivor[high] = low
for label in survivor:
kept = root(label)
if label != kept:
field[np.asarray(mask) == label] = kept
return field
[docs]
def remove_labels(mask: np.ndarray, labels: Iterable[int]) -> np.ndarray:
"""Take objects out of a mask without renumbering what is left.
:param mask: the label image.
:param labels: the ids to erase.
:returns: a NEW mask with those ids set to background.
"""
field = np.asarray(mask).copy()
for label in labels:
field[field == int(label)] = 0
return field
[docs]
def describe_objects(mask: np.ndarray) -> List[Dict[str, Any]]:
"""What can be said about every object without a trained head at all.
:param mask: the label image.
:returns: one row per object -- ``label``, ``area``, ``touches_border``
and ``split_with`` (the ids it may be one object with).
"""
field = _label_mask(mask)
candidates = split_candidates(field)
partners: Dict[int, List[int]] = {}
for first, second in candidates:
partners.setdefault(first, []).append(second)
partners.setdefault(second, []).append(first)
rows = []
for label in _labels_of(field):
label = int(label)
rows.append({
"label": label,
"area": int(np.count_nonzero(field == label)),
"touches_border": touches_border(field, label),
"split_with": sorted(partners.get(label, [])),
})
return rows
def _crops_for(image: np.ndarray, mask: np.ndarray, labels: Sequence[int], *,
channels: Sequence[int], size: Optional[int],
shape: str, padding: int) -> List[Optional[np.ndarray]]:
"""Cut every object out, in the order ``labels`` gives them."""
array = np.asarray(image)
if array.ndim == 2:
array = array[:, :, None]
return [crop_object(array, np.asarray(mask), int(label),
channels=channels, shape=shape, size=size,
padding=padding)
for label in labels]
[docs]
def classify_objects(image: np.ndarray, mask: np.ndarray, *,
head: Optional[Any] = None,
channels: Optional[Sequence[int]] = None,
size: Optional[int] = None,
shape: str = "bounding_box",
padding: int = 0,
remove_artifacts: bool = False,
merge_split: bool = False,
artifact_class: str = "artifact") -> Dict[str, Any]:
"""Classify every object of a mask that already exists.
:param image: the field, ``(H, W)`` or ``(H, W, planes)``.
:param mask: its labels. Not modified; any mask returned is a new one.
:param head: something with ``predict(crops)`` returning a class per
crop, ``(class, probability)`` pairs, or dictionaries containing
``class``, optional ``probability`` and optional regression ``count``.
None reports geometry and crops without making predictions.
:param channels: which planes the head sees, in order. Defaults to
every plane the image has.
:param size: resize each crop to ``size x size`` for the head.
:param shape: ``bounding_box`` or ``object``; see
:func:`spacr.crop_source.crop_object`.
:param padding: pixels of context around each object.
:param remove_artifacts: return a mask with whatever the head called
an artifact erased. OFF by default: a false artifact is a deleted
object, so this is a decision the caller makes rather than a
default they inherit.
:param merge_split: merge pairs :func:`split_candidates` finds, excluding
any pair with either object touching the field edge. This optional
geometry rule can mistake adjacent objects for a segmentation split.
:param artifact_class: which class name means "not a real object", for
``remove_artifacts``. A parameter rather than a constant because a
head trained on somebody else's labels may spell it differently,
and a removal that silently matches nothing is worse than one that
cannot be configured.
:returns: dictionary containing ``objects`` (per-label rows), ``crops``
(label-keyed arrays), a new ``mask``, ``merged`` pairs, ``removed``
ids and ``removed_count``. Rows and crops describe objects after
merging but before removal, retaining evidence for deleted objects.
``label_map`` maps every original id to its final surviving id,
or zero for a removed object. Default mask values and dtype are
unchanged. Geometry candidates are not validated biological splits.
:raises ValueError: invalid image/mask geometry, crop settings or head
output, or artifact removal requested without a head.
"""
source = _label_mask(mask)
array = np.asarray(image)
if (array.ndim not in (2, 3) or array.shape[:2] != source.shape
or (array.ndim == 3 and array.shape[2] == 0)):
raise ValueError("image shape must match the 2-D mask, with channels last")
planes = array.shape[2] if array.ndim == 3 else 1
try:
wanted = ([operator.index(channel) for channel in channels]
if channels is not None else list(range(planes)))
padding = operator.index(padding)
size = operator.index(size) if size is not None else None
except TypeError as exc:
raise ValueError("channels, padding and size must be integers") from exc
if not wanted or any(channel < 0 or channel >= planes for channel in wanted):
raise ValueError("channels must select existing image planes")
if padding < 0 or (size is not None and size <= 0):
raise ValueError("padding must be nonnegative and size must be positive")
if shape not in ("bounding_box", "object"):
raise ValueError("crop shape must be bounding_box or object")
if remove_artifacts and head is None:
raise ValueError("artifact removal requires a classification head")
field = source.copy()
merged: List[Tuple[int, int]] = []
if merge_split:
merged = [pair for pair in split_candidates(field)
if not (touches_border(field, pair[0])
or touches_border(field, pair[1]))]
field = merge_halves(field, merged)
rows = describe_objects(field)
labels = [row["label"] for row in rows]
crops = _crops_for(array, field, labels, channels=wanted, size=size,
shape=shape, padding=padding)
if head is not None and labels:
for row, verdict in zip(rows, _predictions(head, crops)):
row.update(verdict)
removed: List[int] = []
if remove_artifacts:
removed = [row["label"] for row in rows
if row.get("class") == artifact_class]
if removed:
field = remove_labels(field, removed)
original_ids, first_pixels = np.unique(source, return_index=True)
label_map = {int(label): int(field.flat[index])
for label, index in zip(original_ids, first_pixels) if label}
return {"objects": rows, "crops": dict(zip(labels, crops)),
"mask": field, "merged": merged, "removed": removed,
"removed_count": len(removed), "label_map": label_map}
def _predictions(head: Any, crops: Sequence[Optional[np.ndarray]]):
"""One ``{"class": ..., "probability": ...}`` per crop.
An object the cropper could not cut out -- it is not in the mask any
more, which a merge can do -- is reported with no class rather than
being dropped, so the rows still line up with the objects.
"""
usable = [crop for crop in crops if crop is not None]
answers = list(head.predict(usable)) if usable else []
if len(answers) != len(usable):
raise ValueError(
f"the head answered {len(answers)} of {len(usable)} crops; a "
f"classification has to line up with the objects it is about")
answered = iter(answers)
for crop in crops:
if crop is None:
yield {"class": None, "probability": None}
continue
answer = next(answered)
if isinstance(answer, Mapping):
if "class" not in answer:
raise ValueError("a head prediction must contain a class")
verdict = {"class": answer["class"],
"probability": answer.get("probability")}
if "count" in answer:
count = float(answer["count"])
verdict["count"] = count
verdict["class"] = bin_parasite_count(count)
elif isinstance(answer, tuple) and len(answer) == 2:
verdict = {"class": answer[0], "probability": answer[1]}
else:
verdict = {"class": answer, "probability": None}
if verdict["probability"] is not None:
probability = float(verdict["probability"])
if not math.isfinite(probability) or not 0 <= probability <= 1:
raise ValueError("probability must be finite and between 0 and 1")
verdict["probability"] = probability
yield verdict
[docs]
def segment_and_classify(backend: Any, image: np.ndarray, *,
eval_kwargs: Optional[Mapping[str, Any]] = None,
**classify_kwargs) -> Dict[str, Any]:
"""Segment a field with any backend, then classify what it found.
The backend is anything with Cellpose's ``eval`` -- Cellpose-SAM,
Cellpose 3, DINOCell and SAMCell all answer it. This function therefore wraps a model
without knowing which model it has.
:param backend: the segmenter.
:param image: one field.
:param eval_kwargs: passed to the backend's ``eval``.
:param classify_kwargs: passed to :func:`classify_objects`.
:returns: what :func:`classify_objects` returns, with the mask the
backend produced under ``"mask"`` when no option changed it.
"""
output = backend.eval(x=[np.asarray(image)], **dict(eval_kwargs or {}))
masks = output[0] if isinstance(output, tuple) else output
if isinstance(masks, (list, tuple)):
if len(masks) != 1:
raise ValueError("the backend must return one mask for one image")
masks = masks[0]
mask = np.asarray(masks)
if mask.ndim == 3 and mask.shape[0] == 1:
mask = mask[0]
if mask.ndim != 2:
raise ValueError("the backend must return one mask for one image")
return classify_objects(image, mask, **classify_kwargs)
_REAL_SIDE = 24
_REAL_ROLES: Tuple[str, ...] = ("cell", "pathogen", "nucleus")
_REAL_PADDING: Dict[str, float] = {"pathogen": 1.5, "nucleus": 0.5, "cell": 0.3}
_REAL_TYPES: Tuple[str, ...] = ("cell", "nucleus", "pathogen")
def _real_features(crop: np.ndarray) -> np.ndarray:
"""Describe one crop for the real / not-real classifier.
Each channel is scaled between its own 1st and 99th percentiles, so an
8-bit annotation crop and a 16-bit field crop of the same object give
the same features. The scaled crop is shrunk to a small square and
joined with the per-channel mean, spread and bright fraction.
"""
from skimage.transform import resize
array = np.asarray(crop, dtype=np.float32)
if array.ndim == 2:
array = array[:, :, None]
low, high = np.percentile(array, (1, 99), axis=(0, 1))
scaled = np.clip((array - low) / np.maximum(high - low, 1e-6), 0, 1)
small = resize(scaled, (_REAL_SIDE, _REAL_SIDE, scaled.shape[2]),
order=1, anti_aliasing=True, preserve_range=True)
stats = np.concatenate([scaled.mean(axis=(0, 1)), scaled.std(axis=(0, 1)),
(scaled > 0.5).mean(axis=(0, 1))])
return np.concatenate([small.ravel(), stats]).astype(np.float32)
def _read_real_crop(path: str) -> np.ndarray:
"""An annotation crop as ``(H, W, 3)``, CellMask / parasite / Hoechst."""
from PIL import Image
with Image.open(path) as image:
array = np.asarray(image.convert("RGB"))
return array
class _RealObjectHead:
"""A trained real / not-real classifier answering the head protocol.
Each crop gets ``real`` when the probability of being real reaches
``threshold`` and ``not real`` otherwise, with the probability of the
class given.
"""
def __init__(self, bundle: Mapping[str, Any], threshold: float = 0.5):
"""Keep the bundle's estimator and the not-real probability threshold."""
self.estimator = bundle["estimator"]
self.threshold = float(threshold)
if not 0 <= self.threshold <= 1:
raise ValueError("real_object_threshold must be between 0 and 1")
def real_probability(self, crops: Sequence[np.ndarray]) -> np.ndarray:
"""Probability that each crop shows a real object."""
features = np.stack([_real_features(crop) for crop in crops])
classes = list(self.estimator.classes_)
return self.estimator.predict_proba(features)[:, classes.index(1)]
def predict(self, crops: Sequence[np.ndarray]) -> List[Dict[str, Any]]:
"""``{"class", "probability"}`` per crop."""
if not len(crops):
return []
answers = []
for probability in self.real_probability(crops):
real = probability >= self.threshold
answers.append({"class": "real" if real else "not real",
"probability": float(probability if real
else 1 - probability)})
return answers
def _annotated_real_crops(databases: Sequence[str], column: str = "real"):
"""Annotated crops from Annotate's ``png_list`` tables.
:param databases: ``measurements.db`` files whose ``png_list`` has the
annotation column, 1 meaning real and 2 not real.
:param column: the annotation column.
:returns: a frame with ``png_path``, ``real`` (True or False) and
``group``, the plate and well each crop came from, so a split by
group never puts one well on both sides.
:raises ValueError: a database without the column.
"""
import pandas as pd
from .tabular import read_table
frames = []
for database in databases:
frame = read_table(database, table="png_list", report=None)
if column not in frame.columns:
raise ValueError(f"{database} has no {column!r} column in "
f"png_list; annotate it in Annotate first")
calls = pd.to_numeric(frame[column], errors="coerce")
frame = frame.loc[calls.isin([1, 2])].copy()
frame["real"] = calls.loc[frame.index] == 1
names = frame["png_path"].map(lambda p: os.path.splitext(
os.path.basename(str(p)))[0].rsplit("_", 3))
plates = (frame["dataset_plate"].astype(str)
if "dataset_plate" in frame.columns
else names.map(lambda parts: parts[0]))
wells = names.map(lambda parts: parts[1] if len(parts) > 1 else "")
frame["group"] = plates + "|" + wells
frames.append(frame[["png_path", "real", "group"]])
if not frames:
return pd.DataFrame(columns=["png_path", "real", "group"])
return pd.concat(frames, ignore_index=True)
def _real_scorecard(truth: np.ndarray, probability: np.ndarray,
threshold: float) -> Dict[str, Any]:
"""Held-out scores, with "not real" as the class being detected."""
from sklearn.metrics import (balanced_accuracy_score, precision_score,
recall_score, roc_auc_score)
truth = np.asarray(truth, dtype=bool)
called_real = np.asarray(probability) >= threshold
both = truth.any() and (~truth).any()
return {
"n": int(truth.size), "n_not_real": int((~truth).sum()),
"accuracy": float((called_real == truth).mean()) if truth.size else None,
"balanced_accuracy": (float(balanced_accuracy_score(truth, called_real))
if both else None),
"not_real_precision": float(precision_score(
~truth, ~called_real, zero_division=0)),
"not_real_recall": float(recall_score(~truth, ~called_real,
zero_division=0)),
"roc_auc": float(roc_auc_score(truth, probability)) if both else None,
"threshold": float(threshold),
}
def _train_real_classifier(frame, object_type: str, *,
test_fraction: float = 0.2, seed: int = 0,
threshold: float = 0.5) -> Dict[str, Any]:
"""Train one real / not-real classifier and score it on held-out wells.
:param frame: what :func:`_annotated_real_crops` returns.
:param object_type: ``cell``, ``nucleus`` or ``pathogen``.
:param test_fraction: share of the wells held out for the score.
:param seed: the split's seed.
:param threshold: the probability of real below which an object is
called not real when scoring.
:returns: the bundle to save: the estimator refitted on every crop, the
channel roles its crops carry, the padding used to cut field crops
and the held-out ``scorecard``.
:raises ValueError: an unknown object type, one class only, or fewer
than two wells.
"""
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import GroupShuffleSplit
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler
if object_type not in _REAL_TYPES:
raise ValueError(f"object type must be one of {list(_REAL_TYPES)}")
truth = np.asarray(frame["real"], dtype=bool)
groups = np.asarray(frame["group"])
if truth.all() or not truth.any():
raise ValueError("training needs crops called real and not real")
if len(set(groups)) < 2:
raise ValueError("training needs annotated crops from two wells or more")
features = np.stack([_real_features(_read_real_crop(path))
for path in frame["png_path"]])
def estimator():
"""A fresh scaled, class-balanced logistic regression."""
return make_pipeline(StandardScaler(), LogisticRegression(
C=0.1, max_iter=2000, class_weight="balanced"))
split = GroupShuffleSplit(n_splits=1, test_size=test_fraction,
random_state=seed)
train, test = next(split.split(features, truth, groups))
scorecard = {"n_train": int(len(train)),
"wells_held_out": sorted(set(groups[test].tolist()))}
if len(set(truth[train])) == 2:
held_out = estimator().fit(features[train], truth[train].astype(int))
probability = held_out.predict_proba(features[test])[
:, list(held_out.classes_).index(1)]
scorecard.update(_real_scorecard(truth[test], probability, threshold))
else:
scorecard["note"] = "the training wells held one class only; no score"
final = estimator().fit(features, truth.astype(int))
return {"version": 1, "estimator": final, "object_type": object_type,
"roles": _REAL_ROLES, "padding_fraction": _REAL_PADDING[object_type],
"scorecard": scorecard}
def _real_classifier_bundles(location: str) -> Dict[str, Dict[str, Any]]:
"""Load the classifiers a run names, keyed by object type.
:param location: one saved classifier, or a folder holding
``cell.joblib``, ``nucleus.joblib`` and/or ``pathogen.joblib``.
:raises ValueError: nothing loadable there.
"""
import joblib
location = os.path.expanduser(str(location))
if not os.path.exists(location):
from .model_zoo import _ensure_model_file
downloaded = _ensure_model_file(location, kinds=("classifier",))
if downloaded is not None:
location = str(downloaded)
if os.path.isdir(location):
paths = [os.path.join(location, f"{kind}.joblib") for kind in _REAL_TYPES]
paths = [path for path in paths if os.path.isfile(path)]
elif os.path.isfile(location):
paths = [location]
else:
raise ValueError(f"real_object_classifier not found: {location}")
bundles = {}
for path in paths:
bundle = joblib.load(path)
if not isinstance(bundle, Mapping) or bundle.get("object_type") not in _REAL_TYPES:
raise ValueError(f"{path} is not a real / not-real classifier")
bundles[bundle["object_type"]] = dict(bundle)
if not bundles:
raise ValueError(f"no cell, nucleus or pathogen classifier in {location}")
return bundles
def _drop_unreal_objects(src: str, settings: Mapping[str, Any]) -> Dict[str, int]:
"""Erase the objects a real / not-real classifier rejects from each mask.
For every object type with a classifier and a mask folder, each field's
objects are cut from ``stack/<field>.npy`` with the channels the
classifier was trained on, and those called not real are set to
background in ``masks/<type>_mask_stack/<field>.npy``. Remaining ids
are kept. Every verdict goes to ``qc/real_object_filter_<type>.csv``.
:param src: the plate folder holding ``stack`` and ``masks``.
:param settings: reads ``real_object_classifier``,
``real_object_threshold`` and the ``*_channel`` settings.
:returns: objects removed per object type.
"""
import pandas as pd
from .io import _save_array_atomic
from .tabular import write_table
bundles = _real_classifier_bundles(settings["real_object_classifier"])
threshold = float(settings.get("real_object_threshold", 0.5))
removed_counts: Dict[str, int] = {}
for object_type, bundle in bundles.items():
folder = os.path.join(src, "masks", f"{object_type}_mask_stack")
planes = [settings.get(f"{role}_channel") for role in bundle["roles"]]
if not os.path.isdir(folder):
print(f"No {object_type} masks; its real / not-real classifier is skipped.")
continue
if any(plane is None for plane in planes):
print(f"The {object_type} real / not-real classifier needs the "
f"{', '.join(bundle['roles'])} channels; it is skipped.")
continue
head = _RealObjectHead(bundle, threshold)
rows: List[Dict[str, Any]] = []
for name in sorted(f for f in os.listdir(folder)
if f.endswith(".npy") and not f.startswith(".")):
image_path = os.path.join(src, "stack", name)
mask_path = os.path.join(folder, name)
mask = np.load(mask_path)
if mask.ndim != 2 or not os.path.isfile(image_path):
continue
image = np.load(image_path, mmap_mode="r")
if image.ndim == 2:
image = image[:, :, None]
labels, areas = np.unique(mask[mask > 0], return_counts=True)
crops = [crop_object(np.asarray(image), mask, int(label),
channels=planes, padding=int(round(
bundle["padding_fraction"] * math.sqrt(area))))
for label, area in zip(labels, areas)]
verdicts = list(_predictions(head, crops))
dropped = [int(label) for label, verdict in zip(labels, verdicts)
if verdict["class"] == "not real"]
if dropped:
_save_array_atomic(mask_path, remove_labels(mask, dropped))
rows.extend({"field": os.path.splitext(name)[0],
"label": int(label), "class": verdict["class"],
"probability": verdict["probability"]}
for label, verdict in zip(labels, verdicts))
del image, mask, crops
removed_counts[object_type] = sum(row["class"] == "not real" for row in rows)
write_table(pd.DataFrame(rows, columns=["field", "label", "class", "probability"]),
os.path.join(src, "qc", f"real_object_filter_{object_type}.csv"))
print(f"Real / not-real classifier removed {removed_counts[object_type]} "
f"of {len(rows)} {object_type} objects.")
return removed_counts