Source code for spacr.timelapse

"""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)





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_figure(fig, src, figure_number): """Save ``fig`` as ``figure_<figure_number>`` inside a sibling ``results`` folder. The ``.pdf`` extension built here is only a proposal: :func:`spacr.plot.save_figure` rewrites it to the configured figure format. That preference defaults to ``pdf``, so the file on disk is ``figure_1.pdf`` unless another format has been selected. :param fig: matplotlib Figure to persist. :param src: reference path used to derive the parent directory. :param figure_number: integer/string suffix embedded in the filename. :returns: None; the written path is printed. """ source = os.path.dirname(src) results_fldr = os.path.join(source,'results') os.makedirs(results_fldr, exist_ok=True) fig_loc = os.path.join(results_fldr, f'figure_{figure_number}.pdf') fig_loc = save_figure_to_path(fig, fig_loc) print(f'Saved figure:{fig_loc}')
[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)
[docs] def create_results_figure(): """Create the standard 3-panel QC results figure layout. Arrangement is PCA (top-left), XGBoost (top-right) and Histogram (bottom spanning both columns). :returns: tuple ``(fig, ax_pca, ax_xgb, ax_hist)``. """ with figure_style(theme_target()): fig = Figure(figsize=(7, 6), dpi=100) gs = fig.add_gridspec(2, 2, height_ratios=[2, 1]) ax_pca = fig.add_subplot(gs[0, 0]) ax_xgb = fig.add_subplot(gs[0, 1]) ax_hist = fig.add_subplot(gs[1, :]) return fig, ax_pca, ax_xgb, ax_hist
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