"""Time-series tracking, motility analysis, and trajectory utilities."""
import cv2, os, re, glob, random, sqlite3, hashlib
import numpy as np
import pandas as pd
from collections import defaultdict
import matplotlib.pyplot as plt
from matplotlib.figure import Figure
import matplotlib as mpl
from IPython.display import display
from .figures.style import (ROLES, TYPE_SCALE, Palette, figure_style,
reference_line, resolve_ink, theme_target)
from .plot import save_figure as save_figure_to_path
#: THE TWO CONDITIONS EVERY TIMELAPSE PANEL COMPARES, coloured once and never
#: re-mapped: a reader who has learned that blue is infected in the intensity
#: histogram must not find it means something else in the motility plot.
#:
#: They were "red" and "green" at equal weight in eleven figures. That is two
#: full-strength hues where the house style asks for one -- the condition is
#: the claim, the control is the ground -- and it is the one pair a
#: red-green-colour-blind reader cannot separate at all. Infected takes the
#: highlight; uninfected takes the dark grey every control in the published
#: figures takes.
INFECTED_COLOUR = ROLES['highlight']
UNINFECTED_COLOUR = Palette.GREY_DARK
from .openmp_guard import single_threaded_openmp
from IPython.display import Image as ipyimage
from skimage.measure import regionprops_table
from scipy.signal import find_peaks
from scipy.optimize import curve_fit, linear_sum_assignment
try:
from numpy import trapezoid as trapz
except ImportError:
from numpy import trapz
from spacr import schema
from spacr.image_colors import read_image_rgb, rgb_to_cv2
from spacr.utils import _LazyModule, debug
tp = _LazyModule("trackpy")
def _npz_to_movie(arrays, filenames, save_path, fps=10):
"""Write equally sized image frames to a labelled movie without editing them.
:param arrays: nonempty sequence of ``(H, W)`` or ``(H, W, C)`` frames.
uint8 is preserved; uint16 is scaled from 0–65535; finite floats
are clipped to 0–1 and scaled. Grayscale is repeated into RGB,
two channels become red/green, and additional channels after RGB
are omitted. Volumes must be sliced or projected by the caller.
:param filenames: one label per frame, drawn near its bottom edge.
:param save_path: movie path; ``.mp4`` uses mp4v, otherwise XVID.
:param fps: finite positive frame rate, default 10.
:returns: None. The writer is released even if frame encoding raises.
:raises ValueError: invalid dimensions, data type, labels or frame rate,
detected before opening the writer.
:raises OSError: the video writer cannot open the output.
"""
if len(arrays) == 0 or len(arrays) != len(filenames):
raise ValueError("a movie needs frames and one filename per frame")
fps = float(fps)
if not np.isfinite(fps) or fps <= 0:
raise ValueError("fps must be finite and positive")
frame_shape = np.asarray(arrays[0]).shape[:2]
for index, frame in enumerate(arrays):
frame = np.asarray(frame)
if (frame.ndim not in (2, 3) or not all(frame.shape)
or frame.shape[:2] != frame_shape):
raise ValueError(
f"frame {index} must be a 2-D image with optional channels, "
"with the same height and width as every other frame")
if frame.dtype.kind == 'f':
if not np.isfinite(frame).all():
raise ValueError(f"frame {index} contains nonfinite pixels")
elif frame.dtype not in (np.dtype('uint8'), np.dtype('uint16')):
raise ValueError(f"frame {index} must use uint8, uint16 or floats")
save_path = os.fspath(save_path)
fourcc = cv2.VideoWriter_fourcc(*'XVID')
if save_path.lower().endswith('.mp4'):
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
height, width = frame_shape
out = cv2.VideoWriter(save_path, fourcc, fps, (width, height))
try:
if not out.isOpened():
raise OSError(f"video writer could not open {save_path}")
for i, frame in enumerate(arrays):
frame = np.asarray(frame)
if frame.dtype.kind == 'f':
frame = np.clip(frame, 0, 1)
frame = (frame * 255).astype(np.uint8)
elif frame.dtype == np.uint16:
frame = cv2.convertScaleAbs(frame, alpha=(255.0/65535.0))
if frame.ndim == 2 or frame.shape[2] == 1:
frame = cv2.cvtColor(frame, cv2.COLOR_GRAY2RGB)
elif frame.shape[2] == 2:
rgb_frame = np.zeros((height, width, 3), dtype=np.uint8)
rgb_frame[..., 0] = frame[..., 0]
rgb_frame[..., 1] = frame[..., 1]
frame = rgb_frame
else:
frame = np.array(frame[..., :3], copy=True, order='C')
cv2.putText(frame, str(filenames[i]), (10, height - 20),
cv2.FONT_HERSHEY_SIMPLEX, 0.5, (255, 255, 255),
1, cv2.LINE_AA)
out.write(rgb_to_cv2(frame))
finally:
out.release()
print(f"Movie saved to {save_path}")
def _scmovie(folder_paths):
"""
Generate movies from a collection of PNG images in the given folder paths.
Args:
folder_paths (list): List of folder paths containing PNG images.
Returns:
None
"""
folder_paths = list(set(folder_paths))
for folder_path in folder_paths:
movie_path = os.path.join(folder_path, 'movies')
os.makedirs(movie_path, exist_ok=True)
filename_regex = re.compile(r'(\w+)_(\w+)_(\w+)_(\d+)_(\d+).png')
grouped_images = defaultdict(list)
for filename in os.listdir(folder_path):
if filename.endswith('.png'):
match = filename_regex.match(filename)
if match:
plate, well, field, time, object_number = match.groups()
key = (plate, well, field, object_number)
grouped_images[key].append((int(time), os.path.join(folder_path, filename)))
for key, images in grouped_images.items():
images = sorted(images, key=lambda x: x[0])
_, image_paths = zip(*images)
max_height = max_width = 0
for image_path in image_paths:
image = read_image_rgb(image_path)
h, w, _ = image.shape
max_height, max_width = max(max_height, h), max(max_width, w)
plate, well, field, object_number = key
output_filename = f"{plate}_{well}_{field}_{object_number}.mp4"
output_path = os.path.join(movie_path, output_filename)
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
video = cv2.VideoWriter(output_path, fourcc, 10, (max_width, max_height))
for image_path in image_paths:
image = read_image_rgb(image_path)
h, w, _ = image.shape
padded_image = np.zeros((max_height, max_width, 3), dtype=np.uint8)
padded_image[:h, :w, :] = image
video.write(rgb_to_cv2(padded_image))
video.release()
def _sort_key(file_path):
r"""
Returns a sort key for the given file path based on the pattern '(\d+)_([A-Z]\d+)_(\d+)_(\d+).npy'.
The sort key is a tuple containing the plate, well, field, and time values extracted from the file path.
If the file path does not match the pattern, a default sort key is returned to sort the file as "earliest" or "lowest".
Args:
file_path (str): The file path to extract the sort key from.
Returns:
tuple: The sort key tuple containing the plate, well, field, and time values.
"""
match = re.search(
r'(\d+)_([A-Z]\d+)_(\d+)_(\d+)\.npy$',
os.path.basename(file_path))
if match:
plate, well, field, time = match.groups()
return (plate, well, field, int(time))
else:
return ('', '', '', 0)
def _masks_to_gif(masks, gif_folder, name, filenames, object_type):
"""
Converts a sequence of masks into a GIF file.
Args:
masks (list): List of masks representing the sequence.
gif_folder (str): Path to the folder where the GIF file will be saved.
name (str): Name of the GIF file.
filenames (list): List of filenames corresponding to each mask in the sequence.
object_type (str): Type of object represented by the masks.
Returns:
None
"""
from .io import _save_mask_timelapse_as_gif
def _display_gif(path):
"""Read ``path`` into an IPython image, display it, and return ``None``."""
with open(path, 'rb') as file:
display(ipyimage(file.read()))
highest_label = max(np.max(mask) for mask in masks)
random_colors = np.random.rand(highest_label + 1, 4)
random_colors[:, 3] = 1
random_colors[0] = [0, 0, 0, 1]
cmap = plt.cm.colors.ListedColormap(random_colors)
norm = plt.cm.colors.Normalize(vmin=0, vmax=highest_label)
save_path_gif = os.path.join(gif_folder, f'timelapse_masks_{object_type}_{name}.gif')
_save_mask_timelapse_as_gif(masks, None, save_path_gif, cmap, norm, filenames)
def _timelapse_masks_to_gif(folder_path, mask_channels, object_types):
"""
Converts a sequence of masks into a timelapse GIF file.
Args:
folder_path (str): The path to the folder containing the mask files.
mask_channels (list): List of channel indices to extract masks from.
object_types (list): List of object types corresponding to each mask channel.
Returns:
None
"""
master_folder = os.path.dirname(folder_path)
gif_folder = os.path.join(master_folder, 'movies', 'gif')
os.makedirs(gif_folder, exist_ok=True)
paths = glob.glob(os.path.join(folder_path, '*.npy'))
paths.sort(key=_sort_key)
organized_files = {}
for file in paths:
match = re.search(r'(\d+)_([A-Z]\d+)_(\d+)_\d+.npy', os.path.basename(file))
if match:
plate, well, field = match.groups()
key = (plate, well, field)
if key not in organized_files:
organized_files[key] = []
organized_files[key].append(file)
for key, file_list in organized_files.items():
name = f'{key[0]}_{key[1]}_{key[2]}'
for i, mask_channel in enumerate(mask_channels):
object_type = object_types[i]
mask_arrays = []
for file in file_list:
array = np.load(file)
mask_arrays.append(array[:, :, mask_channel])
mask_arrays_np = np.array(mask_arrays)
filenames = [os.path.basename(f) for f in file_list]
_masks_to_gif(mask_arrays_np, gif_folder, name, filenames, object_type)
def _relabel_masks_based_on_tracks(masks, tracks, mode='btrack'):
"""
Relabels the masks based on the tracks DataFrame.
Args:
masks (ndarray): Input masks array with shape (num_frames, height, width).
tracks (DataFrame): DataFrame containing track information.
mode (str, optional): Mode for relabeling. Defaults to 'btrack'.
Returns:
ndarray: Relabeled masks array with the same shape and dtype as the input masks.
"""
relabeled_masks = np.zeros(masks.shape, dtype=masks.dtype)
for frame_number in range(masks.shape[0]):
frame_tracks = tracks[tracks['frame'] == frame_number]
mapping = dict(zip(frame_tracks['original_label'], frame_tracks['track_id']))
current_mask = masks[frame_number, :, :]
for original_label, new_label in mapping.items():
relabeled_masks[frame_number][current_mask == original_label] = new_label
return relabeled_masks
def _require_2d_frames(masks, caller):
"""Stop a volumetric stack from reaching a 2-D-only tracking helper.
Every tracking path in this module reasons in 2-D: the feature table below
renames ``centroid-0``/``centroid-1`` to y/x, the track visualiser plots
two coordinates, and the motility assay measures displacement in the image
plane. Handed a ``(T, Z, Y, X)`` stack, they fail in two different and
equally unhelpful ways:
* :func:`_prepare_for_tracking` and :func:`_relabelled_stack_to_tracks_df`
raise out of skimage -- "Property eccentricity is not implemented for 3D
images", "too many values to unpack" -- which names neither z nor spaCR
and reads to the user as a corrupt mask;
* :func:`_track_by_iou` does not fail at all. Its overlap arithmetic is
dimension-agnostic, so it happily links volumes and returns a table of
perfectly plausible tracks. It returns one just as happily when the
leading axis is z rather than t, in which case the "tracks" are objects
stacked on top of each other and the trajectory is fiction.
The second is the reason this guard exists. See the 4D (Beta) half of
:mod:`spacr.zstack` for the volumetric equivalents.
:param masks: the label stack, ``(T, Y, X)`` or a sequence of 2-D frames.
:param caller: name used in the message.
:raises ValueError: when any frame is not 2-D.
"""
frames = masks if isinstance(masks, (list, tuple)) else np.asarray(masks)
ndims = {int(np.ndim(frame)) for frame in frames}
if ndims and ndims != {2}:
try:
shape = np.asarray(masks).shape
except ValueError:
shape = f'{len(frames)} frames of mixed shape'
raise ValueError(
f"{caller} needs a (T, Y, X) stack of 2-D frames and got frames of "
f"{sorted(ndims)} dimension(s) (stack shape {shape}). Every tracker "
f"in spacr.timelapse links in the image plane only, so a volume "
f"handed to one is either an error further down or -- for the IoU "
f"linker, whose overlap arithmetic works in any dimension -- a "
f"table of plausible tracks that may have been linked along z "
f"instead of along t. Use the 4D (Beta) half of spacr.zstack "
f"(segment_4d / track_4d), which requires the axis order to be "
f"stated, or collapse z before tracking with "
f"spacr.zstack.project_labels and accept that objects separated "
f"only in z merge."
)
def _prepare_for_tracking(mask_array):
"""Convert 2-D label frames into the region table used by trackers.
:param mask_array: label masks ordered by frame; every frame must be 2-D.
:returns: one row per labelled region with frame, centroid, mass, source
label, bounding box, and eccentricity.
:raises ValueError: if any frame is not two-dimensional.
"""
_require_2d_frames(mask_array, '_prepare_for_tracking')
frames = []
for t, frame in enumerate(mask_array):
props = regionprops_table(
frame,
properties=('label', 'centroid', 'area', 'bbox', 'eccentricity')
)
df = pd.DataFrame(props)
df = df.rename(columns={
'centroid-0': 'y',
'centroid-1': 'x',
'area': 'mass',
'label': 'original_label'
})
df['frame'] = t
frames.append(df[['frame','y','x','mass','original_label',
'bbox-0','bbox-1','bbox-2','bbox-3','eccentricity']])
return pd.concat(frames, ignore_index=True)
[docs]
def link_by_iou(mask_prev, mask_next, iou_threshold=0.1):
"""Match labels between two consecutive frames using IoU and Hungarian assignment.
:param mask_prev: labelled mask from the previous frame.
:param mask_next: labelled mask from the next frame.
:param iou_threshold: minimum IoU required to accept a match. Default ``0.1``.
:returns: list of ``(label_prev, label_next)`` matches above the threshold.
"""
labels_prev = np.unique(mask_prev)[1:]
labels_next = np.unique(mask_next)[1:]
bool_prev = {L: mask_prev==L for L in labels_prev}
bool_next = {L: mask_next==L for L in labels_next}
cost = np.ones((len(labels_prev), len(labels_next)), dtype=float)
for i, L1 in enumerate(labels_prev):
m1 = bool_prev[L1]
for j, L2 in enumerate(labels_next):
m2 = bool_next[L2]
inter = np.logical_and(m1, m2).sum()
union = np.logical_or(m1, m2).sum()
cost[i, j] = 1 - inter/union
row_ind, col_ind = linear_sum_assignment(cost)
matches = []
for i, j in zip(row_ind, col_ind):
if cost[i,j] <= 1 - iou_threshold:
matches.append((labels_prev[i], labels_next[j]))
return matches
def _find_optimal_search_range(features, initial_search_range=500, increment=10, max_attempts=49, memory=3):
"""
Find the optimal search range for linking features.
Args:
features (list): List of features to be linked.
initial_search_range (int, optional): Initial search range. Defaults to 500.
increment (int, optional): Increment value for reducing the search range. Defaults to 10.
max_attempts (int, optional): Maximum number of attempts to find the optimal search range. Defaults to 49.
memory (int, optional): Memory parameter for linking features. Defaults to 3.
Returns:
int: The optimal search range for linking features.
"""
optimal_search_range = initial_search_range
for attempt in range(max_attempts):
try:
tp.link(features, search_range=optimal_search_range, memory=memory)
print(f"Success with search_range={optimal_search_range}")
return optimal_search_range
except Exception:
optimal_search_range -= increment
print(f'Retrying with displacement value: {optimal_search_range}', end='\r', flush=True)
min_range = initial_search_range-(max_attempts*increment)
if optimal_search_range <= min_range:
print(f'timelapse_displacement={optimal_search_range} is too high. Lower timelapse_displacement or set to None for automatic thresholding.')
return optimal_search_range
def _remove_objects_from_first_frame(masks, percentage=10):
"""
Removes a specified percentage of objects from the first frame of a sequence of masks.
Parameters:
masks (ndarray): Sequence of masks representing the frames.
percentage (int): Percentage of objects to remove from the first frame.
Returns:
ndarray: Sequence of masks with objects removed from the first frame.
"""
first_frame = masks[0]
unique_labels = np.unique(first_frame[first_frame != 0])
num_labels_to_remove = max(1, int(len(unique_labels) * (percentage / 100)))
labels_to_remove = random.sample(list(unique_labels), num_labels_to_remove)
for label in labels_to_remove:
masks[0][first_frame == label] = 0
return masks
def _track_by_iou(masks, iou_threshold=0.1):
"""
Build a track table by linking masks frame→frame via IoU.
Returns a DataFrame with columns [frame, original_label, track_id].
The 2-D guard is load-bearing here rather than defensive: the overlap
arithmetic below is dimension-agnostic and will link a (T, Z, Y, X) stack
without complaint, including one whose leading axis is z rather than t.
See :func:`_require_2d_frames`.
"""
_require_2d_frames(masks, '_track_by_iou')
n_frames = masks.shape[0]
labels0 = np.unique(masks[0])[1:]
next_track = 1
track_map = {}
for L in labels0:
track_map[(0, L)] = next_track
next_track += 1
for t in range(1, n_frames):
prev, curr = masks[t-1], masks[t]
matches = link_by_iou(prev, curr, iou_threshold=iou_threshold)
used_curr = set()
for L_prev, L_curr in matches:
tid = track_map[(t-1, L_prev)]
track_map[(t, L_curr)] = tid
used_curr.add(L_curr)
for L in np.unique(curr)[1:]:
if L not in used_curr:
track_map[(t, L)] = next_track
next_track += 1
records = []
for (frame, label), tid in track_map.items():
records.append({'frame': frame, 'original_label': label, 'track_id': tid})
return pd.DataFrame(records)
def _facilitate_trackin_with_adaptive_removal(masks, search_range=None, max_attempts=5, memory=3, min_mass=50, track_by_iou=False):
"""
Facilitates object tracking with deterministic initial filtering and
trackpy’s nearest-neighbour linking.
Args:
masks (np.ndarray): integer‐labeled masks (frames × H × W).
search_range (int|None): max displacement; if None, auto‐computed.
max_attempts (int): how many times to retry with smaller search_range.
memory (int): trackpy memory parameter.
min_mass (float): drop any object in frame 0 with area < min_mass.
Returns:
masks, features_df, tracks_df
Raises:
RuntimeError if linking fails after max_attempts.
"""
features = _prepare_for_tracking(masks)
f0 = features[features['frame'] == 0]
valid = f0.loc[f0['mass'] >= min_mass, 'original_label'].unique()
masks[0] = np.where(np.isin(masks[0], valid), masks[0], 0)
features = _prepare_for_tracking(masks)
if search_range is None:
a99 = f0['mass'].quantile(0.99)
search_range = max(1, int(2 * np.sqrt(a99)))
for attempt in range(1, max_attempts + 1):
try:
if track_by_iou:
tracks_df = _track_by_iou(masks, iou_threshold=0.1)
else:
tracks_df = tp.link_df(features, search_range=search_range, memory=memory)
print(f"Linked on attempt {attempt} with search_range={search_range}")
return masks, features, tracks_df
except Exception as e:
search_range = max(1, int(search_range * 0.8))
print(f"Attempt {attempt} failed ({e}); reducing search_range to {search_range}")
raise RuntimeError(
f"Failed to track after {max_attempts} attempts; last search_range={search_range}"
)
def _trackpy_track_cells(src, name, batch_filenames, object_type, masks, timelapse_displacement, timelapse_memory, timelapse_remove_transient, plot, save, mode, track_by_iou):
"""
Track cells using the Trackpy library.
Args:
src (str): The source file path.
name (str): The name of the track.
batch_filenames (list): List of batch filenames.
object_type (str): The type of object to track.
masks (list | np.ndarray): the frames of one field, either as a
list of 2-D label images (what ``spacr.object`` hands over,
because that is what ``CellposeModel.eval`` returns for a list
of images) or as a (T, H, W) array. Coerced to an array below.
timelapse_displacement (int): The displacement for timelapse tracking.
timelapse_memory (int): The memory for timelapse tracking.
timelapse_remove_transient (bool): Whether to remove transient objects in timelapse tracking.
plot (bool): Whether to plot the tracks.
save (bool): Whether to save the tracks.
mode (str): The mode of tracking.
Returns:
list: The mask stack.
"""
from .plot import _visualize_and_save_timelapse_stack_with_tracks
from .utils import _masks_to_masks_stack
print(f'Tracking objects with trackpy')
masks = np.asarray(masks)
if timelapse_displacement is None:
features = _prepare_for_tracking(masks)
timelapse_displacement = _find_optimal_search_range(features, initial_search_range=500, increment=10, max_attempts=49, memory=3)
if timelapse_displacement is None:
timelapse_displacement = 50
masks, features, tracks_df = _facilitate_trackin_with_adaptive_removal(masks, search_range=timelapse_displacement, max_attempts=100, memory=timelapse_memory, track_by_iou=track_by_iou)
if 'particle' not in tracks_df.columns:
tracks_df = tracks_df.rename(columns={'track_id': 'particle'})
tracks_df = tracks_df.merge(features[['frame', 'original_label', 'x', 'y']], on=['frame', 'original_label'], how='left', validate='many_to_one')
tracks_df['particle'] += 1
if timelapse_remove_transient:
tracks_df_filter = tp.filter_stubs(tracks_df, len(masks))
else:
tracks_df_filter = tracks_df.copy()
tracks_df_filter = tracks_df_filter.rename(columns={'particle': 'track_id'})
print(f'Removed {len(tracks_df)-len(tracks_df_filter)} objects that were not present in all frames')
masks = _relabel_masks_based_on_tracks(masks, tracks_df_filter)
tracks_path = os.path.join(os.path.dirname(src), 'tracks')
os.makedirs(tracks_path, exist_ok=True)
tracks_df_filter.to_csv(os.path.join(tracks_path, f'trackpy_tracks_{object_type}_{name}.csv'), index=False)
if plot or save:
_visualize_and_save_timelapse_stack_with_tracks(masks, tracks_df_filter, save, src, name, plot, batch_filenames, object_type, mode)
mask_stack = _masks_to_masks_stack(masks)
return mask_stack
def _trackastra_track_cells(src, name, batch_filenames, object_type, masks, images=None,
timelapse_remove_transient=False, plot=False, save=False,
mode='trackastra', model_name='general_2d', device='automatic',
linking_mode='greedy'):
"""Track objects across frames with Trackastra, a transformer-based tracker.
Trackastra tops the Cell Tracking Challenge leaderboard and, unlike the
other two backends, has no hyperparameters to tune: trackpy needs a
search_range and a memory, and btrack needs a motion-model config, both of
which have to be re-guessed per dataset. It also links divisions natively,
which matters for a replication assay where a parasite splitting into two
is signal rather than a tracking error.
It consumes exactly what spaCR already has — the raw intensity stack and
the label masks — so nothing upstream changes.
:param src: run folder; the tracks CSV lands in ``<dirname(src)>/tracks``,
with the mother's track id in ``parent_track_id`` (0 for none) when
Trackastra reports its division links.
:param name: batch name used in the output filename.
:param batch_filenames: filenames of the frames, for the track visualiser.
:param object_type: 'cell' / 'nucleus' / 'pathogen' / 'organelle'.
:param masks: (T, Y, X) integer label stack.
:param images: (T, Y, X) intensity stack. Trackastra uses appearance as
well as geometry; if omitted the masks are used as a stand-in, which
works but discards the appearance signal.
:param timelapse_remove_transient: drop tracks not present in every frame.
:param model_name: pretrained Trackastra model, e.g. 'general_2d'.
:param device: 'automatic', 'cpu', or a CUDA device string.
:param linking_mode: 'greedy' (fast) or 'ilp' (optimal, needs the extra).
:returns: the relabelled mask stack, ids consistent across frames.
:raises RuntimeError: if trackastra is not installed, naming the fix.
"""
from .plot import _visualize_and_save_timelapse_stack_with_tracks
from .utils import _masks_to_masks_stack
try:
from trackastra.model import Trackastra
from trackastra.tracking import graph_to_ctc
except ImportError as exc:
raise RuntimeError(
"timelapse_mode='trackastra' needs the trackastra package, which is "
"not installed. Install it with `pip install trackastra` (BSD-3, "
"PyTorch-only), or choose timelapse_mode='trackpy' / 'btrack' / 'iou'."
) from exc
masks = np.asarray(masks)
if masks.ndim != 3:
raise ValueError(f"masks must be a (T, Y, X) stack, got shape {masks.shape}")
if masks.shape[0] < 2:
print(f"Trackastra: only {masks.shape[0]} frame(s) for {object_type}; nothing to link.")
return _masks_to_masks_stack(masks)
imgs = np.asarray(images) if images is not None else masks.astype(np.float32)
if imgs.shape != masks.shape:
raise ValueError(
f"images shape {imgs.shape} does not match masks shape {masks.shape}")
model = Trackastra.from_pretrained(model_name, device=device)
track_graph = model.track(imgs, masks, mode=linking_mode)
ctc_df, masks_tracked = graph_to_ctc(track_graph, masks, outdir=None)
tracks_df = _trackastra_graph_to_tracks_df(track_graph, masks_tracked)
if isinstance(ctc_df, pd.DataFrame) and {'label', 'parent'} <= set(ctc_df.columns):
mothers = dict(zip(ctc_df['label'].astype(int), ctc_df['parent'].astype(int)))
tracks_df = _native_lineage_columns(tracks_df, mothers, 'trackastra')
if timelapse_remove_transient:
n_frames = masks_tracked.shape[0]
keep = tracks_df.groupby('track_id')['frame'].nunique() == n_frames
kept_ids = set(keep[keep].index)
before = len(tracks_df)
tracks_df = tracks_df[tracks_df['track_id'].isin(kept_ids)].copy()
print(f'Removed {before - len(tracks_df)} objects that were not present in all frames')
masks_tracked = _relabel_masks_based_on_tracks(masks_tracked, tracks_df)
tracks_path = os.path.join(os.path.dirname(src), 'tracks')
os.makedirs(tracks_path, exist_ok=True)
tracks_df.to_csv(
os.path.join(tracks_path, f'trackastra_tracks_{object_type}_{name}.csv'),
index=False)
if plot or save:
_visualize_and_save_timelapse_stack_with_tracks(
masks_tracked, tracks_df, save, src, name, plot,
batch_filenames, object_type, mode)
return _masks_to_masks_stack(masks_tracked)
def _timeflows_track_cells(src, name, batch_filenames, object_type, masks, images=None,
timelapse_remove_transient=False, plot=False, save=False,
mode='timeflows', model_path=None, device=None,
min_successor=0.5, max_distance=1.0, net=None):
"""Track objects with Timeflows, spaCR's experimental temporal Cellpose.
Timeflows predicts, for every pixel of an object in frame t, where that
object's centre is in frame t+1, plus whether it has a successor at all.
Those votes are assigned to the next frame's masks, and the links are
stitched into whole-movie track ids that are identical on every re-run
(:func:`spacr.timeflows_model._stitch_links`). A daughter after a division
starts a new track; parent/child lineage is not recorded yet.
It is opt-in and experimental: on the held-out movies measured so far it
does not beat overlap linking ('iou'), so it is never the default. It
needs a trained checkpoint (``timeflows_model``); none is downloaded.
:param src: run folder; the tracks CSV lands in ``<dirname(src)>/tracks``.
:param name: batch name used in the output filename.
:param batch_filenames: filenames of the frames, for the track visualiser.
:param object_type: 'cell' / 'nucleus' / 'pathogen' / 'organelle'.
:param masks: (T, Y, X) integer label stack.
:param images: (T, Y, X) or (T, Y, X, C) intensity stack; required,
because the model reads the images, not only the masks.
:param timelapse_remove_transient: drop tracks not present in every frame.
:param model_path: the Timeflows checkpoint to load.
:param device: 'cuda', 'cpu' or None for CUDA when available. On CPU the
encoder runs in float32, which is much faster there than bfloat16.
:param min_successor: successor probability an object needs to be linked.
:param max_distance: the furthest link, in the object's own diameters.
:param net: an already loaded network; used instead of ``model_path``.
:returns: the relabelled mask stack, ids consistent across frames.
:raises ValueError: no checkpoint or images, or mismatched shapes.
"""
from .plot import _visualize_and_save_timelapse_stack_with_tracks
from .qt.i18n import tr
from .utils import _masks_to_masks_stack
from . import timeflows_model
masks = np.asarray(masks)
if masks.ndim != 3:
raise ValueError(tr("Timeflows needs a (T, Y, X) mask stack, got shape {shape}.",
shape=masks.shape))
if masks.shape[0] < 2:
print(tr("Timeflows: only {count} frame(s) for {object_type}; nothing to link.",
count=masks.shape[0], object_type=object_type))
return _masks_to_masks_stack(masks)
if images is None:
raise ValueError(tr("timelapse_mode='timeflows' needs the image stack, not only the masks."))
images = np.asarray(images)
if images.shape[:3] != masks.shape:
raise ValueError(tr("Image stack shape {images} does not match mask stack shape {masks}.",
images=images.shape, masks=masks.shape))
if net is None:
if not model_path or not os.path.isfile(str(model_path)):
raise ValueError(tr("timelapse_mode='timeflows' needs timeflows_model set to a trained "
"Timeflows checkpoint file; none was found at {path}.",
path=model_path))
if device is None:
import torch
device = 'cuda' if torch.cuda.is_available() else 'cpu'
net = timeflows_model._load_timeflows(
str(model_path), device=device,
precision='float32' if str(device).startswith('cpu') else 'checkpoint')
masks_tracked, _links = timeflows_model._track_movie(
net, list(images), masks, device=device or 'cpu',
min_successor=min_successor, max_distance=max_distance)
tracks_df = _relabelled_stack_to_tracks_df(masks_tracked)
if timelapse_remove_transient and not tracks_df.empty:
n_frames = masks_tracked.shape[0]
keep = tracks_df.groupby('track_id')['frame'].nunique() == n_frames
kept_ids = set(keep[keep].index)
before = len(tracks_df)
tracks_df = tracks_df[tracks_df['track_id'].isin(kept_ids)].copy()
print(tr("Removed {count} objects that were not present in all frames",
count=before - len(tracks_df)))
masks_tracked = np.where(np.isin(masks_tracked, list(kept_ids)), masks_tracked, 0)
tracks_path = os.path.join(os.path.dirname(src), 'tracks')
os.makedirs(tracks_path, exist_ok=True)
from .tabular import write_table
write_table(tracks_df, os.path.join(
tracks_path, f'timeflows_tracks_{object_type}_{name}.csv'))
if plot or save:
_visualize_and_save_timelapse_stack_with_tracks(
masks_tracked, tracks_df, save, src, name, plot,
batch_filenames, object_type, mode)
return _masks_to_masks_stack(masks_tracked)
def _sam2_new_seeds(masks, followed, next_id):
"""Objects spaCR segmented that SAM2 is not following, as new seeds.
An object on frame t is new when less than
``_SAM2_NEW_OBJECT_COVER`` of it lies under SAM2's objects on that
frame. A new object overlapping one found new on the frame before is
the same object and is not seeded again, so each is seeded once, on the
first frame it was found.
:param masks: spaCR's ``T x H x W`` per-frame labels.
:param followed: SAM2's ``T x H x W`` labels from the seeds so far.
:param next_id: the first id free for a new object.
:returns: ``{frame: H x W labels}`` of the new objects, ids from
``next_id`` up.
"""
from ._segmentation_backends import (_SAM2_NEW_OBJECT_COVER,
_SAM2_SAME_NEW_OBJECT_IOU)
seeds = {}
previous = None
for frame in range(1, len(masks)):
labels = np.asarray(masks[frame])
covered = np.asarray(followed[frame]) > 0
new = np.zeros(labels.shape, dtype=np.int32)
seeded = np.zeros(labels.shape, dtype=np.int32)
for label in np.unique(labels):
if not label:
continue
region = labels == label
area = int(region.sum())
if (covered & region).sum() >= _SAM2_NEW_OBJECT_COVER * area:
continue
new[region] = label
same = False
if previous is not None:
for other in np.unique(previous[region]):
if not other:
continue
before = previous == other
union = (before | region).sum()
if union and (before & region).sum() / union >= _SAM2_SAME_NEW_OBJECT_IOU:
same = True
break
if not same:
seeded[region] = next_id
next_id += 1
if seeded.any():
seeds[frame] = seeded
previous = new
return seeds
def _sam2_track_cells(src, name, batch_filenames, object_type, masks, images=None,
timelapse_remove_transient=False, plot=False, save=False,
mode='sam2', model=None, device=None, propagate=None):
"""Segment and track in one step with SAM2's video predictor.
spaCR's masks on the first frame with objects seed SAM2, which follows
each object through the movie with a memory of its appearance, so each
keeps one id. Objects spaCR finds later that SAM2 is not following --
cells entering the field, or daughters of a division -- are seeded on
the frame they were first found and the movie is followed once more
with every seed. The masks returned are SAM2's, not spaCR's per-frame
masks. A daughter after a division starts a new track, and the track
it grew out of is written as its ``parent_track_id`` (0 for none), so
lineage trees can be drawn from the tracks table.
SAM2 runs in an environment of its own, installed from the Model Zoo.
:param src: run folder; the tracks CSV lands in ``<dirname(src)>/tracks``.
:param name: batch name used in the output filename.
:param batch_filenames: filenames of the frames, for the track visualiser.
:param object_type: 'cell' / 'nucleus' / 'pathogen' / 'organelle'.
:param masks: (T, Y, X) integer label stack from spaCR's segmentation.
:param images: (T, Y, X) or (T, Y, X, C) intensity stack; required,
because SAM2 follows appearance. With channels, the first is used.
:param timelapse_remove_transient: drop tracks not present in every frame.
:param model: a SAM2.1 checkpoint name; the smallest when None.
:param device: 'cuda', 'cpu' or None for the backend's own choice.
:param propagate: the SAM2 call; the backend's own when None.
:returns: the relabelled mask stack, ids consistent across frames.
:raises ValueError: no images, or mismatched shapes.
:raises ImportError: SAM2 is not installed.
"""
from .plot import _visualize_and_save_timelapse_stack_with_tracks
from .qt.i18n import tr
from .utils import _masks_to_masks_stack
from ._segmentation_backends import _sam2_frames, _sam2_propagate
masks = np.asarray(masks)
if masks.ndim != 3:
raise ValueError(tr("SAM2 needs a (T, Y, X) mask stack, got shape {shape}.",
shape=masks.shape))
if images is None:
raise ValueError(tr("timelapse_mode='sam2' needs the image stack, not only the masks."))
images = np.asarray(images)
if images.shape[:3] != masks.shape:
raise ValueError(tr("Image stack shape {images} does not match mask stack shape {masks}.",
images=images.shape, masks=masks.shape))
occupied = [t for t in range(masks.shape[0]) if masks[t].any()]
if not occupied:
print(tr("SAM2: no {object_type} objects in any frame; nothing to follow.",
object_type=object_type))
return _masks_to_masks_stack(masks)
propagate = propagate or _sam2_propagate
frames = _sam2_frames(images)
start = occupied[0]
seeds = {start: masks[start].astype(np.int32)}
masks_tracked, reply = propagate(frames, seeds, model=model, device=device)
later = masks.copy()
later[:start + 1] = 0
extra = _sam2_new_seeds(later, masks_tracked, int(masks[start].max()) + 1)
if extra:
seeds.update(extra)
masks_tracked, reply = propagate(frames, seeds, model=model, device=device)
print(tr("SAM2 followed {count} {object_type} objects through {frames} frames "
"in {seconds:.1f} s on {device}.",
count=reply.get('objects'), object_type=object_type,
frames=masks.shape[0], seconds=float(reply.get('seconds') or 0.0),
device=reply.get('device')))
tracks_df = _relabelled_stack_to_tracks_df(masks_tracked)
if not tracks_df.empty:
tracks_df = _native_lineage_columns(
tracks_df, _sam2_parent_links(masks_tracked), 'sam2')
if timelapse_remove_transient and not tracks_df.empty:
n_frames = masks_tracked.shape[0]
keep = tracks_df.groupby('track_id')['frame'].nunique() == n_frames
kept_ids = set(keep[keep].index)
before = len(tracks_df)
tracks_df = tracks_df[tracks_df['track_id'].isin(kept_ids)].copy()
print(tr("Removed {count} objects that were not present in all frames",
count=before - len(tracks_df)))
masks_tracked = np.where(np.isin(masks_tracked, list(kept_ids)), masks_tracked, 0)
tracks_path = os.path.join(os.path.dirname(src), 'tracks')
os.makedirs(tracks_path, exist_ok=True)
from .tabular import write_table
write_table(tracks_df, os.path.join(
tracks_path, f'sam2_tracks_{object_type}_{name}.csv'))
if plot or save:
_visualize_and_save_timelapse_stack_with_tracks(
masks_tracked, tracks_df, save, src, name, plot,
batch_filenames, object_type, mode)
return _masks_to_masks_stack(masks_tracked)
_SAM2_PARENT_OVERLAP = 0.3
def _sam2_parent_links(masks_tracked):
"""The track each later-starting SAM2 track grew out of.
A track that starts after the first frame is a daughter of the track
that covers at least ``_SAM2_PARENT_OVERLAP`` of its first mask on the
frame before (the most-covering one when several do). A track that
starts where nothing was before, such as a cell entering the field, has
no parent.
:param masks_tracked: SAM2's ``T x H x W`` labels, ids consistent
across frames.
:returns: ``{daughter_track_id: mother_track_id}``, 0 for no parent.
"""
stack = np.asarray(masks_tracked)
seen = set(int(v) for v in np.unique(stack[0]) if v) if len(stack) else set()
links = {track: 0 for track in seen}
for frame in range(1, len(stack)):
labels = stack[frame]
before = stack[frame - 1]
for track in np.unique(labels):
track = int(track)
if not track or track in seen:
continue
seen.add(track)
region = labels == track
under = before[region]
under = under[(under > 0) & (under != track)]
parent = 0
if under.size:
ids, counts = np.unique(under, return_counts=True)
best = int(np.argmax(counts))
if counts[best] >= _SAM2_PARENT_OVERLAP * int(region.sum()):
parent = int(ids[best])
links[track] = parent
return links
def _relabelled_stack_to_tracks_df(masks_tracked):
"""Flatten a tracker's relabelled label stack into spaCR's tracks table.
Downstream consumers (the track visualiser, the motility assay) expect
``frame`` / ``track_id`` / ``x`` / ``y`` / ``original_label``, which is the
layout trackpy and btrack already emit. Centroids are recomputed from the
relabelled stack so they always agree with the masks actually returned,
rather than trusting whatever table the tracker handed back — Trackastra's
graph and Ultrack's ``to_tracks_layer`` both carry their own coordinates,
and both can disagree with the exported segmentation after post-processing.
Every backend that returns a stack whose ids are already consistent across
frames (Trackastra, Ultrack) shares this function, so their CSVs are
byte-for-byte comparable.
The columns are also the ones the 4D (Beta) table in :mod:`spacr.zstack`
reproduces verbatim (``zstack.BASE_TRACK_COLUMNS``), so a volumetric run's
table drops into the same consumers with z and volume merely appended.
:param masks_tracked: the relabelled (T, Y, X) stack.
:returns: DataFrame with one row per object per frame.
:raises ValueError: when handed a volumetric stack; see
:func:`_require_2d_frames`.
"""
from skimage.measure import regionprops
_require_2d_frames(masks_tracked, '_relabelled_stack_to_tracks_df')
rows = []
for t, frame in enumerate(np.asarray(masks_tracked)):
for region in regionprops(frame):
cy, cx = region.centroid
rows.append({
'frame': t,
'track_id': int(region.label),
'original_label': int(region.label),
'x': float(cx),
'y': float(cy),
})
df = pd.DataFrame(rows, columns=['frame', 'track_id', 'original_label', 'x', 'y'])
if df.empty:
return df
return df.sort_values(['track_id', 'frame']).reset_index(drop=True)
def _trackastra_graph_to_tracks_df(track_graph, masks_tracked):
"""Flatten a Trackastra track graph into spaCR's tracks-table layout.
The graph itself carries no coordinate spaCR trusts, so this is just
:func:`_relabelled_stack_to_tracks_df` over the relabelled stack; the
wrapper is kept because it names the caller's intent.
:param track_graph: the graph returned by ``Trackastra.track``; unused, the
coordinates come from the stack.
:param masks_tracked: the relabelled (T, Y, X) stack.
:returns: DataFrame with one row per object per frame.
"""
return _relabelled_stack_to_tracks_df(masks_tracked)
def _ultrack_set(section, attr, value, setting_name):
"""Apply one spaCR setting onto an Ultrack config section, loudly.
Ultrack's ``MainConfig`` is a nest of pydantic models whose field names have
moved between releases, and spaCR deliberately does not pin a version.
Assigning to a field that no longer exists would surface as a pydantic
error naming only Ultrack's own field, which tells the user nothing about
which spaCR knob to change; this translates it back.
:param section: one of ``config.data_config`` / ``segmentation_config`` /
``linking_config`` / ``tracking_config``.
:param attr: the Ultrack field name to assign.
:param value: the already-coerced value to assign.
:param setting_name: the spaCR setting this came from, for the message.
:raises RuntimeError: if the installed Ultrack has no such field.
"""
if not hasattr(section, attr):
raise RuntimeError(
f"the installed ultrack's {type(section).__name__} has no '{attr}' "
f"field, so the {setting_name} setting cannot be applied. Install a "
"supported release with `pip install spacr[ultrack]`, or choose "
"timelapse_mode='trackastra' / 'trackpy' / 'btrack' / 'iou'."
)
setattr(section, attr, value)
def _ultrack_labels_to_contours(ultrack_utils):
"""Return the installed Ultrack's labels -> (foreground, contours) converter.
Ultrack renamed ``labels_to_edges`` to ``labels_to_contours`` in the same
release that renamed ``detection``/``edges`` to ``foreground``/``contours``.
Both names are still in the wild, so resolve whichever is present rather
than importing one and breaking on the other.
:param ultrack_utils: the imported ``ultrack.utils`` module.
:returns: the converter, called as ``fn(list_of_label_stacks, sigma=...)``.
:raises RuntimeError: if neither name is available.
"""
for attr in ('labels_to_contours', 'labels_to_edges'):
fn = getattr(ultrack_utils, attr, None)
if fn is not None:
return fn
raise RuntimeError(
"the installed ultrack exposes neither ultrack.utils.labels_to_contours "
"nor ultrack.utils.labels_to_edges, so spaCR cannot turn the Cellpose "
"labels into Ultrack's input. Install a supported release with "
"`pip install spacr[ultrack]`, or choose timelapse_mode='trackastra' / "
"'trackpy' / 'btrack' / 'iou'."
)
def _ultrack_track_kwargs(track_fn):
"""Map spaCR's (foreground, contours) pair onto the installed Ultrack's names.
``ultrack.track`` took ``detection=``/``edges=`` before the 2024 rename and
``foreground=``/``contours=`` after it. Reading the signature keeps the one
adapter working against both instead of pinning spaCR to a release.
:param track_fn: the imported ``ultrack.track``.
:returns: ``(foreground_kwarg, contours_kwarg, accepts_images)``.
:raises RuntimeError: if the signature matches neither generation.
"""
import inspect
params = inspect.signature(track_fn).parameters
accepts_images = 'images' in params
if 'foreground' in params and 'contours' in params:
return 'foreground', 'contours', accepts_images
if 'detection' in params and 'edges' in params:
return 'detection', 'edges', accepts_images
raise RuntimeError(
"the installed ultrack's track() accepts neither foreground/contours "
"nor detection/edges, so spaCR cannot drive it. Install a supported "
"release with `pip install spacr[ultrack]`, or choose "
"timelapse_mode='trackastra' / 'trackpy' / 'btrack' / 'iou'."
)
def _native_lineage_columns(tracks, parents, source):
"""Keep native roots and filtered-out parents distinct from unknown links."""
tracks = tracks.copy()
present = set(tracks['track_id'].astype(int))
normalized = {}
for child, value in parents.items():
child = int(child)
if child not in present:
continue
values = value if isinstance(value, (list, tuple, set, np.ndarray)) else [value]
candidates = set()
for parent in values:
if parent is None or pd.isna(parent):
continue
number = int(parent)
if number != float(parent):
raise ValueError(f'{source}: non-integer parent for track {child}')
if number > 0 and number != child:
candidates.add(number)
if len(candidates) > 1:
raise ValueError(f'{source}: track {child} has multiple parents; '
'a division lineage requires one parent')
parent = next(iter(candidates), 0)
normalized[child] = parent if parent in present else 0
tracks['parent_track_id'] = tracks['track_id'].map(normalized).fillna(0).astype(int)
tracks['parent_track_id_source'] = [source if int(t) in normalized else ''
for t in tracks['track_id']]
return tracks
def _ultrack_track_cells(src, name, batch_filenames, object_type, masks, images=None,
timelapse_remove_transient=False, plot=False, save=False,
mode='ultrack', max_distance=25.0, division_weight=-0.1,
contour_sigma=0.0, n_workers=1):
"""Track objects across frames with Ultrack, a global-optimisation tracker.
Ultrack (Nature Methods 2025) is the other Cell Tracking Challenge leader,
and it is kept alongside Trackastra rather than instead of it because the
two fail in different directions. Trackastra scores candidate pairs of
already-segmented objects with a transformer and then links them; Ultrack
enumerates candidate segmentations from a contour map and solves a single
integer program over segmentation *and* linking together. That joint solve
is what makes it the stronger choice on densely packed monolayers and on 3D
stacks, where the correct object boundary is ambiguous until you know the
track — a confluent, heavily infected well is exactly that case. Trackastra
stays the better zero-config generalist, so it stays the default.
Like the Trackastra adapter this consumes what spaCR already has: the
Cellpose label stack, converted to Ultrack's (foreground, contours) input
by ``ultrack.utils.labels_to_contours``, so the user never has to produce a
contour map separately. The intensity stack is forwarded when available —
Ultrack uses it for appearance features during linking.
Ultrack persists its candidate hypotheses and the solved links in a
database (sqlite by default) under ``config.data_config.working_dir``. That
store is pointed at a temporary directory that is deleted when this returns,
so the user's run folder gains the tracks CSV and nothing else.
:param src: run folder; the tracks CSV lands in ``<dirname(src)>/tracks``.
:param name: batch name used in the output filename.
:param batch_filenames: filenames of the frames, for the track visualiser.
:param object_type: 'cell' / 'nucleus' / 'pathogen' / 'organelle'.
:param masks: (T, Y, X) integer label stack.
:param images: (T, Y, X) intensity stack, used for appearance features.
:param timelapse_remove_transient: drop tracks not present in every frame.
:param max_distance: largest displacement in pixels Ultrack will link across.
:param division_weight: solver cost of splitting one track into two.
:param contour_sigma: Gaussian smoothing applied when deriving contours
from the labels; 0 disables smoothing.
:param n_workers: worker processes for the segmentation and linking passes.
:returns: the relabelled mask stack, ids consistent across frames.
:raises RuntimeError: if ultrack is not installed, naming the fix.
The Napari lineage graph maps a child track ID to its parent track IDs.
Missing entries are native roots, never invitations to infer nearby
parents.
"""
from .plot import _visualize_and_save_timelapse_stack_with_tracks
from .utils import _masks_to_masks_stack
try:
from ultrack import MainConfig, track, to_tracks_layer, tracks_to_zarr
from ultrack import utils as ultrack_utils
except ImportError as exc:
raise RuntimeError(
"timelapse_mode='ultrack' needs the ultrack package, which is not "
"installed. Install it with `pip install spacr[ultrack]` (it brings "
"its own ILP solver and database backend), or choose "
"timelapse_mode='trackastra' / 'trackpy' / 'btrack' / 'iou'."
) from exc
masks = np.asarray(masks)
if masks.ndim != 3:
raise ValueError(f"masks must be a (T, Y, X) stack, got shape {masks.shape}")
imgs = np.asarray(images) if images is not None else None
if imgs is not None and imgs.shape != masks.shape:
raise ValueError(
f"images shape {imgs.shape} does not match masks shape {masks.shape}")
if masks.shape[0] < 2:
print(f"Ultrack: only {masks.shape[0]} frame(s) for {object_type}; nothing to link.")
return _masks_to_masks_stack(masks)
labels_to_contours = _ultrack_labels_to_contours(ultrack_utils)
fg_kwarg, contours_kwarg, accepts_images = _ultrack_track_kwargs(track)
import pathlib
import shutil
import tempfile
work_dir = tempfile.mkdtemp(prefix='spacr_ultrack_')
try:
config = MainConfig()
_ultrack_set(config.data_config, 'working_dir', pathlib.Path(work_dir),
'temporary Ultrack data store')
_ultrack_set(config.linking_config, 'max_distance', float(max_distance),
'ultrack_max_distance')
_ultrack_set(config.tracking_config, 'division_weight', float(division_weight),
'ultrack_division_weight')
for section in (config.segmentation_config, config.linking_config):
_ultrack_set(section, 'n_workers', int(n_workers), 'ultrack_n_workers')
sigma = float(contour_sigma)
foreground, contours = labels_to_contours([masks], sigma=sigma if sigma > 0 else None)
track_kwargs = {fg_kwarg: foreground, contours_kwarg: contours}
if accepts_images and imgs is not None:
track_kwargs['images'] = [imgs]
track(config, **track_kwargs)
tracks_table, lineage = to_tracks_layer(config)
masks_tracked = np.asarray(tracks_to_zarr(config, tracks_table))
finally:
shutil.rmtree(work_dir, ignore_errors=True)
tracks_df = _relabelled_stack_to_tracks_df(masks_tracked)
if timelapse_remove_transient:
n_frames = masks_tracked.shape[0]
keep = tracks_df.groupby('track_id')['frame'].nunique() == n_frames
kept_ids = set(keep[keep].index)
before = len(tracks_df)
tracks_df = tracks_df[tracks_df['track_id'].isin(kept_ids)].copy()
print(f'Removed {before - len(tracks_df)} objects that were not present in all frames')
masks_tracked = _relabel_masks_based_on_tracks(masks_tracked, tracks_df)
if lineage is not None:
parents = {int(child): parent for child, parent in lineage.items()}
tracks_df = _native_lineage_columns(
tracks_df, {int(t): parents.get(int(t), [])
for t in tracks_df['track_id'].unique()}, 'ultrack')
tracks_path = os.path.join(os.path.dirname(src), 'tracks')
os.makedirs(tracks_path, exist_ok=True)
tracks_df.to_csv(
os.path.join(tracks_path, f'ultrack_tracks_{object_type}_{name}.csv'),
index=False)
if plot or save:
_visualize_and_save_timelapse_stack_with_tracks(
masks_tracked, tracks_df, save, src, name, plot,
batch_filenames, object_type, mode)
return _masks_to_masks_stack(masks_tracked)
def _filter_short_tracks(df, min_length=5):
"""Filter out tracks that are shorter than min_length.
Args:
df (pandas.DataFrame): The input DataFrame containing track information.
min_length (int, optional): The minimum length of tracks to keep. Defaults to 5.
Returns:
pandas.DataFrame: The filtered DataFrame with only tracks longer than min_length.
"""
track_lengths = df.groupby('track_id').size()
long_tracks = track_lengths[track_lengths >= min_length].index
return df[df['track_id'].isin(long_tracks)]
#: What :func:`spacr.utils._map_wells` puts in every slot when it cannot read a
#: name at all. It does not raise, so the identity columns of the track table
#: carry this string and ``wellID`` has to agree with them.
_UNPARSED_KEY = 'error'
def _track_well_ids(row_ids, column_ids, name, logger):
"""Return one ``wellID`` per parsed ``(rowID, columnID)`` pair.
``wellID`` used to be ``file_name.str.split('_').str[1]``, which copied
whatever spelling the file used through while the ``rowID``/``columnID``
written beside it were canonicalised — a track table could name well
``'a1'`` next to ``'r1'``/``'c1'``. Composing it from the row and column
:func:`spacr.utils._map_wells` has already parsed makes the three agree.
(It does *not* repair a plate id containing an underscore:
:func:`spacr.schema.parse_field_stem` splits a field stem left to right, so
``'exp_plate1_A01_3'`` puts ``'plate1'`` in the well slot there too.)
What it must not do is raise. :func:`spacr.schema.well_id` deliberately
refuses two pairs ``_map_wells`` produces on names this pipeline has always
accepted, and an uncaught :class:`~spacr.schema.KeyParseError` here aborts a
whole timelapse run from the middle of a batch:
* a **positional well**. ``'plate1_5_3'`` is plate1, well 5, field 3;
:func:`spacr.schema.parse_well` passes a bare well number through into
*both* slots, so ``rowID == columnID == '5'`` and there is no
``A01``-style name to build for it. The well *is* ``'5'`` — which is what
the old positional split wrote too — so that is what is written.
* the ``'error'`` sentinel ``_map_wells`` returns for a name it cannot read.
Every other identity column on the row already carries it; ``wellID``
joins them rather than taking the run down, and the batch name is logged.
Anything else :func:`~spacr.schema.well_id` refuses (a well naming column
``0``, say) is a *malformed* well rather than an unnamed one, and is
re-raised with the batch name and the offending pair in the message — the
bare ``KeyParseError: cannot build a well name from column 'c0'`` that
surfaces from four frames down names neither.
:param row_ids: parsed ``rowID`` per row of the track table.
:param column_ids: parsed ``columnID`` per row, aligned to ``row_ids``.
:param name: npz batch name the ids were parsed from, for the messages.
:param logger: logger to report an unreadable batch name on.
:returns: list of ``wellID`` strings, one per pair.
:raises spacr.schema.KeyParseError: for a pair that is a well and is
malformed, naming the batch and the pair.
"""
well_ids = []
warned_unparsed = False
for row, column in zip(row_ids, column_ids):
if schema.is_positional_pair(row, column):
if str(row) == _UNPARSED_KEY and not warned_unparsed:
logger.warning(
"Could not read plate/well/field out of %r; the track "
"table for this batch carries %r in every identity "
"column, including wellID.", name, _UNPARSED_KEY)
warned_unparsed = True
well_ids.append(str(row))
continue
try:
well_ids.append(schema.well_id(row, column))
except schema.KeyParseError as error:
raise schema.KeyParseError(
f"cannot name the well of track batch {name!r}: {error} "
f"_map_wells read that name as row {row!r} / column "
f"{column!r}. Rename the batch to '<plate>_<well>_<field>' "
f"with a well of the form 'A01' (or a bare well number), so "
f"the track table names the well its own rowID/columnID "
f"name.") from error
return well_ids
@debug(enabled=True)
def _btrack_track_cells(src, name, batch_filenames, object_type, plot, save, masks_3D, mode, timelapse_remove_transient, radius=100, n_jobs=10, batch_list=None, optimizer_time_limit_s=120, optimizer_mip_gap=0.01, run_optimization=True, max_objects_for_optimization=20000):
"""
Track cells using the btrack library.
Args:
src (str): The source file path.
name (str): The name of the track (npz batch name).
batch_filenames (list[str]): Filenames for frames in this batch.
object_type (str): The type of object to track (cell, nucleus, pathogen).
plot (bool): Whether to plot the tracks.
save (bool): Whether to save plots.
masks_3D (ndarray or list): 3D label array of masks with shape (T, Y, X),
or list of 2D (Y, X) label arrays (one per frame).
mode (str): The tracking mode (unused here but kept for API consistency).
timelapse_remove_transient (bool): Whether to remove short tracks.
radius (int or None, optional): Max search radius (pixels). If None,
it is set automatically to image_width / 20. Defaults to 100.
n_jobs (int, optional): Number of workers for object extraction. Defaults to 10.
batch_list (list or None, optional): List of intensity images used by Cellpose.
Currently not used by btrack (tracking is shape-based here), but kept
for API compatibility and possible future use.
optimizer_time_limit_s (float or None): Time limit for GLPK in seconds
(used only when global optimisation is actually run).
optimizer_mip_gap (float or None): Relative MIP gap for GLPK (0.01 = 1%).
run_optimization (bool): If False, skip global optimisation entirely.
max_objects_for_optimization (int or None): If not None, skip global
optimisation when the number of objects exceeds this threshold.
Returns:
ndarray: The relabelled mask stack (same shape as masks_3D) where labels
are track IDs.
"""
import os
import logging
import numpy as np
import pandas as pd
try:
import btrack
from btrack import datasets as btrack_datasets
from btrack.constants import BayesianUpdates
except (ImportError, OSError) as exc:
raise ImportError(
"btrack tracking requires the optional btrack dependency. "
"Install it with `pip install 'spacr[btrack]'`, or select "
"trackpy, Trackastra, or IoU tracking instead."
) from exc
from .plot import _visualize_and_save_timelapse_stack_with_tracks
from .utils import _masks_to_masks_stack, _map_wells
logger = logging.getLogger(__name__)
logger.debug(
"Entering _btrack_track_cells: name=%s, object_type=%s, mode=%s",
name,
object_type,
mode,
)
logger.debug("src=%s", src)
logger.debug(
"timelapse_remove_transient=%s, radius=%s, n_jobs=%s, "
"optimizer_time_limit_s=%s, optimizer_mip_gap=%s, run_optimization=%s, "
"max_objects_for_optimization=%s",
timelapse_remove_transient,
radius,
n_jobs,
optimizer_time_limit_s,
optimizer_mip_gap,
run_optimization,
max_objects_for_optimization,
)
logger.debug("masks_3D type: %s", type(masks_3D))
if isinstance(masks_3D, list):
logger.debug("masks_3D is a list with length=%d", len(masks_3D))
if len(masks_3D) == 0:
raise ValueError("masks_3D is an empty list; nothing to track.")
masks_3D = [np.asarray(m) for m in masks_3D]
shapes = {m.shape for m in masks_3D}
logger.debug("Unique mask shapes in list: %s", shapes)
if len(shapes) != 1:
raise ValueError(
f"All masks must have the same shape; got shapes={shapes}."
)
masks_3D = np.stack(masks_3D, axis=0)
logger.debug("Stacked masks_3D into ndarray with shape %s", masks_3D.shape)
else:
masks_3D = np.asarray(masks_3D)
logger.debug("masks_3D array shape: %s", masks_3D.shape)
if masks_3D.ndim != 3:
raise ValueError(
f"masks_3D must be 3D (T, Y, X); got shape {masks_3D.shape}"
)
n_frames, height, width = masks_3D.shape
logger.debug(
"Parsed geometry: n_frames=%d, height=%d, width=%d",
n_frames,
height,
width,
)
if radius is None:
radius = max(1, width // 20)
logger.debug(
"radius was None; automatically set radius=%d (width/20)", radius
)
FEATURES = [
"area",
"major_axis_length",
"minor_axis_length",
"orientation",
"solidity",
]
TRACKING_UPDATES = ["motion", "visual"]
logger.debug("Converting segmentation to btrack objects...")
from .resource_log import _guard_workers
n_jobs = _guard_workers('mask', n_jobs, int(np.asarray(masks_3D[0]).nbytes)
if len(masks_3D) else 0)
objects = btrack.utils.segmentation_to_objects(
masks_3D,
properties=tuple(FEATURES),
num_workers=n_jobs,
)
n_objects = len(objects)
logger.info("Extracted %d objects for tracking.", n_objects)
if n_objects == 0:
columns = (
"track_id", "frame", "x", "y", "original_label", "file_name",
"plateID", "rowID", "columnID", "fieldID", "prcf", "wellID",
)
final_df = pd.DataFrame({
column: pd.Series(dtype=float if column in {
"track_id", "frame", "x", "y", "original_label"
} else "object")
for column in columns
})
masks = [np.zeros_like(frame) for frame in masks_3D]
tracks_path = os.path.join(os.path.dirname(src), "tracks")
os.makedirs(tracks_path, exist_ok=True)
out_csv = os.path.join(
tracks_path, f"btrack_tracks_{object_type}_{name}.csv")
final_df.to_csv(out_csv, index=False)
if plot or save:
_visualize_and_save_timelapse_stack_with_tracks(
masks, final_df, save, src, name, plot, batch_filenames,
object_type, mode,
)
return _masks_to_masks_stack(masks)
CONFIG_FILE = btrack_datasets.cell_config()
with btrack.BayesianTracker() as tracker:
tracker.configure(CONFIG_FILE)
tracker.update_method = BayesianUpdates.APPROXIMATE
tracker.max_search_radius = radius
tracker.features = FEATURES
tracker.append(objects)
tracker.volume = ((0, width), (0, height))
logger.debug(
"Tracker volume set to x=(0,%d), y=(0,%d); update_method=%s",
width,
height,
tracker.update_method,
)
logger.debug("Starting tracking...")
try:
tracker.track(tracking_updates=TRACKING_UPDATES)
except TypeError:
logger.debug(
"tracker.track(tracking_updates=...) not supported; "
"falling back to tracker.tracking_updates + track(step_size=100)."
)
tracker.tracking_updates = [u.upper() for u in TRACKING_UPDATES]
tracker.track(step_size=100)
logger.info(
"Tracking complete. Number of tracks before optimisation: %d",
len(tracker.tracks),
)
do_optimize = bool(run_optimization)
if max_objects_for_optimization is not None and n_objects > max_objects_for_optimization:
logger.warning(
"Skipping btrack global optimisation: %d objects > "
"max_objects_for_optimization=%d. Using pre-optimisation tracks.",
n_objects,
max_objects_for_optimization,
)
do_optimize = False
if do_optimize and len(tracker.tracks) > 0:
glpk_options = {}
if optimizer_time_limit_s is not None and optimizer_time_limit_s > 0:
glpk_options["tm_lim"] = int(optimizer_time_limit_s * 1000)
if optimizer_mip_gap is not None and optimizer_mip_gap > 0:
glpk_options["mip_gap"] = float(optimizer_mip_gap)
try:
if glpk_options:
logger.info(
"Running GLPK optimisation with options: %s", glpk_options
)
tracker.optimize(
backend="glpk",
options={"options": glpk_options},
)
else:
logger.info("Running GLPK optimisation with default options.")
tracker.optimize(backend="glpk")
logger.info(
"Optimisation complete. Number of tracks after optimisation: %d",
len(tracker.tracks),
)
except Exception as e:
logger.warning(
"btrack global optimisation failed or stalled (%s). "
"Using pre-optimisation tracks instead.",
e,
exc_info=True,
)
tracks = tracker.tracks
native_parents = {int(track.ID): track.parent for track in tracks
if hasattr(track, 'parent')}
track_data = []
for track in tracks:
for t, x, y, z in zip(track.t, track.x, track.y, track.z):
track_data.append(
{
"track_id": track.ID,
"frame": t,
"x": x,
"y": y,
"z": z,
}
)
tracks_df = pd.DataFrame(track_data)
logger.debug("tracks_df shape: %s", tracks_df.shape)
if timelapse_remove_transient and not tracks_df.empty:
logger.debug("Removing transient tracks with min_length=%d", n_frames)
tracks_df = _filter_short_tracks(tracks_df, min_length=n_frames)
logger.debug("tracks_df shape after filtering: %s", tracks_df.shape)
if tracks_df.empty:
logger.warning(
"btrack produced no usable tracks for %s (%s); "
"returning an untracked mask stack.", name, object_type,
)
tracks_df = pd.DataFrame(
{col: pd.Series(dtype=float) for col in ("track_id", "frame", "x", "y", "z")}
)
logger.debug("Preparing objects_df from masks_3D...")
objects_df = _prepare_for_tracking(masks_3D)
logger.debug("objects_df shape: %s", objects_df.shape)
tracks_df["x"] = tracks_df["x"].round(2)
tracks_df["y"] = tracks_df["y"].round(2)
objects_df["x"] = objects_df["x"].round(2)
objects_df["y"] = objects_df["y"].round(2)
logger.debug("Merging tracks_df and objects_df on ['frame', 'x', 'y']...")
merged_df = pd.merge(
tracks_df,
objects_df,
on=["frame", "x", "y"],
how="inner",
validate="many_to_many",
)
logger.debug("merged_df shape: %s", merged_df.shape)
final_df = merged_df[["track_id", "frame", "x", "y", "original_label"]].copy()
final_df = _native_lineage_columns(final_df, native_parents, 'btrack')
if final_df.empty:
logger.warning(
"No tracks remained after filtering/merging for %s; "
"writing an empty track table.", name
)
for col in ('file_name', 'plateID', 'rowID', 'columnID', 'fieldID', 'prcf', 'wellID'):
final_df[col] = pd.Series(dtype='object')
else:
try:
final_df['file_name'] = name
final_df[['plateID', 'rowID', 'columnID', 'fieldID', 'prcf']] = (final_df['file_name'].apply(lambda fname: pd.Series(_map_wells(fname, timelapse=False))))
final_df['wellID'] = _track_well_ids(
final_df['rowID'], final_df['columnID'], name, logger)
except IndexError:
logger.warning("Failed to parse plate, well, field from name: %s", name)
logger.debug("Relabelling masks based on tracks...")
masks = _relabel_masks_based_on_tracks(masks_3D, final_df)
tracks_path = os.path.join(os.path.dirname(src), "tracks")
os.makedirs(tracks_path, exist_ok=True)
out_csv = os.path.join(tracks_path, f"btrack_tracks_{object_type}_{name}.csv")
logger.debug("Saving track table to %s", out_csv)
from .tabular import write_table
write_table(final_df, out_csv)
if plot or save:
logger.debug("Generating visualisation (plot=%s, save=%s)...", plot, save)
_visualize_and_save_timelapse_stack_with_tracks(
masks,
final_df,
save,
src,
name,
plot,
batch_filenames,
object_type,
mode,
)
mask_stack = _masks_to_masks_stack(masks)
logger.debug(
"Finished _btrack_track_cells. mask_stack shape: %s",
getattr(mask_stack, "shape", None),
)
return mask_stack
from ._lineage_trees import (
_LINEAGE_SEGMENT_STATS,
_lineage_explicit_parents,
_lineage_inferred_parents,
_lineage_segments,
_lineage_colour_values,
_sibling_correlation,
_lineage_statistics,
_lineage_newick,
_lineage_tree_figure,
_lineage_calibrate_time,
_lineage_hour_statistics,
_lineage_trees_from_tracks,
_run_lineage_step,
)
_EVENT_BACKGROUND = 'none'
_EVENT_CROP = 16
_EVENT_SHAPE_COLUMNS = ('area', 'eccentricity', 'solidity')
def _event_frame_features(mask_stack, images=None, crop=_EVENT_CROP):
"""Shape and intensity of every tracked object in every frame, with crops.
:param mask_stack: ``(T, H, W)`` labels, each object labelled with its
track id.
:param images: ``(T, H, W)`` or ``(T, H, W, C)`` intensities, or None.
:param crop: side of the stored crop; the crop covers twice this many
pixels around the object's centre and is averaged down two-fold.
:returns: ``(features, crops)``: one row per object and frame with
``frame``, ``track_id``, ``x``, ``y``, ``area``, ``eccentricity``,
``solidity`` and, with images, ``intensity_mean_c<k>`` and
``intensity_max_c<k>``; and a ``(rows, C, crop, crop)`` float16
array in the same order, or None without images. Intensities are
divided by each channel's 99.5th percentile over the movie.
"""
masks = np.asarray(mask_stack)
imgs = None
if images is not None:
imgs = np.asarray(images, dtype=np.float32)
if imgs.ndim == 3:
imgs = imgs[..., None]
scale = np.percentile(imgs.reshape(-1, imgs.shape[-1]), 99.5, axis=0)
imgs = imgs / np.where(scale > 0, scale, 1.0)
rows, crops = [], []
for t in range(masks.shape[0]):
lab = masks[t].astype(np.int32)
if not lab.any():
continue
properties = ['label', 'centroid', 'area', 'eccentricity', 'solidity']
if imgs is not None:
properties += ['intensity_mean', 'intensity_max']
props = pd.DataFrame(regionprops_table(
lab, intensity_image=None if imgs is None else imgs[t],
properties=properties))
props = props.rename(columns={'label': 'track_id', 'centroid-0': 'y',
'centroid-1': 'x'})
props = props.rename(columns=lambda c: re.sub(
r'^(intensity_(?:mean|max))-(\d+)$', r'\1_c\2', c))
props.insert(0, 'frame', t)
rows.append(props)
if imgs is not None:
half = crop
padded = np.pad(imgs[t], ((half, half), (half, half), (0, 0)))
for y, x in zip(props['y'], props['x']):
r, c = int(round(y)) + half, int(round(x)) + half
patch = padded[r - half:r + half, c - half:c + half]
patch = patch.reshape(crop, 2, crop, 2, -1).mean(axis=(1, 3))
crops.append(np.moveaxis(patch, -1, 0).astype(np.float16))
if not rows:
return pd.DataFrame(columns=['frame', 'track_id', 'x', 'y']), None
features = pd.concat(rows, ignore_index=True)
return features, (np.stack(crops) if crops else None)
def _event_track_table(tracks, features=None, radius=30.0, partners=None):
"""Per-frame inputs of the event detector for one field's tracks.
Joins the tracks with the frame features of
:func:`_event_frame_features` where there are any and adds what the
tracks themselves say: speed, the log change in area, whether the track
starts or ends in this frame (other than at the movie's first or last
frame), how many other tracks end within ``radius`` pixels in the same
frame and how many new tracks start within ``radius`` pixels in the
next frame.
:param tracks: tracks table with ``frame``, ``track_id``, ``x`` and ``y``.
:param features: optional frame features keyed by ``frame`` and
``track_id``.
:param radius: neighbourhood radius in pixels.
:param partners: optional ``{object_type: tracks}`` of other tracked
objects in the same field, for events across two object types
(parasites leaving or entering a host). For each, the table gains
``<object>_near`` (its objects within ``radius`` pixels in this
frame), ``<object>_near_change`` (the change from the previous
frame), ``<object>_starts_near`` (its tracks starting within
``radius`` in the next frame) and ``<object>_ends_near`` (its tracks
ending within ``radius`` in this frame).
:returns: the table, sorted by track and frame.
"""
df = tracks.dropna(subset=['track_id', 'frame', 'x', 'y']).copy()
df['track_id'] = df['track_id'].astype(int)
df['frame'] = df['frame'].astype(int)
if features is not None and len(features):
extra = features.drop(columns=[c for c in ('x', 'y') if c in features.columns])
extra = extra.astype({'frame': int, 'track_id': int})
keep = [c for c in df.columns if c in ('frame', 'track_id', 'x', 'y',
'parent_track_id')]
df = df[keep].merge(extra, on=['frame', 'track_id'], how='left')
df = df.drop_duplicates(['track_id', 'frame']).sort_values(['track_id', 'frame'])
first, last = df['frame'].min(), df['frame'].max()
group = df.groupby('track_id')
step = np.hypot(group['x'].diff(), group['y'].diff()) / group['frame'].diff()
df['speed'] = step.fillna(0.0)
if 'area' in df.columns:
df['log_area_change'] = np.log(df['area'].clip(lower=1)).groupby(
df['track_id']).diff().fillna(0.0)
starts = group['frame'].transform('min')
ends = group['frame'].transform('max')
df['track_starts'] = ((df['frame'] == starts) & (df['frame'] > first)).astype(float)
df['track_ends'] = ((df['frame'] == ends) & (df['frame'] < last)).astype(float)
born = df[df['track_starts'] > 0]
dying = df[df['track_ends'] > 0]
new_near, end_near = [], []
for frame, x, y, track in zip(df['frame'], df['x'], df['y'], df['track_id']):
b = born[(born['frame'] == frame + 1) & (born['track_id'] != track)]
d = dying[(dying['frame'] == frame) & (dying['track_id'] != track)]
new_near.append(int((np.hypot(b['x'] - x, b['y'] - y) <= radius).sum()))
end_near.append(int((np.hypot(d['x'] - x, d['y'] - y) <= radius).sum()))
df['new_tracks_near'] = new_near
df['ended_tracks_near'] = end_near
for name, other in sorted((partners or {}).items()):
df = _event_partner_columns(df, other, name, radius)
return df.reset_index(drop=True)
def _event_partner_columns(df, other, name, radius):
"""Add the neighbourhood counts of another object type's tracks.
See ``partners`` in :func:`_event_track_table`.
"""
o = other.dropna(subset=['track_id', 'frame', 'x', 'y']).copy()
o['frame'] = o['frame'].astype(int)
span = o.groupby('track_id')['frame']
o_first, o_last = o['frame'].min(), o['frame'].max()
starts = o[(o['frame'] == span.transform('min')) & (o['frame'] > o_first)]
ends = o[(o['frame'] == span.transform('max')) & (o['frame'] < o_last)]
by_frame = {f: g[['x', 'y']].to_numpy(float) for f, g in o.groupby('frame')}
s_frame = {f: g[['x', 'y']].to_numpy(float) for f, g in starts.groupby('frame')}
e_frame = {f: g[['x', 'y']].to_numpy(float) for f, g in ends.groupby('frame')}
empty = np.zeros((0, 2))
def count(points, x, y):
"""Number of ``points`` within ``radius`` of ``(x, y)``."""
return int((np.hypot(points[:, 0] - x, points[:, 1] - y) <= radius).sum())
near, born, gone = [], [], []
for frame, x, y in zip(df['frame'], df['x'], df['y']):
near.append(count(by_frame.get(frame, empty), x, y))
born.append(count(s_frame.get(frame + 1, empty), x, y))
gone.append(count(e_frame.get(frame, empty), x, y))
df[f'{name}_near'] = near
df[f'{name}_near_change'] = df[f'{name}_near'].groupby(df['track_id']).diff().fillna(0.0)
df[f'{name}_starts_near'] = born
df[f'{name}_ends_near'] = gone
return df
def _event_columns(table):
"""The numeric inputs of the detector found in a track table.
:param table: from :func:`_event_track_table`.
:returns: column names, in a fixed order.
"""
fixed = ['speed', 'track_starts', 'track_ends', 'new_tracks_near',
'ended_tracks_near', 'log_area_change', *_EVENT_SHAPE_COLUMNS]
found = [c for c in fixed if c in table.columns]
found += sorted(c for c in table.columns if c.startswith('intensity_'))
found += sorted(c for c in table.columns if c not in found and c.endswith(
('_near', '_near_change', '_starts_near', '_ends_near'))
and c not in ('new_tracks_near', 'ended_tracks_near'))
return found
def _event_windows(table, columns, window, mean, std, crops=None):
"""Cut every track into windows centred on each of its frames.
:param table: from :func:`_event_track_table`, one field.
:param columns: detector inputs; one missing from the table reads 0.
:param window: frames per window (odd; an even value is raised by one).
:param mean: per-column mean used to standardise.
:param std: per-column spread used to standardise.
:param crops: optional ``{(frame, track_id): (C, P, P) array}``.
:returns: ``(inputs, crop_windows, index)``: a ``(n, len(columns) + 1,
window)`` float32 array whose last row marks frames where the track
is present; a ``(n, window, C, P, P)`` array or None; and the
``track_id`` and ``frame`` of each window's centre.
"""
half = int(window) // 2
window = 2 * half + 1
values = np.zeros((len(table), len(columns)), dtype=np.float32)
for k, name in enumerate(columns):
if name in table.columns:
values[:, k] = pd.to_numeric(table[name], errors='coerce').to_numpy(dtype=np.float32)
values = (values - np.asarray(mean, dtype=np.float32)) / np.asarray(std, dtype=np.float32)
values = np.nan_to_num(values)
sample = next(iter(crops.values())) if crops else None
inputs, crop_windows, index = [], [], []
frames_all = table['frame'].to_numpy()
for track, rows in table.groupby('track_id', sort=False).indices.items():
frames = frames_all[rows]
start, stop = frames.min() - half, frames.max() + half
dense = np.zeros((stop - start + 1, len(columns) + 1), dtype=np.float32)
dense[frames - start, :-1] = values[rows]
dense[frames - start, -1] = 1.0
view = np.lib.stride_tricks.sliding_window_view(dense, window, axis=0)
inputs.append(view[frames - frames.min()])
index.extend((int(track), int(f)) for f in frames)
if sample is not None:
stack = np.zeros((stop - start + 1,) + sample.shape, dtype=np.float32)
for f in range(start, stop + 1):
patch = crops.get((f, int(track)))
if patch is not None:
stack[f - start] = patch
crop_windows.append(np.stack([stack[f - frames.min():f - frames.min() + window]
for f in frames]))
if not inputs:
return (np.zeros((0, len(columns) + 1, window), np.float32), None,
pd.DataFrame(columns=['track_id', 'frame']))
return (np.concatenate(inputs).astype(np.float32),
np.concatenate(crop_windows) if crop_windows else None,
pd.DataFrame(index, columns=['track_id', 'frame']))
def _event_network(n_inputs, n_classes, channels=0, video_width=0):
"""The event classifier: a small image encoder and a temporal convolution.
Each frame's crop, when there are crops, is encoded by two convolutions
to 16 numbers that join the frame's track features; two temporal
convolutions over the window, pooled by mean and maximum, feed a linear
layer with one output per class.
:param n_inputs: track inputs per frame, the presence mark included.
:param n_classes: event classes, background included.
:param channels: crop channels, 0 for none.
:param video_width: frozen clip feature width, zero for the original encoder.
:returns: a module called as ``net(inputs, crops=None, video=None)``.
"""
import torch
from torch import nn
class _EventNet(nn.Module):
"""Temporal convolutional classifier of track windows."""
def __init__(self):
"""Build temporal features and the optional per-frame crop encoder."""
super().__init__()
self.encoder = None
width = n_inputs
if channels and not video_width:
self.encoder = nn.Sequential(
nn.Conv2d(channels, 8, 3, padding=1), nn.ReLU(),
nn.MaxPool2d(2), nn.Conv2d(8, 16, 3, padding=1), nn.ReLU(),
nn.AdaptiveAvgPool2d(1), nn.Flatten())
width += 16
self.temporal = nn.Sequential(
nn.Conv1d(width, 32, 3, padding=1), nn.ReLU(),
nn.Conv1d(32, 32, 3, padding=1), nn.ReLU())
self.head = nn.Linear(64 + video_width, n_classes)
def forward(self, inputs, crops=None, video=None):
"""Class scores for a batch of windows."""
if self.encoder is not None and crops is not None:
b, w = crops.shape[:2]
code = self.encoder(crops.reshape((b * w,) + crops.shape[2:]))
inputs = torch.cat([inputs, code.reshape(b, w, -1).transpose(1, 2)], dim=1)
hidden = self.temporal(inputs)
pooled = torch.cat([hidden.mean(dim=2), hidden.amax(dim=2)], dim=1)
if video_width:
if video is None or video.shape != (len(inputs), video_width):
raise ValueError("The event classifier requires matching pretrained clip features")
pooled = torch.cat([pooled, video], dim=1)
return self.head(pooled)
return _EventNet()
def _event_labels(index, field, annotations, classes, radius=1):
"""Training label of each window: the event at its centre, else background.
A window within ``radius`` frames of an annotated event is labelled with
it; windows one frame further out are given weight 0, so the detector
is not taught that the frames just around an event are background.
:param index: ``track_id`` and ``frame`` of each window's centre.
:param field: the field these windows come from.
:param annotations: ``field``, ``track_id``, ``frame``, ``event``.
:param classes: class names, background first.
:param radius: frames either side labelled as the event.
:returns: ``(labels, weights)`` arrays.
"""
labels = np.zeros(len(index), dtype=np.int64)
weights = np.ones(len(index), dtype=np.float32)
ann = annotations[annotations['field'] == field]
code = {name: k for k, name in enumerate(classes)}
lookup = index.reset_index(drop=True)
for track, frame, event in zip(ann['track_id'], ann['frame'], ann['event']):
near = (lookup['track_id'] == int(track)).to_numpy()
offset = np.abs(lookup['frame'].to_numpy() - int(frame))
labels[near & (offset <= radius)] = code[event]
weights[near & (offset == radius + 1) & (labels == 0)] = 0.0
return labels, weights
def _event_train(samples, classes, columns, *, window, channels=0,
epochs=25, seed=0, video_width=0):
"""Train the event classifier on windows from annotated fields.
:param samples: list of ``(inputs, crops, labels, weights)`` per field.
:param classes: class names, background first.
:param columns: the track inputs, stored with the model.
:param window: frames per window.
:param channels: crop channels, 0 for none.
:param epochs: passes over the windows.
:param video_width: pretrained clip width; samples then include a fifth
array of frozen features. Only the event classifier is trained.
:param seed: random seed. Training uses at most four CPU threads, which
is faster for a network this small than many.
:returns: the model as a dict: ``state`` (weights), ``classes``,
``columns``, ``mean``, ``std``, ``window`` and ``channels``.
"""
import torch
torch.manual_seed(seed)
rng = np.random.default_rng(seed)
inputs = np.concatenate([s[0] for s in samples])
crops = (np.concatenate([s[1] for s in samples]).astype(np.float32)
if channels and not video_width else None)
videos = (np.concatenate([s[4] for s in samples]).astype(np.float32)
if video_width else None)
labels = np.concatenate([s[2] for s in samples])
weights = np.concatenate([s[3] for s in samples])
counts = np.bincount(labels[weights > 0], minlength=len(classes)).astype(float)
class_weight = np.where(counts > 0, (counts.sum() / np.maximum(counts, 1)) ** 0.5, 0.0)
net = _event_network(inputs.shape[1], len(classes), channels, video_width)
optimiser = torch.optim.Adam(net.parameters(), lr=3e-3, weight_decay=1e-4)
loss_fn = torch.nn.CrossEntropyLoss(
weight=torch.tensor(class_weight, dtype=torch.float32), reduction='none')
x_all = torch.from_numpy(inputs)
c_all = torch.from_numpy(crops) if crops is not None else None
v_all = torch.from_numpy(videos) if videos is not None else None
y_all = torch.from_numpy(labels)
w_all = torch.from_numpy(weights)
net.train()
threads = torch.get_num_threads()
torch.set_num_threads(min(threads, 4))
try:
for _ in range(int(epochs)):
order = rng.permutation(len(labels))
for start in range(0, len(order), 256):
pick = torch.from_numpy(order[start:start + 256])
scores = net(x_all[pick], None if c_all is None else c_all[pick],
None if v_all is None else v_all[pick])
loss = ((loss_fn(scores, y_all[pick]) * w_all[pick]).sum()
/ w_all[pick].sum().clamp(min=1.0))
optimiser.zero_grad()
loss.backward()
optimiser.step()
finally:
torch.set_num_threads(threads)
return {'state': {k: v.detach().clone() for k, v in net.state_dict().items()},
'classes': list(classes), 'columns': list(columns),
'window': int(window), 'channels': int(channels),
'video_width': int(video_width)}
def _event_probabilities(model, inputs, crops=None, video=None):
"""Class probabilities of each window.
:param model: from :func:`_event_train` or :func:`_event_load_model`.
:param inputs: windows from :func:`_event_windows`.
:param crops: their crops, or None.
:param video: matching frozen clip features for a pretrained event model.
:returns: ``(n, classes)`` array.
"""
import torch
video_width = int(model.get('video_width', 0))
if video_width and (video is None or video.shape != (len(inputs), video_width)):
raise ValueError("The saved event model requires matching pretrained clip features")
net = _event_network(inputs.shape[1], len(model['classes']), model['channels'], video_width)
net.load_state_dict(model['state'])
net.eval()
out = []
with torch.no_grad():
for start in range(0, len(inputs), 1024):
x = torch.from_numpy(inputs[start:start + 1024])
c = None
if model['channels'] and crops is not None and not video_width:
c = torch.from_numpy(crops[start:start + 1024].astype(np.float32))
v = None if video is None else torch.from_numpy(
video[start:start + 1024].astype(np.float32))
out.append(torch.softmax(net(x, c, v), dim=1).numpy())
return np.concatenate(out) if out else np.zeros((0, len(model['classes'])))
def _event_peaks(index, probabilities, classes, threshold=0.5, tolerance=2):
"""Time-stamp events from per-frame class probabilities.
An event is a frame where its class's probability reaches ``threshold``
and is the highest of that track within ``tolerance`` frames. A class
that can happen to a cell only once (its name contains ``death`` or
``lysis``) keeps only its most probable peak per track: a dead cell
stays in the field and would otherwise be called dead again and again.
:param index: ``track_id`` and ``frame`` of each window.
:param probabilities: from :func:`_event_probabilities`.
:param classes: class names, background first.
:param threshold: smallest probability kept.
:param tolerance: frames either side a peak must beat.
:returns: one row per event: ``track_id``, ``frame``, ``event`` and
``probability``.
"""
events = []
frame = index.reset_index(drop=True).assign(_row=np.arange(len(index)))
for track, rows in frame.groupby('track_id'):
rows = rows.sort_values('frame')
frames = rows['frame'].to_numpy()
for k, name in enumerate(classes[1:], start=1):
p = probabilities[rows['_row'].to_numpy(), k]
for i in np.flatnonzero(p >= threshold):
near = np.abs(frames - frames[i]) <= tolerance
if p[i] >= p[near].max() and not (
(p[near] == p[i]) & (frames[near] < frames[i])).any():
events.append({'track_id': int(track), 'frame': int(frames[i]),
'event': name, 'probability': float(p[i])})
found = pd.DataFrame(events, columns=['track_id', 'frame', 'event', 'probability'])
once = found['event'].map(_event_is_terminal).astype(bool)
if once.any():
best = found[once].sort_values(['probability', 'frame'], ascending=[False, True])
best = best.drop_duplicates(['track_id', 'event'])
found = pd.concat([found[~once], best]).sort_values(['track_id', 'frame'])
found = found.reset_index(drop=True)
return found
def _event_is_terminal(name):
"""Whether an event class can happen to one tracked object only once.
:param name: the class name.
:returns: True for names containing ``death`` or ``lysis``.
"""
text = str(name).lower()
return 'death' in text or 'lysis' in text
def _event_scores(detected, annotations, classes, tolerance=2):
"""Precision, recall and timing error of detected against annotated events.
A detection matches the nearest unmatched annotation of the same field,
track and class within ``tolerance`` frames, most probable detection
first.
:param detected: ``field``, ``track_id``, ``frame``, ``event``,
``probability``.
:param annotations: ``field``, ``track_id``, ``frame``, ``event``.
:param classes: event classes to score.
:param tolerance: largest timing error of a match, in frames.
:returns: one row per class and an ``all`` row: ``annotated``,
``detected``, ``true_positives``, ``precision``, ``recall``, ``f1``,
``mean_abs_timing_error`` and ``max_abs_timing_error`` in frames.
"""
rows, all_errors, totals = [], [], np.zeros(3, dtype=int)
for name in classes:
det = detected[detected['event'] == name].sort_values('probability', ascending=False)
ann = annotations[annotations['event'] == name]
free = {key: sorted(g['frame'].astype(int)) for key, g in
ann.groupby(['field', 'track_id'])}
errors = []
for field, track, frame in zip(det['field'], det['track_id'], det['frame']):
pool = free.get((field, int(track)), [])
if pool:
best = min(pool, key=lambda a: abs(a - int(frame)))
if abs(best - int(frame)) <= tolerance:
pool.remove(best)
errors.append(abs(best - int(frame)))
counts = np.array([len(ann), len(det), len(errors)])
totals += counts
all_errors += errors
rows.append(_event_score_row(name, counts, errors))
rows.append(_event_score_row('all', totals, all_errors))
return pd.DataFrame(rows)
def _event_score_row(name, counts, errors):
"""One row of :func:`_event_scores`."""
annotated, detected, hits = (int(c) for c in counts)
precision = hits / detected if detected else np.nan
recall = hits / annotated if annotated else np.nan
f1 = (2 * precision * recall / (precision + recall)
if hits else (0.0 if annotated or detected else np.nan))
return {'event': name, 'annotated': annotated, 'detected': detected,
'true_positives': hits, 'precision': precision, 'recall': recall,
'f1': f1,
'mean_abs_timing_error': float(np.mean(errors)) if errors else np.nan,
'max_abs_timing_error': float(np.max(errors)) if errors else np.nan}
def _event_read_annotations(path):
"""Read an annotated events table.
:param path: a table with ``field``, ``track_id``, ``frame`` and
``event`` columns, optionally ``object`` and tracker provenance.
:returns: the table with the event names stripped and lower-cased.
:raises ValueError: a missing column or repeated track frame.
"""
from .tabular import read_table
from .qt.i18n import tr
ann = read_table(path, canonicalise=False, report=None)
if 'field' not in ann.columns and 'fieldID' in ann.columns:
ann = ann.rename(columns={'fieldID': 'field'})
missing = {'field', 'track_id', 'frame', 'event'} - set(ann.columns)
if missing:
raise ValueError(f"Setting: timelapse_events_annotations {path} lacks "
f"the column(s) {', '.join(sorted(missing))}.")
ann = ann.dropna(subset=['field', 'track_id', 'frame', 'event']).copy()
ann['field'] = ann['field'].astype(str)
ann['track_id'] = ann['track_id'].astype(int)
ann['frame'] = ann['frame'].astype(int)
ann['event'] = ann['event'].astype(str).str.strip().str.lower()
ann = ann[ann['event'] != _EVENT_BACKGROUND].reset_index(drop=True)
keys = ann[['field', 'track_id', 'frame']].copy()
keys['object'] = (ann['object'].fillna('').astype(str)
if 'object' in ann else '')
if keys.duplicated().any():
raise ValueError(tr('An annotation table labels one track frame more than once.'))
return ann
def _event_source_hash(path):
"""Hash one tracker export without holding its whole table in memory."""
digest = hashlib.sha256()
with open(path, 'rb') as source:
for chunk in iter(lambda: source.read(1024 * 1024), b''):
digest.update(chunk)
return digest.hexdigest()
def _event_fields(tracks_dir, object_type, prefix):
"""The tracks tables of one object and tracker, by field name.
:returns: ``{field: path}``.
"""
marker = f'{prefix}_tracks_{object_type}_'
found = {}
for path in sorted(glob.glob(os.path.join(tracks_dir, f'{marker}*.csv'))):
found[os.path.splitext(os.path.basename(path))[0][len(marker):]] = path
return found
def _event_field_inputs(tracks_path, radius, partners=None):
"""The detector's track table and crops of one field.
Reads ``events/<stem>_features.csv`` and ``events/<stem>_crops.npz``
beside the tracks table when the tracking step wrote them.
:param partners: optional ``{object_type: tracks path}`` of the other
objects tracked in this field (see :func:`_event_track_table`).
:returns: ``(table, crops)``; ``crops`` maps ``(frame, track_id)`` to
the crop, or is None.
"""
from .tabular import read_table
tracks = read_table(tracks_path, report=None)
stem = os.path.splitext(os.path.basename(tracks_path))[0]
base = os.path.join(os.path.dirname(tracks_path), 'events', stem)
features = crops = None
if os.path.isfile(base + '_features.csv'):
features = read_table(base + '_features.csv', report=None)
if os.path.isfile(base + '_crops.npz'):
with np.load(base + '_crops.npz') as store:
crops = {(int(f), int(t)): c for (f, t), c in
zip(store['index'], store['crops'])}
others = {k: read_table(v, report=None) for k, v in (partners or {}).items()}
return _event_track_table(tracks, features, radius=radius,
partners=others or None), crops
def _event_clip_features(video, crops, index, crop_source, window):
"""Encode missing frozen windows and reuse them across held-out folds.
Cache entries identify the original crop collection, window size, track
and frame. No fitted normalization, annotations or event heads are cached.
:returns: finite B,768 vectors and verified checkpoint provenance.
:raises ValueError: crops or the explicitly declared checkpoint/map are missing.
"""
from ._segmentation_backends import (
_event_video_features, _event_video_identity, _event_video_inputs)
import hashlib
if crops is None:
raise ValueError("VideoMAE event detection requires saved image crops for every field")
_, channels = _event_video_inputs(crops, video.get('channel_map'))
observed = {(int(frame), int(track))
for track, frame in zip(index['track_id'], index['frame'])}
if not observed <= set(crop_source):
raise ValueError("VideoMAE requires an image crop for every observed track frame")
identity = _event_video_identity(video['model_dir'])
cache = video.setdefault('_cache', {})
keys = [(id(crop_source), int(window), tuple(channels),
str(video.get('device', 'auto')), int(track), int(frame),
hashlib.sha256(crops[i].tobytes()).hexdigest())
for i, (track, frame) in enumerate(zip(index['track_id'], index['frame']))]
missing = [i for i, key in enumerate(keys) if key not in cache]
if missing:
features, provenance = _event_video_features(
crops[missing], video['model_dir'], video['channel_map'],
device=video.get('device', 'auto'))
for i, feature in zip(missing, features):
cache[keys[i]] = feature.copy()
video['_provenance'] = provenance
provenance = dict(video.get('_provenance') or {})
if (any(provenance.get(key) != value for key, value in identity.items())
or provenance.get('channel_map') != list(video['channel_map'])):
raise ValueError("Event video feature cache does not match the checkpoint/channel mapping")
features = np.stack([cache[key] for key in keys]) if keys else np.empty((0, 768), np.float32)
return features, provenance
def _event_fit(tables, annotations, *, window=9, epochs=25, seed=0, video=None):
"""Fit the detector on annotated fields.
:param tables: ``{field: (table, crops)}`` from
:func:`_event_field_inputs`.
:param annotations: events of these fields; every event of a field
that has any annotation is taken to be annotated.
:param video: optional pretrained configuration with model_dir,
channel_map and device. None retains the original small CPU encoder.
:returns: the model dict of :func:`_event_train` with ``mean`` and
``std``.
"""
classes = [_EVENT_BACKGROUND] + sorted(pd.unique(annotations['event']))
columns = _event_columns(pd.concat([t for t, _ in tables.values()]))
stacked = pd.concat([t for t, _ in tables.values()])
mean, std = [], []
for name in columns:
v = pd.to_numeric(stacked[name], errors='coerce') if name in stacked else pd.Series([0.0])
mean.append(float(np.nan_to_num(v.mean())))
std.append(float(v.std()) if np.isfinite(v.std()) and v.std() > 0 else 1.0)
use_crops = all(c for _, c in tables.values())
if video is not None and not use_crops:
raise ValueError("VideoMAE event training requires image crops in every annotated field")
channels = next(iter(next(iter(tables.values()))[1].values())).shape[0] if use_crops else 0
samples = []
for field, (table, crops) in tables.items():
x, c, index = _event_windows(table, columns, window, mean, std,
crops if use_crops else None)
y, w = _event_labels(index, field, annotations, classes)
sample = (x, c, y, w)
if video is not None:
features, provenance = _event_clip_features(video, c, index, crops, window)
sample += (features,)
samples.append(sample)
model = _event_train(samples, classes, columns, window=window,
channels=channels, epochs=epochs, seed=seed,
video_width=768 if video is not None else 0)
if video is not None:
model['video'] = {'model_dir': os.path.abspath(os.fspath(video['model_dir'])),
'channel_map': list(video['channel_map']),
'device': video.get('device', 'auto')}
model['video_provenance'] = provenance
model['mean'], model['std'] = mean, std
return model
def _event_detect(model, table, crops=None, *, threshold=0.5, tolerance=2, video=None):
"""Detect and time-stamp events on one field.
:returns: ``track_id``, ``frame``, ``event``, ``probability`` per event.
"""
x, c, index = _event_windows(table, model['columns'], model['window'],
model['mean'], model['std'],
crops if model['channels'] else None)
if not len(index):
return pd.DataFrame(columns=['track_id', 'frame', 'event', 'probability'])
features = None
if model.get('video_width', 0):
video = video if video is not None else dict(model['video'])
features, provenance = _event_clip_features(video, c, index, crops, model['window'])
expected = model['video_provenance']
for key in ('files_sha256', 'revision', 'channel_map', 'frames', 'intensity',
'frame_sampling', 'pooling'):
if provenance.get(key) != expected.get(key):
raise ValueError("Saved event model and current video encoder provenance differ")
p = _event_probabilities(model, x, c, features)
return _event_peaks(index, p, model['classes'], threshold, tolerance)
def _event_cross_validate(tables, annotations, *, window=9, threshold=0.5,
tolerance=2, epochs=25, seed=0, folds=5, video=None):
"""Precision, recall and timing error on held-out annotated data.
With two or more annotated fields each fold holds out whole fields (at
most ``folds`` folds); with one, it holds out every third track.
:returns: ``(scores, detections)``: :func:`_event_scores` of the pooled
held-out detections, and those detections.
"""
fields = list(tables)
detections = []
if len(fields) >= 2:
groups = np.array_split(np.array(fields, dtype=object), min(folds, len(fields)))
for held in groups:
train = {f: tables[f] for f in fields if f not in set(held)}
model = _event_fit(train, annotations[annotations['field'].isin(train)],
window=window, epochs=epochs, seed=seed, video=video)
for field in held:
found = _event_detect(model, *tables[field], threshold=threshold,
tolerance=tolerance, video=video)
detections.append(found.assign(field=field))
else:
field = fields[0]
table, crops = tables[field]
tracks = np.array(sorted(pd.unique(table['track_id'])))
for k in range(3):
held = set(tracks[k::3])
train = {field: (table[~table['track_id'].isin(held)], crops)}
model = _event_fit(train, annotations[~annotations['track_id'].isin(held)],
window=window, epochs=epochs, seed=seed, video=video)
found = _event_detect(model, table[table['track_id'].isin(held)], crops,
threshold=threshold, tolerance=tolerance, video=video)
detections.append(found.assign(field=field))
detected = pd.concat(detections, ignore_index=True) if detections else pd.DataFrame(
columns=['track_id', 'frame', 'event', 'probability', 'field'])
classes = sorted(pd.unique(annotations['event']))
return _event_scores(detected, annotations, classes, tolerance), detected
def _event_correct_divisions(tracks, events, *, mitosis='mitosis',
max_distance=30.0, tolerance=2):
"""Division links taken from detected mitoses.
For every detected mitosis of a mother track, each track that starts
within ``tolerance`` + 1 frames after it, within ``max_distance`` pixels
of the mother's position at the mitosis, is linked to her in
``parent_track_id``. Links the tracker reported itself are kept; a new
track beside no detected mitosis gets no parent, which removes the
divisions that would otherwise be inferred from broken tracks.
:param tracks: one field's tracks table.
:param events: that field's detected events.
:returns: ``(corrected, links)``: the tracks with ``parent_track_id`` and
a table of the links added (``track_id``, ``parent_track_id``,
``mitosis_frame``).
"""
df = tracks.copy()
if 'parent_track_id' not in df.columns:
df['parent_track_id'] = 0
df['parent_track_id'] = pd.to_numeric(df['parent_track_id'], errors='coerce').fillna(0).astype(int)
starts = df.sort_values('frame').groupby('track_id').first()
has_parent = set(starts.index[starts['parent_track_id'] > 0].astype(int))
links = []
for mother, frame in zip(*(events.loc[events['event'] == mitosis, c]
for c in ('track_id', 'frame'))):
where = df[(df['track_id'] == mother) & (df['frame'] <= frame)].sort_values('frame')
if where.empty:
continue
mx, my = where['x'].iloc[-1], where['y'].iloc[-1]
new = starts[(starts['frame'] > frame - 1) & (starts['frame'] <= frame + tolerance + 1)
& (starts.index != mother)]
near = np.hypot(new['x'] - mx, new['y'] - my) <= max_distance
for daughter in new.index[near].astype(int):
if daughter in has_parent:
continue
has_parent.add(daughter)
df.loc[df['track_id'] == daughter, 'parent_track_id'] = int(mother)
links.append({'track_id': daughter, 'parent_track_id': int(mother),
'mitosis_frame': int(frame)})
return df, pd.DataFrame(links, columns=['track_id', 'parent_track_id', 'mitosis_frame'])
def _event_cycles(spans, marks):
"""Split tracks into the intervals between repeated events such as mitoses.
A mitosis is stamped on the mother's last frame before her daughters
appear; when the tracker carries one daughter on under the mother's id,
the track holds several cell cycles. Each interval runs from its first
frame to the frame after its mitosis (the event) or to the track's last
frame (censored), so consecutive intervals add up to the track.
:param spans: one row per track with ``field``, ``track_id``, ``start``
and ``end``.
:param marks: ``{(field, track_id): sorted event frames}``.
:returns: one row per interval with the columns of ``spans`` (``start``
moved to the interval's first frame), ``event_frame``, ``event`` and
``duration`` in frames.
"""
rows = []
for span in spans.to_dict('records'):
start, end = int(span['start']), int(span['end'])
for frame in marks.get((span['field'], span['track_id']), []):
if start <= frame <= end:
rows.append(dict(span, start=start, event_frame=float(frame),
event=1, duration=float(frame + 1 - start)))
start = frame + 1
if start <= end:
rows.append(dict(span, start=start, event_frame=np.nan, event=0,
duration=float(end - start)))
return pd.DataFrame(rows)
def _event_timing(tables, events, conditions=None):
"""Time to each kind of event, per track or, for mitosis, per cell cycle.
For every event but mitosis this is the time from each track's first
frame to its first event of that kind, censored at the track's last
frame. Mitosis repeats and a tracker often carries one daughter on
under her mother's id, so a track is split at each detected mitosis
(:func:`_event_cycles`) and every interval is one cell cycle: the time
from its first frame to its division, or censored where the track ends.
:param tables: ``{field: track table}``.
:param events: detected events with ``field``.
:param conditions: ``name=wells`` entries; fields are otherwise grouped
by well.
:returns: ``{event: objects}`` in the form
:func:`spacr.measure._time_to_event_statistics` reads: one row per
track (per cell cycle for mitosis) with ``duration`` (frames),
``event`` (1 or 0), ``condition`` and ``well``, plus the order of
conditions under key ``'_order'`` per event.
"""
from .measure import _time_to_event_groups
spans = []
for field, table in tables.items():
try:
key = schema.parse_prcf(field)
ids = (key.plateID, key.rowID, key.columnID, key.fieldID)
except Exception:
ids = (field, 'r0', 'c0', 'f0')
span = table.groupby('track_id')['frame'].agg(start='min', end='max').reset_index()
span['field'] = field
span['plateID'], span['rowID'], span['columnID'], span['fieldID'] = ids
spans.append(span)
spans = pd.concat(spans, ignore_index=True)
config = {'time_to_event_group': 'well', 'time_to_event_reference': '',
'time_to_event_conditions': list(conditions or []),
'time_to_event_covariates': []}
result = {}
for name in sorted(pd.unique(events['event'])) if len(events) else []:
found = events[events['event'] == name]
if name == 'mitosis':
marks = {key: sorted(int(f) for f in pd.unique(group['frame']))
for key, group in found.groupby(['field', 'track_id'])}
objects = _event_cycles(spans, marks)
else:
first = (found.groupby(['field', 'track_id'])['frame']
.min().rename('event_frame').reset_index())
objects = spans.merge(first, on=['field', 'track_id'], how='left')
objects['event'] = objects['event_frame'].notna().astype(int)
objects['duration'] = np.where(objects['event'] == 1,
objects['event_frame'] - objects['start'],
objects['end'] - objects['start']).astype(float)
objects['time_unit'] = 'frames'
grouped, order = _time_to_event_groups(objects, config)
if len(grouped):
result[name] = (grouped, order)
return result
def _event_detection(tracks_dir, object_type, prefix, *, annotations=None,
model_path=None, window=9, threshold=0.5, tolerance=2,
conditions=None, max_distance=30.0, epochs=25, plot=True,
partners=(), encoder='small', video_checkpoint=None,
video_channels=None, video_device='auto'):
"""Detect events on every tracked field of a run and write the results.
With ``annotations``, the detector is scored by cross-validation on the
annotated fields (``events/event_detection_scores.csv``), trained on all
of them and saved as ``events/event_model.pt``; otherwise the model at
``model_path`` is used. Every field's events go to
``events/<stem>_events.csv`` and together to ``events/events.csv``. With
a ``mitosis`` class the tracks are re-linked from the detected mitoses
(``events/<stem>_corrected.csv``) and lineage trees are drawn from them
under ``events/lineage``. The time to each kind of event is compared
across conditions with Kaplan-Meier curves and log-rank tests
(``events/event_timing_<event>_*.csv`` and figure).
GUI-authored annotations bind each field to its exact tracker CSV;
an incompatible backend or changed tracker is refused before fitting.
:param tracks_dir: the run's ``tracks`` folder.
:param object_type: the tracked object.
:param prefix: the tracker's file prefix.
:param partners: other tracked object types read together with this
one, field by field, so events across two object types (a parasite
egressing from or invading a host) can be learned.
:param encoder: small CPU image encoder or frozen pretrained videomae.
:param video_checkpoint: verified local official checkpoint folder. A
saved pretrained model may reuse its recorded folder when blank.
:param video_channels: three explicit ordered source-channel indices.
:param video_device: device used by the isolated pretrained encoder.
:returns: dict with ``scores`` (or None), ``events`` and ``paths``.
:raises ValueError: neither annotations nor a model.
"""
import torch
from .qt.i18n import tr
from .measure import _time_to_event_figure, _time_to_event_statistics
from .tabular import write_table
if encoder not in ('small', 'videomae'):
raise ValueError("timelapse_events_encoder must be small or videomae")
video = None
if encoder == 'videomae' and (annotations is not None or video_checkpoint):
if not video_checkpoint or video_channels is None:
raise ValueError("VideoMAE needs timelapse_events_video_checkpoint and explicit video_channels")
video = {'model_dir': video_checkpoint, 'channel_map': video_channels,
'device': video_device}
out = os.path.join(tracks_dir, 'events')
os.makedirs(out, exist_ok=True)
fields = _event_fields(tracks_dir, object_type, prefix)
if not fields:
raise ValueError(f"No {prefix} tracks of {object_type} in {tracks_dir}.")
other = {k: _event_fields(tracks_dir, k, prefix) for k in partners if k != object_type}
inputs = {f: _event_field_inputs(p, max_distance, {
k: v[f] for k, v in other.items() if f in v}) for f, p in fields.items()}
paths, scores = {}, None
if annotations is not None:
ann = _event_read_annotations(annotations) if isinstance(annotations, str) else annotations
if 'object' in ann.columns:
ann = ann[ann['object'].astype(str).isin([object_type, 'nan', ''])]
ann = ann[ann['field'].isin(inputs)]
if ann.empty:
raise ValueError(f"Setting: timelapse_events_annotations names no "
f"event on the {object_type} tracks of this run.")
provenance = {'tracker_backend', 'track_source_sha256'}
if provenance & set(ann.columns):
if not provenance.issubset(ann.columns):
raise ValueError(tr('Event annotations have incomplete tracker provenance.'))
for field, rows in ann.groupby('field'):
backend = rows['tracker_backend'].fillna('').astype(str)
digest = rows['track_source_sha256'].fillna('').astype(str)
if backend.ne('').any() or digest.ne('').any():
expected = _event_source_hash(fields[field])
if not (backend.eq(prefix) & digest.eq(expected)).all():
raise ValueError(tr(
'Event annotations for {field} belong to another tracker CSV.',
field=field))
annotated = {f: inputs[f] for f in pd.unique(ann['field'])}
scores, _ = _event_cross_validate(annotated, ann, window=window,
threshold=threshold, tolerance=tolerance,
epochs=epochs, video=video)
scores['tolerance_frames'] = tolerance
paths['scores'] = write_table(scores, os.path.join(out, 'event_detection_scores.csv'))
model = _event_fit(annotated, ann, window=window, epochs=epochs, video=video)
paths['model'] = os.path.join(out, 'event_model.pt')
torch.save(model, paths['model'])
elif model_path:
model = torch.load(model_path, map_location='cpu', weights_only=True)
if encoder == 'videomae' and not model.get('video_width', 0):
raise ValueError("The selected event model was not trained with VideoMAE features")
if model.get('video_width', 0) and video is None:
video = dict(model['video'])
video['device'] = video_device
else:
raise ValueError("Setting: timelapse_events needs annotated events in "
"timelapse_events_annotations or a trained model in "
"timelapse_events_model.")
found, corrected = [], {}
for field, (table, crops) in inputs.items():
events = _event_detect(model, table, crops, threshold=threshold,
tolerance=tolerance, video=video).assign(field=field)
stem = os.path.splitext(os.path.basename(fields[field]))[0]
write_table(events, os.path.join(out, f'{stem}_events.csv'), canonicalise=False)
found.append(events)
if 'mitosis' in model['classes']:
fixed, _links = _event_correct_divisions(
table[['frame', 'track_id', 'x', 'y'] + (
['parent_track_id'] if 'parent_track_id' in table else [])],
events, max_distance=max_distance, tolerance=tolerance)
corrected[field] = write_table(fixed, os.path.join(out, f'{stem}_corrected.csv'))
events = pd.concat(found, ignore_index=True)
paths['events'] = write_table(events, os.path.join(out, 'events.csv'), canonicalise=False)
for field, path in corrected.items():
_lineage_trees_from_tracks(path, os.path.join(out, 'lineage'),
max_distance=-1.0, plot=plot)
if corrected:
paths['lineage'] = os.path.join(out, 'lineage')
timing = _event_timing({f: t for f, (t, _) in inputs.items()}, events, conditions)
for name, (objects, order) in timing.items():
stats = _time_to_event_statistics(objects, order, {'time_to_event_covariates': []})
for part in ('curves', 'summary', 'tests'):
paths[f'{name}_{part}'] = write_table(
stats[part], os.path.join(out, f'event_timing_{name}_{part}.csv'))
if plot:
fig = _time_to_event_figure(stats['curves'], stats['summary'],
stats['tests'], f'Time to {name}')
paths[f'{name}_figure'] = save_figure_to_path(
fig, os.path.join(out, f'event_timing_{name}.pdf'), close=True)
return {'scores': scores, 'events': events, 'paths': paths}
def _run_event_features_step(src, name, object_type, mask_stack, images, mode, settings):
"""Store the frame features and crops event detection reads, for one field.
Writes ``tracks/events/<tracker>_tracks_<object>_<name>_features.csv``
and ``..._crops.npz`` beside the tracks table. A failure is reported and
does not stop the run.
"""
from .tabular import write_table
prefix = 'trackpy' if mode == 'iou' else mode
out = os.path.join(os.path.dirname(src), 'tracks', 'events')
stem = f'{prefix}_tracks_{object_type}_{name}'
try:
os.makedirs(out, exist_ok=True)
features, crops = _event_frame_features(mask_stack, images)
write_table(features, os.path.join(out, f'{stem}_features.csv'))
if crops is not None:
np.savez_compressed(os.path.join(out, f'{stem}_crops.npz'),
index=features[['frame', 'track_id']].to_numpy(),
crops=crops)
except Exception as exc:
print(f"Event features could not be stored for {name}: {exc}")
def _run_event_detection_step(src, settings):
"""Detect events on the tracks of a finished timelapse run.
Runs :func:`_event_detection` for every object in ``timelapse_objects``
with the ``timelapse_events_*`` settings, reading the other tracked
objects of each field together with it. A failure is reported and does
not stop the run.
:param src: the run's source folder, holding ``tracks``.
:returns: ``{object: result}``.
"""
mode = settings.get('timelapse_mode') or 'trackastra'
prefix = 'trackpy' if mode == 'iou' else mode
tracks_dir = os.path.join(src, 'tracks')
results = {}
objects = list(settings.get('timelapse_objects') or ['cell'])
for object_type in objects:
try:
result = _event_detection(
tracks_dir, object_type, prefix,
annotations=settings.get('timelapse_events_annotations') or None,
model_path=settings.get('timelapse_events_model') or None,
window=int(settings.get('timelapse_events_window') or 9),
threshold=float(settings.get('timelapse_events_threshold') or 0.5),
conditions=settings.get('timelapse_events_conditions') or None,
max_distance=float(settings.get('timelapse_lineage_max_distance') or 30.0),
plot=bool(settings.get('save', True) or settings.get('plot', False)),
partners=[o for o in objects if o != object_type],
encoder=settings.get('timelapse_events_encoder', 'small'),
video_checkpoint=settings.get('timelapse_events_video_checkpoint'),
video_channels=settings.get('timelapse_events_video_channels'),
video_device=settings.get('timelapse_events_video_device', 'auto'))
except Exception as exc:
print(f"Event detection ({object_type}) failed: {exc}")
continue
results[object_type] = result
counts = result['events']['event'].value_counts().to_dict()
print(f"Events ({object_type}): {counts or 'none'}; written to "
f"{os.path.join(tracks_dir, 'events')}")
if result['scores'] is not None:
overall = result['scores'].iloc[-1]
print(f"Held-out precision {overall['precision']:.2f}, recall "
f"{overall['recall']:.2f}, mean timing error "
f"{overall['mean_abs_timing_error']:.2f} frames")
return results
[docs]
def exponential_decay(x, a, b, c):
"""Return ``a * exp(-b * x) + c`` for curve fitting.
The photobleaching model :func:`analyze_calcium_oscillations` fits with
``scipy.optimize.curve_fit``, which reads the three arguments after ``x``
as the free parameters to solve for.
:param x: time points, as a scalar or a NumPy array / pandas Series; a
``Series`` comes back as a ``Series`` on the same index, which is what
lets the caller's ``df[measurement] / exponential_decay(...)`` align
by label.
:param a: amplitude of the decaying term. At ``x == 0`` the result is
``a + c``, not ``a``; ``a == 0`` flattens the curve to the constant ``c``.
:param b: decay rate. The sign is not checked -- a negative ``b`` grows
instead of decaying, and a large ``-b * x`` overflows to ``inf`` with a
``RuntimeWarning`` rather than raising.
:param c: additive offset, and the asymptote as ``x`` grows. Nothing keeps
the curve positive, so a fit with ``c < 0`` crosses zero and the
caller's division by this curve flips sign across the crossing.
:returns: ``numpy.float64`` for scalar ``x``, otherwise the array type of ``x``.
:raises TypeError: when ``x`` is a plain list and ``b`` is a float, because
``-b * x`` is then list arithmetic. An integer ``b`` does not raise: it
silently returns an empty array for ``b >= 0``.
"""
return a * np.exp(-b * x) + c
#: Well-identifier columns every spaCR object table carries, in the spelling
#: :func:`spacr.utils._merge_and_save_to_database` writes them.
_OBJECT_WELL_KEYS = ['plateID', 'rowID', 'columnID', 'fieldID']
#: Spellings of the timepoint column, most canonical first. The measurement
#: writer emits ``timeID``; ``filepaths_to_database`` spells the same thing
#: ``time_id`` in ``png_list``; ``timeid`` was this module's own private
#: spelling and is accepted so an older hand-built table still reads.
_TIME_KEY_ALIASES = ('timeID', 'time_id', 'timeid')
def _resolve_time_key(df):
"""Return the timepoint column ``df`` carries, or ``None`` if it has none.
A non-timelapse measurements database has no timepoint column at all --
``_merge_and_save_to_database`` only adds one when ``timelapse=True`` --
so callers must treat the time key as optional rather than assume it.
:param df: a measurements DataFrame.
:returns: the column name, or ``None`` when the frame has no time axis.
"""
for key in _TIME_KEY_ALIASES:
if key in df.columns:
return key
return None
def _object_group_keys(df, object_key):
"""Return the identifier columns that address one object in ``df``.
``plateID``/``rowID``/``columnID``/``fieldID``, the timepoint column when
the frame has one, then ``object_key``.
:param df: a measurements DataFrame.
:param object_key: the per-object identifier column, e.g. ``'object_label'``
or ``'cell_id'``.
:returns: list of column names.
:raises KeyError: when the frame is missing a required identifier column.
"""
time_key = _resolve_time_key(df)
keys = list(_OBJECT_WELL_KEYS) + ([time_key] if time_key else []) + [object_key]
missing = [key for key in keys if key not in df.columns]
if missing:
raise KeyError(
f"measurements frame is missing identifier column(s) {missing}; "
f"a spaCR object table carries {_OBJECT_WELL_KEYS} plus "
f"{object_key} (and timeID for a timelapse run). "
f"Got: {list(df.columns)[:12]}"
+ (" ..." if len(df.columns) > 12 else ""))
return keys
#: The bleach-correction methods Measure offers, ``'none'`` first.
_BLEACH_METHODS = ('none', 'ratio', 'exponential', 'histogram')
#: The per-object intensity statistics that scale with the illumination and
#: are therefore corrected: levels, sums and percentiles. Spread and shape
#: statistics (``cv``, ``skew``, ``gini`` ...) are ratios or unitless and are
#: left as measured.
_BLEACH_LEVEL_PATTERN = (
r'(?:mean|median|max|min|integrated|mode)_intensity|percentile_\d+')
def _bleach_channel_columns(df, object_type):
"""Return the intensity columns of each channel that bleach correction rescales.
:param df: an object table from ``measurements.db``.
:param object_type: the table's object prefix, e.g. ``'cell'``.
:returns: ``{channel: [column, ...]}`` for every channel whose
``<object>_channel_<n>_mean_intensity`` is present; that column is the
channel's reference trend and is listed first.
"""
pattern = re.compile(
rf'^{re.escape(object_type)}_channel_(\d+)_(?:{_BLEACH_LEVEL_PATTERN})$')
channels = {}
for column in df.columns:
match = pattern.match(str(column))
if match:
channels.setdefault(int(match.group(1)), []).append(column)
result = {}
for channel in sorted(channels):
reference = f'{object_type}_channel_{channel}_mean_intensity'
if reference in channels[channel]:
rest = [c for c in channels[channel] if c != reference]
result[channel] = [reference] + rest
return result
def _bleach_times(values):
"""Timepoint numbers from a ``timeID`` column such as ``t4``, ``t04`` or ``4``.
Measure writes ``timeID`` as text (``t12``), which neither sorts nor
fits as a number; the trailing integer is the frame.
:param values: the timepoint column.
:returns: a float Series on the same index, NaN where no number ends
the value.
"""
series = pd.Series(values)
if pd.api.types.is_numeric_dtype(series):
return series.astype(float)
text = series.astype(str).str.extract(r'(-?\d+(?:\.\d+)?)\s*$')[0]
return pd.to_numeric(text, errors='coerce')
def _bleach_ring_column(df, object_type, channel):
"""The ring-background column Measure wrote for one channel, if any.
:returns: ``<object>_channel_<n>_outside_percentile_50`` or
``..._outside_mean`` when present, else None.
"""
for stat in ('outside_percentile_50', 'outside_mean'):
name = f'{object_type}_channel_{channel}_{stat}'
if name in df.columns:
return name
return None
def _bleach_signal(frame, column, ring=None):
"""One level column, less the object's ring background when there is one.
:param frame: object rows.
:param column: the level column.
:param ring: the ring-background column, or None.
:returns: a float Series.
"""
values = pd.to_numeric(frame[column], errors='coerce').astype(float)
if ring is not None:
values = values - pd.to_numeric(frame[ring], errors='coerce')
return values
def _bleach_trend(frame, time_key, column, ring=None):
"""Return the median of ``column`` at each timepoint, in time order.
The median over every object in the frame is the background trend a
bleaching series shows: one bright or dying object does not move it.
With ``ring``, each object's ring background is subtracted first, so
the trend follows the fluorescence that bleaches rather than the
camera offset under it.
:returns: ``pandas.Series`` indexed by timepoint number, NaN frames
dropped.
"""
values = _bleach_signal(frame, column, ring)
trend = values.groupby(_bleach_times(frame[time_key]).to_numpy()).median().sort_index()
return trend.dropna()
def _fit_bleach_decay(times, trend):
"""Fit ``a * exp(-b * t) + c`` to a background trend.
Time is counted from the first timepoint. ``a`` and ``b`` are held
non-negative so the fit describes a decay, never a growth.
:param times: timepoints, numeric.
:param trend: the trend value at each timepoint.
:returns: ``(a, b, c)``, or ``None`` when there are fewer than three
timepoints or the fit does not converge.
"""
t = np.asarray(times, dtype=float)
y = np.asarray(trend, dtype=float)
if t.size < 3 or not np.all(np.isfinite(y)):
return None
t = t - t.min()
span = float(t.max()) or 1.0
amplitude = max(float(y.max() - y.min()), 1e-12)
try:
params, _ = curve_fit(
exponential_decay, t, y,
p0=[amplitude, 1.0 / span, float(y.min())],
bounds=([0.0, 0.0, -np.inf], [np.inf, np.inf, np.inf]),
maxfev=10000)
except (RuntimeError, ValueError):
return None
if not np.all(np.isfinite(params)):
return None
return tuple(float(p) for p in params)
def _bleach_factors(trend, method):
"""Return the multiplicative correction at each timepoint of one series.
``ratio`` divides each timepoint by its own trend value and multiplies by
the first: the simple ratio method. Bleaching only dims, so a timepoint
whose trend is above the first is left as measured (factor 1) rather
than darkened: a rise of the median is biology, focus or illumination,
and dividing it away would erase it. ``exponential`` rescales by the
fitted decay in place of the measured trend, which ignores
frame-to-frame noise and, being a decay, never darkens either; when
the fit fails or the curve reaches zero it falls back to the ratio and
reports ``ratio_fallback``.
:param trend: ``pandas.Series`` of the background trend by timepoint.
:param method: ``'ratio'`` or ``'exponential'``.
:returns: ``(factors, applied, params)``: a Series of factors by
timepoint, the method actually applied, and the fitted
``(a, b, c)`` or ``None``.
"""
times = np.asarray(trend.index, dtype=float)
params = None
if method == 'exponential':
params = _fit_bleach_decay(times, trend.to_numpy())
if params is not None:
model = exponential_decay(times - times.min(), *params)
if np.all(model > 0):
return (pd.Series(model[0] / model, index=trend.index),
'exponential', params)
method = 'ratio_fallback'
reference = trend.to_numpy(dtype=float)
with np.errstate(divide='ignore', invalid='ignore'):
factors = np.where(reference > 0, np.maximum(reference[0] / reference, 1.0),
np.nan)
return pd.Series(factors, index=trend.index), method, params
def _histogram_match(values, reference):
"""Map ``values`` onto the distribution of ``reference`` rank for rank.
:returns: an array the length of ``values``; NaN stays NaN.
"""
values = np.asarray(values, dtype=float)
reference = np.asarray(reference, dtype=float)
reference = reference[np.isfinite(reference)]
out = np.full(values.shape, np.nan)
finite = np.isfinite(values)
if not finite.any() or reference.size == 0:
return out
ranks = pd.Series(values[finite]).rank(method='average').to_numpy()
quantiles = (ranks - 0.5) / finite.sum()
out[finite] = np.quantile(reference, quantiles)
return out
def _bleach_correct_table(df, object_type, method):
"""Correct every intensity level of one timelapse object table for bleaching.
Each field (plate, row, column, field) is its own bleaching series and
each channel is corrected on its own, in timepoint order (``t2`` before
``t10``). The background trend of a channel is the per-timepoint median
of ``<object>_channel_<n>_mean_intensity``; ``ratio`` and
``exponential`` rescale every level column of that channel by the
factor that brings the trend back to its first timepoint.
``histogram`` instead maps each column at each timepoint onto that
column's distribution at the first timepoint.
When Measure wrote a ring background for the channel
(``outside_percentile_50``, else ``outside_mean``), the trend and the
correction work on the signal above it: each level less the ring
(the ring times the area for an integrated intensity) is corrected and
the ring is added back. A camera offset, which does not bleach, is then
neither counted in the trend nor rescaled, so the ratio between two
objects' signals in one frame is kept. Without a ring the levels are
rescaled whole.
:param df: the object table as Measure wrote it, with a timepoint column.
:param object_type: its object prefix, e.g. ``'cell'``.
:param method: one of ``'ratio'``, ``'exponential'``, ``'histogram'``.
:returns: ``(corrected, fits)``. ``corrected`` has the identifier
columns, every corrected column under its measured name, and
``bleach_correction_method``. ``fits`` has one row per field and
channel: the method applied, the ring column used (``background``,
empty without one), the timepoints, the trend at the first and last
timepoint before and after correction, how far the trend ever rises
above its first value (``trend_peak_rise``, a fraction; bleaching
alone never raises it), and for an exponential fit
``decay_a``, ``decay_b``, ``decay_c`` and ``half_life`` in timepoint
units.
:raises ValueError: an unknown method, or a table with no timepoint column.
"""
if method not in _BLEACH_METHODS[1:]:
raise ValueError(
f"bleach_correction must be one of {list(_BLEACH_METHODS)}, "
f"got {method!r}")
time_key = _resolve_time_key(df)
if time_key is None:
raise ValueError(
f"the {object_type} table has no timepoint column; bleach "
f"correction needs a timelapse measurement")
channels = _bleach_channel_columns(df, object_type)
ids = [c for c in ('object_label', 'cell_id', *_OBJECT_WELL_KEYS,
time_key, 'prcf', 'file_name') if c in df.columns]
corrected = df[ids].copy()
times = _bleach_times(df[time_key])
area_column = f'{object_type}_area'
area = (pd.to_numeric(df[area_column], errors='coerce')
if area_column in df.columns else None)
fields = [k for k in _OBJECT_WELL_KEYS if k in df.columns]
groups = (df.groupby(fields, sort=True, dropna=False).groups if fields
else {(): df.index})
fit_rows = []
for channel, columns in channels.items():
reference = columns[0]
ring_column = _bleach_ring_column(df, object_type, channel)
for column in columns:
corrected[column] = np.nan
for key, index in groups.items():
frame = df.loc[index]
when = times.loc[index]
ring = (pd.to_numeric(frame[ring_column], errors='coerce')
if ring_column else None)
offsets = {}
for column in columns:
if ring is None:
offsets[column] = 0.0
elif column.endswith('_integrated_intensity'):
offsets[column] = (ring * area.loc[index]
if area is not None else None)
else:
offsets[column] = ring
trend = _bleach_trend(frame, time_key, reference, ring_column)
if trend.empty:
continue
key = key if isinstance(key, tuple) else (key,)
row = dict(zip(fields, key), object_type=object_type,
channel=channel, background=ring_column or '',
n_timepoints=int(len(trend)),
trend_first=float(trend.iloc[0]),
trend_last=float(trend.iloc[-1]),
trend_peak_rise=(float(trend.max() / trend.iloc[0] - 1.0)
if trend.iloc[0] > 0 else np.nan),
decay_a=np.nan, decay_b=np.nan, decay_c=np.nan,
half_life=np.nan)
if method == 'histogram':
first = (when == trend.index[0]).to_numpy()
for column in columns:
offset = offsets[column]
values = pd.to_numeric(frame[column], errors='coerce').astype(float)
if offset is not None:
values = values - offset
target = values.to_numpy()[first]
matched = values.copy()
for _, part in values.groupby(when.to_numpy()):
matched.loc[part.index] = _histogram_match(
part.to_numpy(), target)
corrected.loc[index, column] = (
matched + offset if offset is not None else matched)
row['method'] = 'histogram'
else:
factors, applied, params = _bleach_factors(trend, method)
scale = when.map(factors).astype(float)
for column in columns:
values = pd.to_numeric(frame[column], errors='coerce')
offset = offsets[column]
if offset is None:
corrected.loc[index, column] = values * scale
else:
corrected.loc[index, column] = (
(values - offset) * scale + offset)
row['method'] = applied
if params is not None:
a, b, c = params
row.update(decay_a=a, decay_b=b, decay_c=c,
half_life=(np.log(2) / b) if b > 0 else np.inf)
after = corrected.loc[index].copy()
if ring_column:
after[ring_column] = frame[ring_column]
after = _bleach_trend(after, time_key, reference, ring_column)
row['corrected_first'] = float(after.iloc[0]) if len(after) else np.nan
row['corrected_last'] = float(after.iloc[-1]) if len(after) else np.nan
fit_rows.append(row)
corrected['bleach_correction_method'] = method
return corrected, pd.DataFrame(fit_rows)
def _bleach_decay_figure(df, corrected, fits, object_type, time_key,
max_fields=12):
"""Draw each channel's background trend, the fitted decay and the corrected trend.
One panel per channel; one series per field, at most ``max_fields`` of
them. The trend is the signal above the ring background when the fit
used one. The measured trend is drawn as points, the fitted exponential (or,
for the ratio and histogram methods, the measured trend) as a line, and
the corrected trend dashed in the highlight colour.
:returns: a :class:`matplotlib.figure.Figure`.
"""
channels = sorted(fits['channel'].unique()) if not fits.empty else []
fig = Figure(figsize=(3.2 * max(len(channels), 1), 2.8), dpi=100)
from .figures.bundle import _register_figure_data
_register_figure_data(fig, df, x=time_key, y=f"{object_type}_channel_{channels[0]}_mean_intensity" if channels else "", kind="line")
axes = fig.subplots(1, max(len(channels), 1), squeeze=False)[0]
fields = [k for k in _OBJECT_WELL_KEYS if k in df.columns]
for ax, channel in zip(axes, channels):
reference = f'{object_type}_channel_{channel}_mean_intensity'
rows = fits[fits['channel'] == channel].head(max_fields)
for _, row in rows.iterrows():
mask = np.ones(len(df), dtype=bool)
for k in fields:
mask &= (df[k] == row[k]).to_numpy()
ring = row.get('background') or None
raw = _bleach_trend(df[mask], time_key, reference, ring)
after = corrected[mask].copy()
if ring:
after[ring] = df.loc[mask, ring].to_numpy()
fixed = _bleach_trend(after, time_key, reference, ring)
t = np.asarray(raw.index, dtype=float)
ax.plot(t, raw.to_numpy(), 'o', ms=2.5, color=ROLES['data'])
if np.isfinite(row['decay_b']):
grid = np.linspace(t.min(), t.max(), 100)
ax.plot(grid, exponential_decay(
grid - t.min(), row['decay_a'], row['decay_b'],
row['decay_c']), '-', lw=1, color=ROLES['reference'])
else:
ax.plot(t, raw.to_numpy(), '-', lw=0.8,
color=ROLES['reference'])
ax.plot(np.asarray(fixed.index, dtype=float), fixed.to_numpy(),
'--', lw=1, color=ROLES['highlight'])
ax.set_title(f'{object_type} channel {channel}')
ax.set_xlabel(time_key)
above = bool(rows['background'].astype(str).str.len().gt(0).any()) if 'background' in rows else False
ax.set_ylabel('median mean intensity above ring' if above else 'median mean intensity')
fig.tight_layout()
return fig
def _correct_timelapse_bleaching(db_path, method, *, plot=True,
tables=('cell', 'nucleus', 'pathogen',
'cytoplasm')):
"""Correct every timelapse object table in ``db_path`` for photobleaching.
For each object table present, writes
``measurements.db:<object>_bleach_corrected`` (the corrected intensity
levels under their measured names, keyed like the source table, with the
method in ``bleach_correction_method``) and writes the per-field,
per-channel fits of every table to ``measurements.db:bleach_correction``.
The measured tables are left untouched. With ``plot``, the trends and
fitted decay of each table are saved to
``results/bleach_correction/<object>.pdf`` beside the ``measurements``
folder.
:param db_path: a timelapse ``measurements.db``.
:param method: ``'ratio'``, ``'exponential'`` or ``'histogram'``.
:param plot: save the decay figures.
:param tables: the object tables to correct, where present.
:returns: the fits of every table, as one DataFrame.
:raises ValueError: an unknown method, or no table with a timepoint
column and channel intensities.
"""
from .tabular import database_tables, read_table, write_database
present = [t for t in tables if t in database_tables(db_path)]
all_fits = []
for table in present:
df = read_table(db_path, table=table, report=None)
time_key = _resolve_time_key(df)
if time_key is None or not _bleach_channel_columns(df, table):
continue
corrected, fits = _bleach_correct_table(df, table, method)
write_database(corrected, db_path, f'{table}_bleach_corrected',
if_exists='replace', canonicalise=False)
all_fits.append(fits)
if plot and not fits.empty:
root = os.path.dirname(os.path.dirname(os.path.abspath(db_path)))
fig = _bleach_decay_figure(df, corrected, fits, table, time_key)
save_figure_to_path(fig, os.path.join(
root, 'results', 'bleach_correction', f'{table}.pdf'),
close=True)
if not all_fits:
raise ValueError(
f"{db_path} has no timelapse object table with channel "
f"intensities to correct")
fits = pd.concat(all_fits, ignore_index=True)
write_database(fits, db_path, 'bleach_correction', if_exists='replace',
canonicalise=False)
return fits
[docs]
def preprocess_pathogen_data(pathogen_df):
"""Aggregate a per-parasite table to one row per host cell with a parasite count.
Keys on the column names the measurement writer actually emits:
``columnID`` (not ``column_name``), ``timeID`` (not ``timeid``, and absent
altogether outside a timelapse run) and ``cell_id`` -- the child table's
link to its host cell -- not ``pathogen_cell_id``, which no writer has ever
produced. Under the old names every call died with ``KeyError``.
:param pathogen_df: per-parasite measurements DataFrame with plate/well/field/time/cell identifiers.
:returns: DataFrame aggregated to (plate, row, column, field, time, host cell) with a ``parasite_count`` column.
"""
group_keys = _object_group_keys(pathogen_df, 'cell_id')
parasite_counts = pathogen_df.groupby(group_keys).size().reset_index(name='parasite_count')
value_columns = [
col for col in pathogen_df.columns
if col not in group_keys + ['parasite_count']
]
agg_funcs = {}
for col in value_columns:
dtype = pathogen_df[col].dtype
numeric = (pd.api.types.is_numeric_dtype(dtype)
and not pd.api.types.is_bool_dtype(dtype))
agg_funcs[col] = 'mean' if numeric else 'first'
pathogen_agg = pathogen_df.groupby(group_keys).agg(agg_funcs).reset_index()
pathogen_agg = pathogen_agg.merge(parasite_counts, on=group_keys,
validate='one_to_one')
if 'object_label' in pathogen_agg.columns:
pathogen_agg.drop(columns=['object_label'], inplace=True)
pathogen_agg.rename(columns={'cell_id': 'object_label'}, inplace=True)
return pathogen_agg
[docs]
def plot_data(measurement, group, ax, label, marker='o', linestyle='-'):
"""Plot ``delta_<measurement>`` vs ``time`` for one grouped subset onto ``ax``.
:param measurement: base measurement name; the ``delta_`` prefix is added when reading the column.
:param group: DataFrame subset for a single group.
:param ax: matplotlib axis to draw onto.
:param label: legend label for this series.
:param marker: matplotlib marker. Default ``'o'``.
:param linestyle: matplotlib line style. Default ``'-'``.
:returns: None.
"""
ax.plot(group['time'], group['delta_' + measurement], marker=marker, linestyle=linestyle, label=label)
[docs]
def infected_vs_noninfected(result_df, measurement):
"""Plot per-well mean ``delta_<measurement>`` for infected vs uninfected cell groups.
:param result_df: per-cell/time DataFrame keyed by the composed
``plate_row_column_field_object`` column, with ``time``,
``parasite_count`` and the ``delta_<measurement>`` column.
:param measurement: base measurement column name to plot (the ``delta_`` variant is drawn).
:returns: None.
"""
infected_cells_df = result_df[result_df.groupby('plate_row_column_field_object')['parasite_count'].transform('max') > 0]
uninfected_cells_df = result_df[result_df.groupby('plate_row_column_field_object')['parasite_count'].transform('max') == 0]
with figure_style(theme_target()):
fig, axs = plt.subplots(2, 1, figsize=(12, 10), sharex=True)
from .figures.bundle import _register_figure_data
_register_figure_data(fig, lambda: pd.concat([infected_cells_df.assign(group="Infected"), uninfected_cells_df.assign(group="Uninfected")], ignore_index=True), x="time", y="delta_" + measurement, hue="group", kind="line")
for group_id in infected_cells_df['plate_row_column_field_object'].unique():
group = infected_cells_df[infected_cells_df['plate_row_column_field_object'] == group_id]
plot_data(measurement, group, axs[0], 'Infected', marker='x')
for group_id in uninfected_cells_df['plate_row_column_field_object'].unique():
group = uninfected_cells_df[uninfected_cells_df['plate_row_column_field_object'] == group_id]
plot_data(measurement, group, axs[1], 'Uninfected')
axs[0].set_title('Cells Infected at Some Time')
axs[1].set_title('Cells Never Infected')
for ax in axs:
ax.set_xlabel('Time')
ax.set_ylabel('Normalized Delta ' + measurement)
all_timepoints = sorted(result_df['time'].unique())
ax.set_xticks(all_timepoints)
ax.set_xticklabels(all_timepoints, rotation=45, ha="right")
plt.tight_layout()
plt.show()
[docs]
def save_results_dataframe(df, src, results_name):
"""Save ``df`` as ``<results_name>.csv`` inside a sibling ``results`` folder.
:param df: DataFrame to write.
:param src: reference path used to derive the parent directory.
:param results_name: filename stem (no extension).
:returns: None.
"""
source = os.path.dirname(src)
results_fldr = os.path.join(source,'results')
os.makedirs(results_fldr, exist_ok=True)
csv_loc = os.path.join(results_fldr, f'{results_name}.csv')
df.to_csv(csv_loc, index=True)
print(f'Saved results:{csv_loc}')
_PEAK_ID_SHAPE = ("plate_row_column_field_object — 'plate1_r1_c1_f1_o7', or "
"'plate1_r1_c1_f1_t3_o7' when the key carries a timepoint")
def _explode_peak_ids(peak_details_df, caller):
"""Fill the identity columns of a peak-details frame from its ``ID`` key.
``ID`` is the ``prcfo`` :func:`analyze_calcium_oscillations` composes for
each track. It used to be taken apart with ``str.split('_', expand=True)``
and assigned to exactly five positional columns, which is wrong in both
directions:
* **Loudly, but opaquely**, on any key with six tokens — a plate id that
itself contains an underscore (``exp1_plate1_r1_c1_f1_o1``), or a
timelapse key that kept its timepoint (``plate1_r1_c1_f1_t3_o7``). pandas
answers ``ValueError: Columns must be same length as key``, which names
neither the column nor the id that caused it.
* **Silently**, on any key with five tokens that are not
plate/row/column/field/object. A ``prcf`` from a timelapse
(``plate1_r1_c1_f1_t3``) has exactly five, so the *timepoint* landed in
``object_number`` — and ``cells_per_well`` below, which is
``nunique(object_number)``, then counted timepoints instead of cells and
divided the peak count by the wrong number.
:func:`spacr.schema.parse_prcfo` reads the key **right to left**, so the
plate keeps whatever is left over and the timepoint is recognised by its
``t`` prefix rather than by its position. That is the same parser
:func:`analyze_calcium_oscillations` uses on ``prcf`` one screen below, and
the comment there notes the caller was fixed while these two sites were
not.
:param peak_details_df: frame carrying an ``ID`` column; mutated in place.
:param caller: name of the calling function, used in the error message.
:returns: ``peak_details_df``.
:raises spacr.schema.KeyParseError: naming the ids that are not object
keys, rather than letting a positional split invent an identity.
"""
objects, unparseable = [], []
for value in peak_details_df['ID']:
text = str(value).strip()
try:
objects.append(schema.parse_prcfo(text))
continue
except schema.KeyParseError:
pass
head, separator, tail = text.rpartition(schema.KEY_SEPARATOR)
if separator and tail.isdigit():
try:
objects.append(
schema.parse_prcf(head).with_object(schema.object_id(tail)))
continue
except schema.SchemaError:
pass
unparseable.append(text)
if unparseable:
sample = sorted(set(unparseable))[:5]
raise schema.KeyParseError(
f"{caller}: {len(unparseable)} of {len(peak_details_df)} peak "
f"row(s) have an 'ID' that is not an object key, e.g. {sample}. "
f"'ID' must be {_PEAK_ID_SHAPE}, because every summary below is "
f"keyed on the row and column parsed out of it. Splitting these "
f"on '_' by position would put the field id in the object column "
f"and summarise the peaks into a well that does not exist.")
for column in schema.FIELD_KEY_COLUMNS:
peak_details_df[column] = np.array(
[getattr(obj, column) for obj in objects], dtype=object)
peak_details_df['object_number'] = np.array(
[obj.objectID for obj in objects], dtype=object)
return peak_details_df
[docs]
def summarize_per_well(peak_details_df):
"""Aggregate per-object peak details to one summary row per well.
:param peak_details_df: per-object peak DataFrame with an ``ID`` column
encoding ``plate_row_column_field_object`` and peak metrics.
:returns: DataFrame with one row per well including peak counts, unique
cell counts, and per-well means of the numeric metrics.
:raises spacr.schema.KeyParseError: when an ``ID`` is not an object key.
"""
_explode_peak_ids(peak_details_df, 'summarize_per_well')
peak_details_df['well_ID'] = peak_details_df['rowID'] + '_' + peak_details_df['columnID']
filtered_df = peak_details_df[peak_details_df['amplitude'].notna()]
numeric_cols = filtered_df.select_dtypes(include=['number']).columns
summary_df = filtered_df.groupby('well_ID').agg(
peaks_per_well=('ID', 'size'),
unique_IDs_with_amplitude=('ID', 'nunique'),
**{col: (col, 'mean') for col in numeric_cols}
).reset_index()
peak_details_df['_field_object'] = (
peak_details_df['fieldID'].astype(str) + '_'
+ peak_details_df['object_number'].astype(str))
summary_df_2 = peak_details_df.groupby('well_ID').agg(
cells_per_well=('_field_object', 'nunique'),
).reset_index()
summary_df = summary_df.merge(summary_df_2, on='well_ID', how='left',
validate='one_to_one')
summary_df['peaks_per_cell'] = summary_df['peaks_per_well'] / summary_df['cells_per_well']
return summary_df
[docs]
def summarize_per_well_inf_non_inf(peak_details_df):
"""Aggregate per-object peak details per well, split by infection status.
:param peak_details_df: per-object peak DataFrame with an ``ID`` column
encoding ``plate_row_column_field_object``, peak metrics, and an
``infected`` column whose positive values mark infected objects.
:returns: DataFrame with one row per well and infection status (so one
row for a well seen in a single status) of peak counts, cell counts,
and per-well means of numeric metrics.
:raises spacr.schema.KeyParseError: when an ``ID`` is not an object key.
:raises KeyError: when the frame has no ``infected`` column; the pathogen
count must be carried under exactly that name.
"""
_explode_peak_ids(peak_details_df, 'summarize_per_well_inf_non_inf')
peak_details_df['well_ID'] = peak_details_df['rowID'] + '_' + peak_details_df['columnID']
peak_details_df['infected_status'] = peak_details_df['infected'].apply(lambda x: 'infected' if x > 0 else 'non_infected')
numeric_cols = peak_details_df.select_dtypes(include=['number']).columns
peak_details_df['_field_object'] = (
peak_details_df['fieldID'].astype(str) + '_'
+ peak_details_df['object_number'].astype(str))
summary_df = peak_details_df.groupby(['well_ID', 'infected_status']).agg(
cells_per_well=('_field_object', 'nunique'),
peaks_per_well=('ID', 'size'),
**{col: (col, 'mean') for col in numeric_cols}
).reset_index()
summary_df['peaks_per_cell'] = summary_df['peaks_per_well'] / summary_df['cells_per_well']
return summary_df
def _calcium_shared_bleaching(frame, measurement, method):
"""Keep raw levels and normalize shared field corrections by initial signal."""
match = re.fullmatch(rf'(.+)_channel_(\d+)_(?:{_BLEACH_LEVEL_PATTERN})',
str(measurement))
if match is None:
raise ValueError('Calcium bleach correction requires a measured intensity level column')
role, channel = match.group(1), int(match.group(2))
source = frame.reset_index(drop=True).copy()
columns = _bleach_channel_columns(source, role).get(channel, [])
if measurement not in columns:
raise ValueError('Calcium bleach correction requires the selected level and its channel mean')
keys = _object_group_keys(source, 'object_label')
if source.duplicated(keys).any():
raise ValueError('Calcium bleach correction requires unique field/time/object identities')
working = source.copy()
for column in columns:
working[column] = pd.to_numeric(working[column], errors='coerce').replace(
[np.inf, -np.inf], np.nan)
ring_column = _bleach_ring_column(source, role, channel)
if ring_column:
working[ring_column] = pd.to_numeric(working[ring_column], errors='coerce').replace(
[np.inf, -np.inf], np.nan)
corrected, fits = _bleach_correct_table(working, role, method)
absolute = corrected[measurement]
background = (working[ring_column].copy() if ring_column
else pd.Series(np.nan, index=source.index))
if measurement.endswith('_integrated_intensity'):
area = pd.to_numeric(source.get(f'{role}_area',
pd.Series(np.nan, index=source.index)), errors='coerce')
background *= area.where(np.isfinite(area) & (area > 0))
signal = absolute - background
raw_signal = working[measurement] - background
baseline = pd.Series(np.nan, index=source.index)
applied = pd.Series('unavailable', index=source.index, dtype=object)
time_key = _resolve_time_key(source)
times = _bleach_times(source[time_key])
records = []
for key, indices in source.groupby(_OBJECT_WELL_KEYS, sort=True,
dropna=False).groups.items():
when = times.loc[indices]
first = raw_signal.loc[indices][when == when.min()]
initial = first[np.isfinite(first)].median()
if np.isfinite(initial) and initial > 0:
baseline.loc[indices] = initial
selected = fits
if not selected.empty:
selected = selected[selected['channel'] == channel]
for name, value in zip(_OBJECT_WELL_KEYS, key):
selected = selected[selected[name] == value]
record = (selected.iloc[0].to_dict() if len(selected) == 1 else
dict(zip(_OBJECT_WELL_KEYS, key), channel=channel,
method='unavailable'))
applied.loc[indices] = record['method']
record.update(measurement=measurement, requested_method=method,
normalization='initial field median signal',
baseline_signal=(float(initial) if np.isfinite(initial)
and initial > 0 else np.nan),
background=ring_column or '', source_table='cell',
normalized_units='dimensionless')
records.append(record)
with np.errstate(divide='ignore', invalid='ignore', over='ignore'):
normalized = signal / baseline
source['bleach_corrected_' + measurement] = absolute
source['corrected_' + measurement] = normalized.where(np.isfinite(normalized))
source['bleach_background_' + measurement] = background
source['bleach_baseline_' + measurement] = baseline
source['bleach_correction_requested'] = method
source['bleach_correction_method'] = applied
source['bleach_background_column'] = ring_column or ''
source['bleach_normalization'] = 'initial field median signal'
return source, pd.DataFrame(records)
[docs]
def analyze_calcium_oscillations(db_loc, measurement='cell_channel_1_mean_intensity', size_filter='cell_area', fluctuation_threshold=0.25, num_lines=None, peak_height=0.01, pathogen=None, cytoplasm=None, remove_transient=True, verbose=False, transience_threshold=0.9, *, bleach_correction="legacy"):
"""Detect and summarise per-cell calcium oscillation peaks from a measurements DB.
Loads the ``cell`` (and optionally ``pathogen``/``cytoplasm``) tables,
filters transient tracks, detects peaks on the chosen intensity trace,
and writes the per-peak, per-cell and per-well tables to CSV in a
``results`` folder beside the database.
:param db_loc: path to the measurements SQLite database.
:param measurement: intensity column analysed for oscillations.
:param size_filter: object-size column used for gating.
:param fluctuation_threshold: maximum coefficient of variation
(std / mean) of ``size_filter`` within a track; more variable tracks
are dropped.
:param num_lines: cap on the number of traces to plot; plots all when ``None``.
:param peak_height: minimum absolute peak height, passed to
``scipy.find_peaks(height=...)`` on the delta trace.
:param pathogen: optional pathogen table name to join for infection status.
:param cytoplasm: optional cytoplasm table name to join.
:param remove_transient: drop tracks shorter than the transience threshold.
:param verbose: print diagnostic information.
:param transience_threshold: fraction of timepoints a track must span to be retained.
:param bleach_correction: ``legacy`` preserves the original global decay
fit. Opt-in ``ratio``, ``exponential`` or ``histogram`` uses Measure's
per-field correction once on the original table. Raw intensities remain
unchanged; ``bleach_corrected_<measurement>`` stores absolute corrected
levels. ``corrected_<measurement>`` is dimensionless: corrected signal
above the measured outside ring divided by that field's positive,
finite initial median signal. Integrated levels subtract ring times
area. Missing background or an invalid initial baseline remains NaN;
missing intervals are not bridged for peak detection. The actual method,
background and baseline accompany traces; ``bleach_correction_fits.csv``
records field fits. Histogram matching can erase population changes.
:returns: tuple ``(result_df, peak_details_df, fig)`` -- the
photobleach-corrected per-cell traces, the per-peak details and the
summary matplotlib Figure. The per-well summaries are only written to
CSV. Returns ``None`` when the database has no time axis, the decay
legacy fit fails, or no cells pass the filters.
:raises ValueError: an unknown correction method, an unsupported intensity
column for an explicit shared method, or ambiguous field/time/object
identities.
"""
if bleach_correction not in ('legacy', 'ratio', 'exponential', 'histogram'):
raise ValueError('bleach_correction must be legacy, ratio, exponential or histogram')
conn = sqlite3.connect(db_loc, timeout=30)
cell_df = pd.read_sql(f"SELECT * FROM {'cell'}", conn)
merge_keys = _object_group_keys(cell_df, 'object_label')
if pathogen:
pathogen_df = pd.read_sql("SELECT * FROM pathogen", conn)
if 'cell_id' not in pathogen_df.columns:
conn.close()
raise KeyError(
"the 'pathogen' table has no 'cell_id' column, so parasites "
"cannot be attributed to host cells. That link is only written "
"when the field was measured with a cell mask "
"(settings['cell_mask_dim']); re-run measure_crop with one, or "
"call analyze_calcium_oscillations with pathogen=None.")
pathogen_df['cell_id'] = pathogen_df['cell_id'].astype(float).astype('Int64')
pathogen_df = preprocess_pathogen_data(pathogen_df)
cell_df = cell_df.merge(pathogen_df, on=merge_keys, how='left', suffixes=('', '_pathogen'), validate='many_to_one')
cell_df['parasite_count'] = cell_df['parasite_count'].fillna(0)
print(f'After pathogen merge: {len(cell_df)} objects')
if cytoplasm:
cytoplasm_df = pd.read_sql(f"SELECT * FROM {'cytoplasm'}", conn)
cell_df = cell_df.merge(cytoplasm_df, on=merge_keys, how='left', suffixes=('', '_cytoplasm'), validate='many_to_one')
print(f'After cytoplasm merge: {len(cell_df)} objects')
conn.close()
parsed_prcf = [schema.parse_prcf(value) for value in cell_df['prcf']]
for key in schema.FIELD_KEY_COLUMNS:
cell_df[key] = [getattr(field, key) for field in parsed_prcf]
time_key = _resolve_time_key(cell_df)
prcf_times = [field.timeID for field in parsed_prcf]
if time_key is not None:
cell_df['time'] = cell_df[time_key].astype(str).str.extract(r'(\d+)')[0].astype(int)
elif all(t is not None for t in prcf_times) and prcf_times:
cell_df['time'] = [schema.time_index(t) for t in prcf_times]
else:
print(f"No time axis in {db_loc}: the cell table has no timeID column "
"and its prcf values carry no t<N> element, so this is not a "
"timelapse measurement and calcium oscillations cannot be "
"measured.")
return
cell_df['object_number'] = cell_df['object_label']
cell_df['plate_row_column_field_object'] = [
schema.KEY_SEPARATOR.join([field.prc, field.fieldID,
schema.object_id(label)])
for field, label in zip(parsed_prcf, cell_df['object_label'])]
if 'parasite_count' not in cell_df.columns:
cell_df['parasite_count'] = 0
df = cell_df.copy()
bleach_fits = None
if bleach_correction == 'legacy':
try:
params, _ = curve_fit(exponential_decay, df['time'], df[measurement], p0=[max(df[measurement]), 0.01, min(df[measurement])], maxfev=10000)
df['corrected_' + measurement] = df[measurement] / exponential_decay(df['time'], *params)
except RuntimeError as e:
print(f"Curve fitting failed for the entire dataset with error: {e}")
return
else:
df, bleach_fits = _calcium_shared_bleaching(df, measurement, bleach_correction)
bleach_fits['source_database'] = os.path.abspath(db_loc)
if verbose:
print(f'Analyzing: {len(df)} objects')
corrected_dfs = []
peak_details_list = []
total_timepoints = df['time'].nunique()
field_time_positions = {}
if bleach_correction != 'legacy':
for key, field in df.groupby(_OBJECT_WELL_KEYS, sort=False, dropna=False):
observed_times = sorted(field['time'].unique())
field_time_positions[key] = dict(zip(observed_times,
range(len(observed_times))))
size_filter_removed = 0
transience_removed = 0
for unique_id, group in df.groupby('plate_row_column_field_object'):
group = group.sort_values('time')
if remove_transient:
threshold = int(transience_threshold * total_timepoints)
if verbose:
print(f'Group length: {len(group)} Timelapse length: {total_timepoints}, threshold:{threshold}')
if len(group) <= threshold:
transience_removed += 1
if verbose:
print(f'removed group {unique_id} due to transience')
continue
size_diff = group[size_filter].std() / group[size_filter].mean()
if size_diff <= fluctuation_threshold:
if bleach_correction == 'legacy':
group['delta_' + measurement] = group['corrected_' + measurement].diff().fillna(0)
else:
trace = group['corrected_' + measurement]
delta = trace.diff()
field_key = tuple(group[key].iloc[0] for key in _OBJECT_WELL_KEYS)
position = group['time'].map(field_time_positions[field_key])
delta = delta.where(position.diff() == 1)
if len(delta) and np.isfinite(trace.iloc[0]):
delta.iloc[0] = 0.0
group['delta_' + measurement] = delta
corrected_dfs.append(group)
peaks, properties = find_peaks(group['delta_' + measurement], height=peak_height)
group_filtered = group.copy()
group_filtered['delta_' + measurement] = group['delta_' + measurement].clip(lower=0)
above_zero_auc = trapz(y=group_filtered['delta_' + measurement], x=group_filtered['time'])
auc = trapz(y=group['delta_' + measurement], x=group_filtered['time'])
is_infected = (group['parasite_count'] > 0).any()
if is_infected:
is_infected = 1
else:
is_infected = 0
if len(peaks) == 0:
peak_details_list.append({
'ID': unique_id,
'plateID': group['plateID'].iloc[0],
'rowID': group['rowID'].iloc[0],
'columnID': group['columnID'].iloc[0],
'fieldID': group['fieldID'].iloc[0],
'object_number': group['object_number'].iloc[0],
'time': np.nan,
'amplitude': np.nan,
'delta': np.nan,
'AUC': auc,
'AUC_positive': above_zero_auc,
'AUC_peak': np.nan,
'infected': is_infected
})
for i, peak in enumerate(peaks):
amplitude = properties['peak_heights'][i]
peak_time = group['time'].iloc[peak]
pathogen_count_at_peak = group['parasite_count'].iloc[peak]
start_idx = max(peak - 1, 0)
end_idx = min(peak + 1, len(group) - 1)
peak_segment_y = group['delta_' + measurement].iloc[start_idx:end_idx + 1]
peak_segment_x = group['time'].iloc[start_idx:end_idx + 1]
peak_auc = trapz(y=peak_segment_y, x=peak_segment_x)
peak_details_list.append({
'ID': unique_id,
'plateID': group['plateID'].iloc[0],
'rowID': group['rowID'].iloc[0],
'columnID': group['columnID'].iloc[0],
'fieldID': group['fieldID'].iloc[0],
'object_number': group['object_number'].iloc[0],
'time': peak_time,
'amplitude': amplitude,
'delta': group['delta_' + measurement].iloc[peak],
'AUC': auc,
'AUC_positive': above_zero_auc,
'AUC_peak': peak_auc,
'infected': pathogen_count_at_peak
})
else:
size_filter_removed += 1
if verbose:
print(f'Removed {size_filter_removed} objects due to size filter fluctuation')
print(f'Removed {transience_removed} objects due to transience')
if len(corrected_dfs) > 0:
result_df = pd.concat(corrected_dfs)
else:
print("No suitable cells found for analysis")
return
peak_details_df = pd.DataFrame(peak_details_list)
summary_df = summarize_per_well(peak_details_df)
summary_df_inf_non_inf = summarize_per_well_inf_non_inf(peak_details_df)
if bleach_fits is not None:
save_results_dataframe(df=bleach_fits, src=db_loc,
results_name='bleach_correction_fits')
save_results_dataframe(df=peak_details_df, src=db_loc, results_name='peak_details')
save_results_dataframe(df=result_df, src=db_loc, results_name='results')
save_results_dataframe(df=summary_df, src=db_loc, results_name='well_results')
save_results_dataframe(df=summary_df_inf_non_inf, src=db_loc, results_name='well_results_inf_non_inf')
with figure_style(theme_target()):
fig, ax = plt.subplots(figsize=(10, 8))
from .figures.bundle import _register_figure_data
_register_figure_data(fig, result_df, x="time", y="delta_" + measurement, hue="plate_row_column_field_object", kind="line")
sampled_groups = result_df['plate_row_column_field_object'].unique()
if num_lines is not None and 0 < num_lines < len(sampled_groups):
sampled_groups = np.random.choice(sampled_groups, size=num_lines, replace=False)
for group_id in sampled_groups:
group = result_df[result_df['plate_row_column_field_object'] == group_id]
ax.plot(group['time'], group['delta_' + measurement], marker='o', linestyle='-')
ax.set_xticks(sorted(df['time'].unique()))
ax.set_xticklabels(sorted(df['time'].unique()), rotation=45, ha="right")
ax.set_title(f'Normalized Delta of {measurement} Over Time (Corrected for Photobleaching)')
ax.set_xlabel('Time')
ax.set_ylabel('Normalized Delta ' + measurement)
plt.tight_layout()
plt.show()
save_figure(fig, src=db_loc, figure_number=1)
if pathogen:
infected_vs_noninfected(result_df, measurement)
save_figure(fig, src=db_loc, figure_number=2)
infected_cells = result_df[result_df.groupby('plate_row_column_field_object')['parasite_count'].transform('max') > 0]['plate_row_column_field_object'].unique()
noninfected_cells = result_df[result_df.groupby('plate_row_column_field_object')['parasite_count'].transform('max') == 0]['plate_row_column_field_object'].unique()
infected_peaks = peak_details_df[peak_details_df['ID'].isin(infected_cells)]
noninfected_peaks = peak_details_df[peak_details_df['ID'].isin(noninfected_cells)]
avg_inf_peaks_per_cell = len(infected_peaks) / len(infected_cells) if len(infected_cells) > 0 else 0
avg_non_inf_peaks_per_cell = len(noninfected_peaks) / len(noninfected_cells) if len(noninfected_cells) > 0 else 0
print(f'Average number of peaks per infected cell: {avg_inf_peaks_per_cell:.2f}')
print(f'Average number of peaks per non-infected cell: {avg_non_inf_peaks_per_cell:.2f}')
print(f'done')
return result_df, peak_details_df, fig
def _generate_mask_random_cmap(mask):
"""
Generate a random colormap based on the unique labels in the given mask.
Parameters
----------
mask : ndarray
2D label mask. Background must be 0, objects > 0.
Returns
-------
mpl.colors.ListedColormap
Random colormap with a fixed black background (label 0).
"""
unique_labels = np.unique(mask)
num_objects = np.sum(unique_labels != 0)
random_colors = np.random.rand(num_objects + 1, 4)
random_colors[:, 3] = 1.0
random_colors[0, :] = [0.0, 0.0, 0.0, 1.0]
return mpl.colors.ListedColormap(random_colors)
def _make_intensity_motility_panel(
all_df,
infection_col,
track_df,
per_well_tracks,
n_channels,
motility_dir,
pixels_per_um,
seconds_per_frame,
vel_unit,
settings,
label_tag,
):
"""
Make panels for infection and motility.
Behaviour:
- "mask_*" label_tag:
classic panel with
* per-channel mean intensity (infected vs uninfected)
* (optional) pathogen-channel p75 intensity bar
* (optional) pathogen/cytoplasm intensity ratio bar
* all-tracks motility plot (absolute FOV)
* motility origin plots (infected / uninfected)
* optional small QC image (feature importance PNG)
- "adjusted_*" label_tag:
same as mask panel, plus method-specific QC subplots appended:
* histogram (if strategy == "histogram")
* PCA/UMAP/t-SNE embedding (if strategy == "pca"/"umap"/"tsne")
* XGBoost:
- probability separation histogram
- feature-importance barplot
"""
import os
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.image as mpimg
if all_df.empty or track_df.empty or not per_well_tracks:
print(f"[_make_intensity_motility_panel] No data for panel '{label_tag}', skipping.")
return
os.makedirs(motility_dir, exist_ok=True)
key_cols = ["plateID", "wellID", "fieldID", "cellID"]
label_lower = str(label_tag).lower()
is_mask_panel = label_lower.startswith("mask")
is_adjusted_panel = label_lower.startswith("adjusted")
qc_strategy = str(settings.get("infection_intensity_strategy", "none")).lower()
method_label = qc_strategy if qc_strategy else "none"
panel_label = "mask" if is_mask_panel else ("adjusted" if is_adjusted_panel else label_tag)
qc_graphs_enabled = bool(settings.get("infection_intensity_qc_graphs", True))
qc_panel_type = settings.get("infection_intensity_qc_panel_type", None)
qc_panel_path = settings.get("infection_intensity_qc_panel_path", None)
hist_data = settings.get("infection_hist_data", None)
pca_data = settings.get("infection_pca_data", None)
xgb_data = settings.get("infection_xgb_importance", None)
has_pca = pca_data is not None
has_xgb = xgb_data is not None
qc_panel_needed_mask = (
is_mask_panel
and qc_graphs_enabled
and isinstance(qc_panel_path, str)
and qc_panel_path
and os.path.exists(qc_panel_path)
)
qc_axes_count = 0
if is_adjusted_panel and qc_graphs_enabled:
if qc_strategy == "histogram":
qc_axes_count = 1
elif qc_strategy in {"pca", "umap", "tsne"} and has_pca:
qc_axes_count = 1
elif qc_strategy == "xgboost" and has_xgb:
qc_axes_count = 2
origin_xlim = settings.get("motility_xlim", settings.get("motility_origin_xlim"))
origin_ylim = settings.get("motility_ylim", settings.get("motility_origin_ylim"))
if pixels_per_um is not None and pixels_per_um > 0:
coord_scale = 1.0 / float(pixels_per_um)
coord_label_x = "x (µm)"
coord_label_y = "y (µm)"
else:
coord_scale = 1.0
coord_label_x = "x (pixels)"
coord_label_y = "y (pixels)"
pathogen_chan = settings.get("pathogen_channel", None)
def _plot_hist_qc(ax, source):
"""
Draw infected vs uninfected intensity histogram.
`source` can be:
- a dict payload (settings['infection_hist_data']) with keys:
'intensities_inf', 'intensities_uninf', 'bin_edges',
'thr_val' (optional), 'intensity_col' (optional)
- or a DataFrame (df_well), in which case the histogram
is computed on the fly using a reasonable intensity
column and `infection_intensity_n_bins`.
"""
try:
if isinstance(source, dict):
intens_inf = np.asarray(source["intensities_inf"], dtype=float)
intens_uninf = np.asarray(source["intensities_uninf"], dtype=float)
bin_edges = np.asarray(source["bin_edges"], dtype=float)
thr_val = float(source.get("thr_val", np.nan))
intensity_col = source.get("intensity_col", "intensity")
else:
df_vals = source
intensity_col = settings.get("infection_hist_intensity_col", None)
if not intensity_col or intensity_col not in df_vals.columns:
cand_cols = []
if pathogen_chan is not None:
cand_cols.extend(
[
f"cell_mean_intensity_ch{pathogen_chan}",
f"cell_p75_intensity_ch{pathogen_chan}",
f"pathogen_mean_intensity_ch{pathogen_chan}",
]
)
cand_cols.extend(
[c for c in df_vals.columns if c.startswith("cell_mean_intensity_ch")]
)
for c in cand_cols:
if c in df_vals.columns:
intensity_col = c
break
if not intensity_col or intensity_col not in df_vals.columns:
ax.set_visible(False)
return
cell_level = (
df_vals[key_cols + [intensity_col, infection_col]]
.groupby(key_cols, dropna=False)
.agg({intensity_col: "mean", infection_col: "max"})
.reset_index()
)
cell_level = cell_level.replace([np.inf, -np.inf], np.nan)
cell_level = cell_level.dropna(subset=[intensity_col])
if cell_level.empty:
ax.set_visible(False)
return
mask_inf = cell_level[infection_col].astype(bool)
intens_inf = cell_level.loc[mask_inf, intensity_col].to_numpy()
intens_uninf = cell_level.loc[~mask_inf, intensity_col].to_numpy()
all_vals = np.concatenate(
[arr for arr in (intens_inf, intens_uninf) if arr.size]
)
n_bins = int(settings.get("infection_intensity_n_bins", 64) or 64)
bin_edges = np.histogram_bin_edges(all_vals, bins=n_bins)
thr_val = settings.get(
"infection_hist_thr_val",
settings.get("infection_intensity_threshold", np.nan),
)
thr_val = float(thr_val) if thr_val is not None else np.nan
ax.hist(
intens_uninf,
bins=bin_edges,
alpha=0.5,
color=UNINFECTED_COLOUR,
label="Uninfected",
)
ax.hist(
intens_inf,
bins=bin_edges,
alpha=0.5,
color=INFECTED_COLOUR,
label="Infected",
)
if np.isfinite(thr_val):
reference_line(ax, x=thr_val)
ax.set_xlabel(intensity_col)
ax.set_ylabel("Count")
ax.set_title("Pathogen-channel intensity\n(adjusted labels)")
ax.legend(fontsize=TYPE_SCALE["legend"], frameon=False)
except Exception as e:
print(f"[_make_intensity_motility_panel] Histogram payload invalid: {e}")
ax.set_visible(False)
def _plot_pca_qc(ax, pdata):
"""Plot ``pdata``'s 2-D embedding on ``ax``, or hide it, then return ``None``."""
import numpy as np
try:
coords = np.asarray(pdata["coords"], dtype=float)
labels = np.asarray(pdata["labels"], dtype=bool)
except Exception as e:
print(f"[_make_intensity_motility_panel] PCA/embedding payload invalid: {e}")
ax.set_visible(False)
return
if coords.ndim != 2 or coords.shape[1] < 2:
ax.set_visible(False)
return
method_label = str(pdata.get("method_label", "PCA"))
x = coords[:, 0]
y = coords[:, 1]
ax.scatter(
x[~labels],
y[~labels],
s=5,
alpha=0.4,
color=UNINFECTED_COLOUR,
label="Uninfected",
)
ax.scatter(
x[labels],
y[labels],
s=5,
alpha=0.4,
color=INFECTED_COLOUR,
label="Infected",
)
ax.set_xlabel(f"{method_label} 1")
ax.set_ylabel(f"{method_label} 2")
ax.set_title(f"{method_label} of features\n(adjusted labels)")
ax.legend(fontsize=7)
def _plot_xgb_importance_qc(ax, xdata):
"""Plot ``xdata`` importances on ``ax``, or hide it, then return ``None``."""
try:
feat_names = xdata["feature_names"]
feat_vals = xdata["feature_importances"]
except Exception as e:
print(f"[_make_intensity_motility_panel] XGB importance payload invalid: {e}")
ax.set_visible(False)
return
if not feat_names:
ax.set_visible(False)
return
y_pos = np.arange(len(feat_names))
ax.barh(y_pos, feat_vals)
ax.set_yticks(y_pos)
ax.set_yticklabels(feat_names, fontsize=7)
ax.invert_yaxis()
ax.set_xlabel("Importance (gain)")
ax.set_title("XGBoost feature importance")
def _plot_xgb_prob_qc(ax, df_prob):
"""
Per-cell probability distribution by adjusted infection label.
Uses settings['infection_xgb_proba_column'] if available,
otherwise falls back through a few common column names.
"""
prob_col_candidates = []
cfg_col = settings.get("infection_xgb_proba_column", None)
if isinstance(cfg_col, str) and cfg_col:
prob_col_candidates.append(cfg_col)
prob_col_candidates.extend(
[
"infection_prob",
"infection_xgb_proba",
"xgb_prob",
]
)
prob_col = None
for c in prob_col_candidates:
if c in df_prob.columns:
prob_col = c
break
if prob_col is None:
ax.set_visible(False)
return
cell_probs = (
df_prob[key_cols + [prob_col, infection_col]]
.groupby(key_cols, dropna=False)
.agg({prob_col: "mean", infection_col: "max"})
.reset_index()
)
cell_probs = cell_probs.replace([np.inf, -np.inf], np.nan)
cell_probs = cell_probs.dropna(subset=[prob_col])
if cell_probs.empty:
ax.set_visible(False)
return
mask_inf = cell_probs[infection_col].astype(bool)
probs_inf = cell_probs.loc[mask_inf, prob_col].to_numpy()
probs_uninf = cell_probs.loc[~mask_inf, prob_col].to_numpy()
bins = np.linspace(0.0, 1.0, 21)
if probs_uninf.size:
ax.hist(
probs_uninf,
bins=bins,
alpha=0.5,
color=UNINFECTED_COLOUR,
label="Uninfected",
)
if probs_inf.size:
ax.hist(
probs_inf,
bins=bins,
alpha=0.5,
color=INFECTED_COLOUR,
label="Infected",
)
ax.set_xlabel("XGBoost infection probability")
ax.set_ylabel("Cells")
ax.set_title("Probability separation (adjusted labels)")
ax.legend(fontsize=TYPE_SCALE["legend"], frameon=False)
def _plot_inf_uninf_bar(ax, df_vals, value_col, title, ylabel):
"""
Helper to plot infected vs uninfected distributions for the given column,
using violin plots (with mean markers) instead of barplots.
"""
cell_level = (
df_vals[key_cols + [value_col, infection_col]]
.groupby(key_cols, dropna=False)
.agg({value_col: "mean", infection_col: "max"})
.reset_index()
)
cell_level = cell_level.replace([np.inf, -np.inf], np.nan)
cell_level = cell_level.dropna(subset=[value_col])
if cell_level.empty:
ax.set_visible(False)
return
mask_inf = cell_level[infection_col].astype(bool)
vals_inf = cell_level.loc[mask_inf, value_col].to_numpy()
vals_uninf = cell_level.loc[~mask_inf, value_col].to_numpy()
data = []
positions = []
colors = []
labels_xtick = []
pos = 0
if vals_inf.size:
data.append(vals_inf)
positions.append(pos)
colors.append(INFECTED_COLOUR)
labels_xtick.append("Inf")
pos += 1
if vals_uninf.size:
data.append(vals_uninf)
positions.append(pos)
colors.append(UNINFECTED_COLOUR)
labels_xtick.append("Uninf")
vp = ax.violinplot(
data,
positions=positions,
widths=0.6,
showmeans=False,
showmedians=False,
showextrema=False,
)
ink = resolve_ink(theme_target())
for body, color in zip(vp["bodies"], colors):
body.set_facecolor(color)
body.set_edgecolor(ink)
body.set_alpha(0.6)
means = [float(np.nanmean(d)) for d in data]
ax.scatter(positions, means, color=ink, s=10, zorder=3)
ax.set_xticks(positions)
ax.set_xticklabels(labels_xtick)
flat = np.concatenate(data)
if np.nanmin(flat) >= 0:
ymin, ymax = ax.get_ylim()
ax.set_ylim(bottom=0, top=ymax)
ax.set_title(title)
ax.set_ylabel(ylabel)
if not {"plateID", "wellID"}.issubset(all_df.columns):
print(
"[_make_intensity_motility_panel] Missing 'plateID'/'wellID' columns; "
"cannot make per-well panels."
)
return
unique_wells = (
all_df[["plateID", "wellID"]]
.dropna()
.drop_duplicates()
.to_records(index=False)
)
for plate_id, well_id in unique_wells:
df_well = all_df[
(all_df["plateID"] == plate_id) & (all_df["wellID"] == well_id)
]
track_df_well = track_df[
(track_df["plateID"] == plate_id) & (track_df["wellID"] == well_id)
]
well_tracks = []
for tracks in per_well_tracks.values():
for tr in tracks:
if tr.get("plateID") == plate_id and tr.get("wellID") == well_id:
well_tracks.append(tr)
if df_well.empty or track_df_well.empty or not well_tracks:
print(
f"[_make_intensity_motility_panel] No data for plate={plate_id}, "
f"well={well_id}; skipping."
)
continue
available_channels = [
ch
for ch in range(n_channels)
if f"cell_mean_intensity_ch{ch}" in df_well.columns
]
if not available_channels:
print(
f"[_make_intensity_motility_panel] No cell_mean_intensity_ch* "
f"columns for plate={plate_id}, well={well_id}; skipping."
)
continue
has_p75_path = False
has_rel_int = False
if pathogen_chan is not None:
p75_col = f"cell_p75_intensity_ch{pathogen_chan}"
if p75_col in df_well.columns:
has_p75_path = True
path_col = f"pathogen_mean_intensity_ch{pathogen_chan}"
cyto_col = f"cytoplasm_mean_intensity_ch{pathogen_chan}"
if path_col in df_well.columns and cyto_col in df_well.columns:
has_rel_int = True
extra_int_plots = (1 if has_p75_path else 0) + (1 if has_rel_int else 0)
n_int_plots = len(available_channels) + extra_int_plots
n_cols = n_int_plots + 3 + (1 if qc_panel_needed_mask else 0) + qc_axes_count
with figure_style(theme_target()):
fig, axes = plt.subplots(1, n_cols, figsize=(4 * n_cols, 4))
from .figures.bundle import _register_figure_data
_register_figure_data(fig, df_well, x=infection_col, y=f"cell_mean_intensity_ch{available_channels[0]}" if len(available_channels) else "", kind="bar")
axes = np.array(axes).ravel()
axis_idx = 0
for ch in available_channels:
col_int = f"cell_mean_intensity_ch{ch}"
ax = axes[axis_idx]
axis_idx += 1
_plot_inf_uninf_bar(
ax,
df_well,
value_col=col_int,
title=f"Ch {ch} mean",
ylabel="Mean cell intensity",
)
if pathogen_chan is not None and ch == pathogen_chan:
if has_p75_path:
ax_p75 = axes[axis_idx]
axis_idx += 1
p75_col = f"cell_p75_intensity_ch{pathogen_chan}"
_plot_inf_uninf_bar(
ax_p75,
df_well,
value_col=p75_col,
title=f"Ch {pathogen_chan} p75",
ylabel="Cell p75 intensity",
)
if has_rel_int:
ax_rel = axes[axis_idx]
axis_idx += 1
path_col = f"pathogen_mean_intensity_ch{pathogen_chan}"
cyto_col = f"cytoplasm_mean_intensity_ch{pathogen_chan}"
df_ratio = df_well[
key_cols + [path_col, cyto_col, infection_col]
].copy()
df_ratio["rel_intensity"] = df_ratio[path_col] / df_ratio[
cyto_col
].replace(0, np.nan)
_plot_inf_uninf_bar(
ax_rel,
df_ratio,
value_col="rel_intensity",
title=f"Ch {pathogen_chan} pathogen/cytoplasm",
ylabel="Intensity ratio",
)
def _plot_all_tracks(ax):
"""Plot tracks on ``ax`` in absolute pixel or calibrated coordinates; return ``None``."""
xs_all = []
ys_all = []
n_inf_tr = 0
n_uninf_tr = 0
for tr in well_tracks:
x_px = np.asarray(tr["x_px"], dtype=float)
y_px = np.asarray(tr["y_px"], dtype=float)
if x_px.size < 2:
continue
x = x_px * coord_scale
y = y_px * coord_scale
infected_tr = bool(tr.get("infected", False))
color = INFECTED_COLOUR if infected_tr else UNINFECTED_COLOUR
ax.plot(x, y, color=color, alpha=0.15, linewidth=0.5)
ax.scatter(x[-1], y[-1], color=color, s=5)
xs_all.append(x)
ys_all.append(y)
if infected_tr:
n_inf_tr += 1
else:
n_uninf_tr += 1
if not xs_all:
ax.set_visible(False)
return
xs_all = np.concatenate(xs_all)
ys_all = np.concatenate(ys_all)
ax.set_aspect("equal", "box")
ax.set_xlabel(coord_label_x)
ax.set_ylabel(coord_label_y)
x_margin = 0.05 * (xs_all.max() - xs_all.min() + 1e-9)
y_margin = 0.05 * (ys_all.max() - ys_all.min() + 1e-9)
ax.set_xlim(xs_all.min() - x_margin, xs_all.max() + x_margin)
ax.set_ylim(ys_all.min() - y_margin, ys_all.max() + y_margin)
mask_inf = track_df_well["infected"].astype(bool)
v_inf = track_df_well.loc[mask_inf, "velocity"].to_numpy()
v_uninf = track_df_well.loc[~mask_inf, "velocity"].to_numpy()
mean_inf_v = float(np.nanmean(v_inf)) if v_inf.size else np.nan
mean_uninf_v = float(np.nanmean(v_uninf)) if v_uninf.size else np.nan
txt_lines = []
txt_lines.append(f"Infected ({mean_inf_v:.2f} {vel_unit})")
txt_lines.append(f"Uninfected ({mean_uninf_v:.2f} {vel_unit})")
if pixels_per_um is not None and pixels_per_um > 0:
txt_lines.append(f"1 µm = {pixels_per_um:.2f} px")
if seconds_per_frame is not None:
txt_lines.append(f"1 frame = {seconds_per_frame:.0f} s")
txt = "\n".join(txt_lines)
ax.text(
0.98,
0.02,
txt,
transform=ax.transAxes,
ha="right",
va="bottom",
fontsize=TYPE_SCALE["annotation"],
color=resolve_ink(theme_target()),
)
ax_all = axes[axis_idx]
axis_idx += 1
_plot_all_tracks(ax_all)
def _plot_origin(ax, want_infected: bool):
"""Plot the selected group on ``ax`` relative to its origins; return ``None``."""
n_tr = 0
color = INFECTED_COLOUR if want_infected else UNINFECTED_COLOUR
for tr in well_tracks:
if bool(tr.get("infected", False)) != want_infected:
continue
x_px = np.asarray(tr["x_px"], dtype=float)
y_px = np.asarray(tr["y_px"], dtype=float)
if x_px.size < 2:
continue
x = (x_px - x_px[0]) * coord_scale
y = (y_px - y_px[0]) * coord_scale
ax.plot(x, y, color=color, alpha=0.15, linewidth=0.5)
ax.scatter(x[-1], y[-1], color=color, s=5)
n_tr += 1
ax.set_aspect("equal", "box")
ax.set_xlabel(coord_label_x)
ax.set_ylabel(coord_label_y)
if origin_xlim is not None and len(origin_xlim) == 2:
ax.set_xlim(origin_xlim)
if origin_ylim is not None and len(origin_ylim) == 2:
ax.set_ylim(origin_ylim)
mask = track_df_well["infected"].astype(bool)
if not want_infected:
mask = ~mask
v = track_df_well.loc[mask, "velocity"].to_numpy()
mean_v = float(np.nanmean(v)) if v.size else np.nan
label = "Infected" if want_infected else "Uninfected"
ax.set_title(f"{label}\n(n={n_tr}, v={mean_v:.2f} {vel_unit})")
ax_inf = axes[axis_idx]
axis_idx += 1
_plot_origin(ax_inf, True)
ax_uninf = axes[axis_idx]
axis_idx += 1
_plot_origin(ax_uninf, False)
if qc_panel_needed_mask and axis_idx < len(axes):
ax_qc = axes[axis_idx]
axis_idx += 1
try:
img = mpimg.imread(qc_panel_path)
ax_qc.imshow(img)
ax_qc.axis("off")
tmap = {
"histogram": "Intensity histogram",
"pca": "PCA/UMAP clustering",
"xgboost": "XGBoost feature importance",
}
ttl = tmap.get(str(qc_panel_type).lower(), "Infection QC")
ax_qc.set_title(ttl, fontsize=9)
except Exception as e:
print(
f"[_make_intensity_motility_panel] Could not embed QC plot "
f"from {qc_panel_path}: {e}"
)
ax_qc.set_visible(False)
if is_adjusted_panel and qc_graphs_enabled:
if qc_strategy == "histogram" and axis_idx < len(axes):
ax_hist = axes[axis_idx]
axis_idx += 1
src = hist_data if hist_data is not None else df_well
_plot_hist_qc(ax_hist, src)
elif qc_strategy in {"pca", "umap", "tsne"} and has_pca and axis_idx < len(axes):
ax_pca = axes[axis_idx]
axis_idx += 1
_plot_pca_qc(ax_pca, pca_data)
elif qc_strategy == "xgboost" and has_xgb:
ax_prob = axes[axis_idx]
axis_idx += 1
_plot_xgb_prob_qc(ax_prob, df_well)
ax_xgb = axes[axis_idx]
axis_idx += 1
_plot_xgb_importance_qc(ax_xgb, xgb_data)
meta_tag = f"{plate_id}_{well_id}"
fig.suptitle(
f"Infection panel – {panel_label} labels – method={method_label}\n{meta_tag}",
fontsize=10,
)
fig.tight_layout(rect=[0, 0, 1, 0.90])
if is_adjusted_panel:
out_name = f"{meta_tag}_{method_label}_adjusted.pdf"
elif is_mask_panel:
out_name = f"{meta_tag}.pdf"
else:
out_name = f"{meta_tag}_{label_tag}_{method_label}.pdf"
out_path = os.path.join(motility_dir, out_name)
out_path = save_figure_to_path(fig, out_path)
plt.close(fig)
print(
f"[summarise_tracks_from_merged] Saved per-well intensity+motility panel "
f"({panel_label}, method={method_label}) for plate={plate_id}, well={well_id} "
f"to {out_path}"
)
def _infer_plate_well_meta_tag(df):
"""
Infer a compact 'plate_well' tag for filenames from a DataFrame that has
plateID / wellID columns.
Examples
--------
plate1 + A02 -> 'plate1_A02'
plate1 + many -> 'plate1_MULTI_WELLS'
many + A02 -> 'MULTI_PLATES_A02'
many + many -> 'MULTI_PLATES_MULTI_WELLS'
"""
plates = sorted(df["plateID"].dropna().unique()) if "plateID" in df.columns else []
wells = sorted(df["wellID"].dropna().unique()) if "wellID" in df.columns else []
if len(plates) == 1 and len(wells) == 1:
return f"{plates[0]}_{wells[0]}"
elif len(plates) == 1 and len(wells) > 1:
return f"{plates[0]}_MULTI_WELLS"
elif len(plates) > 1 and len(wells) == 1:
return f"MULTI_PLATES_{wells[0]}"
else:
return "MULTI_PLATES_MULTI_WELLS"
def _compute_cell_mean_intensity_per_channel(
mask_stack,
intensity_stack,
channel_index,
):
"""
Compute per-frame, per-cell mean intensity for a given channel.
Parameters
----------
mask_stack : ndarray
Label image stack of shape (T, Y, X) for cells (track_id labels).
intensity_stack : ndarray
Intensity stack of shape (T, Y, X, C).
channel_index : int
Channel index in intensity_stack to use.
Returns
-------
DataFrame
Columns: ['frame', 'track_id', f'cell_mean_intensity_ch{channel_index}']
"""
import numpy as np
import pandas as pd
if intensity_stack is None:
print(
f"[cell_mean_intensity] channel {channel_index}: "
"intensity_stack is None, skipping."
)
return pd.DataFrame(
columns=["frame", "track_id", f"cell_mean_intensity_ch{channel_index}"]
)
if channel_index is None or channel_index < 0 or channel_index >= intensity_stack.shape[-1]:
print(
f"[cell_mean_intensity] channel {channel_index}: "
"invalid channel index for intensity_stack, skipping."
)
return pd.DataFrame(
columns=["frame", "track_id", f"cell_mean_intensity_ch{channel_index}"]
)
T = mask_stack.shape[0]
dfs = []
col_name = f"cell_mean_intensity_ch{channel_index}"
for frame in range(T):
labels = mask_stack[frame]
if not np.any(labels):
continue
intensity_image = intensity_stack[frame, :, :, channel_index]
props_table = regionprops_table(
labels,
intensity_image=intensity_image,
properties=("label", "mean_intensity"),
)
frame_df = pd.DataFrame(props_table)
frame_df = frame_df.rename(
columns={
"label": "track_id",
"mean_intensity": col_name,
}
)
frame_df["frame"] = frame
dfs.append(frame_df)
if not dfs:
print(
f"[cell_mean_intensity] channel {channel_index}: "
f"no objects found in any of {T} frames."
)
return pd.DataFrame(columns=["frame", "track_id", col_name])
out_df = pd.concat(dfs, ignore_index=True)
n_rows = out_df.shape[0]
n_frames_detected = out_df["frame"].nunique()
n_objs = out_df["track_id"].nunique()
print(
f"[cell_mean_intensity] channel {channel_index}: "
f"frames_with_objects={n_frames_detected}/{T}, "
f"unique_track_id={n_objs}, rows={n_rows}"
)
return out_df
def _reorient_merged_array(arr, n_channels, max_extra_masks=3):
"""
Ensure merged array has shape (planes, H, W) with planes as the first axis.
Handles both (planes, H, W) and (H, W, planes) layouts by detecting which
axis likely corresponds to the small "planes" dimension (~n_channels + masks).
"""
import numpy as np
if arr.ndim != 3:
raise ValueError(
f"_reorient_merged_array expected 3D array, got ndim={arr.ndim}"
)
target_min = n_channels
target_max = n_channels + max_extra_masks
shape = arr.shape
plane_axis = None
for ax, dim in enumerate(shape):
if target_min <= dim <= target_max:
plane_axis = ax
break
if plane_axis is None:
plane_axis = int(np.argmin(shape))
if plane_axis != 0:
arr = np.moveaxis(arr, plane_axis, 0)
planes, H, W = arr.shape
return arr, planes, H, W
def _parse_merged_filename(fname):
"""
Parse a merged .npy filename of the form:
plate_well_field_time.npy
Returns a dict with:
plateID, wellID, rowID, columnID, fieldID, timeID, prcf, prcft, filename
.. warning::
**These are not the canonical spaCR keys and must never be written to
a database or joined on.** This is a grouping-and-sorting helper for
the motility/track summaries, and it deliberately does not go through
:mod:`spacr.schema`:
* ``rowID`` here is the row *letter* (``'B'``), not ``'r2'``;
* ``columnID`` is an ``int`` (``3``), not ``'c3'``;
* ``timeID`` is an ``int``, because the callers sort on it and
``'t10' < 't2'`` as a string;
* ``prcf`` here is ``plate_well_field``, **not** the canonical
``plate_row_column_field``, and ``prcft`` appends the bare integer
timepoint.
Only ``plateID`` / ``wellID`` / ``fieldID`` / ``timeID`` / ``filename``
are read by any caller (``_process_merged_group``,
``summarise_tracks_from_merged`` and the Qt motility preview all group
on ``(plateID, wellID, fieldID)`` and sort on ``timeID``). If you need
a real key from one of these names, call
:func:`spacr.schema.parse_field_stem`, which returns the same identity
the measurement tables carry.
"""
from ._merged_names import parse_merged_filename
return parse_merged_filename(fname)
def _compute_parent_child_overlaps(
parent_masks,
child_masks,
parent_label_col,
child_label_col,
):
"""
For each frame, find which child labels overlap which parent labels.
Returns columns: 'frame', parent_label_col, child_label_col
"""
T = parent_masks.shape[0]
records = []
for frame in range(T):
p = parent_masks[frame]
c = child_masks[frame]
m = (p > 0) & (c > 0)
if not np.any(m):
continue
p_flat = p[m].ravel()
c_flat = c[m].ravel()
pairs = np.stack([p_flat, c_flat], axis=1)
unique_pairs = np.unique(pairs, axis=0)
for parent_label, child_label in unique_pairs:
records.append(
{
"frame": frame,
parent_label_col: int(parent_label),
child_label_col: int(child_label),
}
)
if not records:
return pd.DataFrame(columns=["frame", parent_label_col, child_label_col])
return pd.DataFrame.from_records(records)
def _summarise_child_features_per_parent(
overlaps_df,
child_props_df,
parent_label_col,
child_label_col,
count_col_name,
):
"""
Summarise child object features per parent object.
- Counts distinct children -> count_col_name
- Aggregates numeric child features per parent:
* '*area*' -> sum
* '*intensity*' -> mean
* '*dist*'/'*distance*' -> min
* everything else -> mean
"""
if overlaps_df.empty or child_props_df.empty:
return pd.DataFrame(columns=["frame", parent_label_col, count_col_name])
df = overlaps_df.merge(child_props_df, on=["frame", child_label_col], how="left", validate="many_to_one")
if df.empty:
return pd.DataFrame(columns=["frame", parent_label_col, count_col_name])
group_cols = ["frame", parent_label_col]
counts = (
df.groupby(group_cols)[child_label_col]
.nunique()
.reset_index()
.rename(columns={child_label_col: count_col_name})
)
numeric_cols = [
c
for c in df.columns
if c not in group_cols + [child_label_col]
and pd.api.types.is_numeric_dtype(df[c])
]
if not numeric_cols:
return counts
def _agg_for_feature(col_name: str) -> str:
"""Return sum for area, minimum for distance, otherwise mean."""
name = col_name.lower()
if "area" in name:
return "sum"
if "intensity" in name:
return "mean"
if "distance" in name or "dist" in name:
return "min"
return "mean"
agg_dict = {c: _agg_for_feature(c) for c in numeric_cols}
agg_df = df.groupby(group_cols).agg(agg_dict).reset_index()
summary = agg_df.merge(counts, on=group_cols, how="left",
validate="one_to_one")
return summary
def _load_intensity_stack_from_merged(
src,
filenames,
n_channels,
height,
width,
dtype=np.float32,
):
"""
Load intensity channels from merged/*.npy into a (T, H, W, C) stack.
Supports merged arrays stored either as (planes, H, W) or (H, W, planes).
The first n_channels planes are intensities and any remaining planes are masks.
"""
import os
import numpy as np
merged_dir = os.path.join(src, "merged")
T = len(filenames)
if not os.path.isdir(merged_dir) or n_channels is None or n_channels <= 0:
return np.zeros((T, height, width, 0), dtype=dtype)
stack = np.zeros((T, height, width, n_channels), dtype=dtype)
for t, fn in enumerate(filenames):
base = os.path.splitext(os.path.basename(fn))[0]
candidates = [
os.path.join(merged_dir, base + ".npy"),
os.path.join(merged_dir, fn),
os.path.join(merged_dir, fn + ".npy"),
]
arr = None
for path in candidates:
if os.path.exists(path):
arr = np.load(path)
break
if arr is None or arr.ndim != 3:
continue
try:
arr, planes, H_img, W_img = _reorient_merged_array(
arr, n_channels=n_channels
)
except ValueError:
continue
if H_img != height or W_img != width:
print(
f"[_load_intensity_stack_from_merged] Skipping {fn}: "
f"reoriented size=({planes}, {H_img}, {W_img}), "
f"expected H={height}, W={width}"
)
continue
use_planes = min(n_channels, planes)
if use_planes <= 0:
continue
img = arr[:use_planes].transpose(1, 2, 0)
C = img.shape[2]
stack[t, :, :, :C] = img
return stack
def _load_masks_from_merged(
src,
filenames,
n_channels,
height,
width,
nucleus_chan=None,
pathogen_chan=None,
dtype=None,
):
"""
Load cell / nucleus / pathogen masks from merged/*.npy.
Supports merged arrays stored either as (planes, H, W) or (H, W, planes).
Layout per merged array after reorientation (planes, H, W):
0 .. n_channels-1 → intensity channels
n_channels → cell_mask (always present)
n_channels + 1 (optional) → nucleus_mask or pathogen_mask
n_channels + 2 (optional) → pathogen_mask (when both nuc+pathogen exist)
The exact interpretation of mask planes depends on whether
`nucleus_chan` and/or `pathogen_chan` are None.
"""
import os
import numpy as np
if dtype is None:
dtype = np.int32
merged_dir = os.path.join(src, "merged")
T = len(filenames)
cell_masks = np.zeros((T, height, width), dtype=dtype)
nucleus_masks = np.zeros((T, height, width), dtype=dtype)
pathogen_masks = np.zeros((T, height, width), dtype=dtype)
if not os.path.isdir(merged_dir):
return cell_masks, nucleus_masks, pathogen_masks
for t, fn in enumerate(filenames):
base = os.path.splitext(os.path.basename(fn))[0]
candidates = [
os.path.join(merged_dir, base + ".npy"),
os.path.join(merged_dir, fn),
os.path.join(merged_dir, fn + ".npy"),
]
arr = None
for path in candidates:
if os.path.exists(path):
arr = np.load(path)
break
if arr is None or arr.ndim != 3:
continue
try:
arr, planes, H_img, W_img = _reorient_merged_array(
arr, n_channels=n_channels
)
except ValueError:
continue
if H_img != height or W_img != width:
print(
f"[_load_masks_from_merged] Skipping {fn}: "
f"reoriented size=({planes}, {H_img}, {W_img}), "
f"expected H={height}, W={width}"
)
continue
if planes <= n_channels:
continue
n_masks = planes - n_channels
cell_masks[t] = arr[n_channels].astype(dtype)
if n_masks >= 2:
if nucleus_chan is not None and pathogen_chan is None:
nucleus_masks[t] = arr[n_channels + 1].astype(dtype)
elif nucleus_chan is None and pathogen_chan is not None:
pathogen_masks[t] = arr[n_channels + 1].astype(dtype)
elif nucleus_chan is not None and pathogen_chan is not None:
nucleus_masks[t] = arr[n_channels + 1].astype(dtype)
if n_masks >= 3 and pathogen_chan is not None:
pathogen_masks[t] = arr[n_channels + 2].astype(dtype)
return cell_masks, nucleus_masks, pathogen_masks
def _compute_regionprops_stack(
mask_stack,
intensity_stack,
channel_index,
object_prefix,
label_as_track_id=False,
):
"""
Compute regionprops over a (T, Y, X) label stack.
Parameters
----------
mask_stack : ndarray
Label image stack of shape (T, Y, X).
intensity_stack : ndarray or None
Intensity stack of shape (T, Y, X, C) or None.
channel_index : int or None
Channel index in intensity_stack to use for intensity props.
object_prefix : str
Prefix for column names ("cell", "nucleus", "pathogen", "cytoplasm").
label_as_track_id : bool
If True, rename 'label' to 'track_id',
otherwise to f"{object_prefix}_label".
Returns
-------
DataFrame
One row per object per frame with prefixed column names.
"""
import numpy as np
import pandas as pd
T, H, W = mask_stack.shape
use_intensity = (
intensity_stack is not None
and channel_index is not None
and 0 <= channel_index < intensity_stack.shape[-1]
)
geom_props = [
"label",
"area",
"bbox_area",
"equivalent_diameter",
"perimeter",
"perimeter_crofton",
"solidity",
"centroid",
]
intensity_props = [
"max_intensity",
"mean_intensity",
"min_intensity",
]
props = geom_props + intensity_props if use_intensity else geom_props
label_col_name = "track_id" if label_as_track_id else f"{object_prefix}_label"
dfs = []
for frame in range(T):
labels = mask_stack[frame]
if not np.any(labels):
continue
if use_intensity:
intensity_image = intensity_stack[frame, :, :, channel_index]
props_table = regionprops_table(
labels,
intensity_image=intensity_image,
properties=props,
)
else:
props_table = regionprops_table(labels, properties=props)
frame_df = pd.DataFrame(props_table)
frame_df = frame_df.rename(columns={"label": label_col_name})
frame_df["frame"] = frame
feature_cols = [
c for c in frame_df.columns if c not in ("frame", label_col_name)
]
frame_df = frame_df.rename(
columns={c: f"{object_prefix}_{c}" for c in feature_cols}
)
dfs.append(frame_df)
if not dfs:
print(f"[regionprops] {object_prefix}: no objects found in any of {T} frames.")
return pd.DataFrame(columns=["frame", label_col_name])
out_df = pd.concat(dfs, ignore_index=True)
n_rows = out_df.shape[0]
n_frames_detected = out_df["frame"].nunique()
n_objs = out_df[label_col_name].nunique()
print(
f"[regionprops] {object_prefix}: frames_with_objects="
f"{n_frames_detected}/{T}, unique_{label_col_name}={n_objs}, rows={n_rows}"
)
return out_df
def _motility_role_bleaching(props, masks, images, role, role_channel, times, method):
"""Correct individual objects before aggregation, retaining their raw levels.
A five-pixel ring excludes all foreground labels of the current role.
Missing ring pixels yield NaN, never an assumed zero camera background.
Return corrected properties, fit rows, and columns added only by this path.
"""
from scipy.ndimage import binary_dilation
if props.empty:
return props.copy(), pd.DataFrame(), set()
result = props.copy()
original_columns = set(result.columns)
label_key = 'track_id' if role == 'cell' else f'{role}_label'
measured = pd.DataFrame({'timeID': [times[int(t)] for t in result['frame']]})
measured[f'{role}_area'] = result[f'{role}_area'].to_numpy()
for channel in range(images.shape[-1]):
means, backgrounds = [], []
for frame, label in zip(result['frame'], result[label_key]):
labels = masks[int(frame)]
inside = labels == label
ring = binary_dilation(inside, iterations=5) & (labels == 0)
plane = images[int(frame), :, :, channel]
values = plane[inside]
values = values[np.isfinite(values)]
outside = plane[ring]
outside = outside[np.isfinite(outside)]
means.append(float(np.mean(values)) if values.size else np.nan)
backgrounds.append(float(np.median(outside)) if outside.size else np.nan)
mean_key = f'{role}_mean_intensity_ch{channel}'
if mean_key not in result:
result[mean_key] = means
ring_key = f'{role}_channel_{channel}_outside_percentile_50'
result[ring_key] = backgrounds
measured[ring_key] = backgrounds
mapping = {}
for column in list(result):
match = re.fullmatch(rf'{role}_(mean|max|min|p[0-9]+)_intensity_ch([0-9]+)', column)
if match:
stat, channel = match.groups()
stat = f'percentile_{int(stat[1:])}' if stat.startswith('p') else stat + '_intensity'
mapping[column] = f'{role}_channel_{channel}_{stat}'
elif role_channel is not None and column in {
f'{role}_mean_intensity', f'{role}_max_intensity', f'{role}_min_intensity'}:
mapping[column] = f'{role}_channel_{role_channel}_{column[len(role) + 1:]}'
for column, canonical in mapping.items():
measured[canonical] = result[column].to_numpy()
corrected, fits = _bleach_correct_table(measured, role, method)
for column, canonical in mapping.items():
if column in original_columns:
result['raw_' + column] = result[column]
result[column] = corrected[canonical].to_numpy()
added = set(result) - original_columns
return result, fits, added
def _motility_raw_frame(corrected, added):
"""Recover the unchanged measured levels before saving the original table."""
raw = corrected.copy()
for column in list(raw):
if column.startswith('raw_'):
raw[column[4:]] = raw[column]
return raw.drop(columns=[c for c in added if c in raw])
def _motility_smooth_corrected(frame, max_displacement, track_outlier_zscore):
"""Keep unknown corrected levels unknown when geometric glitches are smoothed."""
keys = ['plateID', 'wellID', 'fieldID', 'cellID', 'frame']
levels = [c for c in frame if '_intensity' in c]
missing = frame.set_index(keys)[levels].isna()
result = _smooth_tracks_and_features(frame, max_displacement, track_outlier_zscore)
positions = pd.MultiIndex.from_frame(result[keys])
for column in levels:
result.loc[missing[column].reindex(positions).to_numpy(), column] = np.nan
return result
def _process_merged_group(args):
"""
Worker: process one (plate, well, field) group of merged .npy files.
Returns per-cell-per-frame DataFrame with:
- metadata
- cell features
- aggregated nucleus / pathogen / cytoplasm features
- per-channel cell mean intensities (cell_mean_intensity_ch{c})
An optional seventh argument selects bleach correction. That path returns
(corrected, raw, fits), correcting individual roles before aggregation.
Empty/unusable groups retain the historical empty-DataFrame result.
"""
import numpy as np
import pandas as pd
import os
(
src,
file_basenames,
n_channels,
cell_chan,
nucleus_chan,
pathogen_chan,
) = args[:6]
bleach_method = args[6] if len(args) > 6 else 'none'
if bleach_method not in _BLEACH_METHODS:
raise ValueError(f'Unknown motility bleach correction: {bleach_method!r}')
if not file_basenames:
print("[_process_merged_group] Empty file_basenames list.")
return pd.DataFrame()
merged_dir = os.path.join(src, "merged")
metas = []
for bn in file_basenames:
meta = _parse_merged_filename(bn)
metas.append(meta)
metas_sorted = sorted(metas, key=lambda m: m["timeID"])
sorted_basenames = [m["filename"] for m in metas_sorted]
key = (
metas_sorted[0]["plateID"],
metas_sorted[0]["wellID"],
metas_sorted[0]["fieldID"],
)
print(f"[_process_merged_group] Start group {key}, files={len(sorted_basenames)}")
first_path = os.path.join(merged_dir, sorted_basenames[0])
first_arr_raw = np.load(first_path)
if first_arr_raw.ndim != 3:
print(
f"[_process_merged_group] First array for group {key} is not 3D, "
"skipping."
)
return pd.DataFrame()
try:
first_arr, planes, H, W = _reorient_merged_array(
first_arr_raw, n_channels=n_channels
)
except ValueError:
print(
f"[_process_merged_group] Group {key}: could not reorient first array "
f"with shape={first_arr_raw.shape}, skipping."
)
return pd.DataFrame()
base_dtype = first_arr.dtype
print(
f"[_process_merged_group] Group {key}: first array original_shape="
f"{first_arr_raw.shape}, reoriented_shape=({planes}, {H}, {W}), "
f"dtype={base_dtype}"
)
intensity_stack = _load_intensity_stack_from_merged(
src=src,
filenames=sorted_basenames,
n_channels=n_channels,
height=H,
width=W,
dtype=base_dtype,
)
cell_masks, nucleus_masks, pathogen_masks = _load_masks_from_merged(
src=src,
filenames=sorted_basenames,
n_channels=n_channels,
height=H,
width=W,
nucleus_chan=nucleus_chan,
pathogen_chan=pathogen_chan,
dtype=np.int32,
)
T = cell_masks.shape[0]
if T == 0 or not np.any(cell_masks):
print(f"[_process_merged_group] Group {key}: no cell masks found, skipping.")
return pd.DataFrame()
print(
f"[_process_merged_group] Group {key}: frames={T}, "
f"any_nucleus={np.any(nucleus_masks)}, any_pathogen={np.any(pathogen_masks)}"
)
has_nucleus = np.any(nucleus_masks)
has_pathogen = np.any(pathogen_masks)
cytoplasm_masks = None
if has_nucleus or has_pathogen:
cytoplasm_masks = cell_masks.copy()
if has_nucleus:
cytoplasm_masks[nucleus_masks > 0] = 0
if has_pathogen:
cytoplasm_masks[pathogen_masks > 0] = 0
cell_props_df = _compute_regionprops_stack(
mask_stack=cell_masks,
intensity_stack=intensity_stack,
channel_index=cell_chan,
object_prefix="cell",
label_as_track_id=True,
)
nucleus_props_df = _compute_regionprops_stack(
mask_stack=nucleus_masks,
intensity_stack=intensity_stack,
channel_index=nucleus_chan,
object_prefix="nucleus",
label_as_track_id=False,
)
pathogen_props_df = _compute_regionprops_stack(
mask_stack=pathogen_masks,
intensity_stack=intensity_stack,
channel_index=pathogen_chan,
object_prefix="pathogen",
label_as_track_id=False,
)
cytoplasm_props_df = pd.DataFrame()
if cytoplasm_masks is not None and np.any(cytoplasm_masks):
cytoplasm_props_df = _compute_regionprops_stack(
mask_stack=cytoplasm_masks,
intensity_stack=intensity_stack,
channel_index=cell_chan,
object_prefix="cytoplasm",
label_as_track_id=False,
)
percentile_dfs_cell = []
percentile_dfs_nucleus = []
percentile_dfs_pathogen = []
percentile_dfs_cytoplasm = []
for ch in range(n_channels):
df_p = _compute_intensity_percentiles_per_channel(
mask_stack=cell_masks,
intensity_stack=intensity_stack,
channel_index=ch,
object_prefix="cell",
label_as_track_id=True,
)
if not df_p.empty:
percentile_dfs_cell.append(df_p)
if np.any(nucleus_masks):
df_p_n = _compute_intensity_percentiles_per_channel(
mask_stack=nucleus_masks,
intensity_stack=intensity_stack,
channel_index=ch,
object_prefix="nucleus",
label_as_track_id=False,
)
if not df_p_n.empty:
percentile_dfs_nucleus.append(df_p_n)
if np.any(pathogen_masks):
df_p_pa = _compute_intensity_percentiles_per_channel(
mask_stack=pathogen_masks,
intensity_stack=intensity_stack,
channel_index=ch,
object_prefix="pathogen",
label_as_track_id=False,
)
if not df_p_pa.empty:
percentile_dfs_pathogen.append(df_p_pa)
if cytoplasm_masks is not None and np.any(cytoplasm_masks):
df_p_cy = _compute_intensity_percentiles_per_channel(
mask_stack=cytoplasm_masks,
intensity_stack=intensity_stack,
channel_index=ch,
object_prefix="cytoplasm",
label_as_track_id=False,
)
if not df_p_cy.empty:
percentile_dfs_cytoplasm.append(df_p_cy)
if percentile_dfs_cell:
tmp = percentile_dfs_cell[0]
for df_p in percentile_dfs_cell[1:]:
tmp = tmp.merge(df_p, on=["frame", "track_id"], how="outer",
validate="one_to_one")
cell_props_df = cell_props_df.merge(
tmp, on=["frame", "track_id"], how="left", validate="one_to_one"
)
if np.any(nucleus_masks) and not nucleus_props_df.empty and percentile_dfs_nucleus:
tmp = percentile_dfs_nucleus[0]
for df_p in percentile_dfs_nucleus[1:]:
tmp = tmp.merge(df_p, on=["frame", "nucleus_label"], how="outer",
validate="one_to_one")
nucleus_props_df = nucleus_props_df.merge(
tmp, on=["frame", "nucleus_label"], how="left", validate="one_to_one"
)
if np.any(pathogen_masks) and not pathogen_props_df.empty and percentile_dfs_pathogen:
tmp = percentile_dfs_pathogen[0]
for df_p in percentile_dfs_pathogen[1:]:
tmp = tmp.merge(df_p, on=["frame", "pathogen_label"], how="outer",
validate="one_to_one")
pathogen_props_df = pathogen_props_df.merge(
tmp, on=["frame", "pathogen_label"], how="left", validate="one_to_one"
)
if (
cytoplasm_masks is not None
and np.any(cytoplasm_masks)
and not cytoplasm_props_df.empty
and percentile_dfs_cytoplasm
):
tmp = percentile_dfs_cytoplasm[0]
for df_p in percentile_dfs_cytoplasm[1:]:
tmp = tmp.merge(df_p, on=["frame", "cytoplasm_label"], how="outer",
validate="one_to_one")
cytoplasm_props_df = cytoplasm_props_df.merge(
tmp, on=["frame", "cytoplasm_label"], how="left", validate="one_to_one"
)
if cell_props_df.empty:
print(f"[_process_merged_group] Group {key}: cell_props_df empty, skipping.")
return pd.DataFrame()
per_channel_intensity_dfs = []
for ch in range(n_channels):
df_ch = _compute_cell_mean_intensity_per_channel(
mask_stack=cell_masks,
intensity_stack=intensity_stack,
channel_index=ch,
)
if not df_ch.empty:
per_channel_intensity_dfs.append(df_ch)
cell_intensity_df = None
if per_channel_intensity_dfs:
cell_intensity_df = per_channel_intensity_dfs[0]
for df_ch in per_channel_intensity_dfs[1:]:
cell_intensity_df = cell_intensity_df.merge(
df_ch,
on=["frame", "track_id"],
how="outer",
validate="one_to_one",
)
added_cols = [
c
for c in cell_intensity_df.columns
if c.startswith("cell_mean_intensity_ch")
]
print(
f"[_process_merged_group] Group {key}: added per-channel cell "
f"intensity columns: {added_cols}"
)
bleach_fits = []
bleach_added = set()
if bleach_method != 'none':
if cell_intensity_df is not None:
cell_props_df = cell_props_df.merge(
cell_intensity_df, on=['frame', 'track_id'], validate='one_to_one')
cell_intensity_df = None
role_frames = []
for props, masks, role, channel in (
(cell_props_df, cell_masks, 'cell', cell_chan),
(nucleus_props_df, nucleus_masks, 'nucleus', nucleus_chan),
(pathogen_props_df, pathogen_masks, 'pathogen', pathogen_chan),
(cytoplasm_props_df, cytoplasm_masks, 'cytoplasm', cell_chan)):
corrected, fits, added = _motility_role_bleaching(
props, masks, intensity_stack, role, channel,
[m['timeID'] for m in metas_sorted], bleach_method)
role_frames.append(corrected)
bleach_added.update(added)
if not fits.empty:
fits = fits.assign(plateID=key[0], wellID=key[1], fieldID=key[2])
bleach_fits.append(fits)
cell_props_df, nucleus_props_df, pathogen_props_df, cytoplasm_props_df = role_frames
nucleus_summary = None
if has_nucleus:
overlaps_cn = _compute_parent_child_overlaps(
parent_masks=cell_masks,
child_masks=nucleus_masks,
parent_label_col="track_id",
child_label_col="nucleus_label",
)
if not overlaps_cn.empty and not nucleus_props_df.empty:
nucleus_summary = _summarise_child_features_per_parent(
overlaps_df=overlaps_cn,
child_props_df=nucleus_props_df,
parent_label_col="track_id",
child_label_col="nucleus_label",
count_col_name="n_nuclei",
)
print(
f"[_process_merged_group] Group {key}: nucleus_summary rows="
f"{len(nucleus_summary)}"
)
pathogen_summary = None
if has_pathogen:
overlaps_cp = _compute_parent_child_overlaps(
parent_masks=cell_masks,
child_masks=pathogen_masks,
parent_label_col="track_id",
child_label_col="pathogen_label",
)
if not overlaps_cp.empty and not pathogen_props_df.empty:
pathogen_summary = _summarise_child_features_per_parent(
overlaps_df=overlaps_cp,
child_props_df=pathogen_props_df,
parent_label_col="track_id",
child_label_col="pathogen_label",
count_col_name="n_pathogens",
)
print(
f"[_process_merged_group] Group {key}: pathogen_summary rows="
f"{len(pathogen_summary)}"
)
cytoplasm_summary = None
if (
cytoplasm_masks is not None
and np.any(cytoplasm_masks)
and not cytoplasm_props_df.empty
):
overlaps_cc = _compute_parent_child_overlaps(
parent_masks=cell_masks,
child_masks=cytoplasm_masks,
parent_label_col="track_id",
child_label_col="cytoplasm_label",
)
if not overlaps_cc.empty:
cytoplasm_summary = _summarise_child_features_per_parent(
overlaps_df=overlaps_cc,
child_props_df=cytoplasm_props_df,
parent_label_col="track_id",
child_label_col="cytoplasm_label",
count_col_name="n_cytoplasm",
)
print(
f"[_process_merged_group] Group {key}: cytoplasm_summary rows="
f"{len(cytoplasm_summary)}"
)
enriched_df = cell_props_df.copy()
if nucleus_summary is not None and not nucleus_summary.empty:
enriched_df = enriched_df.merge(
nucleus_summary,
on=["frame", "track_id"],
how="left",
validate="one_to_one",
)
if pathogen_summary is not None and not pathogen_summary.empty:
enriched_df = enriched_df.merge(
pathogen_summary,
on=["frame", "track_id"],
how="left",
validate="one_to_one",
)
if cytoplasm_summary is not None and not cytoplasm_summary.empty:
enriched_df = enriched_df.merge(
cytoplasm_summary,
on=["frame", "track_id"],
how="left",
validate="one_to_one",
)
if cell_intensity_df is not None:
enriched_df = enriched_df.merge(
cell_intensity_df,
on=["frame", "track_id"],
how="left",
validate="one_to_one",
)
meta_records = []
for local_frame_idx, meta in enumerate(metas_sorted):
rec = {"frame": local_frame_idx}
rec.update(meta)
meta_records.append(rec)
meta_df = pd.DataFrame(meta_records)
enriched_df = enriched_df.merge(meta_df, on="frame", how="left",
validate="many_to_one")
enriched_df["cellID"] = enriched_df["track_id"]
n_tracks = (
enriched_df[["plateID", "wellID", "fieldID", "cellID"]]
.drop_duplicates()
.shape[0]
)
print(
f"[_process_merged_group] Group {key}: enriched_df rows={len(enriched_df)}, "
f"unique_tracks={n_tracks}"
)
if bleach_method != 'none':
import json
raw = _motility_raw_frame(enriched_df, bleach_added)
enriched_df = enriched_df.drop(columns=[c for c in enriched_df if c.startswith('raw_')])
enriched_df['bleach_correction_method'] = bleach_method
enriched_df['bleach_background_policy'] = '5px_role_label_zero_median'
fits = pd.concat(bleach_fits, ignore_index=True) if bleach_fits else pd.DataFrame()
if fits.empty:
fits = pd.DataFrame(columns=['object_type', 'channel', 'method'])
fits['source_files'] = json.dumps(sorted_basenames)
fits['background_policy'] = '5px_role_label_zero_median'
fits['time_unit'] = 'source_frame'
return enriched_df, raw, fits
return enriched_df
def _smooth_tracks_and_features(df, max_displacement=50.0, track_outlier_zscore=3.0):
"""
Smooth cell tracks and a small set of scalar features.
- Fixes single-frame "teleport" glitches in centroid position.
- Optionally drops tracks with impossible jumps.
- Smooths a subset of scalar cell_* features using a z-score heuristic.
"""
import numpy as np
if df.empty:
print("[_smooth_tracks_and_features] Input DataFrame is empty.")
return df
n_rows_before = len(df)
n_tracks_before = df[["plateID", "wellID", "fieldID", "cellID"]].drop_duplicates().shape[0]
df = df.sort_values(
["plateID", "wellID", "fieldID", "cellID", "frame"]
).reset_index(drop=True)
y_col = "cell_centroid-0"
x_col = "cell_centroid-1"
if y_col not in df.columns or x_col not in df.columns:
print("[_smooth_tracks_and_features] Centroid columns missing, nothing to smooth.")
return df
drop_indices = set()
updates = {}
candidate_cols = [
"cell_area",
"cell_bbox_area",
"cell_equivalent_diameter",
"cell_perimeter",
"cell_perimeter_crofton",
"cell_solidity",
"cell_mean_intensity",
"cell_max_intensity",
"cell_min_intensity",
]
cell_feature_cols = [c for c in candidate_cols if c in df.columns]
for col in [y_col, x_col] + cell_feature_cols:
if col in df.columns and not np.issubdtype(df[col].dtype, np.floating):
df[col] = df[col].astype(float)
grouped = df.groupby(["plateID", "wellID", "fieldID", "cellID"], sort=False)
n_tracks_processed = 0
n_tracks_dropped = 0
n_glitches_fixed = 0
for (plateID, wellID, fieldID, cellID), g in grouped:
idx = g.index.to_numpy()
if len(idx) < 2:
continue
n_tracks_processed += 1
y = g[y_col].to_numpy(dtype=float, copy=True)
x = g[x_col].to_numpy(dtype=float, copy=True)
n = len(idx)
glitch_frames = set()
if n >= 3:
for i_local in range(1, n - 1):
y_prev, y_curr, y_next = y[i_local - 1], y[i_local], y[i_local + 1]
x_prev, x_curr, x_next = x[i_local - 1], x[i_local], x[i_local + 1]
d_prev = np.hypot(y_curr - y_prev, x_curr - x_prev)
d_next = np.hypot(y_curr - y_next, x_curr - x_next)
d_neigh = np.hypot(y_next - y_prev, x_next - x_prev)
if (
d_prev > max_displacement
and d_next > max_displacement
and d_neigh <= max_displacement
):
glitch_frames.add(i_local)
for i_local in glitch_frames:
n_glitches_fixed += 1
y_new = 0.5 * (y[i_local - 1] + y[i_local + 1])
x_new = 0.5 * (x[i_local - 1] + x[i_local + 1])
y[i_local] = y_new
x[i_local] = x_new
for col in cell_feature_cols:
s = g[col].to_numpy(dtype=float)
s_new = 0.5 * (s[i_local - 1] + s[i_local + 1])
updates.setdefault(col, {})[idx[i_local]] = s_new
drop_track = False
for i_local in range(1, n):
d = np.hypot(y[i_local] - y[i_local - 1], x[i_local] - x[i_local - 1])
if (
d > max_displacement
and i_local not in glitch_frames
and (i_local - 1) not in glitch_frames
):
drop_track = True
break
if drop_track:
n_tracks_dropped += 1
drop_indices.update(idx.tolist())
continue
for i_local, global_idx in enumerate(idx):
if y[i_local] != g[y_col].iloc[i_local]:
updates.setdefault(y_col, {})[global_idx] = y[i_local]
if x[i_local] != g[x_col].iloc[i_local]:
updates.setdefault(x_col, {})[global_idx] = x[i_local]
if len(idx) < 3 or not cell_feature_cols:
continue
for col in cell_feature_cols:
s = g[col].to_numpy(dtype=float)
if np.all(~np.isfinite(s)):
continue
mean = np.nanmean(s)
std = np.nanstd(s)
if not np.isfinite(std) or std == 0:
continue
z = (s - mean) / std
for i_local in range(1, n - 1):
if not np.isfinite(z[i_local]) or abs(z[i_local]) <= track_outlier_zscore:
continue
if (
abs(z[i_local - 1]) <= track_outlier_zscore / 2
and abs(z[i_local + 1]) <= track_outlier_zscore / 2
):
new_val = 0.5 * (s[i_local - 1] + s[i_local + 1])
updates.setdefault(col, {})[idx[i_local]] = new_val
for col, mapping in updates.items():
df.loc[list(mapping.keys()), col] = list(mapping.values())
if drop_indices:
df = df.drop(index=list(drop_indices)).reset_index(drop=True)
n_rows_after = len(df)
n_tracks_after = df[["plateID", "wellID", "fieldID", "cellID"]].drop_duplicates().shape[0]
print(
"[_smooth_tracks_and_features] rows_before="
f"{n_rows_before}, rows_after={n_rows_after}, "
f"tracks_before={n_tracks_before}, tracks_after={n_tracks_after}, "
f"tracks_processed={n_tracks_processed}, tracks_dropped={n_tracks_dropped}, "
f"glitches_fixed={n_glitches_fixed}"
)
return df
def _debug_plot_merged_planes(src, sample_filename, n_channels, nucleus_chan, pathogen_chan, out_dir):
"""
Debug-plot a single merged .npy file.
The plot is saved as a PDF and contains:
- one panel per raw intensity channel (normalized 2–98 percent)
- one panel per mask plane (random colormap)
- one panel showing merged intensity channels with all masks overlaid
using a random colormap with alpha=0.6.
"""
import os
import numpy as np
import matplotlib.pyplot as plt
import matplotlib as mpl
merged_path = os.path.join(src, "merged", sample_filename)
if not os.path.isfile(merged_path):
print(f"[_debug_plot_merged_planes] File not found: {merged_path}")
return
arr = np.load(merged_path)
original_shape = arr.shape
if arr.ndim == 3:
if arr.shape[-1] != n_channels and arr.shape[0] == n_channels:
planes = arr
else:
planes = np.moveaxis(arr, -1, 0)
elif arr.ndim == 4:
if arr.shape[-1] >= n_channels:
planes = np.moveaxis(arr[0], -1, 0)
else:
planes = arr.reshape(-1, arr.shape[-2], arr.shape[-1])
else:
planes = arr
reoriented_shape = planes.shape
print(
f"[_debug_plot_merged_planes] Sample '{sample_filename}': "
f"original_shape={original_shape}, reoriented_shape={reoriented_shape}"
)
if planes.ndim != 3:
print(
f"[_debug_plot_merged_planes] Expected 3D array after reorientation, "
f"got shape={planes.shape}; skipping."
)
return
n_planes = planes.shape[0]
if n_planes < n_channels:
n_channels = n_planes
intensity_planes = planes[:n_channels].astype(float)
mask_planes = planes[n_channels:]
n_masks = mask_planes.shape[0]
norm_intensity = []
for ch_idx in range(n_channels):
p = intensity_planes[ch_idx].astype(float)
lo = np.percentile(p, 2)
hi = np.percentile(p, 98)
if hi <= lo:
p_norm = np.zeros_like(p, dtype=float)
else:
p_norm = np.clip((p - lo) / (hi - lo), 0.0, 1.0)
norm_intensity.append(p_norm)
norm_intensity = np.asarray(norm_intensity)
if norm_intensity.size == 0:
print("[_debug_plot_merged_planes] No intensity channels to plot; skipping.")
return
H, W = norm_intensity[0].shape
merged_rgb = np.zeros((H, W, 3), dtype=float)
merged_rgb[..., 0] = norm_intensity[0]
if n_channels >= 2:
merged_rgb[..., 1] = norm_intensity[1]
if n_channels >= 3:
merged_rgb[..., 2] = norm_intensity[2]
combined_mask = None
if n_masks > 0:
combined_mask = np.zeros((H, W), dtype=int)
offset = 0
for m in mask_planes:
m_int = m.astype(int)
if m_int.max() <= 0:
continue
nonzero = m_int > 0
combined_mask[nonzero] = m_int[nonzero] + offset
offset += int(m_int.max())
if offset == 0:
combined_mask = None
extra = 1 if combined_mask is not None else 0
n_cols = n_channels + n_masks + extra
with figure_style(theme_target()):
fig, axes = plt.subplots(
1,
n_cols,
figsize=(3 * n_cols, 3),
dpi=150,
squeeze=False,
)
from .figures.bundle import _register_figure_data
_register_figure_data(fig, None, kind="mask", title="Merged planes")
axes = axes[0]
col_idx = 0
for ch_idx in range(n_channels):
ax = axes[col_idx]
col_idx += 1
ax.imshow(norm_intensity[ch_idx], cmap="gray")
ax.set_title(f"Ch {ch_idx} (2–98% norm)")
ax.axis("off")
for m_idx in range(n_masks):
ax = axes[col_idx]
col_idx += 1
mask_plane = mask_planes[m_idx]
try:
random_cmap = _generate_mask_random_cmap(mask_plane)
except NameError:
unique_labels = np.unique(mask_plane)
unique_labels = unique_labels[unique_labels != 0]
n_labels = len(unique_labels)
rng = np.random.default_rng(seed=42)
colors = np.ones((n_labels + 1, 4))
colors[1:, :3] = rng.random((n_labels, 3))
random_cmap = mpl.colors.ListedColormap(colors)
ax.imshow(mask_plane, cmap=random_cmap, interpolation="nearest")
ax.set_title(f"Mask {m_idx}")
ax.axis("off")
if combined_mask is not None:
ax = axes[col_idx]
try:
merged_cmap = _generate_mask_random_cmap(combined_mask)
except NameError:
unique_labels = np.unique(combined_mask)
unique_labels = unique_labels[unique_labels != 0]
n_labels = len(unique_labels)
rng = np.random.default_rng(seed=123)
colors = np.ones((n_labels + 1, 4))
colors[1:, :3] = rng.random((n_labels, 3))
merged_cmap = mpl.colors.ListedColormap(colors)
ax.imshow(merged_rgb)
ax.imshow(combined_mask, cmap=merged_cmap, alpha=0.6, interpolation="nearest")
ax.set_title("Merged channels + masks")
ax.axis("off")
fig.tight_layout()
base = os.path.splitext(sample_filename)[0]
out_path = os.path.join(out_dir, f"merged_planes_{base}.pdf")
out_path = save_figure_to_path(fig, out_path, bbox_inches="tight")
plt.close(fig)
print(
f"[_debug_plot_merged_planes] Saved merged plane debug figure to {out_path}"
)
def _infection_qc_pca_clustering(
all_df,
settings,
infection_col,
pathogen_chan,
motility_dir,
):
"""
Embedding-based infection intensity QC (PCA/UMAP/t-SNE).
Steps
-----
1. Aggregate to per-cell features (cell_* columns).
2. Select morphology + pathogen-channel intensity features.
3. Optionally transform features to improve structure:
- log1p on intensity-like features (if enabled)
- up-weight pathogen-channel features (if configured)
4. Embed to 2D using PCA / UMAP / t-SNE
(controlled by settings['infection_intensity_strategy']:
'pca', 'umap', or 'tsne').
For UMAP and t-SNE, an internal hyperparameter search is
performed (if enabled) to maximize a separation score based on:
- distance between the two clusters in the embedding
- how well clusters separate GT infected vs GT uninfected.
5. KMeans clustering (2 clusters) in the embedded space.
6. Define "ground-truth" subsets based on pathogen-channel intensity:
- uninfected_gt: lowest 25% of intensities in the UNINFECTED group
- infected_gt : highest 25% of intensities in the INFECTED group
7. Assign clusters as "infected" vs "uninfected" based on which
ground-truth class dominates each cluster.
8. Depending on infection_intensity_mode:
- 'relabel': adjusted_infected = cluster assignment.
- 'remove' : drop cells whose original label disagrees with the
cluster assignment.
9. Map adjusted_infected back to all_df (frame level).
Hyperparameter search
---------------------
UMAP:
- Enabled if settings.get('infection_pca_umap_search', True) is True.
- Grid (overrideable via settings):
settings['infection_pca_umap_n_neighbors_grid'] (default [5, 10, 15, 30])
settings['infection_pca_umap_min_dist_grid'] (default [0.0, 0.05, 0.1, 0.3])
- Score = centroid_distance * the absolute GT-infected fraction difference between clusters.
t-SNE:
- Enabled if settings.get('infection_pca_tsne_search', True) is True.
- Grid (overrideable via settings):
settings['infection_pca_tsne_perplexity_grid'] (default [15.0, 30.0, 45.0])
settings['infection_pca_tsne_learning_rate_grid'] (default [200.0, 500.0])
(perplexity is clamped to < (n_samples-1)/3.)
- Same score definition as for UMAP.
PCA “structure” helpers
-----------------------
- Optional log1p transform on intensity-like features
(if settings.get('infection_pca_log_intensity', True) is True).
- Optional up-weighting of pathogen-channel features after
standardization:
settings['infection_pca_pathogen_weight'] (default 1.0)
Side effects
------------
- settings['infection_pca_data'] with:
{
'coords': coords (n_cells x 2),
'labels': adjusted_infected (bool),
'cluster_labels': cluster_ids (0 or 1),
'method_label': 'PCA' / 'UMAP' / 't-SNE',
'infected_cluster': int,
'uninfected_cluster': int,
'initial_infected_frac_infected_cluster': float (0-1),
'initial_infected_frac_uninfected_cluster': float (0-1),
'gt_sep_score': float,
'silhouette_score': float or None,
'centroid_distance': float,
'embedding_params': dict (e.g. {'n_neighbors': 15, 'min_dist': 0.1}),
}
- settings['infection_intensity_qc_panel_type'] = 'pca'
- settings['infection_intensity_qc_panel_path'] = None
"""
import os
import pandas as pd
from sklearn.preprocessing import StandardScaler
from sklearn.cluster import KMeans
from sklearn.metrics import silhouette_score
from sklearn.decomposition import PCA
try:
from sklearn.manifold import TSNE
except Exception:
TSNE = None
try:
from .utils import umap # type: ignore
umap.UMAP # noqa: B018
except Exception:
umap = None
source_all_df = all_df
def _evaluate_embedding(coords, cluster_labels, y_orig, gt_uninf, gt_inf):
"""
Compute:
- infected/uninfected cluster mapping using GT sets
- GT separation score
- silhouette score
- centroid distance between the two clusters
- original infected fractions in each cluster
- overall score = centroid_distance * GT separation
"""
frac_inf_gt = []
for k in (0, 1):
mask_k = cluster_labels == k
n_inf_k = int(np.sum(mask_k & gt_inf))
n_uninf_k = int(np.sum(mask_k & gt_uninf))
tot_k = n_inf_k + n_uninf_k
if tot_k == 0:
frac_inf_gt.append(0.5)
else:
frac_inf_gt.append(n_inf_k / float(tot_k))
infected_cluster = int(np.argmax(frac_inf_gt))
uninfected_cluster = 1 - infected_cluster
gt_sep_score = abs(frac_inf_gt[0] - frac_inf_gt[1])
centroids = []
for k in (0, 1):
mask_k = cluster_labels == k
if mask_k.any():
centroids.append(coords[mask_k].mean(axis=0))
else:
centroids.append(np.zeros(coords.shape[1], dtype=float))
centroid_distance = float(np.linalg.norm(centroids[0] - centroids[1]))
sil = None
if coords.shape[0] > 10 and len(np.unique(cluster_labels)) > 1:
try:
sil = float(silhouette_score(coords, cluster_labels))
except Exception:
sil = None
mask_inf_cluster = cluster_labels == infected_cluster
mask_uninf_cluster = cluster_labels == uninfected_cluster
frac_inf_infected_cluster = (
float(y_orig[mask_inf_cluster].mean()) if mask_inf_cluster.any() else 0.0
)
frac_inf_uninfected_cluster = (
float(y_orig[mask_uninf_cluster].mean()) if mask_uninf_cluster.any() else 0.0
)
score = centroid_distance * gt_sep_score
return {
"score": score,
"infected_cluster": infected_cluster,
"uninfected_cluster": uninfected_cluster,
"gt_sep_score": gt_sep_score,
"silhouette_score": sil,
"centroid_distance": centroid_distance,
"frac_inf_infected_cluster": frac_inf_infected_cluster,
"frac_inf_uninfected_cluster": frac_inf_uninfected_cluster,
}
def _search_umap(X_scaled, y_orig, gt_uninf, gt_inf, settings_local):
"""Return coordinates, labels, statistics, and parameters from UMAP."""
random_state = int(settings_local.get("infection_pca_random_state", 0))
do_search = bool(settings_local.get("infection_pca_umap_search", True))
if not do_search:
n_neighbors = int(settings_local.get("infection_pca_umap_n_neighbors", 15))
min_dist = float(settings_local.get("infection_pca_umap_min_dist", 0.1))
reducer = umap.UMAP(
n_components=2,
random_state=random_state,
n_neighbors=n_neighbors,
min_dist=min_dist,
)
coords = reducer.fit_transform(X_scaled)
kmeans = KMeans(
n_clusters=2, random_state=random_state, n_init="auto"
)
cluster_labels = kmeans.fit_predict(coords)
stats = _evaluate_embedding(coords, cluster_labels, y_orig, gt_uninf, gt_inf)
return coords, cluster_labels, stats, {"n_neighbors": n_neighbors, "min_dist": min_dist}
nn_grid = settings_local.get(
"infection_pca_umap_n_neighbors_grid", [5, 10, 15, 30]
)
md_grid = settings_local.get(
"infection_pca_umap_min_dist_grid", [0.0, 0.05, 0.1, 0.3]
)
best = None
for nn in nn_grid:
for md in md_grid:
try:
reducer = umap.UMAP(
n_components=2,
random_state=random_state,
n_neighbors=int(nn),
min_dist=float(md),
)
coords = reducer.fit_transform(X_scaled)
kmeans = KMeans(
n_clusters=2, random_state=random_state, n_init="auto"
)
cluster_labels = kmeans.fit_predict(coords)
stats = _evaluate_embedding(
coords, cluster_labels, y_orig, gt_uninf, gt_inf
)
if (best is None) or (stats["score"] > best["stats"]["score"]):
best = {
"coords": coords,
"cluster_labels": cluster_labels,
"stats": stats,
"params": {"n_neighbors": int(nn), "min_dist": float(md)},
}
except Exception as e:
print(
f"[infection_intensity_qc:PCA] UMAP trial failed for "
f"n_neighbors={nn}, min_dist={md}: {e}"
)
continue
if best is None:
raise RuntimeError("UMAP hyperparameter search failed for all trials.")
return best["coords"], best["cluster_labels"], best["stats"], best["params"]
def _search_tsne(X_scaled, y_orig, gt_uninf, gt_inf, settings_local):
"""Return coordinates, labels, statistics, and parameters from t-SNE."""
random_state = int(settings_local.get("infection_pca_random_state", 0))
do_search = bool(settings_local.get("infection_pca_tsne_search", True))
n_samples = X_scaled.shape[0]
max_perp = max(5.0, (n_samples - 1) / 3.0)
def _run_tsne(perplexity, learning_rate):
"""Return coordinates, labels, and statistics for one t-SNE pair."""
tsne = TSNE(
n_components=2,
random_state=random_state,
init="pca",
learning_rate=learning_rate,
perplexity=perplexity,
)
coords_ = tsne.fit_transform(X_scaled)
kmeans_ = KMeans(
n_clusters=2, random_state=random_state, n_init="auto"
)
cluster_labels_ = kmeans_.fit_predict(coords_)
stats_ = _evaluate_embedding(
coords_, cluster_labels_, y_orig, gt_uninf, gt_inf
)
return coords_, cluster_labels_, stats_
if not do_search:
base_perp = float(settings_local.get("infection_pca_tsne_perplexity", 30.0))
perplexity = min(base_perp, max_perp)
if perplexity <= 0:
perplexity = max_perp
coords, cluster_labels, stats = _run_tsne(perplexity, learning_rate="auto")
return coords, cluster_labels, stats, {"perplexity": perplexity, "learning_rate": "auto"}
perp_grid = settings_local.get(
"infection_pca_tsne_perplexity_grid", [15.0, 30.0, 45.0]
)
lr_grid = settings_local.get(
"infection_pca_tsne_learning_rate_grid", [200.0, 500.0]
)
perp_candidates = [
float(p) for p in perp_grid if float(p) < max_perp and float(p) > 0
]
if not perp_candidates:
perp_candidates = [min(30.0, max_perp)]
best = None
for perp in perp_candidates:
for lr in lr_grid:
try:
coords, cluster_labels, stats = _run_tsne(perp, float(lr))
if (best is None) or (stats["score"] > best["stats"]["score"]):
best = {
"coords": coords,
"cluster_labels": cluster_labels,
"stats": stats,
"params": {"perplexity": float(perp), "learning_rate": float(lr)},
}
except Exception as e:
print(
f"[infection_intensity_qc:PCA] t-SNE trial failed for "
f"perplexity={perp}, learning_rate={lr}: {e}"
)
continue
if best is None:
raise RuntimeError("t-SNE hyperparameter search failed for all trials.")
return best["coords"], best["cluster_labels"], best["stats"], best["params"]
if all_df.empty:
print("[infection_intensity_qc:PCA] all_df is empty; skipping embedding QC.")
return all_df, infection_col
if infection_col not in all_df.columns:
print(
f"[infection_intensity_qc:PCA] infection_col {infection_col!r} missing; "
"skipping embedding QC."
)
return all_df, infection_col
mode = str(settings.get("infection_intensity_mode", "relabel")).lower()
if mode not in {"relabel", "remove"}:
print(
f"[infection_intensity_qc:PCA] Unsupported mode={mode!r}; "
"expected 'relabel' or 'remove'. Skipping embedding QC."
)
return all_df, infection_col
strategy = str(settings.get("infection_intensity_strategy", "pca")).lower()
if strategy in {"pca", "umap", "tsne"}:
embed_method = strategy
else:
embed_method = "pca"
settings["infection_pca_method"] = embed_method
key_cols = ["plateID", "wellID", "fieldID", "cellID"]
for col in key_cols:
if col not in all_df.columns:
raise KeyError(
f"[infection_intensity_qc:PCA] Required column {col!r} not in all_df."
)
cols_to_drop = [
c
for c in all_df.columns
if c == "adjusted_infected" or c.startswith("adjusted_infected_")
]
if cols_to_drop:
all_df = all_df.drop(columns=cols_to_drop)
pathogen_token = f"ch{pathogen_chan}".lower()
feature_candidates = [
c
for c in all_df.columns
if c.startswith("cell_")
and c != "cellID"
and (
"ch" not in c.lower()
or pathogen_token in c.lower()
)
]
if infection_col in key_cols or infection_col in feature_candidates:
role = (
"a grouping key"
if infection_col in key_cols
else "a cell_* feature"
)
print(
f"[infection_intensity_qc:PCA] infection_col {infection_col!r} is also "
f"{role}; it cannot be both the call and what the call is made from. "
"Skipping embedding QC."
)
return all_df, infection_col
candidate_frame = schema.coerce_model_feature_types(
all_df.loc[:, feature_candidates],
extra_features=feature_candidates,
)
numeric_cols = [
c
for c in feature_candidates
if pd.api.types.is_numeric_dtype(candidate_frame[c])
]
if not numeric_cols:
print(
"[infection_intensity_qc:PCA] No numeric cell_* features found; "
"skipping embedding QC."
)
return all_df, infection_col
tmp = all_df[key_cols + [infection_col]].copy()
for column in numeric_cols:
tmp[column] = candidate_frame[column]
tmp.replace([np.inf, -np.inf], np.nan, inplace=True)
group = tmp.groupby(key_cols, observed=True)
cell_level = group[numeric_cols].median(numeric_only=True).reset_index()
inf_any = group[infection_col].max().reset_index()
cell_level = cell_level.merge(inf_any, on=key_cols, how="left", suffixes=("", "_y"), validate="one_to_one")
cell_level[infection_col] = cell_level[infection_col].fillna(0).astype(bool)
intensity_col = None
if pathogen_chan is not None:
cand_int = [
f"cell_p95_intensity_ch{pathogen_chan}",
f"cell_max_intensity_ch{pathogen_chan}",
f"cell_mean_intensity_ch{pathogen_chan}",
]
for c in cand_int:
if c in cell_level.columns:
intensity_col = c
break
if intensity_col is None:
print(
"[infection_intensity_qc:PCA] No pathogen-channel cell_* intensity column "
"found; skipping embedding QC."
)
return all_df, infection_col
morph_cols = [
c
for c in numeric_cols
if c.startswith("cell_") and ("ch" not in c.lower())
]
path_cols = [
c
for c in numeric_cols
if c.startswith("cell_") and f"ch{pathogen_chan}" in c.lower()
]
feature_cols = sorted(set(morph_cols + path_cols))
if intensity_col not in feature_cols and intensity_col in cell_level.columns:
feature_cols.append(intensity_col)
clean_feature_cols = []
for c in feature_cols:
s = cell_level[c]
if s.notna().sum() < 10:
continue
if s.nunique(dropna=True) <= 1:
continue
clean_feature_cols.append(c)
feature_cols = clean_feature_cols
if not feature_cols:
print(
"[infection_intensity_qc:PCA] No usable morphology + pathogen features; "
"skipping embedding QC."
)
return all_df, infection_col
log_intensity = bool(settings.get("infection_pca_log_intensity", True))
cell_for_X = cell_level.copy()
if log_intensity:
for c in feature_cols:
cl = c.lower()
if ("intensity" in cl) or ("p75" in cl) or ("p95" in cl) or ("max" in cl):
vals = cell_for_X[c].to_numpy(dtype=float, copy=True)
finite = np.isfinite(vals)
if finite.any() and np.nanmin(vals[finite]) >= 0:
vals[finite] = np.log1p(vals[finite])
cell_for_X[c] = vals
X = cell_for_X[feature_cols].to_numpy(dtype=float)
y_orig = cell_level[infection_col].astype(bool).to_numpy()
finite_counts = np.isfinite(X).sum(axis=1)
mask_rows = finite_counts > 0
X = X[mask_rows]
cell_level = cell_level.loc[mask_rows].reset_index(drop=True)
y_orig = y_orig[mask_rows]
for j in range(X.shape[1]):
col = X[:, j]
m = np.isfinite(col)
med = np.nanmedian(col[m])
col[~m] = med
X[:, j] = col
max_cells = int(settings.get("infection_pca_max_cells", 50000))
if X.shape[0] > max_cells:
rng = np.random.default_rng(0)
idx = rng.choice(np.arange(X.shape[0]), size=max_cells, replace=False)
X = X[idx]
cell_level = cell_level.iloc[idx].reset_index(drop=True)
y_orig = y_orig[idx]
intens = cell_level[intensity_col].to_numpy(dtype=float)
mask_finite_int = np.isfinite(intens)
intens = intens[mask_finite_int]
y_int = y_orig[mask_finite_int]
if intens.size < 40 or np.sum(y_int) < 10 or np.sum(~y_int) < 10:
print(
"[infection_intensity_qc:PCA] Not enough cells with finite intensity in "
"both infected/uninfected for ground-truth definition; "
"skipping embedding QC."
)
return all_df, infection_col
inf_vals = intens[y_int]
uninf_vals = intens[~y_int]
thr_uninf = float(np.nanpercentile(uninf_vals, 25.0))
thr_inf = float(np.nanpercentile(inf_vals, 75.0))
intens_full = cell_level[intensity_col].to_numpy(dtype=float)
mask_finite_full = np.isfinite(intens_full)
gt_uninf = mask_finite_full & (~y_orig) & (intens_full <= thr_uninf)
gt_inf = mask_finite_full & (y_orig) & (intens_full >= thr_inf)
n_gt_uninf = int(gt_uninf.sum())
n_gt_inf = int(gt_inf.sum())
print(
"[infection_intensity_qc:PCA] Ground-truth sets: "
f"uninfected_gt={n_gt_uninf}, infected_gt={n_gt_inf} "
f"(thr_uninf={thr_uninf:.3f}, thr_inf={thr_inf:.3f})."
)
if n_gt_uninf < 10 or n_gt_inf < 10:
print(
"[infection_intensity_qc:PCA] Very small ground-truth subsets; "
"embedding QC may be unstable."
)
scaler = StandardScaler()
X_scaled = scaler.fit_transform(X)
path_weight = float(settings.get("infection_pca_pathogen_weight", 1.0))
if path_weight != 1.0 and path_cols:
path_idx = [feature_cols.index(c) for c in feature_cols if c in path_cols]
if path_idx:
X_scaled[:, path_idx] *= path_weight
random_state = int(settings.get("infection_pca_random_state", 0))
method_label = "PCA"
embedding_params = {}
coords = None
cluster_labels = None
eval_stats = None
if embed_method == "umap" and umap is not None:
coords, cluster_labels, eval_stats, embedding_params = _search_umap(
X_scaled, y_orig, gt_uninf, gt_inf, settings
)
method_label = "UMAP"
print(
"[infection_intensity_qc:PCA] UMAP best params: "
f"{embedding_params}, score={eval_stats['score']:.4f}"
)
elif embed_method == "tsne" and TSNE is not None:
coords, cluster_labels, eval_stats, embedding_params = _search_tsne(
X_scaled, y_orig, gt_uninf, gt_inf, settings
)
method_label = "t-SNE"
print(
"[infection_intensity_qc:PCA] t-SNE best params: "
f"{embedding_params}, score={eval_stats['score']:.4f}"
)
else:
if embed_method in {"umap", "tsne"}:
print(
f"[infection_intensity_qc:PCA] Requested method={embed_method!r} "
"not available; falling back to PCA."
)
pca = PCA(
n_components=2,
random_state=random_state,
)
coords = pca.fit_transform(X_scaled)
kmeans = KMeans(
n_clusters=2,
random_state=random_state,
n_init="auto",
)
cluster_labels = kmeans.fit_predict(coords)
eval_stats = _evaluate_embedding(coords, cluster_labels, y_orig, gt_uninf, gt_inf)
method_label = "PCA"
embedding_params = {}
infected_cluster = int(eval_stats["infected_cluster"])
uninfected_cluster = int(eval_stats["uninfected_cluster"])
gt_sep_score = float(eval_stats["gt_sep_score"])
sil_score = eval_stats["silhouette_score"]
centroid_distance = float(eval_stats["centroid_distance"])
frac_inf_infected_cluster = float(eval_stats["frac_inf_infected_cluster"])
frac_inf_uninfected_cluster = float(eval_stats["frac_inf_uninfected_cluster"])
min_gt_sep = float(settings.get("infection_pca_min_gt_separation", 0.2))
min_sil = float(settings.get("infection_pca_min_silhouette", 0.05))
if gt_sep_score < min_gt_sep or (sil_score is not None and sil_score < min_sil):
print(
"[infection_intensity_qc:PCA] WARNING: weak cluster structure "
f"(gt_sep_score={gt_sep_score:.3f}, silhouette={sil_score}). "
"To improve separation you can try:\n"
" - Tightening infection ground-truth thresholds (e.g. more extreme percentiles)\n"
" - Reducing noise features, especially non-morphology/non-pathogen\n"
" - Adjusting UMAP/t-SNE grids to favor more local structure\n"
" - Increasing infection_pca_pathogen_weight to emphasize pathogen features."
)
print(
"[infection_intensity_qc:PCA] Cluster infected fractions (original labels): "
f"infected_cluster={frac_inf_infected_cluster:.3f}, "
f"uninfected_cluster={frac_inf_uninfected_cluster:.3f}, "
f"centroid_distance={centroid_distance:.3f}, gt_sep={gt_sep_score:.3f}."
)
cluster_infected = (cluster_labels == infected_cluster)
removed_ids = set()
if mode == "relabel":
adjusted = cluster_infected.astype(bool)
n_changed = int((adjusted != y_orig).sum())
print(
"[infection_intensity_qc:PCA] Relabel mode: adjusted infection labels for "
f"{n_changed} cells based on {method_label} clusters."
)
cell_level["adjusted_infected"] = adjusted.astype(bool)
else:
consistent = cluster_infected == y_orig
to_remove = ~consistent
if to_remove.any():
removed = cell_level.loc[to_remove, key_cols]
removed_ids = {
(r["plateID"], r["wellID"], r["fieldID"], r["cellID"])
for _, r in removed.iterrows()
}
cell_level = cell_level.loc[consistent].copy()
cluster_infected = cluster_infected[consistent]
y_orig = y_orig[consistent]
coords = coords[consistent]
cluster_labels = cluster_labels[consistent]
print(
"[infection_intensity_qc:PCA] Remove mode: removed "
f"{len(removed_ids)} cells with cluster vs label disagreement."
)
cell_level["adjusted_infected"] = y_orig.astype(bool)
if all_df is source_all_df:
all_df = all_df.copy(deep=False)
for col in key_cols:
all_df[col] = all_df[col].astype(cell_level[col].dtype)
all_df = all_df.merge(
cell_level[key_cols + ["adjusted_infected"]],
on=key_cols,
how="left",
validate="m:1",
)
if removed_ids:
mask_drop = all_df.apply(
lambda r: (r["plateID"], r["wellID"], r["fieldID"], r["cellID"])
in removed_ids,
axis=1,
)
all_df = all_df.loc[~mask_drop].reset_index(drop=True)
mask_missing = all_df["adjusted_infected"].isna()
if mask_missing.any():
all_df.loc[mask_missing, "adjusted_infected"] = (
all_df.loc[mask_missing, infection_col].astype(bool)
)
all_df["adjusted_infected"] = all_df["adjusted_infected"].astype(bool)
infection_col = "adjusted_infected"
settings["infection_pca_data"] = {
"coords": coords,
"labels": cell_level["adjusted_infected"].astype(bool).to_numpy(),
"cluster_labels": cluster_labels,
"method_label": method_label,
"infected_cluster": int(infected_cluster),
"uninfected_cluster": int(uninfected_cluster),
"initial_infected_frac_infected_cluster": frac_inf_infected_cluster,
"initial_infected_frac_uninfected_cluster": frac_inf_uninfected_cluster,
"gt_sep_score": gt_sep_score,
"silhouette_score": sil_score,
"centroid_distance": centroid_distance,
"embedding_params": embedding_params,
}
settings["infection_intensity_qc_panel_type"] = "pca"
settings["infection_intensity_qc_panel_path"] = None
try:
if motility_dir is not None:
import matplotlib.pyplot as plt
os.makedirs(motility_dir, exist_ok=True)
with figure_style(theme_target()):
fig, ax = plt.subplots(figsize=(4, 4))
from .figures.bundle import _register_figure_data
_register_figure_data(fig, lambda: pd.DataFrame({"pc_1": np.asarray(coords)[:, 0], "pc_2": np.asarray(coords)[:, 1], "cluster": np.asarray(cluster_labels).astype(str)}), x="pc_1", y="pc_2", hue="cluster", kind="scatter")
mask_uninf_cluster_plot = cluster_labels == uninfected_cluster
mask_inf_cluster_plot = cluster_labels == infected_cluster
ax.scatter(
coords[mask_uninf_cluster_plot, 0],
coords[mask_uninf_cluster_plot, 1],
s=2,
alpha=0.6,
color=UNINFECTED_COLOUR,
label=(
f"Uninfected cluster "
f"({frac_inf_uninfected_cluster*100:.1f}% infected at start)"
),
)
ax.scatter(
coords[mask_inf_cluster_plot, 0],
coords[mask_inf_cluster_plot, 1],
s=2,
alpha=0.6,
color=INFECTED_COLOUR,
label=(
f"Infected cluster "
f"({frac_inf_infected_cluster*100:.1f}% infected at start)"
),
)
ax.set_xlabel(f"{method_label} 1")
ax.set_ylabel(f"{method_label} 2")
title = f"{method_label} infection QC"
if embedding_params:
param_str = ", ".join(
f"{k}={v}" for k, v in embedding_params.items()
)
title += f"\n{param_str}"
if sil_score is not None:
title += f"\nGT-sep={gt_sep_score:.2f}, sil={sil_score:.2f}"
else:
title += f"\nGT-sep={gt_sep_score:.2f}"
ax.set_title(title)
ax.legend(fontsize=7, loc="best")
out_png = os.path.join(
motility_dir, f"infection_{embed_method}_qc_embedding.png"
)
out_png = save_figure_to_path(fig, out_png, bbox_inches="tight")
plt.close(fig)
except Exception as e:
print(f"[infection_intensity_qc:PCA] Failed to save embedding QC plot: {e}")
return all_df, infection_col
def _apply_infection_intensity_qc(
all_df,
settings,
infection_col,
pathogen_chan,
motility_dir,
):
"""
Dispatch to different infection QC strategies based on
settings['infection_intensity_strategy'] and
settings['infection_intensity_qc_scope'].
Strategies
----------
'histogram' / 'hist' / 'histagram'
1D intensity histogram thresholding
'pca' / 'umap' / 'tsne'
PCA/UMAP/TSNE + clustering (via _infection_qc_pca_clustering)
'xgboost' / 'xgb'
Supervised XGBoost classifier on extreme intensities
Scope
-----
settings['infection_intensity_qc_scope']:
'combined' (default)
Run QC once on all_df (old behaviour).
'plate'
Run QC separately per plateID.
'well'
Run QC separately per (plateID, wellID).
'none' / 'off'
Skip QC entirely and return original labels.
If settings['infection_intensity_qc'] is False or pathogen_chan is None,
this function is a no-op and returns the input as-is.
Returns
-------
all_df : DataFrame
Frame-level measurements with possibly updated 'infection_col'.
infection_col : str
Name of the column in all_df that encodes the (possibly adjusted)
infection status.
"""
import os
import pandas as pd
settings["infection_hist_data"] = None
settings["infection_pca_data"] = None
settings["infection_xgb_importance"] = None
settings["infection_intensity_qc_panel_type"] = None
settings["infection_intensity_qc_panel_path"] = None
infection_intensity_qc = bool(settings.get("infection_intensity_qc", False))
if (not infection_intensity_qc) or (pathogen_chan is None):
print("[infection_intensity_qc] QC disabled or no pathogen channel; skipping.")
return all_df, infection_col
os.makedirs(motility_dir, exist_ok=True)
strategy = str(settings.get("infection_intensity_strategy", "histogram")).lower()
if strategy in {"hist", "histogram", "histagram"}:
qc_func = _infection_qc_histogram
elif strategy in {"xgboost", "xgb"}:
qc_func = _infection_qc_xgboost
elif strategy in {"pca", "umap", "tsne"}:
qc_func = _infection_qc_pca_clustering
else:
print(
"[infection_intensity_qc] Unknown strategy "
f"{strategy!r}; falling back to 'histogram'."
)
qc_func = _infection_qc_histogram
scope = str(settings.get("infection_intensity_qc_scope", "combined") or "combined").lower()
if scope in {"none", "off"}:
return all_df, infection_col
if scope in {"combined", "global", "all"}:
local_settings = dict(settings)
df_qc, inf_col_out = qc_func(
all_df=all_df,
settings=local_settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
settings["infection_hist_data"] = local_settings.get("infection_hist_data")
settings["infection_pca_data"] = local_settings.get("infection_pca_data")
settings["infection_xgb_importance"] = local_settings.get("infection_xgb_importance")
settings["infection_intensity_qc_panel_type"] = local_settings.get(
"infection_intensity_qc_panel_type"
)
settings["infection_intensity_qc_panel_path"] = local_settings.get(
"infection_intensity_qc_panel_path"
)
if "adjusted_infected" in df_qc.columns:
if df_qc["adjusted_infected"].isna().any():
df_qc["adjusted_infected"] = df_qc["adjusted_infected"].fillna(
df_qc[infection_col]
)
try:
df_qc["adjusted_infected"] = df_qc["adjusted_infected"].astype(int)
except Exception:
df_qc["adjusted_infected"] = df_qc["adjusted_infected"].astype(bool)
inf_col_out = "adjusted_infected"
return df_qc, inf_col_out
if scope in {"plate", "per_plate", "plateid"}:
group_cols = ["plateID"]
elif scope in {"well", "per_well"}:
group_cols = ["plateID", "wellID"]
else:
print(
f"[_apply_infection_intensity_qc] Unknown scope={scope!r}; "
"using 'combined' behaviour."
)
local_settings = dict(settings)
df_qc, inf_col_out = qc_func(
all_df=all_df,
settings=local_settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
settings["infection_hist_data"] = local_settings.get("infection_hist_data")
settings["infection_pca_data"] = local_settings.get("infection_pca_data")
settings["infection_xgb_importance"] = local_settings.get("infection_xgb_importance")
settings["infection_intensity_qc_panel_type"] = local_settings.get(
"infection_intensity_qc_panel_type"
)
settings["infection_intensity_qc_panel_path"] = local_settings.get(
"infection_intensity_qc_panel_path"
)
if "adjusted_infected" in df_qc.columns:
if df_qc["adjusted_infected"].isna().any():
df_qc["adjusted_infected"] = df_qc["adjusted_infected"].fillna(
df_qc[infection_col]
)
try:
df_qc["adjusted_infected"] = df_qc["adjusted_infected"].astype(int)
except Exception:
df_qc["adjusted_infected"] = df_qc["adjusted_infected"].astype(bool)
inf_col_out = "adjusted_infected"
return df_qc, inf_col_out
if not set(group_cols).issubset(all_df.columns):
print(
f"[_apply_infection_intensity_qc] Requested scope={scope!r} but "
f"missing grouping columns {group_cols}; falling back to combined QC."
)
local_settings = dict(settings)
df_qc, inf_col_out = qc_func(
all_df=all_df,
settings=local_settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
settings["infection_hist_data"] = local_settings.get("infection_hist_data")
settings["infection_pca_data"] = local_settings.get("infection_pca_data")
settings["infection_xgb_importance"] = local_settings.get("infection_xgb_importance")
settings["infection_intensity_qc_panel_type"] = local_settings.get(
"infection_intensity_qc_panel_type"
)
settings["infection_intensity_qc_panel_path"] = local_settings.get(
"infection_intensity_qc_panel_path"
)
if "adjusted_infected" in df_qc.columns:
if df_qc["adjusted_infected"].isna().any():
df_qc["adjusted_infected"] = df_qc["adjusted_infected"].fillna(
df_qc[infection_col]
)
try:
df_qc["adjusted_infected"] = df_qc["adjusted_infected"].astype(int)
except Exception:
df_qc["adjusted_infected"] = df_qc["adjusted_infected"].astype(bool)
inf_col_out = "adjusted_infected"
return df_qc, inf_col_out
parts = []
any_adjusted = False
first_payload_settings = None
for g_key, df_group in all_df.groupby(group_cols, sort=False):
if df_group.empty:
continue
local_settings = dict(settings)
df_group_qc, inf_col_group = qc_func(
all_df=df_group,
settings=local_settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
if "adjusted_infected" in df_group_qc.columns and df_group_qc["adjusted_infected"].notna().any():
any_adjusted = True
if first_payload_settings is None:
first_payload_settings = local_settings
parts.append(df_group_qc)
if not parts:
return all_df, infection_col
all_df_qc = pd.concat(parts, axis=0, ignore_index=True)
settings["infection_hist_data"] = first_payload_settings.get("infection_hist_data")
settings["infection_pca_data"] = first_payload_settings.get("infection_pca_data")
settings["infection_xgb_importance"] = first_payload_settings.get("infection_xgb_importance")
settings["infection_intensity_qc_panel_type"] = first_payload_settings.get(
"infection_intensity_qc_panel_type"
)
settings["infection_intensity_qc_panel_path"] = first_payload_settings.get(
"infection_intensity_qc_panel_path"
)
if any_adjusted and "adjusted_infected" in all_df_qc.columns:
if all_df_qc["adjusted_infected"].isna().any():
all_df_qc["adjusted_infected"] = all_df_qc["adjusted_infected"].fillna(
all_df_qc[infection_col]
)
try:
all_df_qc["adjusted_infected"] = all_df_qc["adjusted_infected"].astype(int)
except Exception:
all_df_qc["adjusted_infected"] = all_df_qc["adjusted_infected"].astype(bool)
infection_col_out = "adjusted_infected"
else:
infection_col_out = infection_col
return all_df_qc, infection_col_out
def _compute_velocities_and_well_summary(
all_df,
settings,
infection_col,
pixels_per_um,
seconds_per_frame,
):
"""
Compute per-track velocities, straightness and per-well motility summary.
Returns
-------
track_df : DataFrame
per_well_tracks : dict[(plateID, wellID) -> list of track dicts]
well_summary_df : DataFrame
vel_unit : str
"""
import numpy as np
import pandas as pd
y_col = "cell_centroid-0"
x_col = "cell_centroid-1"
track_df = pd.DataFrame()
well_summary_df = pd.DataFrame()
per_well_tracks = {}
vel_unit = "px/frame"
if y_col not in all_df.columns or x_col not in all_df.columns:
print(
"[summarise_tracks_from_merged] Centroid columns missing; "
"motility summary and plots will not be generated."
)
return track_df, per_well_tracks, well_summary_df, vel_unit
gtracks = all_df.groupby(["plateID", "wellID", "fieldID", "cellID"])
track_records = []
for (plateID, wellID, fieldID, cellID), g in gtracks:
g = g.sort_values("frame")
x_px = g[x_col].to_numpy(dtype=float)
y_px = g[y_col].to_numpy(dtype=float)
if len(x_px) < 2:
continue
dx = np.diff(x_px)
dy = np.diff(y_px)
d = np.hypot(dx, dy)
if d.size == 0 or not np.isfinite(d).any():
continue
v_px = float(np.nanmean(d))
path_length = float(np.nansum(d))
net_dx = float(x_px[-1] - x_px[0])
net_dy = float(y_px[-1] - y_px[0])
net_disp = float(np.hypot(net_dx, net_dy))
if path_length > 0 and np.isfinite(net_disp):
straightness = net_disp / path_length
else:
straightness = np.nan
infected_track = bool(g[infection_col].any())
track_records.append(
{
"plateID": plateID,
"wellID": wellID,
"fieldID": fieldID,
"cellID": cellID,
"infected": infected_track,
"v_px_per_frame": v_px,
"straightness": straightness,
}
)
key_well = (plateID, wellID)
per_well_tracks.setdefault(key_well, []).append(
{
"plateID": plateID,
"wellID": wellID,
"fieldID": fieldID,
"cellID": cellID,
"infected": infected_track,
"x_px": x_px,
"y_px": y_px,
"v_px_per_frame": v_px,
"straightness": straightness,
}
)
if not track_records:
print(
"[summarise_tracks_from_merged] No tracks with >=2 frames; "
"skipping motility summary and plots."
)
return track_df, per_well_tracks, well_summary_df, vel_unit
track_df = pd.DataFrame(track_records)
use_physical_units = (
pixels_per_um is not None and seconds_per_frame is not None
)
if use_physical_units:
pixels_per_um = float(pixels_per_um)
seconds_per_frame = float(seconds_per_frame)
factor = (1.0 / pixels_per_um) * (60.0 / seconds_per_frame)
vel_unit = "µm/min"
else:
factor = 1.0
vel_unit = "px/frame"
track_df["velocity"] = track_df["v_px_per_frame"] * factor
track_df["velocity_unit"] = vel_unit
straightness_threshold = float(
settings.get("straightness_threshold", 0.95)
)
drop_straight_tracks = bool(settings.get("drop_straight_tracks", False))
n_tracks_before = track_df.shape[0]
n_high = int((track_df["straightness"] >= straightness_threshold).sum())
print(
"[summarise_tracks_from_merged] Straightness metric: "
f"{n_high} of {n_tracks_before} tracks have straightness "
f">= {straightness_threshold:.2f} "
"(net displacement / path length)."
)
if drop_straight_tracks and n_high > 0:
drop_mask = track_df["straightness"] >= straightness_threshold
dropped = track_df.loc[
drop_mask, ["plateID", "wellID", "fieldID", "cellID"]
].copy()
drop_keys = set(
zip(
dropped["plateID"],
dropped["wellID"],
dropped["fieldID"],
dropped["cellID"],
)
)
track_df = track_df.loc[~drop_mask].reset_index(drop=True)
print(
"[summarise_tracks_from_merged] Straightness filter "
f"removed {n_high} overly straight tracks "
f"(threshold={straightness_threshold:.2f})."
)
for well_key, track_list in list(per_well_tracks.items()):
filtered_list = [
tr
for tr in track_list
if (
tr["plateID"],
tr["wellID"],
tr["fieldID"],
tr["cellID"],
)
not in drop_keys
]
if filtered_list:
per_well_tracks[well_key] = filtered_list
else:
del per_well_tracks[well_key]
if track_df.empty:
print(
"[summarise_tracks_from_merged] No tracks left after "
"straightness filtering; skipping motility summary and plots."
)
return track_df, per_well_tracks, well_summary_df, vel_unit
well_records = []
for (plateID, wellID), g in track_df.groupby(["plateID", "wellID"]):
n_tracks_well = len(g)
n_inf_well = int(g["infected"].sum())
n_uninf_well = n_tracks_well - n_inf_well
mean_all = float(g["velocity"].mean()) if n_tracks_well > 0 else np.nan
mean_inf = (
float(g.loc[g["infected"], "velocity"].mean())
if n_inf_well > 0
else np.nan
)
mean_uninf = (
float(g.loc[~g["infected"], "velocity"].mean())
if n_uninf_well > 0
else np.nan
)
well_records.append(
dict(
plateID=plateID,
wellID=wellID,
n_tracks=n_tracks_well,
n_infected_tracks=n_inf_well,
n_uninfected_tracks=n_uninf_well,
mean_velocity_all=mean_all,
mean_velocity_infected=mean_inf,
mean_velocity_uninfected=mean_uninf,
velocity_unit=vel_unit,
)
)
well_summary_df = pd.DataFrame(well_records)
print(
"[summarise_tracks_from_merged] Computed per-track velocities "
f"in units: {vel_unit}"
)
return track_df, per_well_tracks, well_summary_df, vel_unit
#: Tables in ``measurements.db`` that spaCR itself writes and reads. The
#: motility summary writes its table with ``if_exists='replace'``, so letting
#: the free-text ``db_table_name`` setting name one of these would drop it and
#: everything measured into it.
RESERVED_DB_TABLE_NAMES = (
'cell',
'cytoplasm',
'nucleus',
'object_counts',
'organelle',
'pathogen',
'png_list',
'run_status',
'settings',
)
def _validate_db_table_name(db_table_name):
"""Refuse a ``db_table_name`` that would overwrite a spaCR-owned table.
:param db_table_name: the value of ``settings['db_table_name']``.
:returns: ``db_table_name`` unchanged when it is safe to write.
:raises ValueError: when it names one of :data:`RESERVED_DB_TABLE_NAMES`.
"""
if str(db_table_name).strip().lower() in RESERVED_DB_TABLE_NAMES:
raise ValueError(
f"settings['db_table_name'] = {db_table_name!r} names a table spaCR "
"owns. The motility summary writes that table with "
"if_exists='replace', which would drop the existing "
f"{db_table_name!r} table and every measurement in it. Choose a "
"name of your own, e.g. the default "
"'timelapse_object_measurements'. Reserved names: "
+ ', '.join(RESERVED_DB_TABLE_NAMES) + '.')
return db_table_name
def _save_measurements_and_well_summary(
all_df,
well_summary_df,
src,
db_table_name,
):
"""
Save per-frame measurements and well-level motility summary to SQLite.
Returns (measurements_dir, db_path).
:raises ValueError: when ``db_table_name`` names a spaCR-owned table; see
:func:`_validate_db_table_name`.
"""
import os
import sqlite3
_validate_db_table_name(db_table_name)
measurements_dir = os.path.join(src, "measurements")
os.makedirs(measurements_dir, exist_ok=True)
db_path = os.path.join(measurements_dir, "measurements.db")
with sqlite3.connect(db_path, timeout=30) as conn:
all_df.to_sql(db_table_name, conn, if_exists="replace", index=False)
print(
f"[summarise_tracks_from_merged] Saved measurements to "
f"{db_path} (table='{db_table_name}')"
)
if not well_summary_df.empty:
well_table_name = db_table_name + "_well_motility"
well_summary_df.to_sql(
well_table_name,
conn,
if_exists="replace",
index=False,
)
print(
f"[summarise_tracks_from_merged] Saved well-level motility "
f"summary to {db_path} (table='{well_table_name}')"
)
else:
print(
"[summarise_tracks_from_merged] No well-level motility "
"summary table was created."
)
return measurements_dir, db_path
def _feature_velocity_correlations(all_df, track_df, measurements_dir):
"""
Correlate per-track velocity with median per-track features (all / infected / uninfected).
Saves CSV to measurements_dir/velocity_feature_correlations.csv
"""
import numpy as np
import os
import pandas as pd
if track_df.empty:
return
try:
group_cols = ["plateID", "wellID", "fieldID", "cellID"]
numeric_cols = all_df.select_dtypes(include=[np.number]).columns.tolist()
for col_rm in ("frame", "timeID", "cellID"):
if col_rm in numeric_cols:
numeric_cols.remove(col_rm)
if not numeric_cols:
print(
"[summarise_tracks_from_merged] No numeric feature columns "
"available for correlation analysis."
)
return
agg_features = (
all_df[group_cols + numeric_cols]
.groupby(group_cols, dropna=False)
.median()
.reset_index()
)
track_features = track_df.merge(agg_features, on=group_cols, how="left",
validate="many_to_one")
exclude_cols = set(
group_cols
+ ["infected", "v_px_per_frame", "velocity", "velocity_unit"]
)
candidate_cols = [
c
for c in track_features.columns
if c not in exclude_cols
and np.issubdtype(track_features[c].dtype, np.number)
]
if not candidate_cols:
print(
"[summarise_tracks_from_merged] No numeric feature columns "
"available for correlation analysis."
)
return
def _corr_subset(mask, label):
"""Return correlations for at least five finite-velocity rows, else ``None``."""
sub = track_features.loc[mask, candidate_cols + ["velocity"]].copy()
sub = sub[np.isfinite(sub["velocity"])]
if sub.shape[0] < 5:
print(
"[summarise_tracks_from_merged] "
f"Not enough tracks for correlation ({label})."
)
return None
corr_series = sub.corr(method="pearson")["velocity"].drop("velocity")
corr_df = (
corr_series.rename("pearson_r")
.to_frame()
.reset_index()
.rename(columns={"index": "feature"})
)
corr_df["n_tracks"] = sub.shape[0]
corr_df["group"] = label
return corr_df
mask_all = np.isfinite(track_features["velocity"])
results = []
res_all = _corr_subset(mask_all, "all")
if res_all is not None:
results.append(res_all)
mask_inf = track_features["infected"].astype(bool) & mask_all
res_inf = _corr_subset(mask_inf, "infected")
if res_inf is not None:
results.append(res_inf)
mask_uninf = (~track_features["infected"].astype(bool)) & mask_all
res_uninf = _corr_subset(mask_uninf, "uninfected")
if res_uninf is not None:
results.append(res_uninf)
if not results:
return
corr_all = pd.concat(results, ignore_index=True)
corr_all["abs_pearson_r"] = corr_all["pearson_r"].abs()
corr_all = corr_all.sort_values(
["group", "abs_pearson_r"], ascending=[True, False]
)
corr_out = os.path.join(measurements_dir, "velocity_feature_correlations.csv")
corr_all.to_csv(corr_out, index=False)
print(
"[summarise_tracks_from_merged] Saved velocity–feature "
f"correlations to {corr_out}"
)
except Exception as e:
print(
"[summarise_tracks_from_merged] Feature–velocity correlation "
f"analysis failed with error: {e}"
)
def _make_intensity_sanity_plots(all_df, infection_col, n_channels, motility_dir):
"""
Per-channel intensity sanity-check plots (infected vs uninfected).
"""
import numpy as np
import os
import matplotlib.pyplot as plt
if all_df.empty:
return
keys = ["plateID", "wellID", "fieldID", "cellID"]
os.makedirs(motility_dir, exist_ok=True)
for ch in range(n_channels):
col_int = f"cell_mean_intensity_ch{ch}"
if col_int not in all_df.columns:
continue
cell_level_int = (
all_df[keys + [col_int, infection_col]]
.groupby(keys, dropna=False)
.agg(
{
col_int: "mean",
infection_col: "max",
}
)
.reset_index()
)
cell_level_int = cell_level_int.replace([np.inf, -np.inf], np.nan)
cell_level_int = cell_level_int.dropna(subset=[col_int])
if cell_level_int.empty:
print(
f"[summarise_tracks_from_merged] No data for intensity "
f"channel {ch}, skipping sanity plot."
)
continue
mask_inf = cell_level_int[infection_col].astype(bool)
vals_inf = cell_level_int.loc[mask_inf, col_int].to_numpy()
vals_uninf = cell_level_int.loc[~mask_inf, col_int].to_numpy()
mean_inf = float(np.nanmean(vals_inf)) if vals_inf.size else np.nan
std_inf = (
float(np.nanstd(vals_inf, ddof=1)) if vals_inf.size > 1 else np.nan
)
mean_uninf = float(np.nanmean(vals_uninf)) if vals_uninf.size else np.nan
std_uninf = (
float(np.nanstd(vals_uninf, ddof=1))
if vals_uninf.size > 1
else np.nan
)
x_pos = np.arange(2)
heights = [mean_inf, mean_uninf]
errors = [std_inf, std_uninf]
with figure_style(theme_target()):
fig_ch, ax_ch = plt.subplots(figsize=(4, 4))
from .figures.bundle import _register_figure_data
_register_figure_data(fig_ch, {"infected": vals_inf, "uninfected": vals_uninf}, x="group", y="intensity", kind="bar_strip")
ax_ch.bar(
x_pos,
heights,
yerr=errors,
capsize=5,
color=[INFECTED_COLOUR, UNINFECTED_COLOUR],
alpha=0.7,
)
ax_ch.set_xticks(x_pos)
ax_ch.set_xticklabels(["Infected", "Uninfected"])
ax_ch.set_ylabel(f"Mean cell intensity (channel {ch})")
ax_ch.set_title(f"Intensity vs infection – channel {ch}")
ax_ch.set_ylim(bottom=0)
plt.tight_layout()
out_ch = os.path.join(
motility_dir, f"intensity_channel{ch}_infected_vs_uninfected.png"
)
out_ch = save_figure_to_path(fig_ch, out_ch)
plt.close(fig_ch)
print(
f"[summarise_tracks_from_merged] Saved intensity sanity plot "
f"for channel {ch} to {out_ch}"
)
def _track_frame(groups, origin=False):
"""One tidy row per track point: ``x``, ``y``, ``track`` and ``infected``.
:param groups: iterables of track dicts with ``x_px``, ``y_px`` and
``infected``.
:param origin: shift every track to start at the origin.
:returns: the frame, empty when there are no tracks.
"""
import pandas as pd
rows = []
number = 0
for tracks in groups:
for track in tracks:
x = np.asarray(track["x_px"], dtype=float)
y = np.asarray(track["y_px"], dtype=float)
if origin and len(x):
x, y = x - x[0], y - y[0]
rows.append(pd.DataFrame({"x": x, "y": y, "track": number,
"infected": bool(track["infected"])}))
number += 1
return (pd.concat(rows, ignore_index=True) if rows
else pd.DataFrame(columns=["x", "y", "track", "infected"]))
def _make_motility_plots(
track_df,
per_well_tracks,
well_summary_df,
motility_dir,
pixels_per_um,
seconds_per_frame,
vel_unit,
settings,
):
"""
Motility plots (combined + per-well) with compact text box.
Axis control via settings:
- motility_xlim / motility_ylim: applied to absolute-coordinate plots
- motility_origin_xlim / motility_origin_ylim: applied to origin plots
"""
import numpy as np
import os
import matplotlib.pyplot as plt
from matplotlib import patches
if track_df.empty or not per_well_tracks:
print(
"[summarise_tracks_from_merged] No per-track velocities available; "
"motility plots were not generated."
)
return
def _fmt_vel(val):
"""Return finite ``val`` to two decimals, or ``n/a`` when non-finite."""
return "n/a" if not np.isfinite(val) else f"{val:.2f}"
def _apply_axis_limits(ax, xlim, ylim):
"""Apply valid two-value limits to ``ax`` and return ``None``."""
if xlim is not None and len(xlim) == 2:
ax.set_xlim(float(xlim[0]), float(xlim[1]))
if ylim is not None and len(ylim) == 2:
ax.set_ylim(float(ylim[0]), float(ylim[1]))
abs_xlim = settings.get("motility_xlim", None)
abs_ylim = settings.get("motility_ylim", None)
origin_xlim = settings.get("motility_origin_xlim", None)
origin_ylim = settings.get("motility_origin_ylim", None)
if pixels_per_um is not None:
unit_line1 = f"1 µm = {float(pixels_per_um):.2f} px"
coord_label_x = "x (µm)"
coord_label_y = "y (µm)"
coord_scale = 1.0 / float(pixels_per_um)
else:
unit_line1 = "1 µm = ? px"
coord_label_x = "x (pixels)"
coord_label_y = "y (pixels)"
coord_scale = 1.0
if seconds_per_frame is not None:
unit_line2 = f"1 frame = {float(seconds_per_frame):g} s"
else:
unit_line2 = "1 frame = ? s"
box_x0 = 0.64
box_y0 = 0.69
box_width = 0.30
box_height = 0.23
text_x = box_x0 + 0.02
y_top = box_y0 + box_height - 0.03
line_spacing = 0.07
fontsize_main = 8
fontsize_units = 7
os.makedirs(motility_dir, exist_ok=True)
with figure_style(theme_target()):
fig_all, ax_all = plt.subplots(figsize=(6, 6))
from .figures.bundle import _register_figure_data
_register_figure_data(fig_all, lambda: _track_frame(per_well_tracks.values()), x="x", y="y", hue="infected", kind="scatter")
for tracks in per_well_tracks.values():
for tr in tracks:
x = tr["x_px"] * coord_scale
y = tr["y_px"] * coord_scale
infected_track = tr["infected"]
color = INFECTED_COLOUR if infected_track else UNINFECTED_COLOUR
ax_all.plot(x, y, color=color, alpha=0.2, linewidth=0.5)
ax_all.scatter(x[-1], y[-1], color=color, s=5)
vel_all = track_df["velocity"].to_numpy()
vel_inf = track_df.loc[track_df["infected"], "velocity"].to_numpy()
vel_uninf = track_df.loc[~track_df["infected"], "velocity"].to_numpy()
mean_vel_all = float(np.nanmean(vel_all)) if vel_all.size else np.nan
mean_vel_inf = float(np.nanmean(vel_inf)) if vel_inf.size else np.nan
mean_vel_uninf = float(np.nanmean(vel_uninf)) if vel_uninf.size else np.nan
print(
"[summarise_tracks_from_merged] Velocity stats "
f"({vel_unit}): all={mean_vel_all:.3f} "
f"(n={vel_all.size} tracks with >=2 frames), "
f"infected={mean_vel_inf:.3f} (n={vel_inf.size}), "
f"uninfected={mean_vel_uninf:.3f} (n={vel_uninf.size})"
)
ax_all.set_aspect("equal", "box")
ax_all.set_xlabel(coord_label_x)
ax_all.set_ylabel(coord_label_y)
_apply_axis_limits(ax_all, abs_xlim, abs_ylim)
bbox_all = patches.FancyBboxPatch(
(box_x0, box_y0),
box_width,
box_height,
transform=ax_all.transAxes,
facecolor="none",
edgecolor="none",
boxstyle="round,pad=0.02",
alpha=0.0,
)
ax_all.add_patch(bbox_all)
ax_all.text(
text_x,
y_top,
f"Infected ({_fmt_vel(mean_vel_inf)} {vel_unit})",
color=INFECTED_COLOUR,
transform=ax_all.transAxes,
fontsize=fontsize_main,
va="top",
)
ax_all.text(
text_x,
y_top - line_spacing,
f"Uninfected ({_fmt_vel(mean_vel_uninf)} {vel_unit})",
color=UNINFECTED_COLOUR,
transform=ax_all.transAxes,
fontsize=fontsize_main,
va="top",
)
ax_all.text(
text_x,
y_top - 2 * line_spacing,
unit_line1,
color=resolve_ink(theme_target()),
transform=ax_all.transAxes,
fontsize=fontsize_units,
va="top",
)
ax_all.text(
text_x,
y_top - 3 * line_spacing,
unit_line2,
color=resolve_ink(theme_target()),
transform=ax_all.transAxes,
fontsize=fontsize_units,
va="top",
)
plt.tight_layout()
out_png_all = os.path.join(motility_dir, "motility_all_tracks.png")
out_png_all = save_figure_to_path(fig_all, out_png_all)
plt.close(fig_all)
print(
f"[summarise_tracks_from_merged] Saved combined motility plot to "
f"{out_png_all}"
)
well_summary_map = {}
if not well_summary_df.empty:
for _, row in well_summary_df.iterrows():
well_summary_map[(row["plateID"], row["wellID"])] = row
for (plateID, wellID), tracks in per_well_tracks.items():
with figure_style(theme_target()):
fig_w, ax_w = plt.subplots(figsize=(6, 6))
from .figures.bundle import _register_figure_data
_register_figure_data(fig_w, lambda: _track_frame([tracks]), x="x", y="y", hue="infected", kind="scatter", well=f"{plateID}_{wellID}")
has_infected = False
has_uninfected = False
for tr in tracks:
x = tr["x_px"] * coord_scale
y = tr["y_px"] * coord_scale
infected_track = tr["infected"]
color = INFECTED_COLOUR if infected_track else UNINFECTED_COLOUR
if infected_track:
has_infected = True
else:
has_uninfected = True
ax_w.plot(x, y, color=color, alpha=0.2, linewidth=0.5)
ax_w.scatter(x[-1], y[-1], color=color, s=5)
ax_w.set_aspect("equal", "box")
ax_w.set_xlabel(coord_label_x)
ax_w.set_ylabel(coord_label_y)
_apply_axis_limits(ax_w, abs_xlim, abs_ylim)
mean_inf_w = np.nan
mean_uninf_w = np.nan
summary_row = well_summary_map.get((plateID, wellID))
if summary_row is not None:
mean_inf_w = summary_row["mean_velocity_infected"]
mean_uninf_w = summary_row["mean_velocity_uninfected"]
bbox_w = patches.FancyBboxPatch(
(box_x0, box_y0),
box_width,
box_height,
transform=ax_w.transAxes,
facecolor="none",
edgecolor="none",
boxstyle="round,pad=0.02",
alpha=0.0,
)
ax_w.add_patch(bbox_w)
ax_w.text(
text_x,
y_top,
f"Infected ({_fmt_vel(mean_inf_w)} {vel_unit})",
color=INFECTED_COLOUR,
transform=ax_w.transAxes,
fontsize=fontsize_main,
va="top",
)
ax_w.text(
text_x,
y_top - line_spacing,
f"Uninfected ({_fmt_vel(mean_uninf_w)} {vel_unit})",
color=UNINFECTED_COLOUR,
transform=ax_w.transAxes,
fontsize=fontsize_main,
va="top",
)
ax_w.text(
text_x,
y_top - 2 * line_spacing,
unit_line1,
color=resolve_ink(theme_target()),
transform=ax_w.transAxes,
fontsize=fontsize_units,
va="top",
)
ax_w.text(
text_x,
y_top - 3 * line_spacing,
unit_line2,
color=resolve_ink(theme_target()),
transform=ax_w.transAxes,
fontsize=fontsize_units,
va="top",
)
plt.tight_layout()
out_well = os.path.join(
motility_dir, f"motility_{plateID}_{wellID}_all_tracks.png"
)
out_well = save_figure_to_path(fig_w, out_well)
plt.close(fig_w)
print(
f"[summarise_tracks_from_merged] Saved per-well motility plot "
f"to {out_well}"
)
if has_infected:
with figure_style(theme_target()):
fig_inf, ax_inf = plt.subplots(figsize=(6, 6))
from .figures.bundle import _register_figure_data
_register_figure_data(fig_inf, lambda: _track_frame([[tr for tr in tracks if tr["infected"]]], origin=True), x="x", y="y", kind="scatter")
for tr in tracks:
if not tr["infected"]:
continue
x = (tr["x_px"] - tr["x_px"][0]) * coord_scale
y = (tr["y_px"] - tr["y_px"][0]) * coord_scale
ax_inf.plot(x, y, color=INFECTED_COLOUR, alpha=0.2, linewidth=0.5)
ax_inf.scatter(x[-1], y[-1], color=INFECTED_COLOUR, s=5)
ax_inf.set_aspect("equal", "box")
ax_inf.set_xlabel(coord_label_x)
ax_inf.set_ylabel(coord_label_y)
_apply_axis_limits(ax_inf, origin_xlim, origin_ylim)
plt.tight_layout()
out_inf = os.path.join(
motility_dir, f"motility_{plateID}_{wellID}_infected_origin.png"
)
out_inf = save_figure_to_path(fig_inf, out_inf)
plt.close(fig_inf)
print(
f"[summarise_tracks_from_merged] Saved per-well infected "
f"origin plot to {out_inf}"
)
if has_uninfected:
with figure_style(theme_target()):
fig_uninf, ax_uninf = plt.subplots(figsize=(6, 6))
from .figures.bundle import _register_figure_data
_register_figure_data(fig_uninf, lambda: _track_frame([[tr for tr in tracks if not tr["infected"]]], origin=True), x="x", y="y", kind="scatter")
for tr in tracks:
if tr["infected"]:
continue
x = (tr["x_px"] - tr["x_px"][0]) * coord_scale
y = (tr["y_px"] - tr["y_px"][0]) * coord_scale
ax_uninf.plot(x, y, color=UNINFECTED_COLOUR, alpha=0.2, linewidth=0.5)
ax_uninf.scatter(x[-1], y[-1], color=UNINFECTED_COLOUR, s=5)
ax_uninf.set_aspect("equal", "box")
ax_uninf.set_xlabel(coord_label_x)
ax_uninf.set_ylabel(coord_label_y)
_apply_axis_limits(ax_uninf, origin_xlim, origin_ylim)
plt.tight_layout()
out_uninf = os.path.join(
motility_dir, f"motility_{plateID}_{wellID}_uninfected_origin.png"
)
out_uninf = save_figure_to_path(fig_uninf, out_uninf)
plt.close(fig_uninf)
print(
f"[summarise_tracks_from_merged] Saved per-well uninfected "
f"origin plot to {out_uninf}"
)
def _select_infection_feature_columns(all_df, pathogen_chan):
"""
Select numeric feature columns for infection QC:
- numeric columns
- drop obvious IDs / motility metrics
- drop centroid (coordinate) features
- keep intensity features only for the pathogen channel
- drop near-constant or almost-empty columns at cell level
"""
import numpy as np
numeric_cols = schema.model_feature_columns(
all_df,
allow_unknown=True,
)
exclude = {
"frame",
"timeID",
"cellID",
"n_pathogens",
"v_px_per_frame",
"velocity",
"straightness",
}
exclude |= {c for c in numeric_cols if c.endswith("_idx")}
exclude |= {c for c in numeric_cols if "centroid" in c.lower()}
if pathogen_chan is not None:
for c in numeric_cols:
if "intensity_ch" in c:
try:
digits = "".join(ch for ch in c.split("ch")[-1] if ch.isdigit())
if digits != "":
ch_idx = int(digits)
if ch_idx != pathogen_chan:
exclude.add(c)
except Exception:
pass
feature_cols = [c for c in numeric_cols if c not in exclude]
if not feature_cols:
return []
key_cols = ["plateID", "wellID", "fieldID", "cellID"]
agg_cols = list(feature_cols)
cell_level = (
all_df[key_cols + agg_cols]
.groupby(key_cols, dropna=False)
.median()
.reset_index()
)
filtered = []
for c in agg_cols:
arr = cell_level[c].to_numpy(dtype=float)
finite = np.isfinite(arr)
if finite.sum() < 10:
continue
if np.nanstd(arr[finite]) < 1e-6:
continue
filtered.append(c)
return filtered
def _compute_intensity_percentiles_per_channel(
mask_stack,
intensity_stack,
channel_index,
object_prefix,
percentiles=(1, 5, 10, 25, 75, 95, 99),
label_as_track_id=False,
):
"""
Compute per-frame, per-object intensity percentiles for a given channel.
Parameters
----------
mask_stack : ndarray
Label image stack of shape (T, Y, X).
intensity_stack : ndarray
Intensity stack of shape (T, Y, X, C).
channel_index : int
Channel index in intensity_stack.
object_prefix : str
Prefix for column names ("cell", "nucleus", "pathogen", "cytoplasm").
percentiles : tuple of int
Percentiles to compute (0–100).
label_as_track_id : bool
If True, rename 'label' -> 'track_id'; otherwise
'label' -> f"{object_prefix}_label".
Returns
-------
DataFrame
Columns: ['frame', label_col, f'{object_prefix}_pXX_intensity_ch{channel_index}', ...]
"""
import numpy as np
import pandas as pd
if intensity_stack is None:
return pd.DataFrame(
columns=["frame", "track_id" if label_as_track_id else f"{object_prefix}_label"]
)
if channel_index is None or channel_index < 0 or channel_index >= intensity_stack.shape[-1]:
return pd.DataFrame(
columns=["frame", "track_id" if label_as_track_id else f"{object_prefix}_label"]
)
T = mask_stack.shape[0]
dfs = []
label_col_name = "track_id" if label_as_track_id else f"{object_prefix}_label"
perc = np.array(percentiles, dtype=float)
for frame in range(T):
labels = mask_stack[frame]
if not np.any(labels):
continue
intensity_image = intensity_stack[frame, :, :, channel_index]
obj_labels = np.unique(labels)
obj_labels = obj_labels[obj_labels > 0]
if obj_labels.size == 0:
continue
records = []
for lab in obj_labels:
mask = labels == lab
vals = intensity_image[mask]
vals = vals[np.isfinite(vals)]
if vals.size == 0:
continue
pvals = np.percentile(vals, perc)
rec = {"frame": frame, label_col_name: int(lab)}
for p, v in zip(perc, pvals):
col_name = f"{object_prefix}_p{int(p):02d}_intensity_ch{channel_index}"
rec[col_name] = float(v)
records.append(rec)
if records:
dfs.append(pd.DataFrame.from_records(records))
if not dfs:
return pd.DataFrame(columns=["frame", label_col_name])
out_df = pd.concat(dfs, ignore_index=True)
return out_df
def _make_adjusted_qc_panel(
all_df,
infection_col,
motility_dir,
settings,
label_tag,
):
"""
Build a QC results panel for adjusted labels using 3 subplots:
- top-left: PCA (adjusted_infected)
- top-right: XGBoost feature importance
- bottom: pathogen-channel intensity histogram
Uses payloads stored in `settings` by the QC functions:
settings["infection_hist_data"]
settings["infection_pca_data"]
settings["infection_xgb_importance"]
"""
import os
import numpy as np
import matplotlib.pyplot as plt
os.makedirs(motility_dir, exist_ok=True)
meta_tag = _infer_plate_well_meta_tag(all_df)
fig, ax_pca, ax_xgb, ax_hist = create_results_figure()
from .figures.bundle import _register_figure_data
_register_figure_data(fig, lambda: {"infected": (settings.get("infection_hist_data") or {}).get("intensities_inf", []), "uninfected": (settings.get("infection_hist_data") or {}).get("intensities_uninf", [])}, x="group", y="intensity", kind="hist")
hist_data = settings.get("infection_hist_data") or {}
vals_inf = np.asarray(hist_data.get("intensities_inf", []), dtype=float)
vals_uninf = np.asarray(hist_data.get("intensities_uninf", []), dtype=float)
bin_edges = np.asarray(hist_data.get("bin_edges", []), dtype=float)
thr_val = hist_data.get("thr_val", None)
pathogen_chan = hist_data.get("pathogen_chan", None)
do_log = bool(hist_data.get("log_transform", False))
if vals_inf.size + vals_uninf.size > 0 and bin_edges.size > 0:
ax_hist.hist(
vals_uninf,
bins=bin_edges,
alpha=0.5,
color=UNINFECTED_COLOUR,
label="Uninfected",
)
ax_hist.hist(
vals_inf,
bins=bin_edges,
alpha=0.5,
color=INFECTED_COLOUR,
label="Infected",
)
if thr_val is not None:
ax_hist.axvline(
thr_val,
linestyle=(0, (4, 3)),
linewidth=0.6,
color=ROLES["reference"],
label=f"thr={thr_val:.2f}",
)
if pathogen_chan is not None:
if do_log:
ax_hist.set_xlabel(f"log10 intensity (channel {pathogen_chan})")
else:
ax_hist.set_xlabel(f"Intensity (channel {pathogen_chan})")
else:
ax_hist.set_xlabel("Intensity")
ax_hist.set_ylabel("Cell count")
ax_hist.set_title("Pathogen-channel intensity histogram")
ax_hist.legend(loc="best", frameon=False)
else:
ax_hist.text(
0.5,
0.5,
"No histogram data",
ha="center",
va="center",
transform=ax_hist.transAxes,
)
ax_hist.axis("off")
pca_data = settings.get("infection_pca_data") or {}
coords = pca_data.get("coords", None)
labels = pca_data.get("labels", None)
method_label = pca_data.get("method_label", "PCA")
if coords is not None and labels is not None:
coords = np.asarray(coords, dtype=float)
labels = np.asarray(labels, dtype=bool)
if coords.ndim == 2 and coords.shape[0] == labels.shape[0] and coords.shape[1] >= 2:
x = coords[:, 0]
y = coords[:, 1]
ax_pca.scatter(
x[~labels],
y[~labels],
s=8,
c=UNINFECTED_COLOUR,
alpha=0.5,
label="Uninfected",
)
ax_pca.scatter(
x[labels],
y[labels],
s=8,
c=INFECTED_COLOUR,
alpha=0.5,
label="Infected",
)
ax_pca.set_xlabel("component 1")
ax_pca.set_ylabel("component 2")
ax_pca.set_title(f"{method_label} embedding")
ax_pca.legend(loc="best", frameon=False)
else:
ax_pca.text(
0.5,
0.5,
"No PCA data",
ha="center",
va="center",
transform=ax_pca.transAxes,
)
ax_pca.axis("off")
else:
ax_pca.text(
0.5,
0.5,
"No PCA data",
ha="center",
va="center",
transform=ax_pca.transAxes,
)
ax_pca.axis("off")
xgb_data = settings.get("infection_xgb_importance") or {}
feat_names = xgb_data.get("feature_names") or []
feat_vals = xgb_data.get("feature_importances") or []
if feat_names and feat_vals and len(feat_names) == len(feat_vals):
feat_names = list(feat_names)
feat_vals = np.asarray(feat_vals, dtype=float)
y_pos = np.arange(len(feat_names))
ax_xgb.barh(y_pos, feat_vals)
ax_xgb.set_yticks(y_pos)
ax_xgb.set_yticklabels(feat_names)
ax_xgb.invert_yaxis()
ax_xgb.set_xlabel("Importance (gain)")
ax_xgb.set_title("XGBoost feature importance")
else:
ax_xgb.text(
0.5,
0.5,
"No XGBoost importance data",
ha="center",
va="center",
transform=ax_xgb.transAxes,
)
ax_xgb.axis("off")
fig.suptitle(
f"Infection QC panel – {label_tag} labels\n{meta_tag}",
fontsize=10,
)
fig.tight_layout(rect=[0, 0, 1, 0.94])
out_name = f"infection_qc_panel_{label_tag}_{meta_tag}.png"
out_path = os.path.join(motility_dir, out_name)
out_path = save_figure_to_path(fig, out_path)
plt.close(fig)
print(
f"[summarise_tracks_from_merged] Saved infection QC results panel "
f"({label_tag}) to {out_path}"
)
def _load_measurements_from_db(db_path, db_table_name):
"""
Load per-cell measurements from an existing SQLite database.
Parameters
----------
db_path : str
Path to the SQLite database file (measurements.db).
db_table_name : str
Name of the table that stores per-cell measurements.
Returns
-------
pandas.DataFrame
DataFrame with the measurements, or an empty DataFrame if the
database/table is missing or unreadable.
"""
import os
import sqlite3
import pandas as pd
if not os.path.isfile(db_path):
return pd.DataFrame()
conn = sqlite3.connect(db_path, timeout=30)
try:
query = f"SELECT * FROM {db_table_name}"
df = pd.read_sql_query(query, conn)
except Exception as e:
print(
"[summarise_tracks_from_merged] Could not load existing measurements "
f"from {db_path} (table='{db_table_name}'): {e}"
)
df = pd.DataFrame()
finally:
conn.close()
return df
def _infection_qc_histogram(
all_df,
settings,
infection_col,
pathogen_chan,
motility_dir,
):
"""Refine mask infection labels from pathogen-channel intensity.
The mutable ``settings`` mapping receives the histogram payload, panel
type, and optional panel path. The returned pair is the updated frame
and infection-column name; insufficient data preserves the input pair.
"""
import matplotlib.pyplot as plt
import os
import numpy as np
cand_cols = [
f"cell_p95_intensity_ch{pathogen_chan}",
f"cell_mean_intensity_ch{pathogen_chan}",
]
intensity_col = None
for c in cand_cols:
if c in all_df.columns:
intensity_col = c
break
settings["infection_hist_data"] = None
if intensity_col is None:
print(
f"[infection_intensity_qc] None of {cand_cols} found; "
f"skipping intensity-based relabelling."
)
settings["infection_intensity_qc_panel_type"] = "histogram"
settings["infection_intensity_qc_panel_path"] = None
return all_df, infection_col
cols_to_drop = [
c
for c in all_df.columns
if c == "adjusted_infected" or c.startswith("adjusted_infected_")
]
if cols_to_drop:
all_df = all_df.drop(columns=cols_to_drop)
key_cols = ["plateID", "wellID", "fieldID", "cellID"]
cell_level = (
all_df[key_cols + [intensity_col, infection_col]]
.groupby(key_cols, dropna=False)
.agg({intensity_col: "mean", infection_col: "max"})
.reset_index()
)
cell_level = cell_level.replace([np.inf, -np.inf], np.nan)
cell_level = cell_level.dropna(subset=[intensity_col])
if len(cell_level) < 20 or cell_level[intensity_col].nunique() < 2:
print(
"[infection_intensity_qc] Too few cells or no intensity variation; "
"skipping intensity-based relabelling."
)
settings["infection_intensity_qc_panel_type"] = "histogram"
settings["infection_intensity_qc_panel_path"] = None
return all_df, infection_col
intensities = cell_level[intensity_col].to_numpy(dtype=float)
mask_labels = cell_level[infection_col].to_numpy(dtype=bool)
do_log = bool(settings.get("infection_intensity_log", False))
if do_log:
eps = np.nanmax([np.nanmin(intensities[intensities > 0]) * 0.5, 1e-6])
intensities = np.log10(intensities + eps)
n_bins = int(settings.get("infection_intensity_n_bins", 64))
n_bins = max(10, min(n_bins, 256))
counts_all, bin_edges = np.histogram(intensities, bins=n_bins)
counts_inf, _ = np.histogram(intensities[mask_labels], bins=bin_edges)
denom = np.maximum(counts_all, 1)
frac_inf = counts_inf.astype(float) / denom.astype(float)
target_frac = float(settings.get("infection_intensity_frac_infected", 0.7))
target_frac = max(0.5, min(target_frac, 0.95))
hist_pct = float(settings.get("infection_hist_percentile", 25.0))
hist_pct = max(0.0, min(hist_pct, 100.0))
thresh_idx = None
for i, frac in enumerate(frac_inf):
if frac >= target_frac:
thresh_idx = i
break
if thresh_idx is None:
thr_val = float(np.nanpercentile(intensities, hist_pct))
print(
"[infection_intensity_qc] Could not find bin with infected ≥ "
f"{target_frac:.2f}; using {hist_pct:.1f}th percentile of all cells "
f"({thr_val:.2f}) as threshold."
)
else:
thr_val = float(bin_edges[thresh_idx])
print(
"[infection_intensity_qc] Automatic intensity threshold at first bin "
f"where infected ≥ {target_frac:.2f}: {thr_val:.2f} (bin {thresh_idx})"
)
cell_level["intensity_positive"] = intensities >= thr_val
mode = str(settings.get("infection_intensity_mode", "relabel")).lower()
if mode not in {"relabel", "remove"}:
mode = "relabel"
removed_ids = None
if mode == "relabel":
cell_level["adjusted_infected"] = cell_level["intensity_positive"].astype(bool)
n_changed = int(
(
cell_level["adjusted_infected"]
!= cell_level[infection_col].astype(bool)
).sum()
)
print(
"[infection_intensity_qc] Adjusted infection labels for "
f"{n_changed} cells (mode=relabel)."
)
else:
consistent = (
cell_level[infection_col].astype(bool)
== cell_level["intensity_positive"].astype(bool)
)
removed = cell_level.loc[
~consistent, ["plateID", "wellID", "fieldID", "cellID"]
]
removed_ids = set(
zip(
removed["plateID"],
removed["wellID"],
removed["fieldID"],
removed["cellID"],
)
)
cell_level = cell_level.loc[consistent].copy()
cell_level["adjusted_infected"] = cell_level["intensity_positive"].astype(bool)
print(
"[infection_intensity_qc] Removed "
f"{len(removed_ids)} cells with conflicting mask vs intensity labels "
"(mode=remove)."
)
all_df = all_df.merge(
cell_level[key_cols + ["adjusted_infected"]],
on=key_cols,
how="left",
validate="many_to_one",
)
if removed_ids:
mask_keep = ~all_df.apply(
lambda r: (r["plateID"], r["wellID"], r["fieldID"], r["cellID"])
in removed_ids,
axis=1,
)
all_df = all_df.loc[mask_keep].reset_index(drop=True)
adjusted = all_df["adjusted_infected"].astype("boolean")
fallback = all_df[infection_col].astype("boolean")
all_df["adjusted_infected"] = adjusted.fillna(fallback).fillna(False).astype(bool)
infection_col = "adjusted_infected"
make_graphs = bool(settings.get("infection_intensity_qc_graphs", True))
meta_tag = _infer_plate_well_meta_tag(all_df)
hist_path = None
vals_inf = intensities[mask_labels]
vals_uninf = intensities[~mask_labels]
hist_payload = {
"intensities_inf": vals_inf,
"intensities_uninf": vals_uninf,
"bin_edges": bin_edges,
"thr_val": thr_val,
"pathogen_chan": pathogen_chan,
"log_transform": do_log,
"intensity_col": intensity_col,
}
settings["infection_hist_data"] = hist_payload
if make_graphs:
os.makedirs(motility_dir, exist_ok=True)
with figure_style(theme_target()):
fig_h, ax_h = plt.subplots(figsize=(6, 4))
from .figures.bundle import _register_figure_data
_register_figure_data(fig_h, {"infected": vals_inf, "uninfected": vals_uninf}, x="group", y="intensity", kind="hist")
ax_h.hist(
vals_uninf,
bins=bin_edges,
alpha=0.5,
color=UNINFECTED_COLOUR,
label="Uninfected (mask-based)",
)
ax_h.hist(
vals_inf,
bins=bin_edges,
alpha=0.5,
color=INFECTED_COLOUR,
label="Infected (mask-based)",
)
ax_h.axvline(
thr_val,
linestyle=(0, (4, 3)),
linewidth=0.6,
color=ROLES["reference"],
label=f"Threshold = {thr_val:.1f}",
)
if do_log:
ax_h.set_xlabel(f"log10 intensity metric (channel {pathogen_chan})")
else:
ax_h.set_xlabel(f"Intensity metric (channel {pathogen_chan})")
ax_h.set_ylabel("Cell count")
ax_h.set_title(
f"Pathogen-channel intensity histogram (thr={thr_val:.1f}, "
f"{hist_pct:.1f}th pct fallback)"
)
ax_h.legend(loc="best", frameon=False)
hist_filename = f"infection_intensity_histogram_{meta_tag}.png"
hist_path = os.path.join(motility_dir, hist_filename)
fig_h.tight_layout()
hist_path = save_figure_to_path(fig_h, hist_path, fmt="png")
plt.close(fig_h)
print(f"[infection_intensity_qc] Saved histogram to: {hist_path}")
else:
print(
"[infection_intensity_qc] infection_intensity_qc_graphs=False; "
"skipping histogram plot."
)
settings["infection_intensity_qc_panel_type"] = "histogram"
settings["infection_intensity_qc_panel_path"] = hist_path
return all_df, infection_col
@single_threaded_openmp('XGBoost infection QC')
def _infection_qc_xgboost(all_df, settings, infection_col, pathogen_chan, motility_dir):
"""
Use an XGBoost classifier to refine infection calling based on per-object features.
Key behaviour
-------------
- Training:
* per-cell (or per-object) medians across frames
* {tracked_object}_* morphology + {tracked_object}_pathogen-channel intensity features
* training labels from pathogen-channel intensity extremes:
- bottom 25% of UNINFECTED → strong negatives
- top 25% of INFECTED → strong positives
* training data curated per well:
- wells with both classes in the extreme set:
- if both classes have >= infection_xgb_min_cells_per_class examples:
→ balanced sampling per class within that well
- else (small wells with both classes):
→ keep all extreme examples from that well
- wells with only one class in the extreme set are skipped
- if no well has both classes in the extreme set → skip XGBoost QC
- Prediction:
* get P(infected) for all cells
* "mode":
- 'relabel': start from original labels, override only when model
is confident; ambiguous are still dropped if requested.
- 'remove' : drop strong label–model disagreements AND ambiguous if
requested.
- Ambiguous band removal:
* if infection_xgb_drop_ambiguous is True (default),
drop cells with proba in [infection_xgb_ambiguous_low,
infection_xgb_ambiguous_high]
(defaults: 0.25 and 0.75).
Additionally, this function stores three QC payloads in `settings`:
settings["infection_hist_data"] : dict for histogram panel
settings["infection_pca_data"] : dict for PCA panel
settings["infection_xgb_importance"] : dict for feature-importance panel
Returns
-------
all_df, infection_col='adjusted_infected'
"""
import re
import numpy as np
import pandas as pd
try:
import xgboost as xgb
except ImportError:
print("[_infection_qc_xgboost] XGBoost not installed; using histogram QC.")
return _infection_qc_histogram(
all_df=all_df,
settings=settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
settings["infection_hist_data"] = None
settings["infection_pca_data"] = None
settings["infection_xgb_importance"] = None
source_all_df = all_df
cols_to_drop = [
c
for c in all_df.columns
if c == "adjusted_infected"
or c.startswith("adjusted_infected_")
or c == "infection_prob"
or c.startswith("infection_prob_")
]
if cols_to_drop:
all_df = all_df.drop(columns=cols_to_drop)
if all_df is source_all_df:
all_df = all_df.copy(deep=False)
orig_infection_col = infection_col
if pathogen_chan is None:
print("[_infection_qc_xgboost] pathogen_chan is None; using histogram QC.")
return _infection_qc_histogram(
all_df=all_df,
settings=settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
if "n_pathogens" in all_df.columns:
all_df["n_pathogens"] = all_df["n_pathogens"].fillna(0)
if orig_infection_col not in all_df.columns:
infect_like = [c for c in all_df.columns if "infect" in c.lower()]
if infect_like:
new_col = infect_like[0]
print(
f"[_infection_qc_xgboost] Column {orig_infection_col!r} not found; "
f"using {new_col!r} instead."
)
orig_infection_col = new_col
elif "n_pathogens" in all_df.columns:
orig_infection_col = "_infected_from_n_pathogens"
all_df[orig_infection_col] = (all_df["n_pathogens"] > 0).astype(int)
print(
"[_infection_qc_xgboost] Column 'infected' not found; created "
f"{orig_infection_col!r} from 'n_pathogens > 0'."
)
else:
print(
"[_infection_qc_xgboost] No infection label and no 'n_pathogens'; "
"using histogram QC instead."
)
return _infection_qc_histogram(
all_df=all_df,
settings=settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
key_cols = ["plateID", "wellID", "fieldID", "cellID"]
for col in key_cols:
if col not in all_df.columns:
raise KeyError(f"[_infection_qc_xgboost] Required column {col!r} not in all_df.")
tracked_object = str(settings.get("tracked_object", "cell")).strip().lower()
if tracked_object not in {"cell", "nucleus", "pathogen"}:
print(
f"[_infection_qc_xgboost] Unknown tracked_object={tracked_object!r}; "
"falling back to 'cell'."
)
tracked_object = "cell"
obj_prefix = f"{tracked_object}_"
pattern_obj = re.compile(rf"^{re.escape(obj_prefix)}")
pattern_ch = re.compile(r"ch(\d+)\b")
feature_candidates = []
for column in all_df.columns:
if column == orig_infection_col or not pattern_obj.match(column):
continue
if "centroid" in column.lower():
continue
channel_match = pattern_ch.search(column)
if channel_match and channel_match.group(1) != str(pathogen_chan):
continue
feature_candidates.append(column)
candidate_frame = schema.coerce_model_feature_types(
all_df.loc[:, feature_candidates],
extra_features=feature_candidates,
)
aggregation_df = all_df.copy(deep=False)
for column in feature_candidates:
aggregation_df[column] = candidate_frame[column]
agg_cols = [
c
for c in aggregation_df.columns
if c not in (key_cols + ["frame", "timeID", orig_infection_col])
]
group = aggregation_df.groupby(key_cols, observed=True)
cell_level = group[agg_cols].median(numeric_only=True).reset_index()
infection_any = (
aggregation_df.groupby(key_cols, observed=True)[orig_infection_col]
.max()
.reset_index()
)
cell_level = cell_level.merge(
infection_any, on=key_cols, how="left", suffixes=("", "_y"),
validate="one_to_one",
)
cell_level[orig_infection_col] = (
cell_level[orig_infection_col]
.fillna(0)
.astype(bool)
)
if "n_pathogens" in cell_level.columns:
cell_level["n_pathogens"] = cell_level["n_pathogens"].fillna(0)
intensity_candidates = [
f"{obj_prefix}p95_intensity_ch{pathogen_chan}",
f"{obj_prefix}max_intensity_ch{pathogen_chan}",
f"{obj_prefix}mean_intensity_ch{pathogen_chan}",
]
intensity_col = None
for c in intensity_candidates:
if c in cell_level.columns:
intensity_col = c
break
if intensity_col is None:
print(
"[_infection_qc_xgboost] No pathogen-channel intensity column found for "
f"tracked_object={tracked_object!r} "
f"(tried: {intensity_candidates}); using histogram QC."
)
return _infection_qc_histogram(
all_df=all_df,
settings=settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
numeric_cols = schema.model_feature_columns(cell_level)
feature_cols = []
for c in numeric_cols:
if "centroid" in c.lower():
continue
if not pattern_obj.match(c):
continue
m = pattern_ch.search(c)
if m:
ch_idx = int(m.group(1))
if ch_idx != int(pathogen_chan):
continue
feature_cols.append(c)
else:
feature_cols.append(c)
if intensity_col in cell_level.columns and intensity_col not in feature_cols:
feature_cols.append(intensity_col)
clean_feature_cols = []
for c in feature_cols:
s = cell_level[c]
if s.notna().sum() < 10:
continue
if s.nunique(dropna=True) <= 1:
continue
clean_feature_cols.append(c)
feature_cols = clean_feature_cols
if not feature_cols:
print(
"[_infection_qc_xgboost] No usable "
f"{tracked_object}_* feature columns; using histogram QC."
)
return _infection_qc_histogram(
all_df=all_df,
settings=settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
infected_cells = cell_level[cell_level[orig_infection_col]]
uninfected_cells = cell_level[~cell_level[orig_infection_col]]
if len(infected_cells) < 10 or len(uninfected_cells) < 10:
print(
"[_infection_qc_xgboost] Too few infected or uninfected cells overall; "
"using histogram QC."
)
return _infection_qc_histogram(
all_df=all_df,
settings=settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
inf_int = infected_cells[intensity_col].to_numpy(dtype=float)
uninf_int = uninfected_cells[intensity_col].to_numpy(dtype=float)
inf_int = inf_int[np.isfinite(inf_int)]
uninf_int = uninf_int[np.isfinite(uninf_int)]
if inf_int.size == 0 or uninf_int.size == 0:
print(
"[_infection_qc_xgboost] No finite intensities for infected/uninfected; "
"using histogram QC."
)
return _infection_qc_histogram(
all_df=all_df,
settings=settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
high_thr_inf = np.nanpercentile(inf_int, 75.0)
low_thr_uninf = np.nanpercentile(uninf_int, 25.0)
hi_inf = infected_cells[infected_cells[intensity_col] >= high_thr_inf].copy()
lo_uninf = uninfected_cells[uninfected_cells[intensity_col] <= low_thr_uninf].copy()
print(
"[_infection_qc_xgboost] Extreme-intensity candidates: "
f"infected={len(hi_inf)}, uninfected={len(lo_uninf)} "
f"(tracked_object={tracked_object}, intensity_col={intensity_col}, "
f"low_thr_uninf={low_thr_uninf:.3f}, high_thr_inf={high_thr_inf:.3f})"
)
hi_inf["xgb_label"] = 1
lo_uninf["xgb_label"] = 0
train_candidates = pd.concat([hi_inf, lo_uninf], axis=0)
min_per_class = int(settings.get("infection_xgb_min_cells_per_class", 10))
rng_seed = settings.get("infection_xgb_random_state", 42)
try:
rng = np.random.default_rng(int(rng_seed))
except Exception:
rng = np.random.default_rng(42)
train_idx_list = []
y_train_list = []
wells_used = set()
wells_single_class = []
grouped = train_candidates.groupby(["plateID", "wellID"], observed=True)
for (plate_id, well_id), df_w in grouped:
pos_df = df_w[df_w["xgb_label"] == 1]
neg_df = df_w[df_w["xgb_label"] == 0]
n_pos = len(pos_df)
n_neg = len(neg_df)
if n_pos == 0 or n_neg == 0:
wells_single_class.append((plate_id, well_id))
continue
pos_idx = pos_df.index.to_numpy(copy=True)
neg_idx = neg_df.index.to_numpy(copy=True)
if n_pos >= min_per_class and n_neg >= min_per_class:
n_per_class = min(n_pos, n_neg)
if n_pos > n_per_class:
pos_sel = rng.choice(pos_idx, size=n_per_class, replace=False)
else:
pos_sel = pos_idx
if n_neg > n_per_class:
neg_sel = rng.choice(neg_idx, size=n_per_class, replace=False)
else:
neg_sel = neg_idx
else:
pos_sel = pos_idx
neg_sel = neg_idx
train_idx_list.extend(pos_sel.tolist())
y_train_list.extend([1] * pos_sel.size)
train_idx_list.extend(neg_sel.tolist())
y_train_list.extend([0] * neg_sel.size)
wells_used.add((plate_id, well_id))
if not wells_used:
print(
"[_infection_qc_xgboost] No wells with both infected and uninfected "
"extreme-intensity examples; skipping XGBoost QC."
)
return _infection_qc_histogram(
all_df=all_df,
settings=settings,
infection_col=infection_col,
pathogen_chan=pathogen_chan,
motility_dir=motility_dir,
)
train_idx = np.array(train_idx_list, dtype=int)
y_train = np.array(y_train_list, dtype=int)
n_pos_train = int((y_train == 1).sum())
n_neg_train = int((y_train == 0).sum())
print(
"[_infection_qc_xgboost] Training set (after per-well curation): "
f"wells used={len(wells_used)}, positives={n_pos_train}, "
f"negatives={n_neg_train}, min_per_class={min_per_class}."
)
if wells_single_class:
print(
"[_infection_qc_xgboost] Wells skipped due to single class in extreme set: "
+ ", ".join([f"{p}_{w}" for (p, w) in wells_single_class])
)
X_all = cell_level[feature_cols].to_numpy(dtype=float, copy=True)
for j in range(X_all.shape[1]):
col = X_all[:, j]
mask = np.isfinite(col)
if not mask.any():
X_all[:, j] = 0.0
else:
med = np.nanmedian(col[mask])
col[~mask] = med
X_all[:, j] = col
X_train = X_all[train_idx]
if X_train.shape[1] > 1:
corr = np.corrcoef(X_train, rowvar=False)
corr_thr = float(settings.get("infection_xgb_corr_threshold", 0.95))
keep = np.ones(corr.shape[0], dtype=bool)
for i in range(corr.shape[0]):
if not keep[i]:
continue
for j in range(i + 1, corr.shape[0]):
if keep[j] and abs(corr[i, j]) >= corr_thr:
keep[j] = False
removed = [f for f, k in zip(feature_cols, keep) if not k]
if removed:
print(
"[_infection_qc_xgboost] Removing highly correlated features "
f"(>|{corr_thr:.2f}|): " + ", ".join(removed)
)
feature_cols = [f for f, k in zip(feature_cols, keep) if k]
X_all = X_all[:, keep]
X_train = X_train[:, keep]
used_feature_cols = feature_cols
print(
f"[_infection_qc_xgboost] Using {len(used_feature_cols)} "
f"{tracked_object}_* features:"
)
print(" " + ", ".join(used_feature_cols))
dtrain = xgb.DMatrix(X_train, label=y_train, feature_names=used_feature_cols)
params = {
"objective": "binary:logistic",
"eval_metric": "logloss",
"max_depth": int(settings.get("infection_xgb_max_depth", 3)),
"eta": float(settings.get("infection_xgb_learning_rate", 0.1)),
"subsample": float(settings.get("infection_xgb_subsample", 0.8)),
"colsample_bytree": float(settings.get("infection_xgb_colsample_bytree", 0.8)),
"lambda": float(settings.get("infection_xgb_reg_lambda", 1.0)),
"alpha": 0.0,
"verbosity": 0,
"nthread": int(settings.get("infection_xgb_n_jobs", -1)),
}
num_round = int(settings.get("infection_xgb_n_estimators", 200))
bst = xgb.train(params, dtrain, num_boost_round=num_round)
dall = xgb.DMatrix(X_all, feature_names=used_feature_cols)
probs = bst.predict(dall)
prob_thr = float(settings.get("infection_xgb_proba_threshold", 0.5))
margin = float(settings.get("infection_xgb_margin", 0.0))
margin = max(0.0, min(margin, 0.49))
orig_arr = cell_level[orig_infection_col].astype(int).to_numpy()
pred_arr = (probs >= prob_thr).astype(int)
mode = str(settings.get("infection_intensity_mode", "relabel")).lower()
if mode not in {"relabel", "remove"}:
mode = "relabel"
removed_ids = set()
ambiguous_ids = set()
if mode == "relabel":
adjusted = orig_arr.copy()
hi_conf = probs >= (prob_thr + margin)
lo_conf = probs <= (prob_thr - margin)
adjusted[hi_conf] = 1
adjusted[lo_conf] = 0
n_changed = int((adjusted != orig_arr).sum())
print(
"[_infection_qc_xgboost] Relabel mode: adjusted infection labels for "
f"{n_changed} cells (prob_thr={prob_thr:.2f}, margin={margin:.2f})."
)
cell_level["adjusted_infected"] = adjusted
cell_level["infection_prob"] = probs
else:
if margin > 0:
ambig = np.abs(probs - prob_thr) < margin
else:
ambig = np.zeros_like(probs, dtype=bool)
disagree = pred_arr != orig_arr
to_remove = disagree & ~ambig
removed = cell_level.loc[to_remove, key_cols]
removed_ids = {
(r["plateID"], r["wellID"], r["fieldID"], r["cellID"])
for _, r in removed.iterrows()
}
cell_level = cell_level.loc[~to_remove].copy()
kept_idx = (~to_remove).nonzero()[0]
adjusted = orig_arr[kept_idx].copy()
pred_kept = pred_arr[kept_idx]
ambig_kept = ambig[kept_idx]
adjusted[~ambig_kept] = pred_kept[~ambig_kept]
cell_level["adjusted_infected"] = adjusted
cell_level["infection_prob"] = probs[~to_remove]
print(
"[_infection_qc_xgboost] Remove mode: removed "
f"{len(removed_ids)} cells with strong model vs label disagreement "
f"(prob_thr={prob_thr:.2f}, margin={margin:.2f})."
)
drop_amb = bool(settings.get("infection_xgb_drop_ambiguous", True))
amb_low = float(settings.get("infection_xgb_ambiguous_low", 0.25))
amb_high = float(settings.get("infection_xgb_ambiguous_high", 0.75))
amb_low = max(0.0, min(amb_low, 1.0))
amb_high = max(0.0, min(amb_high, 1.0))
if amb_low > amb_high:
amb_low, amb_high = amb_high, amb_low
if drop_amb and "infection_prob" in cell_level.columns:
amb_mask = (
(cell_level["infection_prob"] >= amb_low)
& (cell_level["infection_prob"] <= amb_high)
)
if amb_mask.any():
amb = cell_level.loc[amb_mask, key_cols]
ambiguous_ids = {
(r["plateID"], r["wellID"], r["fieldID"], r["cellID"])
for _, r in amb.iterrows()
}
cell_level = cell_level.loc[~amb_mask].copy()
print(
"[_infection_qc_xgboost] Dropped "
f"{len(ambiguous_ids)} cells with ambiguous XGBoost probability "
f"in [{amb_low:.2f}, {amb_high:.2f}]."
)
for col in key_cols:
all_df[col] = all_df[col].astype(cell_level[col].dtype)
all_df = all_df.merge(
cell_level[key_cols + ["adjusted_infected", "infection_prob"]],
on=key_cols,
how="left",
validate="m:1",
)
ids_to_remove = set()
if removed_ids:
ids_to_remove |= removed_ids
if ambiguous_ids:
ids_to_remove |= ambiguous_ids
if ids_to_remove:
mask_drop = all_df.apply(
lambda r: (r["plateID"], r["wellID"], r["fieldID"], r["cellID"])
in ids_to_remove,
axis=1,
)
all_df = all_df.loc[~mask_drop].reset_index(drop=True)
mask_missing = all_df["adjusted_infected"].isna()
if mask_missing.any():
all_df.loc[mask_missing, "adjusted_infected"] = (
all_df.loc[mask_missing, orig_infection_col].astype(int)
)
all_df["adjusted_infected"] = all_df["adjusted_infected"].astype(int)
infection_col = "adjusted_infected"
try:
n_inf = int(all_df[infection_col].sum())
n_uninf = int((1 - all_df[infection_col]).sum())
print(
f"[_infection_qc_xgboost] Final infection counts (frame-level): "
f"infected={n_inf}, uninfected={n_uninf}"
)
except Exception:
pass
try:
intens = cell_level[intensity_col].to_numpy(dtype=float)
labels_adj = cell_level["adjusted_infected"].astype(bool).to_numpy()
mask_fin = np.isfinite(intens)
intens = intens[mask_fin]
labels_adj = labels_adj[mask_fin]
vals_inf = intens[labels_adj]
vals_uninf = intens[~labels_adj]
if intens.size >= 10:
n_bins = int(settings.get("infection_intensity_n_bins", 64))
n_bins = max(10, min(n_bins, 256))
_, bin_edges = np.histogram(intens, bins=n_bins)
thr_val = float(0.5 * (low_thr_uninf + high_thr_inf))
hist_payload = {
"intensities_inf": vals_inf,
"intensities_uninf": vals_uninf,
"bin_edges": bin_edges,
"thr_val": thr_val,
"pathogen_chan": pathogen_chan,
"log_transform": False,
"intensity_col": intensity_col,
"tracked_object": tracked_object,
}
settings["infection_hist_data"] = hist_payload
from sklearn.decomposition import PCA
from sklearn.preprocessing import StandardScaler
X_panel = cell_level[used_feature_cols].to_numpy(
dtype=float, copy=True
)
for j in range(X_panel.shape[1]):
col = X_panel[:, j]
m = np.isfinite(col)
if not m.any():
X_panel[:, j] = 0.0
else:
med = np.nanmedian(col[m])
col[~m] = med
X_panel[:, j] = col
scaler = StandardScaler()
X_scaled_panel = scaler.fit_transform(X_panel)
panel_components = min(2, X_scaled_panel.shape[1])
pca = PCA(
n_components=panel_components,
random_state=int(settings.get("infection_pca_random_state", 0)),
)
coords = pca.fit_transform(X_scaled_panel)
if panel_components == 1:
coords = np.column_stack(
[coords[:, 0], np.zeros(coords.shape[0], dtype=float)]
)
labels_adj_panel = cell_level["adjusted_infected"].astype(bool).to_numpy()
pca_payload = {
"coords": coords,
"labels": labels_adj_panel,
"method_label": "PCA",
"tracked_object": tracked_object,
}
settings["infection_pca_data"] = pca_payload
except Exception as e:
print(f"[_infection_qc_xgboost] Could not compute histogram/PCA payloads: {e}")
try:
importance_dict = bst.get_score(importance_type="gain") or {}
feat_names = used_feature_cols
feat_vals = [importance_dict.get(f, 0.0) for f in feat_names]
sorted_pairs = sorted(
zip(feat_names, feat_vals),
key=lambda x: x[1],
reverse=True,
)
feat_names = [p[0] for p in sorted_pairs]
feat_vals = [p[1] for p in sorted_pairs]
top_k = int(settings.get("infection_xgb_top_features", 20))
feat_names = feat_names[:top_k]
feat_vals = feat_vals[:top_k]
settings["infection_xgb_importance"] = {
"feature_names": feat_names,
"feature_importances": feat_vals,
"tracked_object": tracked_object,
}
if feat_names:
print("[_infection_qc_xgboost] Top XGBoost features (gain):")
for name, val in zip(feat_names, feat_vals):
print(f" {name}: {val:.4g}")
except Exception as e:
print(f"[_infection_qc_xgboost] Could not compute feature importances: {e}")
settings["infection_intensity_qc_panel_type"] = "xgboost"
settings["infection_intensity_qc_panel_path"] = None
return all_df, infection_col
[docs]
def automated_motility_assay(settings):
"""End-to-end merged-npy pipeline for cell/pathogen motility and infection QC.
Reads ``merged/*.npy`` frames, builds per-cell measurements, cleans and
persists them to SQLite, computes per-track velocities, generates
intensity + motility QC panels (mask-based and, optionally, XGBoost /
histogram / PCA / UMAP / t-SNE adjusted labels), and writes a
well-level motility summary.
:param settings: dict of assay settings; see
``get_automated_motility_assay_default_settings`` for keys
including ``src``, ``db_table_name``, ``n_jobs``,
``max_displacement``, ``track_outlier_zscore``, ``infection_intensity_qc``,
``infection_intensity_strategy``, ``infection_intensity_mode``,
``infection_xgb_drop_ambiguous``, ``infection_xgb_ambiguous_low``,
``infection_xgb_ambiguous_high``, ``infection_xgb_proba_column``,
``infection_hist_percentile``, ``make_mask_panel``,
``make_adjusted_panel``, ``motility_xlim``, ``motility_ylim``,
``motility_origin_xlim``, ``motility_origin_ylim``, and
``reuse_existing_measurements``. Optional ``bleach_correction`` defaults
to ``none``; ``ratio``, ``exponential`` or ``histogram`` correct each
field/channel and object role before child aggregation and infection
QC, using a five-pixel label-zero background ring (missing rings give
NaN). Histogram matching can erase biological changes. Opt-in runs
require recomputation rather than reuse of cached measurement rows.
Original intensity measurements remain raw; separate tables with
``_bleach_corrected`` and ``_bleach_fits`` suffixes and a fits CSV
record the corrected levels, source filenames and method.
:returns: the per-cell measurements DataFrame carrying the final
(QC-adjusted) labels. Measurements and summary tables are also written
to ``measurements/measurements.db`` and the QC panels saved under
``src``.
:raises ValueError: when ``settings['db_table_name']`` names a spaCR-owned
table; see :func:`_validate_db_table_name`.
Pre-QC corrected values are persisted separately; the cached originals stay
raw.
"""
import numpy as np
import pandas as pd
import os
from multiprocessing import cpu_count
from .resource_log import _parallel_pool as Pool
import sqlite3
from .settings import get_automated_motility_assay_default_settings
settings = get_automated_motility_assay_default_settings(settings)
src = settings["src"]
db_table_name = _validate_db_table_name(settings["db_table_name"])
n_jobs = settings["n_jobs"]
max_displacement = settings["max_displacement"]
track_outlier_zscore = settings["track_outlier_zscore"]
reuse_existing = settings.get("reuse_existing_measurements", True)
bleach_method = settings.get('bleach_correction', 'none')
if bleach_method not in _BLEACH_METHODS:
raise ValueError(f'Unknown motility bleach correction: {bleach_method!r}')
measurements_dir = os.path.join(src, "measurements")
db_path = os.path.join(measurements_dir, "measurements.db")
if bleach_method != 'none' and reuse_existing and os.path.isfile(db_path):
from pathlib import Path
with sqlite3.connect(Path(db_path).resolve().as_uri() + '?mode=ro', uri=True, timeout=30.0) as conn:
exists = conn.execute(
'SELECT 1 FROM sqlite_master WHERE type="table" AND name=?',
(db_table_name,)).fetchone()
if exists:
raise ValueError(
'Motility bleach correction needs original per-object backgrounds; '
'cached rows cannot verify them. Set reuse_existing_measurements=False '
'to recompute from the original merged arrays.')
os.makedirs(measurements_dir, exist_ok=True)
bleach_raw = None
bleach_fits = None
all_df = None
loaded_from_db = False
if reuse_existing and os.path.exists(db_path):
try:
print(
f"[summarise_tracks_from_merged] Attempting to reuse existing "
f"measurements from {db_path} (table='{db_table_name}')."
)
with sqlite3.connect(db_path, timeout=30) as conn:
all_df = pd.read_sql_query(f"SELECT * FROM {db_table_name}", conn)
if (
all_df is not None
and not all_df.empty
and {"plateID", "wellID", "fieldID", "cellID", "frame"}.issubset(
all_df.columns
)
):
n_frames_db = all_df["frame"].nunique()
n_tracks_db = (
all_df[["plateID", "wellID", "fieldID", "cellID"]]
.drop_duplicates()
.shape[0]
)
print(
"[summarise_tracks_from_merged] Loaded ORIGINAL measurements "
f"from DB: shape={all_df.shape}, frames={n_frames_db}, "
f"tracks={n_tracks_db}. Skipping regionprops/intensity "
"computation from merged .npy files."
)
loaded_from_db = True
else:
print(
"[summarise_tracks_from_merged] Loaded table is empty or missing "
"required columns; recomputing from merged .npy files."
)
all_df = None
except Exception as e:
print(
"[summarise_tracks_from_merged] Failed to reuse existing measurements "
f"({e}); recomputing from merged .npy files."
)
all_df = None
merged_dir = os.path.join(src, "merged")
if not os.path.isdir(merged_dir):
raise FileNotFoundError(f"No merged directory at: {merged_dir}")
all_files = [f for f in os.listdir(merged_dir) if f.endswith(".npy")]
if not all_files:
raise FileNotFoundError(f"No .npy files found in {merged_dir}")
print(
f"[summarise_tracks_from_merged] Found {len(all_files)} merged .npy files "
f"in {merged_dir}"
)
groups = {}
for fname in all_files:
meta = _parse_merged_filename(fname)
key = (meta["plateID"], meta["wellID"], meta["fieldID"])
groups.setdefault(key, []).append(fname)
print(
"[summarise_tracks_from_merged] Number of (plate, well, field) groups: "
f"{len(groups)}"
)
cell_chan = settings.get("cell_channel", None)
nucleus_chan = settings.get("nucleus_channel", None)
pathogen_chan = settings.get("pathogen_channel", None)
channels_list = settings.get("channels", [])
pixels_per_um = settings.get("pixels_per_um", None)
seconds_per_frame = settings.get("seconds_per_frame", None)
n_channels = len(channels_list) if isinstance(channels_list, (list, tuple)) else None
if n_channels is None or n_channels <= 0:
raise ValueError(
"settings['channels'] must be a non-empty list of channels used "
"in merged arrays."
)
print(
f"[summarise_tracks_from_merged] Channels={channels_list}, "
f"cell_chan={cell_chan}, nucleus_chan={nucleus_chan}, "
f"pathogen_chan={pathogen_chan}"
)
motility_dir = os.path.join(src, "motility_plots")
os.makedirs(motility_dir, exist_ok=True)
sample_filename = sorted(all_files)[0]
print(
"[summarise_tracks_from_merged] Debug plotting planes for sample file: "
f"{sample_filename}"
)
_debug_plot_merged_planes(
src=src,
sample_filename=sample_filename,
n_channels=n_channels,
nucleus_chan=nucleus_chan,
pathogen_chan=pathogen_chan,
out_dir=motility_dir,
)
if not loaded_from_db:
worker_args = []
for key, file_basenames in groups.items():
args = (src, file_basenames, n_channels, cell_chan, nucleus_chan, pathogen_chan)
worker_args.append(args if bleach_method == 'none' else args + (bleach_method,))
if n_jobs is None:
n_jobs = max(cpu_count() - 1, 1)
if worker_args and worker_args[0][1]:
from .resource_log import _array_file_nbytes, _guard_workers
first = worker_args[0][1][0]
n_jobs = _guard_workers('motility', n_jobs, len(worker_args[0][1]) * _array_file_nbytes(
first if os.path.isabs(str(first)) else os.path.join(src, str(first))))
print(f"[summarise_tracks_from_merged] Using n_jobs={n_jobs}")
if n_jobs == 1:
dfs = [_process_merged_group(args) for args in worker_args]
else:
with Pool(processes=n_jobs) as pool:
dfs = pool.map(_process_merged_group, worker_args)
if bleach_method != 'none':
results = [result for result in dfs if isinstance(result, tuple)]
if results:
bleach_raw = pd.concat([result[1] for result in results], ignore_index=True)
bleach_raw = _smooth_tracks_and_features(
bleach_raw, max_displacement=max_displacement,
track_outlier_zscore=track_outlier_zscore)
bleach_fits = pd.concat([result[2] for result in results], ignore_index=True)
dfs = [result[0] for result in results]
non_empty = [df for df in dfs if not df.empty]
all_df = (
pd.concat(non_empty, ignore_index=True) if non_empty else pd.DataFrame()
)
if all_df.empty:
raise RuntimeError("No measurements were produced from merged .npy files.")
print(
"[summarise_tracks_from_merged] Combined raw measurements: "
f"shape={all_df.shape}, frames={all_df['frame'].nunique()}"
)
n_tracks_raw = (
all_df[["plateID", "wellID", "fieldID", "cellID"]]
.drop_duplicates()
.shape[0]
)
print(
"[summarise_tracks_from_merged] Unique tracks before smoothing: "
f"{n_tracks_raw}"
)
smooth = (_smooth_tracks_and_features if bleach_method == 'none'
else _motility_smooth_corrected)
all_df = smooth(
all_df,
max_displacement=max_displacement,
track_outlier_zscore=track_outlier_zscore,
)
n_tracks_smoothed = (
all_df[["plateID", "wellID", "fieldID", "cellID"]]
.drop_duplicates()
.shape[0]
)
print(
"[summarise_tracks_from_merged] After smoothing: "
f"shape={all_df.shape}, frames={all_df['frame'].nunique()}, "
f"tracks={n_tracks_smoothed}"
)
else:
n_frames_db = all_df["frame"].nunique()
n_tracks_db = (
all_df[["plateID", "wellID", "fieldID", "cellID"]]
.drop_duplicates()
.shape[0]
)
print(
"[summarise_tracks_from_merged] Reusing ORIGINAL smoothed measurements "
f"from DB: shape={all_df.shape}, frames={n_frames_db}, "
f"tracks={n_tracks_db}"
)
if "infected" in all_df.columns:
all_df["infected"] = all_df["infected"].fillna(False).astype(bool)
elif "n_pathogens" in all_df.columns:
tmp = all_df[["plateID", "wellID", "fieldID", "cellID", "n_pathogens"]].copy()
tmp["n_pathogens"] = tmp["n_pathogens"].fillna(0)
infected = (
tmp.groupby(["plateID", "wellID", "fieldID", "cellID"])["n_pathogens"]
.max()
.gt(0)
)
infected = infected.reset_index()
infected = infected.rename(columns={"n_pathogens": "infected"})
infected["infected"] = infected["infected"].astype(bool)
all_df = all_df.merge(
infected[["plateID", "wellID", "fieldID", "cellID", "infected"]],
on=["plateID", "wellID", "fieldID", "cellID"],
how="left",
validate="many_to_one",
)
all_df["infected"] = all_df["infected"].fillna(False).astype(bool)
else:
all_df["infected"] = False
n_infected_tracks = (
all_df[all_df["infected"]][["plateID", "wellID", "fieldID", "cellID"]]
.drop_duplicates()
.shape[0]
)
n_uninfected_tracks = (
all_df[~all_df["infected"]][["plateID", "wellID", "fieldID", "cellID"]]
.drop_duplicates()
.shape[0]
)
print(
"[summarise_tracks_from_merged] Tracks (mask-based): "
f"infected={n_infected_tracks}, uninfected={n_uninfected_tracks}"
)
all_df_original = all_df.copy(deep=True)
if bleach_raw is not None:
keys = ['plateID', 'wellID', 'fieldID', 'cellID', 'frame']
all_df_original = bleach_raw.drop(columns=['infected'], errors='ignore').merge(
all_df[keys + ['infected']], on=keys, how='inner', validate='one_to_one')
from .tabular import write_database, write_table
write_database(all_df, db_path, db_table_name + '_bleach_corrected',
if_exists='replace')
write_database(bleach_fits, db_path, db_table_name + '_bleach_fits',
if_exists='replace')
write_table(bleach_fits, os.path.join(measurements_dir,
db_table_name + '_bleach_fits.csv'))
if bleach_method == 'histogram':
print('Motility histogram matching removes population intensity changes; '
'do not interpret these values as quantitative bleaching correction.')
infection_col = "infected"
all_df, infection_col = _apply_infection_intensity_qc(
all_df=all_df,
settings=settings,
infection_col=infection_col,
motility_dir=motility_dir,
pathogen_chan=pathogen_chan,
)
if all_df.empty and not all_df_original.empty:
print(
"[summarise_tracks_from_merged] WARNING: infection-intensity QC "
"removed every row; falling back to the original mask labels so "
"the assay and well summary remain usable."
)
all_df = all_df_original.copy(deep=True)
all_df["adjusted_infected"] = (
all_df["infected"].fillna(False).astype(int))
all_df["infection_prob"] = np.nan
infection_col = "adjusted_infected"
if (
settings.get("infection_intensity_qc", False)
and str(settings.get("infection_intensity_strategy", "")).lower() == "xgboost"
and settings.get("infection_xgb_drop_ambiguous", True)
):
low = settings.get("infection_xgb_ambiguous_low", 0.25)
high = settings.get("infection_xgb_ambiguous_high", 0.75)
xgb_proba_col = settings.get("infection_xgb_proba_column", None)
_shipped_default = "infection_xgb_proba"
_chosen_by_user = xgb_proba_col not in (None, "", _shipped_default)
if not _chosen_by_user and xgb_proba_col not in all_df.columns:
cand_cols = [
c
for c in all_df.columns
if "xgb" in c.lower()
and (
"proba" in c.lower()
or "prob" in c.lower()
or "score" in c.lower()
)
]
if not cand_cols:
cand_cols = [
c
for c in all_df.columns
if "infection" in c.lower()
and (
"proba" in c.lower()
or "prob" in c.lower()
or "score" in c.lower()
)
]
if cand_cols:
xgb_proba_col = cand_cols[0]
if xgb_proba_col and xgb_proba_col in all_df.columns:
track_keys = ["plateID", "wellID", "fieldID", "cellID"]
track_scores = (
all_df[track_keys + [xgb_proba_col]]
.groupby(track_keys)[xgb_proba_col]
.mean()
.reset_index()
)
ambiguous = track_scores[
(track_scores[xgb_proba_col] > low)
& (track_scores[xgb_proba_col] < high)
][track_keys]
if not ambiguous.empty:
before = all_df.shape[0]
all_df = all_df.merge(
ambiguous.assign(_ambiguous_flag=1),
on=track_keys,
how="left",
validate="many_to_one",
)
all_df = all_df[all_df["_ambiguous_flag"].isna()].drop(
columns=["_ambiguous_flag"]
)
after = all_df.shape[0]
print(
"[summarise_tracks_from_merged] Dropped "
f"{before - after} rows from {ambiguous.shape[0]} ambiguous "
f"XGBoost tracks ({low} < proba < {high})."
)
else:
print(
"[summarise_tracks_from_merged] WARNING: "
"infection_xgb_drop_ambiguous is True, but no XGBoost "
"probability/score column was found. Skipping ambiguous-track "
"filtering."
)
try:
qc_strategy = str(settings.get("infection_intensity_strategy", "none")).lower()
if qc_strategy in {"", "none", "null"}:
adjusted_basename = f"{db_table_name}_adjusted.csv"
else:
adjusted_basename = f"{db_table_name}_adjusted_{qc_strategy}.csv"
adjusted_csv_path = os.path.join(measurements_dir, adjusted_basename)
all_df.to_csv(adjusted_csv_path, index=False)
print(
"[summarise_tracks_from_merged] Saved ADJUSTED frame-level measurements "
f"to CSV: {adjusted_csv_path}"
)
except Exception as e:
print(
f"[summarise_tracks_from_merged] WARNING: failed to save adjusted CSV "
f"({e})"
)
(
track_df_mask,
per_well_tracks_mask,
well_summary_mask,
vel_unit_mask,
) = _compute_velocities_and_well_summary(
all_df=all_df,
settings=settings,
infection_col="infected",
pixels_per_um=pixels_per_um,
seconds_per_frame=seconds_per_frame,
)
(
track_df,
per_well_tracks,
well_summary_df,
vel_unit,
) = _compute_velocities_and_well_summary(
all_df=all_df,
settings=settings,
infection_col=infection_col,
pixels_per_um=pixels_per_um,
seconds_per_frame=seconds_per_frame,
)
measurements_dir, db_path = _save_measurements_and_well_summary(
all_df=all_df_original,
well_summary_df=well_summary_df,
src=src,
db_table_name=db_table_name,
)
_feature_velocity_correlations(all_df, track_df, measurements_dir)
qc_strategy = str(settings.get("infection_intensity_strategy", "none")).lower()
if settings.get("make_mask_panel", True):
_make_intensity_motility_panel(
all_df=all_df,
infection_col="infected",
track_df=track_df_mask,
per_well_tracks=per_well_tracks_mask,
n_channels=n_channels,
motility_dir=motility_dir,
pixels_per_um=pixels_per_um,
seconds_per_frame=seconds_per_frame,
vel_unit=vel_unit_mask,
settings=settings,
label_tag=f"mask_{qc_strategy}",
)
if (
settings.get("make_adjusted_panel", True)
and infection_col in all_df.columns
and infection_col != "infected"
):
_make_intensity_motility_panel(
all_df=all_df,
infection_col=infection_col,
track_df=track_df,
per_well_tracks=per_well_tracks,
n_channels=n_channels,
motility_dir=motility_dir,
pixels_per_um=pixels_per_um,
seconds_per_frame=seconds_per_frame,
vel_unit=vel_unit,
settings=settings,
label_tag=f"adjusted_{qc_strategy}",
)
return all_df