diff --git a/.github/workflows/prepare_test_data.yaml b/.github/workflows/prepare_test_data.yaml index a72da9d6..e9fc831f 100644 --- a/.github/workflows/prepare_test_data.yaml +++ b/.github/workflows/prepare_test_data.yaml @@ -85,6 +85,17 @@ jobs: # OMAP10 for format v0.x.x curl -o OMAP10_small.zip "https://zenodo.org/api/records/18196366/files-archive" + # ------- + # Stellaromics Pyxa, 100 um cube cropped from the public demo dataset + # https://huggingface.co/datasets/Stellaromics/demo + mkdir -p pyxa_xsmall + for file in cell_assigned_gene_v1.csv cell_by_gene_v1.csv cell_metadata_v1.csv segmentation_geometries_v1.parquet mosaic_3d.ome.zarr.zip; do + curl -L -o "pyxa_xsmall/$file" "https://huggingface.co/datasets/Stellaromics/demo/resolve/main/xsmall/$file" + done + # the zipped OME-Zarr mosaic is extracted in place (it contains a single `mosaic_3d.ome.zarr/` directory) + unzip -q pyxa_xsmall/mosaic_3d.ome.zarr.zip -d pyxa_xsmall + rm pyxa_xsmall/mosaic_3d.ome.zarr.zip + - name: Unzip files run: | cd ./data diff --git a/CHANGELOG.md b/CHANGELOG.md index 6521a5dc..491d20f9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -16,6 +16,7 @@ Release notes for `v0.7.1` and earlier are available on the [Releases][] page. ### Added - `spatialdata_io` ships a `py.typed` marker, so downstream type checkers use its annotations. +- Experimental `pyxa` reader (Stellaromics Pyxa): table (sparse counts, Pyxa Studio clusters/UMAP), transcripts, segmentation shapes, mosaic image (directory or zip), and with `labels=True` 3D cell labels on the mosaic grid. ### Changed diff --git a/README.md b/README.md index adb20e76..3748eb3a 100644 --- a/README.md +++ b/README.md @@ -48,6 +48,16 @@ Contributions for addressing the below limitations are very welcomed. - Only Stereo-seq 7.x is supported, 8.x is not currently supported. https://github.com/scverse/spatialdata-io/issues/161 +## Experimental readers + +Readers without (yet) a public specification for their raw data format live +in `spatialdata_io.experimental` rather than the main technology list above. +No stability guarantees are made for these. + +- Pyxa (Stellaromics): no public format specification yet; validated against + the public [demo dataset](https://huggingface.co/datasets/Stellaromics/demo). + `labels=True` adds 3D cell labels on the mosaic grid. + ## Getting started Please refer to the [documentation][link-docs]. In particular, the diff --git a/pyproject.toml b/pyproject.toml index 9c7135e1..ffddfa91 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -33,6 +33,7 @@ dependencies = [ "joblib", "numpy", "ome-types", + "pillow", "pyarrow", "readfcs", "scanpy", @@ -157,6 +158,7 @@ exclude = "^tests/data/" module = [ "dask_image.*", "h5py.*", + "joblib.*", "multiscale_spatial_image.*", "pyarrow.*", "rasterio.*", diff --git a/src/spatialdata_io/__main__.py b/src/spatialdata_io/__main__.py index f0a5f9f7..5fb0a808 100644 --- a/src/spatialdata_io/__main__.py +++ b/src/spatialdata_io/__main__.py @@ -910,6 +910,61 @@ def macsima_wrapper( sdata.write(output) +@cli.command(name="pyxa") +@_input_output_click_options +@click.option("--dataset-id", type=str, default="pyxa", help="Dataset ID. [default: pyxa]") +@click.option( + "--image", + type=click.Path(exists=True, file_okay=True, dir_okay=True), + default=None, + help="Mosaic OME-Zarr directory or zip, if not in the input directory. [default: found in the input]", +) +@click.option("--no-image", is_flag=True, default=False, help="Leave out the mosaic even if present.") +@click.option( + "--pyxa-studio", + type=click.Path(exists=True, file_okay=True, dir_okay=False), + default=None, + help="Path to a Pyxa Studio export (cluster labels, UMAP) outside the input directory. [default: None]", +) +@click.option( + "--skip", + type=click.Choice(["cell_assigned_gene", "segmentation_geometries", "pyxa_studio"]), + multiple=True, + help="Optional input file to leave out even if present; repeatable. [default: none]", +) +@click.option("--labels", is_flag=True, default=False, help="Rasterize 3D cell labels onto the mosaic's grid.") +@click.option( + "--shapes/--no-shapes", + default=None, + help="Return the polygons as shapes. [default: when read and --labels is not set]", +) +def pyxa_wrapper( + input: str, + output: str, + dataset_id: str = "pyxa", + image: str | None = None, + no_image: bool = False, + pyxa_studio: str | None = None, + skip: tuple[str, ...] = (), + labels: bool = False, + shapes: bool | None = None, +) -> None: + """Pyxa (Stellaromics) conversion to SpatialData.""" + from spatialdata_io.experimental import pyxa + + if no_image and image is not None: + raise click.UsageError("--image and --no-image are mutually exclusive") + inputs: dict[str, str | bool] = dict.fromkeys(skip, False) + if pyxa_studio is not None and "pyxa_studio" not in skip: + inputs["pyxa_studio"] = pyxa_studio + if no_image: + inputs["image"] = False + elif image is not None: + inputs["image"] = image + sdata = pyxa(input, dataset_id=dataset_id, labels=labels, shapes=shapes, **inputs) # type: ignore[arg-type] + sdata.write(output) + + @cli.command(name="generic") @click.option( "--input", diff --git a/src/spatialdata_io/_constants/_constants.py b/src/spatialdata_io/_constants/_constants.py index 46f983b1..5c426d9f 100644 --- a/src/spatialdata_io/_constants/_constants.py +++ b/src/spatialdata_io/_constants/_constants.py @@ -409,3 +409,59 @@ class VisiumHDKeys(ModeEnum): # Cell Segmentation keys CELL_SEG_KEY_HD = "cell_segmentations" NUCLEUS_SEG_KEY_HD = "nucleus_segmentations" + + +class PyxaKeys(ModeEnum): + """Keys for *Pyxa* (Stellaromics) output. + + No public specification exists yet; keys are validated against the public + demo dataset at https://huggingface.co/datasets/Stellaromics/demo. + """ + + # files + CELL_ASSIGNED_GENE_FILE = "cell_assigned_gene_v1.csv" + CELL_BY_GENE_FILE = "cell_by_gene_v1.csv" + CELL_METADATA_FILE = "cell_metadata_v1.csv" + SEGMENTATION_GEOMETRIES_FILE = "segmentation_geometries_v1.parquet" + # Pyxa Studio export: cells that passed Pyxa's filters, with cluster labels and a 3D UMAP + PYXA_STUDIO_FILE = "pyxa_studio_v1.csv" + + # shared columns + CELL_ID = "cell_id" + GENE = "Gene" + X_UM = "X_um" + Y_UM = "Y_um" + Z_UM = "Z_um" + X_PIXELS = "X_pixels" + Y_PIXELS = "Y_pixels" + Z_PIXELS = "Z_pixels" + VOLUME_UM3 = "Volume_um3" + ROI = "ROI" + Z_INDEX = "ZIndex" + BORDER = "Border" + FOV = "FOV" + CLUSTER = "Cluster" + X_UMAP = "X_UMAP" + Y_UMAP = "Y_UMAP" + Z_UMAP = "Z_UMAP" + + # unassigned transcripts have cell_id ending in this suffix, e.g. "Region_-1" + UNASSIGNED_SUFFIX = "_-1" + + # constructed metadata + REGION_KEY = "region" + # per-cell footprint (union of the cell's z-plane polygons), annotated by the table + REGION = "cell_boundaries" + # per-cell, per-z-plane polygons, as stored on disk + CELL_BOUNDARIES_Z = "cell_boundaries_z" + INSTANCE_KEY = "cell_id" + ASSIGNED = "assigned" + MOSAIC_IMAGE = "mosaic_image" + UMAP_KEY = "X_umap" + + # mosaic image, looked up in the Pyxa directory (unzipped or zipped, as on the Hub) + MOSAIC_FILE = "mosaic_3d.ome.zarr" + MOSAIC_ZIP_FILE = "mosaic_3d.ome.zarr.zip" + # 3D cell labels rasterized from the segmentation polygons onto the mosaic's grid + CELL_LABELS = "cell_labels" + LABEL_ID = "label_id" diff --git a/src/spatialdata_io/experimental/__init__.py b/src/spatialdata_io/experimental/__init__.py index 36dd4d8b..395911b6 100644 --- a/src/spatialdata_io/experimental/__init__.py +++ b/src/spatialdata_io/experimental/__init__.py @@ -3,9 +3,11 @@ to_legacy_anndata, ) from spatialdata_io.readers.iss import iss +from spatialdata_io.readers.pyxa import pyxa _readers_technologies = [ "iss", + "pyxa", ] _readers_file_types: list[str] = [ # add experimental readers for new file types here diff --git a/src/spatialdata_io/readers/pyxa.py b/src/spatialdata_io/readers/pyxa.py new file mode 100644 index 00000000..7eb57dd5 --- /dev/null +++ b/src/spatialdata_io/readers/pyxa.py @@ -0,0 +1,1011 @@ +from __future__ import annotations + +import zipfile +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Any, cast + +import anndata as ad +import dask +import dask.array as da +import dask.dataframe as dd +import geopandas as gpd +import joblib +import numpy as np +import pandas as pd +import pyarrow.compute as pc +import pyarrow.parquet as pq +import shapely +import zarr +from joblib.externals.loky import get_reusable_executor +from scipy import sparse +from spatialdata import SpatialData +from spatialdata._logging import logger +from spatialdata.models import Image3DModel, Labels3DModel, PointsModel, ShapesModel, TableModel +from spatialdata.transformations import Scale, Sequence, Translation, set_transformation +from xarray import DataArray, Dataset, DataTree + +from spatialdata_io._constants._constants import PyxaKeys +from spatialdata_io._docs import inject_docs + +__all__ = ["pyxa"] + + +def _validate_columns(df: pd.DataFrame | dd.DataFrame, required: set[str], file_name: str) -> None: + """Raise a clear ``ValueError`` naming the file and any missing required column(s).""" + missing = required - set(df.columns) + if missing: + raise ValueError(f"{file_name} is missing required column(s): {sorted(missing)}") + + +def _get_points(path: Path) -> dd.DataFrame: + ddf = dd.read_csv(path, dtype={PyxaKeys.CELL_ID.value: str}) + _validate_columns( + ddf, + { + PyxaKeys.CELL_ID.value, + PyxaKeys.GENE.value, + PyxaKeys.X_UM.value, + PyxaKeys.Y_UM.value, + PyxaKeys.Z_UM.value, + }, + path.name, + ) + ddf[PyxaKeys.ASSIGNED.value] = ~ddf[PyxaKeys.CELL_ID.value].str.endswith(PyxaKeys.UNASSIGNED_SUFFIX.value) + # PointsModel needs known feature categories; computing them here is one pass over the gene column + ddf[PyxaKeys.GENE.value] = ddf[PyxaKeys.GENE.value].astype("category").cat.as_known() + return ddf + + +def _read_cells(path: Path) -> pd.DataFrame: + return pd.read_csv(path, index_col=PyxaKeys.CELL_ID.value, dtype={PyxaKeys.CELL_ID.value: str}) + + +def _cluster_categorical(values: pd.Series) -> pd.Categorical[str]: + """Cluster labels as string categories, in numeric order when the labels are integers.""" + present = values.dropna() + numeric = pd.to_numeric(present, errors="coerce") + if numeric.notna().all(): + labels = numeric.astype(int).astype(str) + categories = [str(c) for c in sorted(numeric.astype(int).unique())] + else: + labels = present.astype(str) + categories = sorted(labels.unique()) + return pd.Categorical(labels.reindex(values.index), categories=categories) + + +def _get_table( + cell_by_gene_path: Path, + cell_metadata_path: Path, + pyxa_studio_path: Path | None = None, +) -> ad.AnnData: + """Build the cell table from the counts and metadata, optionally with Pyxa Studio's clusters. + + ``pyxa_studio`` lists only the cells that passed Pyxa's filters, so its ``Cluster`` + (categorical) and ``obsm["X_umap"]`` are missing (NaN) for the other cells. + """ + by_gene = _read_cells(cell_by_gene_path) + metadata = _read_cells(cell_metadata_path) + + spatial_cols = [PyxaKeys.X_UM.value, PyxaKeys.Y_UM.value, PyxaKeys.Z_UM.value] + _validate_columns(metadata, set(spatial_cols), cell_metadata_path.name) + + metadata = metadata.loc[by_gene.index] + # Pyxa counts are mostly zeros: a full Region's dense float64 table is several GB + adata = ad.AnnData( + sparse.csr_matrix(by_gene.to_numpy()), + obs=metadata.drop(columns=spatial_cols), + var=pd.DataFrame(index=by_gene.columns), + ) + adata.obsm["spatial"] = metadata[spatial_cols].values + + if pyxa_studio_path is not None: + studio = _read_cells(pyxa_studio_path) + _validate_columns(studio, {PyxaKeys.CLUSTER.value}, pyxa_studio_path.name) + n_missing = int((~adata.obs_names.isin(studio.index)).sum()) + if n_missing: + logger.info( + f"{pyxa_studio_path.name}: {n_missing} of {adata.n_obs} cells were filtered out by Pyxa Studio; " + f"their {PyxaKeys.CLUSTER.value!r} and {PyxaKeys.UMAP_KEY.value!r} are missing" + ) + studio = studio.reindex(adata.obs_names) + adata.obs[PyxaKeys.CLUSTER.value] = _cluster_categorical(studio[PyxaKeys.CLUSTER.value]) + umap_cols = [c for c in (PyxaKeys.X_UMAP.value, PyxaKeys.Y_UMAP.value, PyxaKeys.Z_UMAP.value) if c in studio] + if umap_cols: + adata.obsm[PyxaKeys.UMAP_KEY.value] = studio[umap_cols].to_numpy(dtype=np.float64) + + adata.obs[PyxaKeys.REGION_KEY.value] = pd.Series(PyxaKeys.REGION.value, index=adata.obs_names, dtype="category") + adata.obs[PyxaKeys.CELL_ID.value] = adata.obs_names + # the cell_id column carries the instance key; an index with the same name breaks table joins + adata.obs.index.name = None + return adata + + +def _get_voxel_size(cell_metadata_path: Path, n_rows: int = 10_000) -> tuple[float, float]: + """Infer the (xy, z) voxel size in um of the segmentation polygons from the per-cell metadata. + + Polygons are stored in pixel coordinates (xy) with a z-plane index (``ZIndex``), while points + and the table use micrometers. The metadata carries each cell centroid in both units, related + by a pure scale per axis (no offset), so each size is a least-squares fit through the origin. + A pure scale is fully determined by a few cells, so only the first ``n_rows`` cells are read; + the fit is then checked to reproduce every sampled cell, and an error is raised if it does not. + """ + columns = [ + PyxaKeys.X_UM.value, + PyxaKeys.Y_UM.value, + PyxaKeys.Z_UM.value, + PyxaKeys.X_PIXELS.value, + PyxaKeys.Y_PIXELS.value, + PyxaKeys.Z_PIXELS.value, + ] + metadata = pd.read_csv(cell_metadata_path, usecols=lambda c: c in columns, nrows=n_rows) + _validate_columns(metadata, set(columns), cell_metadata_path.name) + + def fit(um_cols: list[str], px_cols: list[str]) -> float: + um = metadata[um_cols].to_numpy().ravel() + px = metadata[px_cols].to_numpy().ravel() + size = float(np.dot(um, px) / np.dot(px, px)) + residual = np.abs(um - size * px).max() + if residual > 1e-6 * max(np.abs(um).max(), 1.0): + raise ValueError( + f"{cell_metadata_path.name}: {um_cols} and {px_cols} are not related by a pure scale " + f"(max residual {residual:.3g} um with a fitted size of {size:.6g} um/pixel)" + ) + return size + + xy = fit([PyxaKeys.X_UM.value, PyxaKeys.Y_UM.value], [PyxaKeys.X_PIXELS.value, PyxaKeys.Y_PIXELS.value]) + z = fit([PyxaKeys.Z_UM.value], [PyxaKeys.Z_PIXELS.value]) + return xy, z + + +def _polygonal_part(geometry: shapely.Geometry) -> shapely.Geometry: + """Keep only the (multi)polygonal part of a geometry, dropping any lines or points.""" + if isinstance(geometry, shapely.Polygon | shapely.MultiPolygon): + return geometry + polygons = [p for p in shapely.get_parts(geometry) if isinstance(p, shapely.Polygon)] + return shapely.MultiPolygon(polygons) if len(polygons) > 1 else polygons[0] if polygons else shapely.Polygon() + + +def _make_polygonal_valid(geometries: np.ndarray) -> np.ndarray: + """Repair invalid geometries in place of dropping them, keeping each one a (Multi)Polygon. + + Valid geometries are returned untouched. Invalid ones are repaired with + :func:`shapely.make_valid` (``"structure"`` method), which rebuilds the polygon from its + rings, and any zero-area parts produced by the repair (lines, points) are discarded. + """ + geometries = geometries.copy() + invalid = ~shapely.is_valid(geometries) + if invalid.any(): + repaired = shapely.make_valid(geometries[invalid], method="structure", keep_collapsed=False) + geometries[invalid] = [_polygonal_part(g) for g in repaired] + return geometries + + +def _get_shapes(path: Path, xy_size: float, z_size: float) -> gpd.GeoDataFrame: + """Read the per-cell, per-z-plane segmentation polygons and convert them to micrometers. + + xy coordinates are scaled from pixels by ``xy_size``. Since shapes are 2D in spatialdata, z is + stored as a ``Z_um`` column: plane ``k`` spans ``[k, k + 1)`` in ``Z_pixels`` units, so its + centre sits at ``(k + 0.5) * z_size``. Scaling can turn polygons that touch themselves at a + single vertex into self-intersecting ones through floating point rounding, so the scaled + geometries are repaired and then validated. + """ + parquet_file = pq.ParquetFile(path) + chunks = [] + for batch in parquet_file.iter_batches(): + chunk = batch.to_pandas() + chunk["geometry"] = shapely.from_wkb(chunk["geometry"]) + chunks.append(gpd.GeoDataFrame(chunk, geometry="geometry")) + gdf = pd.concat(chunks, ignore_index=True) + gdf = gpd.GeoDataFrame(gdf, geometry="geometry") + + _validate_columns(gdf, {PyxaKeys.CELL_ID.value, PyxaKeys.Z_INDEX.value}, path.name) + + gdf[PyxaKeys.CELL_ID.value] = gdf[PyxaKeys.CELL_ID.value].astype(str) + gdf[PyxaKeys.Z_UM.value] = (gdf[PyxaKeys.Z_INDEX.value] + 0.5) * z_size + + scaled = shapely.transform(gdf.geometry.to_numpy(), lambda coords: coords * xy_size) + n_invalid = int((~shapely.is_valid(scaled)).sum()) + fixed = _make_polygonal_valid(scaled) + if n_invalid: + area_change = np.abs(shapely.area(fixed) - shapely.area(scaled)) / np.maximum(shapely.area(scaled), 1e-12) + logger.info( + f"{path.name}: repaired {n_invalid} invalid polygon(s) after scaling to micrometers " + f"(max relative area change {area_change.max():.2g})" + ) + gdf = gdf.set_geometry(fixed) + + empty = gdf.geometry.is_empty.to_numpy() + if empty.any(): + logger.warning(f"{path.name}: dropping {int(empty.sum())} polygon(s) with no area left after repair") + gdf = gdf[~empty] + if not gdf.geometry.is_valid.all() or not set(gdf.geom_type) <= {"Polygon", "MultiPolygon"}: + raise ValueError(f"{path.name}: segmentation polygons are still invalid or non-polygonal after repair") + + return gdf.reset_index(drop=True) + + +def _get_footprints(planes: gpd.GeoDataFrame) -> gpd.GeoDataFrame: + """Merge each cell's z-plane polygons into one 2D footprint, indexed by cell id. + + A table can only annotate an element with one row per instance, so the per-plane polygons are + kept as a separate element and the table annotates these footprints instead. The union is the + expensive step and is independent per cell; shapely releases the GIL, so cells are merged in a + thread pool. + """ + cell_ids = planes[PyxaKeys.CELL_ID.value].to_numpy() + order = np.argsort(cell_ids, kind="stable") + cell_ids, geometries = cell_ids[order], planes.geometry.to_numpy()[order] + unique_ids, starts = np.unique(cell_ids, return_index=True) + stops = np.append(starts[1:], len(cell_ids)) + + with ThreadPoolExecutor() as executor: + unions = list( + executor.map(lambda bounds: shapely.union_all(geometries[slice(*bounds)]), zip(starts, stops, strict=True)) + ) + + footprints = gpd.GeoDataFrame( + geometry=_make_polygonal_valid(np.array(unions, dtype=object)), + index=pd.Index(unique_ids, name=PyxaKeys.CELL_ID.value), + ) + if not footprints.geometry.is_valid.all(): + raise ValueError("cell footprints are invalid after merging the z-plane polygons") + return footprints + + +def _open_mosaic(path: Path) -> zarr.Group: + """Open a mosaic OME-Zarr read-only, from its directory or from a zip of it (read in place). + + A zip holds either the group at its root or a single top-level ``.ome.zarr/`` directory, + as in the Stellaromics/demo dataset. Top-level entries starting with ``__`` (e.g. the + ``__MACOSX/`` tree of a zip made on macOS) are not the mosaic and are ignored. + """ + if path.suffix != ".zip": + return zarr.open_group(store=str(path), mode="r") + with zipfile.ZipFile(path) as zf: + tops = {top for name in zf.namelist() if not (top := name.split("/", 1)[0]).startswith("__")} + group_path = "" if "zarr.json" in tops or len(tops) != 1 else tops.pop() + return zarr.open_group(store=zarr.storage.ZipStore(path, mode="r"), mode="r", path=group_path) + + +@dataclass(frozen=True) +class _MosaicGrid: + """The mosaic's voxel grid. + + Every level's (z, y, x) shape and OME-NGFF scale, and level 0's frame (scale and translation) in + micrometers. + """ + + shapes: tuple[tuple[int, int, int], ...] + scale: tuple[float, float, float] + translation: tuple[float, float, float] + scales: tuple[tuple[float, float, float], ...] + + @property + def transformation(self) -> Sequence: + axes = ("z", "y", "x") + return Sequence([Scale(list(self.scale), axes=axes), Translation(list(self.translation), axes=axes)]) + + def step(self, level: int) -> tuple[int, int, int]: + """Level-0 voxels per voxel of ``level``, per axis: the OME-NGFF scale ratio to level 0. + + Each axis's ratio must be within ``1e-6`` of a positive integer (nearest-neighbour stride); + the ratio of the array *shapes* is not used, since a cropped or oddly-sized level can round + that ratio to the wrong integer even where the scale ratio itself is exact. + """ + ratios = tuple(s / s0 for s0, s in zip(self.scales[0], self.scales[level], strict=True)) + steps = tuple(round(r) for r in ratios) + for axis, (ratio, step) in enumerate(zip(ratios, steps, strict=True)): + if step < 1 or abs(ratio - step) > 1e-6: + raise ValueError(f"scale{level} axis {axis}: scale ratio to scale0 ({ratio}) is not a positive integer") + return steps # type: ignore[return-value] + + +def _mosaic_grid(path: Path) -> _MosaicGrid: + """The mosaic's level shapes and OME-NGFF scales, and level 0's scale and translation, from its OME-NGFF metadata.""" + group = _open_mosaic(path) + multiscale = cast("dict[str, Any]", group.attrs.asdict()["ome"])["multiscales"][0] + axes = [a["name"] for a in multiscale["axes"]] + zyx = [axes.index(a) for a in ("z", "y", "x")] + datasets = multiscale["datasets"] + shapes = tuple(tuple(int(cast("Any", group[d["path"]]).shape[i]) for i in zyx) for d in datasets) + scales = tuple( + tuple(float(next(t for t in d["coordinateTransformations"] if t["type"] == "scale")["scale"][i]) for i in zyx) + for d in datasets + ) + transforms0 = {t["type"]: t for t in datasets[0]["coordinateTransformations"]} + return _MosaicGrid( + shapes=shapes, # type: ignore[arg-type] + scale=scales[0], # type: ignore[arg-type] + translation=tuple(float(transforms0["translation"]["translation"][i]) for i in zyx), # type: ignore[arg-type] + scales=scales, # type: ignore[arg-type] + ) + + +def _get_image(path: Path) -> DataTree: + """Load every level of an OME-Zarr (OME-NGFF v0.5) mosaic, from a directory or a zip, as a multiscale image. + + The pyramid levels already in the store are opened lazily, not recomputed. As in spatialdata's + own OME-Zarr reader, the ``global`` transformation comes from the full-resolution level and + each coarser level is related to it by the ratio of the array shapes. + """ + group = _open_mosaic(path) + ome = cast("dict[str, Any]", group.attrs.asdict()["ome"]) + multiscale = ome["multiscales"][0] + datasets = multiscale["datasets"] + + # drop the singleton "t" axis, which spatialdata's image models don't model + all_axes = [a["name"] for a in multiscale["axes"]] + t_index = all_axes.index("t") + axes = tuple(a for a in all_axes if a != "t") + arrays = [da.squeeze(da.from_zarr(group[d["path"]]), axis=t_index) for d in datasets] + + transformation = _mosaic_grid(path).transformation + spatial_axes = tuple(a for a in axes if a != "c") + + n_channels = arrays[0].shape[axes.index("c")] + channel_labels = [c.get("label") for c in ome.get("omero", {}).get("channels", [])] + c_coords = channel_labels if len(channel_labels) == n_channels else list(range(n_channels)) + + levels = {} + for i, array in enumerate(arrays): + # coordinates of every level are pixel centres in scale0 pixel units, as spatialdata assigns them + coords: dict[str, Any] = {"c": c_coords} + for ax in spatial_axes: + n0, n = arrays[0].shape[axes.index(ax)], array.shape[axes.index(ax)] + coords[ax] = np.linspace(0, n0, n + 1)[:-1] + n0 / n / 2 + levels[f"scale{i}"] = Dataset({"image": DataArray(array, dims=axes, coords=coords)}) + image = DataTree.from_dict(levels) + set_transformation(image, {"global": transformation}, set_all=True) + Image3DModel.validate(image) + return image + + +# --- 3D cell labels, rasterized lazily from the segmentation polygons ------------------------------ +# The polygons (one per cell per z-plane, in pixel units) are drawn onto the mosaic image's level-0 +# voxel grid, so labels and image overlay voxel for voxel at level 0. Ring decoding (``_read_rings``) +# is eager; drawing (``_draw`` / ``_rasterize_tile``) is one ``dask.delayed`` task per tile and runs +# only when the labels are computed or written. See ``_get_labels`` for how the coarser pyramid +# levels are obtained without redrawing. + +# PIL draws labels into signed 32-bit ("I" mode) images +_MAX_LABEL = 2**31 - 1 + + +def _label_ids(cell_ids: pd.Index) -> tuple[np.ndarray, str]: + """Integer label per cell, and the rule used. + + The trailing integer of each ``cell_id`` (after any prefix, e.g. ``Region_17`` -> 17) when every + cell has one and they are unique, > 0 (0 is background) and < 2^31; otherwise 1..n in table order. + """ + n = len(cell_ids) + digits = pd.Series(np.asarray(cell_ids, dtype=str)).str.extract(r"(\d+)$", expand=False) + if digits.isna().any(): + why = "no trailing integer in some cell_id" + elif (digits.str.len() > 10).any() or (numbers := digits.astype(np.int64)).max() > _MAX_LABEL: + why = "a trailing integer is not below 2^31" + elif numbers.min() < 1: + why = "a trailing integer is 0 (the background)" + elif not numbers.is_unique: + why = "trailing integers are not unique" + else: + return numbers.to_numpy(dtype=np.uint32), "trailing integer of cell_id" + return np.arange(1, n + 1, dtype=np.uint32), f"1..n in table order ({why})" + + +@dataclass(frozen=True) +class _Rings: + """Polygon exterior rings in level-0 voxel index space (voxel centres at integer coordinates).""" + + label: np.ndarray + plane: np.ndarray + length: np.ndarray + coords: np.ndarray + bounds: np.ndarray + + def __len__(self) -> int: + return len(self.label) + + +# Rows decoded to shapely at once, per row group: bounds a worker's peak memory to one batch's +# geometries and their transformed/simplified copies, instead of a whole row group's (up to ~1.36M +# polygons for the largest Region measured), at a small cost in per-batch overhead. +_DECODE_BATCH_ROWS = 65_536 + + +def _rings_from_row_group( + path: Path, + row_group: int, + labels: pd.Series, + grid: _MosaicGrid, + xy_size: float, + z_size: float, + simplify: float, +) -> tuple[dict[str, np.ndarray], dict[str, int]]: + """One parquet row group's rings on the grid, and how many polygon parts were dropped and why. + + Streamed in ``_DECODE_BATCH_ROWS``-row batches, so only one batch's shapely geometries (and their + transformed/simplified copies) are held in memory at a time, not the whole row group's. + """ + sz, sy, sx = grid.scale + tz, ty, tx = grid.translation + nz = grid.shapes[0][0] + batches = pq.ParquetFile(path).iter_batches( + batch_size=_DECODE_BATCH_ROWS, + row_groups=[row_group], + columns=[PyxaKeys.CELL_ID.value, PyxaKeys.Z_INDEX.value, "geometry"], + ) + batch_arrays: list[dict[str, np.ndarray]] = [] + dropped = {"not in the table": 0, "off the mosaic's z range": 0, "empty": 0} + for batch in batches: + cell_id = pc.cast(batch.column(PyxaKeys.CELL_ID.value), "string").to_numpy(zero_copy_only=False) + row = labels.index.get_indexer(cell_id) + zindex = batch.column(PyxaKeys.Z_INDEX.value).to_numpy() + geoms = shapely.from_wkb(batch.column("geometry").to_numpy(zero_copy_only=False)) + poly_parts, part_of = shapely.get_parts(geoms, return_index=True) + rings = shapely.get_exterior_ring(poly_parts) + rings = shapely.transform( + rings, + lambda c: np.column_stack(((c[:, 0] * xy_size - tx) / sx, (c[:, 1] * xy_size - ty) / sy)), + ) + rings = shapely.simplify(rings, simplify) + # plane k holds Z_um = (ZIndex + 0.5) * z_size, the reader's plane centre + plane = np.rint(((zindex[part_of] + 0.5) * z_size - tz) / sz).astype(np.int32) + in_table = row[part_of] >= 0 + on_grid = (plane >= 0) & (plane < nz) + empty = shapely.is_empty(rings) | (shapely.get_num_coordinates(rings) < 3) + keep = in_table & on_grid & ~empty + rings = rings[keep] + coords, ring_of = shapely.get_coordinates(rings, return_index=True) + batch_arrays.append( + { + "label": labels.to_numpy(dtype=np.uint32)[row[part_of][keep]], + "plane": plane[keep], + "length": np.bincount(ring_of, minlength=len(rings)).astype(np.int64), + "coords": coords.astype(np.float32), + "bounds": shapely.bounds(rings).astype(np.float32).reshape(-1, 4), + } + ) + dropped["not in the table"] += int((~in_table).sum()) + dropped["off the mosaic's z range"] += int((in_table & ~on_grid).sum()) + dropped["empty"] += int((in_table & on_grid & empty).sum()) + # drop this batch's shapely/numpy arrays before decoding the next one + del geoms, poly_parts, part_of, rings, plane, in_table, on_grid, empty, keep, coords, ring_of + if batch_arrays: + out = {k: np.concatenate([b[k] for b in batch_arrays]) for k in batch_arrays[0]} + else: + # the row group's batches kept nothing (every polygon filtered out, or no batches at all) + out = { + "label": np.empty(0, dtype=np.uint32), + "plane": np.empty(0, dtype=np.int32), + "length": np.empty(0, dtype=np.int64), + "coords": np.empty((0, 2), dtype=np.float32), + "bounds": np.empty((0, 4), dtype=np.float32), + } + return out, dropped + + +def _read_rings( + path: Path, + labels: pd.Series, + grid: _MosaicGrid, + xy_size: float, + z_size: float, + *, + simplify: float = 0.25, +) -> _Rings: + """Every segmentation polygon's exterior ring on the mosaic's level-0 grid, with its label and plane. + + ``labels`` maps ``cell_id`` to the cell's label. Row groups are decoded in worker processes, one + task per row group across up to the CPU count of workers (``joblib``, preferring processes, so + ``joblib.parallel_config(backend=...)`` can choose another backend); the default ``loky`` workers + are shut down once decoding ends. Although pyarrow and shapely release the GIL, decoding a full + Region's row groups concurrently on threads was measured to contend badly on the allocator (one + process, many threads each doing large shapely/numpy malloc/free traffic) rather than the GIL, + taking ~24x longer wall time and ~3x the peak memory of the same work split across processes. + A single row group runs directly, with no pool. Rings are simplified to ``simplify`` voxels: the + polygons are traced on a finer pixel grid than the mosaic's, and the staircase vertices add + nothing at its resolution. Holes are ignored (exteriors are filled). + """ + n_groups = pq.ParquetFile(path).metadata.num_row_groups + if n_groups <= 1: + results = [_rings_from_row_group(path, i, labels, grid, xy_size, z_size, simplify) for i in range(n_groups)] + else: + n_jobs = min(n_groups, joblib.cpu_count()) + backend, _ = joblib.parallel.get_active_backend(prefer="processes") + results = joblib.Parallel(n_jobs=n_jobs, prefer="processes")( + joblib.delayed(_rings_from_row_group)(path, i, labels, grid, xy_size, z_size, simplify) + for i in range(n_groups) + ) + if n_jobs > 1 and isinstance(backend, joblib.parallel.BACKENDS["loky"]): + # loky keeps its workers (hundreds of MB each after a decode) alive for reuse; nothing else + # here needs them, so release them now rather than holding that memory through the write + get_reusable_executor(reuse=True).shutdown(wait=True) + parts = [r[0] for r in results] + dropped: dict[str, int] = {} + for _, d in results: + for why, count in d.items(): + dropped[why] = dropped.get(why, 0) + count + if parts: + rings = _Rings(**{k: np.concatenate([p[k] for p in parts]) for k in parts[0]}) + else: + # no row groups at all (e.g. a crop with no cells): nothing to concatenate over + rings = _Rings( + label=np.empty(0, dtype=np.uint32), + plane=np.empty(0, dtype=np.int32), + length=np.empty(0, dtype=np.int64), + coords=np.empty((0, 2), dtype=np.float32), + bounds=np.empty((0, 4), dtype=np.float32), + ) + summary = ", ".join(f"{count} {why}" for why, count in dropped.items() if count) + logger.info( + f"{path.name}: {len(rings)} polygon rings on the mosaic grid" + (f"; dropped {summary}" if summary else "") + ) + return rings + + +# one drawing task: a whole number of storage chunks, so tiles never share a chunk +TILE = (32, 1024, 1024) +# storage chunks, as the mosaic's: an inspect window reads only the chunks it covers +CHUNKS = (32, 256, 256) + + +@dataclass(frozen=True) +class _Tile: + """The rings to draw into one tile of one pyramid level.""" + + origin: tuple[int, int, int] + shape: tuple[int, int, int] + step: tuple[int, int] + label: np.ndarray + plane: np.ndarray + offsets: np.ndarray + coords: np.ndarray + + +def _ragged_gather(starts: np.ndarray, lengths: np.ndarray) -> np.ndarray: + """Indices of the concatenated slices ``[s, s + n)`` for each (s, n), in order.""" + new_starts = np.cumsum(lengths) - lengths + return np.arange(int(lengths.sum())) - np.repeat(new_starts, lengths) + np.repeat(starts, lengths) + + +def _plan_tiles( + rings: _Rings, + shape: tuple[int, int, int], + step: tuple[int, int, int], + tile: tuple[int, int, int] = TILE, +) -> list[_Tile]: + """Group the rings of one level (``shape``, ``step`` from level 0) by tile. + + The level keeps the planes a nearest-neighbour stride of level 0 keeps; a ring crossing tiles goes + to each. Within a tile, rings are ordered by label, so where they overlap the higher label wins. + """ + nz, ny, nx = shape + dz, dy, dx = step + tz, ty, tx = tile + keep = np.flatnonzero(rings.plane % dz == 0) + plane = rings.plane[keep] // dz + on = plane < nz + keep, plane = keep[on], plane[on] + # bounds in level-voxel index space. Drawing snaps each vertex to its nearest voxel, so a ring fills + # voxels floor(lo + 0.5)..floor(hi + 0.5): its high side can reach the next tile's first voxel. + # Tiles are assigned from bounds padded by one voxel on the high side (the low side needs none); a + # tile that the ring does not actually reach draws nothing for it. + bounds = rings.bounds[keep] / np.array([dx, dy, dx, dy], dtype=np.float32) + n_ty, n_tx = -(-ny // ty), -(-nx // tx) + y_lo = np.clip(np.floor(bounds[:, 1] / ty), 0, n_ty - 1).astype(np.int64) + y_hi = np.clip(np.floor((bounds[:, 3] + 1) / ty), 0, n_ty - 1).astype(np.int64) + x_lo = np.clip(np.floor(bounds[:, 0] / tx), 0, n_tx - 1).astype(np.int64) + x_hi = np.clip(np.floor((bounds[:, 2] + 1) / tx), 0, n_tx - 1).astype(np.int64) + idx, keys = [], [] + for oy in range(int((y_hi - y_lo).max(initial=0)) + 1): + for ox in range(int((x_hi - x_lo).max(initial=0)) + 1): + hit = np.flatnonzero((y_lo + oy <= y_hi) & (x_lo + ox <= x_hi)) + idx.append(hit) + keys.append(((plane[hit] // tz) * n_ty + y_lo[hit] + oy) * n_tx + x_lo[hit] + ox) + idx_arr, keys_arr = np.concatenate(idx), np.concatenate(keys) + if keys_arr.size == 0: + # no ring falls in this level at all (e.g. an empty _Rings): nothing to draw, no tiles + return [] + order = np.lexsort((rings.label[keep][idx_arr], keys_arr)) + idx_arr, keys_arr = idx_arr[order], keys_arr[order] + + ring = keep[idx_arr] + starts = np.cumsum(rings.length) - rings.length + lengths = rings.length[ring] + coords = rings.coords[_ragged_gather(starts[ring], lengths)] + offsets = np.concatenate([[0], np.cumsum(lengths)]) + tiles = [] + edges = np.flatnonzero(np.diff(keys_arr)) + 1 + for lo, hi in zip(np.concatenate([[0], edges]), np.concatenate([edges, [len(keys_arr)]]), strict=True): + kz, rest = divmod(int(keys_arr[lo]), n_ty * n_tx) + ky, kx = divmod(rest, n_tx) + origin = (kz * tz, ky * ty, kx * tx) + c0, c1 = int(offsets[lo]), int(offsets[hi]) + tiles.append( + _Tile( + origin=origin, + shape=(min(tz, nz - origin[0]), min(ty, ny - origin[1]), min(tx, nx - origin[2])), + step=(dy, dx), + label=rings.label[ring[lo:hi]], + plane=plane[idx_arr[lo:hi]] - origin[0], + offsets=offsets[lo : hi + 1] - c0, + coords=coords[c0:c1], + ) + ) + return tiles + + +def _rasterize_tile(tile: _Tile) -> np.ndarray: + """Fill the tile's rings plane by plane into a ``uint32`` block (0 = background).""" + from PIL import Image, ImageDraw + + depth, height, width = tile.shape + block = np.zeros(tile.shape, dtype=np.uint32) + dy, dx = tile.step + _, y0, x0 = tile.origin + # Snap each vertex to the voxel whose centre is nearest (PIL fills the pixels of an integer polygon, + # vertices included), in level voxel index space, so the snap does not depend on the tile. PIL draws + # a polygon differently when some vertices are negative (it truncates toward zero and clips), so the + # canvas starts at the lowest vertex, not the tile's origin, and is cropped to the tile after drawing: + # a tile then draws exactly what one whole-level draw would, and neighbouring tiles join seamlessly. + xs = np.floor(tile.coords[:, 0] / dx + 0.5).astype(np.int64) + ys = np.floor(tile.coords[:, 1] / dy + 0.5).astype(np.int64) + cx, cy = min(x0, int(xs.min(initial=x0))), min(y0, int(ys.min(initial=y0))) + xy = np.column_stack((xs - cx, ys - cy)) + for plane in np.unique(tile.plane): + image = Image.new("I", (x0 + width - cx, y0 + height - cy)) + draw = ImageDraw.Draw(image) + for r in np.flatnonzero(tile.plane == plane): + draw.polygon(xy[tile.offsets[r] : tile.offsets[r + 1]].ravel().tolist(), fill=int(tile.label[r])) + block[plane] = np.asarray(image, dtype=np.int32)[y0 - cy :, x0 - cx :] + return block + + +def _labels_level( + rings: _Rings, + shape: tuple[int, int, int], + step: tuple[int, int, int], + *, + tile: tuple[int, int, int] = TILE, + chunks: tuple[int, int, int] = CHUNKS, +) -> da.Array: + """One pyramid level as a lazy array: a ``dask.delayed`` drawing task per tile, zeros where no ring falls.""" + return _tiles_array(_plan_tiles(rings, shape, step, tile), shape, tile=tile, chunks=chunks) + + +def _tiles_array( + tiles: list[_Tile], + shape: tuple[int, int, int], + *, + tile: tuple[int, int, int] = TILE, + chunks: tuple[int, int, int] = CHUNKS, +) -> da.Array: + """A level's planned tiles as a lazy array, zeros where no tile is planned.""" + by_origin = {t.origin: t for t in tiles} + nz, ny, nx = shape + tz, ty, tx = tile + blocks = [] + for z in range(0, nz, tz): + rows = [] + for y in range(0, ny, ty): + row = [] + for x in range(0, nx, tx): + size = (min(tz, nz - z), min(ty, ny - y), min(tx, nx - x)) + t = by_origin.get((z, y, x)) + if t is None: + row.append(da.zeros(size, dtype=np.uint32, chunks=size)) + else: + # looked up at call time, so the drawing function can be patched in tests + row.append(da.from_delayed(dask.delayed(_draw)(t), shape=size, dtype=np.uint32)) + rows.append(row) + blocks.append(rows) + return da.block(blocks).rechunk(tuple(min(c, s) for c, s in zip(chunks, shape, strict=True))) + + +def _draw(tile: _Tile) -> np.ndarray: + return _rasterize_tile(tile) + + +def _get_labels(rings: _Rings, grid: _MosaicGrid) -> DataTree: + """The cells as a lazy multiscale ``Labels3DModel`` on the mosaic's grid, one level per mosaic level. + + Level 0 is drawn from the rings, one ``dask.delayed`` tile at a time. Every coarser level is a + nearest-neighbour strided *view* of the level-0 array (matching ``grid.step``), not redrawn + independently, so a level-0 tile's drawing task is shared by every level that needs it: computing or + writing the whole tree draws each level-0 tile at most once. The stride starts at each block's centre + (offset ``step // 2`` per axis, less where that would leave the level short of the mosaic's shape): + the mosaic's own pyramid is a smoothed block average, so centre samples track it best. ``grid.step`` + is the OME-NGFF scale ratio to level 0 (not the array-shape ratio, which a cropped or odd-sized level + can round to the wrong integer). + """ + n0 = grid.shapes[0] + tiles = _plan_tiles(rings, n0, (1, 1, 1)) + level0 = _tiles_array(tiles, n0) + levels = {} + for i, shape in enumerate(grid.shapes): + if i == 0: + array = level0 + else: + step = grid.step(i) + # each coarse voxel takes the level-0 voxel at its block's centre, as the mosaic's pyramid + # (a smoothed block average) centres it there; the offset shrinks where it would run off the end + oz, oy, ox = (max(0, min(d // 2, a - 1 - (b - 1) * d)) for a, b, d in zip(n0, shape, step, strict=True)) + dz, dy, dx = step + strided = level0[oz::dz, oy::dy, ox::dx] + if any(a < b for a, b in zip(strided.shape, shape, strict=True)): + raise ValueError( + f"scale{i}: level-0 stride {step} gives shape {strided.shape}, " + f"shorter than the mosaic's {shape} on some axis" + ) + sz, sy, sx = shape + array = strided[:sz, :sy, :sx].rechunk(tuple(min(c, s) for c, s in zip(CHUNKS, shape, strict=True))) + # coordinates of every level are pixel centres in scale0 pixel units, as spatialdata assigns them + coords = {ax: np.linspace(0, a, b + 1)[:-1] + a / b / 2 for ax, a, b in zip("zyx", n0, shape, strict=True)} + levels[f"scale{i}"] = Dataset({"image": DataArray(array, dims=("z", "y", "x"), coords=coords)}) + tree = DataTree.from_dict(levels) + set_transformation(tree, {"global": grid.transformation}, set_all=True) + Labels3DModel.validate(tree) + logger.info( + f"{PyxaKeys.CELL_LABELS.value}: {len(rings)} rings in {len(tiles)} level-0 tiles, " + f"{len(grid.shapes)} levels planned; drawn when computed or written" + ) + return tree + + +InputPath = str | Path | bool | None + +RequiredPath = str | Path | None + + +def _required(value: RequiredPath) -> str | Path | bool: + """A required file is never skipped; unset means it must be in ``path``.""" + if isinstance(value, bool): + raise TypeError("cell_by_gene and cell_metadata are required; pass a path or leave them unset") + return True if value is None else value + + +def _resolve_input(value: InputPath, path: Path | None, file_name: str) -> Path | None: + """Where to read one Pyxa file from, or ``None`` to skip it. + + ``False`` skips the file. ``None`` reads ``path / file_name`` when that exists and skips it + otherwise; ``True`` requires it. A path reads that file, which must exist. + """ + if value is False: + return None + if value is None or value is True: + candidate = path / file_name if path is not None else None + if candidate is not None and candidate.exists(): + return candidate + if value is True: + raise FileNotFoundError(f"Expected Pyxa output file not found: {candidate or file_name}") + return None + explicit = Path(value) + if not explicit.exists(): + raise FileNotFoundError(f"Expected Pyxa output file not found: {explicit}") + return explicit + + +def _resolve_image(value: InputPath, path: Path | None) -> Path | None: + """Where to read the mosaic from, or ``None`` to skip it. + + Like :func:`_resolve_input`, looking in ``path`` for the unzipped ``mosaic_3d.ome.zarr`` first + and then for ``mosaic_3d.ome.zarr.zip``. An explicit path may be either. + """ + if value is False: + return None + if value is None or value is True: + for name in (PyxaKeys.MOSAIC_FILE.value, PyxaKeys.MOSAIC_ZIP_FILE.value): + if path is not None and (path / name).exists(): + return path / name + if value is True: + raise FileNotFoundError(f"Expected Pyxa mosaic image not found: {PyxaKeys.MOSAIC_FILE.value}(.zip)") + return None + explicit = Path(value) + if not explicit.exists(): + raise FileNotFoundError(f"Expected Pyxa mosaic image not found: {explicit}") + return explicit + + +@inject_docs(px=PyxaKeys) +def pyxa( + path: str | Path | None = None, + dataset_id: str = "pyxa", + *, + cell_by_gene: RequiredPath = None, + cell_metadata: RequiredPath = None, + cell_assigned_gene: InputPath = None, + segmentation_geometries: InputPath = None, + pyxa_studio: InputPath = None, + image: InputPath = None, + shapes: bool | None = None, + labels: bool = False, +) -> SpatialData: + """ + Read *Pyxa* (Stellaromics) output. + + The ``rna`` table is always read, from two required files: + + - ``{px.CELL_BY_GENE_FILE!r}``: Per-cell gene expression counts, the table's ``X``. + - ``{px.CELL_METADATA_FILE!r}``: Per-cell metadata (volume, spatial coordinates), the + table's ``obs`` and ``obsm["spatial"]``. + + Everything else is optional, each read when present: + + - ``{px.CELL_ASSIGNED_GENE_FILE!r}``: Transcript-level gene assignments, as the + ``transcripts`` points. + - ``{px.SEGMENTATION_GEOMETRIES_FILE!r}``: Per-cell segmentation polygons, as shapes, or + with ``labels=True`` as 3D cell labels. With ``labels=False`` the table annotates the + cell footprints only when these are read. + - ``{px.PYXA_STUDIO_FILE!r}``: Pyxa Studio's export of the cells that passed its filters, + adding ``{px.CLUSTER!r}`` (categorical) to the table's ``obs`` and the 3D UMAP as + ``obsm[{px.UMAP_KEY!r}]``. Cells it filtered out keep missing values there. + - ``{px.MOSAIC_FILE!r}`` (or ``{px.MOSAIC_ZIP_FILE!r}``, read in place): the mosaic + OME-Zarr (OME-NGFF v0.5) image, all pyramid levels, as ``{px.MOSAIC_IMAGE!r}``. + + Files are looked up in ``path`` by default. Pass a path for a file to read it from + elsewhere, and for an optional file ``False`` to skip it even if present (e.g. the + transcripts of a large Region) or ``True`` to require it. + + No public specification exists for this format at the time of writing; this + reader is validated against the public demo dataset at + https://huggingface.co/datasets/Stellaromics/demo. + + All elements are returned in micrometers in the ``global`` coordinate + system. Segmentation polygons are stored on disk in pixel units, one polygon + per cell per z-plane (``ZIndex``); the reader converts them to micrometers + (repairing any polygon that the conversion makes invalid) and adds their z + as a ``Z_um`` column (the centre of the z-plane), since shapes are 2D in + spatialdata. Both voxel sizes are inferred from ``{px.CELL_METADATA_FILE!r}``, + whose per-cell centroids are the area-weighted centroids of each cell's + polygons in both units. + + As shapes (by default only with ``labels=False``; see ``shapes``), the polygons are two elements: + + - ``{px.REGION!r}``: one 2D footprint per cell (the union of its z-plane + polygons), indexed by ``cell_id`` and, with ``labels=False``, annotated by the ``rna`` table. + Cells stacked in z have overlapping footprints, so use these for 2D + display and table annotation, not for 2D spatial aggregation (the + transcripts' ``cell_id`` already gives each transcript's cell). + - ``{px.CELL_BOUNDARIES_Z!r}``: the per-cell, per-z-plane polygons, with + ``cell_id``, ``ZIndex`` and ``Z_um`` columns. + + With ``labels=True`` the polygons are instead rasterized into ``{px.CELL_LABELS!r}``, a 3D + labels element on the mosaic's voxel grid (same pyramid levels and transformation as + ``{px.MOSAIC_IMAGE!r}``), and the table annotates it through the integer + ``{px.LABEL_ID!r}``: the trailing integer of each ``cell_id`` (``Region_17`` -> 17) when + those are unique, positive and below 2^31, otherwise 1..n in table order. Holes are + filled, and where two cells overlap on a plane the higher label wins. The labels are + lazy: level 0 is drawn, one task per 32 x 1024 x 1024 tile, when computed or written, and + the coarser levels are strided views of it (nearest neighbour, sampled at each coarse + voxel's block centre), so writing draws each tile once, with dask's default threaded + scheduler (a process scheduler is much slower here, since every drawn tile is pickled + back). The lazy labels' dask graph holds every ring's coordinates (GBs for a full Region) + for as long as the element is alive, and would ship them all to a distributed scheduler. + Decoding the polygons is eager, in worker processes: one task per parquet row group, + across up to the CPU count of workers (threads contend badly on the allocator with this + many large shapely/numpy arrays; ``joblib.parallel_config(backend=...)`` can choose another + backend), each row group streamed in small batches so a worker never holds more than one + batch's geometries at once. That takes about a minute and tens of GB for a full Region of + ~20M polygons. + + Unassigned transcripts (``cell_id`` ending in ``"_-1"``) are kept in the + points element, flagged via an ``assigned`` column, rather than dropped. + + Parameters + ---------- + path + Directory holding Pyxa's output files. ``None`` reads only the files given + explicitly, in which case ``cell_by_gene`` and ``cell_metadata`` must be. + dataset_id + Dataset identifier, currently unused for element naming (reserved for + future multi-sample support). + cell_by_gene, cell_metadata + Required files: ``None`` (default) reads them from ``path``, a path reads that file. + cell_assigned_gene, segmentation_geometries, pyxa_studio, image + Optional files: ``None`` (default) reads one from ``path`` if present, a path reads + that file, ``False`` skips it and ``True`` requires it in ``path``. + + ``image`` may be the mosaic's directory or a zip of it; a zipped mosaic is read in place. + Compute a zipped mosaic with dask's threaded scheduler (the default): a process scheduler + pickles its ``ZipStore``, which reopens the zip in every worker. + shapes + Return the polygons as shapes. ``None`` (default): when the segmentation geometries are + read and ``labels`` is ``False``. + labels + Rasterize the polygons into 3D cell labels on the mosaic's grid (needs the segmentation + geometries and the mosaic); the table then annotates the labels. + + Returns + ------- + :class:`spatialdata.SpatialData` + """ + directory = Path(path) if path is not None else None + if directory is not None and not directory.is_dir(): + raise FileNotFoundError(f"Pyxa output directory not found: {directory}") + by_gene_path = _resolve_input(_required(cell_by_gene), directory, PyxaKeys.CELL_BY_GENE_FILE.value) + metadata_path = _resolve_input(_required(cell_metadata), directory, PyxaKeys.CELL_METADATA_FILE.value) + assigned_gene_path = _resolve_input(cell_assigned_gene, directory, PyxaKeys.CELL_ASSIGNED_GENE_FILE.value) + geometries_path = _resolve_input(segmentation_geometries, directory, PyxaKeys.SEGMENTATION_GEOMETRIES_FILE.value) + studio_path = _resolve_input(pyxa_studio, directory, PyxaKeys.PYXA_STUDIO_FILE.value) + if by_gene_path is None or metadata_path is None: # unreachable: required inputs resolve or raise + raise FileNotFoundError("cell_by_gene and cell_metadata are required") + + image_source = _resolve_image(image, directory) + inputs = [p for p in (by_gene_path, metadata_path, assigned_gene_path, geometries_path, studio_path) if p] + logger.info(f"Reading Pyxa {', '.join(p.name for p in inputs)}") + + if labels: + missing = [ + name + for name, found in (("segmentation_geometries", geometries_path), ("a mosaic image", image_source)) + if found is None + ] + if missing: + raise ValueError( + f"labels=True needs segmentation_geometries and a mosaic image; missing: {', '.join(missing)}" + ) + if shapes and geometries_path is None: + raise FileNotFoundError( + f"Expected Pyxa output file not found: {PyxaKeys.SEGMENTATION_GEOMETRIES_FILE.value} (shapes=True)" + ) + read_shapes = geometries_path is not None and (shapes if shapes is not None else not labels) + + points = {} + if assigned_gene_path is not None: + points["transcripts"] = PointsModel.parse( + _get_points(assigned_gene_path), + coordinates={"x": PyxaKeys.X_UM.value, "y": PyxaKeys.Y_UM.value, "z": PyxaKeys.Z_UM.value}, + feature_key=PyxaKeys.GENE.value, + instance_key=PyxaKeys.CELL_ID.value, + ) + + xy_size, z_size = _get_voxel_size(metadata_path) if (read_shapes or labels) else (1.0, 1.0) + shapes_elements = {} + if read_shapes: + planes = _get_shapes(geometries_path, xy_size, z_size) # type: ignore[arg-type] + shapes_elements[PyxaKeys.REGION.value] = ShapesModel.parse(_get_footprints(planes)) + shapes_elements[PyxaKeys.CELL_BOUNDARIES_Z.value] = ShapesModel.parse(planes) + + adata = _get_table(by_gene_path, metadata_path, studio_path) + labels_elements = {} + if labels: + ids, rule = _label_ids(adata.obs_names) + logger.info(f"{PyxaKeys.LABEL_ID.value}: {rule}") + adata.obs[PyxaKeys.LABEL_ID.value] = ids + adata.obs[PyxaKeys.REGION_KEY.value] = pd.Series( + PyxaKeys.CELL_LABELS.value, index=adata.obs_names, dtype="category" + ) + grid = _mosaic_grid(image_source) # type: ignore[arg-type] + rings = _read_rings( + geometries_path, # type: ignore[arg-type] + pd.Series(ids, index=adata.obs_names), + grid, + xy_size, + z_size, + ) + labels_elements[PyxaKeys.CELL_LABELS.value] = _get_labels(rings, grid) + table = TableModel.parse( + adata, + region=PyxaKeys.CELL_LABELS.value, + region_key=PyxaKeys.REGION_KEY.value, + instance_key=PyxaKeys.LABEL_ID.value, + ) + elif shapes_elements: + table = TableModel.parse( + adata, + region=PyxaKeys.REGION.value, + region_key=PyxaKeys.REGION_KEY.value, + instance_key=PyxaKeys.INSTANCE_KEY.value, + ) + else: + table = TableModel.parse(adata) + + images = {} + if image_source is not None: + images[PyxaKeys.MOSAIC_IMAGE.value] = _get_image(image_source) + + return SpatialData( + points=points, shapes=shapes_elements, labels=labels_elements, tables={"rna": table}, images=images + ) diff --git a/tests/test_pyxa.py b/tests/test_pyxa.py new file mode 100644 index 00000000..27583daf --- /dev/null +++ b/tests/test_pyxa.py @@ -0,0 +1,1246 @@ +import dataclasses +import math +import os +import shutil +import tempfile +import urllib.request +import uuid +import zipfile +from pathlib import Path +from tempfile import TemporaryDirectory +from typing import Any, cast + +import dask.dataframe as dd +import geopandas as gpd +import numpy as np +import pandas as pd +import pyarrow.parquet as pq +import pytest +import shapely +import zarr +from click.testing import CliRunner +from spatialdata import get_extent, match_element_to_table, match_table_to_element, read_zarr +from spatialdata.models import get_table_keys +from spatialdata.transformations import Identity, Scale, Sequence, Translation, get_transformation +from xarray import DataTree + +from spatialdata_io.__main__ import pyxa_wrapper +from spatialdata_io._constants._constants import PyxaKeys +from spatialdata_io.readers import pyxa as pyxa_module +from spatialdata_io.readers.pyxa import ( + _get_footprints, + _get_image, + _get_labels, + _get_points, + _get_shapes, + _get_table, + _get_voxel_size, + _label_ids, + _labels_level, + _make_polygonal_valid, + _mosaic_grid, + _MosaicGrid, + _plan_tiles, + _rasterize_tile, + _read_rings, + _Rings, + _validate_columns, + pyxa, +) + +# See https://github.com/scverse/spatialdata-io/blob/main/.github/workflows/prepare_test_data.yaml for instructions on +# how to download and place the data on disk +DATASETS = ["pyxa_xsmall"] +FIXTURE_DIR = Path("./data") / DATASETS[0] +MOSAIC_DIR = FIXTURE_DIR / "mosaic_3d.ome.zarr" + + +def _ensure_pyxa_xsmall_fixture() -> None: + """Download the xsmall Pyxa fixture if the CI test-data artifact doesn't have it yet. + + Mirrors the `pyxa_xsmall` step of `.github/workflows/prepare_test_data.yaml`. Downloads and + extracts into a uniquely named temp directory, then swaps it into place, so this is safe when + several `pytest -n auto` workers import this module at once. Remove once the shared CI artifact + includes `pyxa_xsmall`. + """ + if FIXTURE_DIR.exists(): + return + base_url = "https://huggingface.co/datasets/Stellaromics/demo/resolve/main/xsmall/" + files = [ + "cell_assigned_gene_v1.csv", + "cell_by_gene_v1.csv", + "cell_metadata_v1.csv", + "segmentation_geometries_v1.parquet", + "mosaic_3d.ome.zarr.zip", + ] + tmp_dir = FIXTURE_DIR.parent / f".{FIXTURE_DIR.name}.tmp-{uuid.uuid4().hex}" + try: + tmp_dir.mkdir(parents=True) + for name in files: + with urllib.request.urlopen(base_url + name, timeout=60) as response, open(tmp_dir / name, "wb") as f: + shutil.copyfileobj(response, f) + zip_path = tmp_dir / "mosaic_3d.ome.zarr.zip" + with zipfile.ZipFile(zip_path) as zf: + zf.extractall(tmp_dir) + zip_path.unlink() + os.replace(tmp_dir, FIXTURE_DIR) + except OSError as err: + shutil.rmtree(tmp_dir, ignore_errors=True) + if FIXTURE_DIR.exists(): + return # another worker downloaded it first + pytest.skip(f"Could not download the pyxa_xsmall fixture: {err}", allow_module_level=True) + + +_ensure_pyxa_xsmall_fixture() + + +TINY_SCALE0 = np.arange(2 * 4 * 4, dtype="uint8").reshape(1, 1, 2, 4, 4) +# a constant rather than a downsampling of scale0, so a test can tell a loaded level from a recomputed one +TINY_SCALE1 = np.full((1, 1, 1, 2, 2), 7, dtype="uint8") + + +def _make_tiny_ome_zarr(path: Path) -> None: + """Build a minimal two-level OME-NGFF v0.5 store, shapes (t=1, c=1, z=2, y=4, x=4) and (1, 1, 1, 2, 2).""" + group = zarr.open_group(store=str(path), mode="w") + for name, data in (("scale0/image", TINY_SCALE0), ("scale1/image", TINY_SCALE1)): + array = group.create_array(name, shape=data.shape, dtype=data.dtype, dimension_names=["t", "c", "z", "y", "x"]) + array[:] = data + group.attrs["ome"] = { + "version": "0.5", + "multiscales": [ + { + "axes": [ + {"name": "t", "type": "time"}, + {"name": "c", "type": "channel"}, + {"name": "z", "type": "space"}, + {"name": "y", "type": "space"}, + {"name": "x", "type": "space"}, + ], + "datasets": [ + { + "path": "scale0/image", + "coordinateTransformations": [ + {"type": "scale", "scale": [1.0, 1.0, 0.5, 0.2, 0.2]}, + {"type": "translation", "translation": [0.0, 0.0, 1.0, 2.0, 3.0]}, + ], + }, + { + "path": "scale1/image", + "coordinateTransformations": [ + {"type": "scale", "scale": [1.0, 1.0, 1.0, 0.4, 0.4]}, + {"type": "translation", "translation": [0.0, 0.0, 1.25, 2.1, 3.1]}, + ], + }, + ], + "name": "image", + } + ], + "omero": {"channels": [{"label": "DAPI"}]}, + } + + +def test_pyxa_keys_filenames() -> None: + assert PyxaKeys.CELL_ASSIGNED_GENE_FILE == "cell_assigned_gene_v1.csv" + assert PyxaKeys.CELL_BY_GENE_FILE == "cell_by_gene_v1.csv" + assert PyxaKeys.CELL_METADATA_FILE == "cell_metadata_v1.csv" + assert PyxaKeys.SEGMENTATION_GEOMETRIES_FILE == "segmentation_geometries_v1.parquet" + assert PyxaKeys.PYXA_STUDIO_FILE == "pyxa_studio_v1.csv" + + +def test_pyxa_keys_columns() -> None: + assert PyxaKeys.CELL_ID == "cell_id" + assert PyxaKeys.GENE == "Gene" + assert PyxaKeys.X_UM == "X_um" + assert PyxaKeys.Y_UM == "Y_um" + assert PyxaKeys.Z_UM == "Z_um" + assert PyxaKeys.VOLUME_UM3 == "Volume_um3" + assert PyxaKeys.ROI == "ROI" + assert PyxaKeys.Z_INDEX == "ZIndex" + assert PyxaKeys.BORDER == "Border" + assert PyxaKeys.FOV == "FOV" + assert PyxaKeys.UNASSIGNED_SUFFIX == "_-1" + assert PyxaKeys.REGION_KEY == "region" + assert PyxaKeys.REGION == "cell_boundaries" + assert PyxaKeys.CELL_BOUNDARIES_Z == "cell_boundaries_z" + assert PyxaKeys.INSTANCE_KEY == "cell_id" + assert PyxaKeys.ASSIGNED == "assigned" + + +def test_validate_columns_passes_when_present() -> None: + df = pd.DataFrame({"cell_id": [1], "Gene": ["A"]}) + _validate_columns(df, {"cell_id", "Gene"}, "test_file.csv") + + +def test_validate_columns_raises_when_missing() -> None: + df = pd.DataFrame({"cell_id": [1]}) + with pytest.raises(ValueError, match=r"test_file\.csv is missing required column\(s\): \['Gene'\]"): + _validate_columns(df, {"cell_id", "Gene"}, "test_file.csv") + + +def test_get_points_keeps_unassigned_transcripts() -> None: + points = _get_points(FIXTURE_DIR / "cell_assigned_gene_v1.csv") + assert isinstance(points, dd.DataFrame) + computed = points.compute() + assert "assigned" in computed.columns + assert (~computed["assigned"]).sum() > 0 + assert computed[~computed["assigned"]]["cell_id"].str.endswith("_-1").all() + assert computed["assigned"].sum() > 0 + + +def test_get_points_has_required_coordinate_columns() -> None: + points = _get_points(FIXTURE_DIR / "cell_assigned_gene_v1.csv") + computed = points.compute() + for col in ("X_um", "Y_um", "Z_um", "Gene", "cell_id"): + assert col in computed.columns + + +def test_pyxa_reader_gene_categories_are_known(caplog: pytest.LogCaptureFixture) -> None: + points = pyxa(FIXTURE_DIR)["transcripts"] + # PointsModel.parse warns (and computes them itself) when the feature categories are unknown + assert "unknown categories" not in caplog.text + assert points["Gene"].cat.known + raw = pd.read_csv(FIXTURE_DIR / "cell_assigned_gene_v1.csv") + assert set(points["Gene"].cat.categories) == set(raw["Gene"]) + + +def test_get_table_matches_raw_values() -> None: + adata = _get_table( + FIXTURE_DIR / "cell_by_gene_v1.csv", + FIXTURE_DIR / "cell_metadata_v1.csv", + ) + raw_by_gene = pd.read_csv(FIXTURE_DIR / "cell_by_gene_v1.csv", index_col="cell_id") + raw_metadata = pd.read_csv(FIXTURE_DIR / "cell_metadata_v1.csv", index_col="cell_id") + + assert adata.n_obs == len(raw_by_gene) + sample_cell = raw_by_gene.index[0] + sample_gene = raw_by_gene.columns[0] + assert adata[sample_cell, sample_gene].to_df().iloc[0, 0] == raw_by_gene.loc[sample_cell, sample_gene] + + assert list(adata.obsm["spatial"][0]) == list(raw_metadata.loc[sample_cell, ["X_um", "Y_um", "Z_um"]]) + assert (adata.obs["region"] == "cell_boundaries").all() + # an index named like the cell_id column breaks spatialdata's table joins + assert adata.obs.index.name is None + + +def test_get_shapes_matches_raw_row_count() -> None: + gdf = _get_shapes(FIXTURE_DIR / "segmentation_geometries_v1.parquet", xy_size=0.114984751, z_size=0.5) + raw = gpd.read_parquet(FIXTURE_DIR / "segmentation_geometries_v1.parquet") + assert len(gdf) == len(raw) + assert all(isinstance(c, str) for c in gdf["cell_id"]) + assert gdf.geometry.is_valid.all() + assert set(gdf.geom_type) <= {"Polygon", "MultiPolygon"} + assert gdf.index.is_unique + + +def test_get_footprints_is_union_of_planes() -> None: + planes = _get_shapes(FIXTURE_DIR / "segmentation_geometries_v1.parquet", xy_size=0.114984751, z_size=0.5) + footprints = _get_footprints(planes) + assert footprints.index.name == "cell_id" + assert footprints.index.is_unique + assert set(footprints.index) == set(planes["cell_id"]) + assert footprints.geometry.is_valid.all() + # every z-plane polygon lies inside its cell's footprint + covered = footprints.loc[planes["cell_id"]].geometry.buffer(1e-9).covers(planes.geometry, align=False) + assert covered.all() + + +def test_get_shapes_converts_to_um() -> None: + gdf = _get_shapes(FIXTURE_DIR / "segmentation_geometries_v1.parquet", xy_size=0.114984751, z_size=0.5) + raw = gpd.read_parquet(FIXTURE_DIR / "segmentation_geometries_v1.parquet") + np.testing.assert_allclose(gdf.total_bounds, raw.total_bounds * 0.114984751) + np.testing.assert_allclose(gdf["Z_um"], (raw["ZIndex"] + 0.5) * 0.5) + + +def test_make_polygonal_valid_fixes_self_intersection() -> None: + bowtie = shapely.Polygon([(0, 0), (2, 2), (2, 0), (0, 2)]) + square = shapely.box(0, 0, 1, 1) + fixed = _make_polygonal_valid(np.array([bowtie, square])) + assert shapely.is_valid(fixed).all() + assert set(shapely.get_type_id(fixed)) <= {shapely.GeometryType.POLYGON, shapely.GeometryType.MULTIPOLYGON} + assert shapely.area(fixed[0]) == pytest.approx(2.0) + assert fixed[1] is square # valid geometries are passed through untouched + + +def test_make_polygonal_valid_drops_non_polygonal_parts() -> None: + # a polygon with a zero-width spike: make_valid returns the square plus a dangling line + spiky = shapely.Polygon([(0, 0), (1, 0), (1, 1), (1, 2), (1, 1), (0, 1)]) + (fixed,) = _make_polygonal_valid(np.array([spiky])) + assert fixed.is_valid + assert fixed.geom_type in {"Polygon", "MultiPolygon"} + assert fixed.area == pytest.approx(1.0) + + +def test_get_voxel_size() -> None: + xy, z = _get_voxel_size(FIXTURE_DIR / "cell_metadata_v1.csv") + assert xy == pytest.approx(0.114984751) + assert z == pytest.approx(0.5) + + +def _write_metadata(path: Path, n: int, xy: float, z: float) -> pd.DataFrame: + rng = np.random.default_rng(0) + px = pd.DataFrame(rng.uniform(1, 1000, size=(n, 3)), columns=["X_pixels", "Y_pixels", "Z_pixels"]) + metadata = px.assign(X_um=px["X_pixels"] * xy, Y_um=px["Y_pixels"] * xy, Z_um=px["Z_pixels"] * z) + metadata.to_csv(path, index=False) + return metadata + + +def test_get_voxel_size_reads_only_a_sample(tmp_path: Path) -> None: + path = tmp_path / "cell_metadata_v1.csv" + metadata = _write_metadata(path, n=5_000, xy=0.25, z=1.5) + # corrupt everything past the fitted sample: the fit must not read these rows + metadata.loc[2_000:, ["X_um", "Y_um", "Z_um"]] = 1e6 + metadata.to_csv(path, index=False) + assert _get_voxel_size(path, n_rows=1_000) == pytest.approx((0.25, 1.5)) + + +def test_get_voxel_size_raises_when_not_a_pure_scale(tmp_path: Path) -> None: + path = tmp_path / "cell_metadata_v1.csv" + metadata = _write_metadata(path, n=100, xy=0.25, z=1.5) + metadata["X_um"] += 10.0 # an offset between pixel and um coordinates + metadata.to_csv(path, index=False) + with pytest.raises(ValueError, match="not related by a pure scale"): + _get_voxel_size(path) + + +def _area_weighted_centroids(shapes: gpd.GeoDataFrame) -> pd.DataFrame: + """Per-cell centroid of the polygon stack, weighting each z-plane polygon by its area.""" + centroids = shapes.geometry.centroid + weighted = ( + pd.DataFrame( + { + "cell_id": shapes["cell_id"].to_numpy(), + "area": shapes.geometry.area.to_numpy(), + "x": (centroids.x * shapes.geometry.area).to_numpy(), + "y": (centroids.y * shapes.geometry.area).to_numpy(), + "z": (shapes["Z_um"] * shapes.geometry.area).to_numpy(), + } + ) + .groupby("cell_id")[["area", "x", "y", "z"]] + .sum() + ) + return weighted[["x", "y", "z"]].div(weighted["area"], axis=0) + + +def test_pyxa_reader_shapes_in_um_with_identity_transform() -> None: + sdata = pyxa(FIXTURE_DIR) + for name in ("cell_boundaries", "cell_boundaries_z"): + assert isinstance(get_transformation(sdata[name], to_coordinate_system="global"), Identity) + assert sdata[name].geometry.is_valid.all() + + +def test_pyxa_reader_shapes_aligned_with_cell_metadata() -> None: + shapes = pyxa(FIXTURE_DIR)["cell_boundaries_z"] + metadata = pd.read_csv(FIXTURE_DIR / "cell_metadata_v1.csv", index_col="cell_id") + # the per-cell metadata centroid is exactly the area-weighted centroid of the cell's polygon stack + centroids = _area_weighted_centroids(shapes) + expected = metadata.loc[centroids.index, ["X_um", "Y_um", "Z_um"]].to_numpy() + np.testing.assert_allclose(centroids.to_numpy(), expected, atol=1e-6) + + +def test_pyxa_reader_builds_valid_sdata() -> None: + sdata = pyxa(FIXTURE_DIR) + + assert "transcripts" in sdata.points + assert "cell_boundaries" in sdata.shapes + assert "cell_boundaries_z" in sdata.shapes + assert "rna" in sdata.tables + + raw_transcripts = pd.read_csv(FIXTURE_DIR / "cell_assigned_gene_v1.csv") + extent = get_extent(sdata["transcripts"]) + assert math.floor(extent["x"][0]) <= math.floor(raw_transcripts["X_um"].min()) + assert math.ceil(extent["x"][1]) >= math.ceil(raw_transcripts["X_um"].max()) + + +def test_pyxa_reader_missing_file_raises() -> None: + with tempfile.TemporaryDirectory() as tmpdir: + with pytest.raises(FileNotFoundError): + pyxa(Path(tmpdir)) + + +def test_get_image_loads_all_scales() -> None: + with tempfile.TemporaryDirectory() as tmpdir: + zarr_path = Path(tmpdir) / "tiny.ome.zarr" + _make_tiny_ome_zarr(zarr_path) + + image = _get_image(zarr_path) + assert isinstance(image, DataTree) + assert list(image.keys()) == ["scale0", "scale1"] + scale0, scale1 = image["scale0"]["image"], image["scale1"]["image"] + assert scale0.dims == ("c", "z", "y", "x") + assert list(scale0.coords["c"].values) == ["DAPI"] + # both levels are read from the store as written, not recomputed from scale0 + np.testing.assert_array_equal(scale0.values, TINY_SCALE0[0]) + np.testing.assert_array_equal(scale1.values, TINY_SCALE1[0]) + + expected = Sequence( + [Scale([0.5, 0.2, 0.2], axes=("z", "y", "x")), Translation([1.0, 2.0, 3.0], axes=("z", "y", "x"))] + ) + affine = get_transformation(image, to_coordinate_system="global").to_affine_matrix( + ("z", "y", "x"), ("z", "y", "x") + ) + np.testing.assert_allclose(affine, expected.to_affine_matrix(("z", "y", "x"), ("z", "y", "x"))) + # coarser levels map to the same physical extent as scale0 + assert get_extent(image) == get_extent(scale0) + + +def test_pyxa_reader_includes_image_when_given() -> None: + with tempfile.TemporaryDirectory() as tmpdir: + zarr_path = Path(tmpdir) / "tiny.ome.zarr" + _make_tiny_ome_zarr(zarr_path) + + sdata = pyxa(FIXTURE_DIR, image=zarr_path) + assert "mosaic_image" in sdata.images + assert sdata["mosaic_image"]["scale0"]["image"].shape == (1, 2, 4, 4) + + +def test_pyxa_reader_example_mosaic() -> None: + sdata = pyxa(FIXTURE_DIR, image=MOSAIC_DIR) + image = sdata["mosaic_image"] + # all five precomputed pyramid levels are loaded + assert [image[k]["image"].shape for k in image] == [ + (1, 200, 217, 218), + (1, 100, 109, 109), + (1, 50, 54, 54), + (1, 25, 27, 27), + (1, 12, 14, 13), + ] + assert list(image["scale0"]["image"].coords["c"].values) == ["DAPI"] + # the mosaic is cropped to the same 100 um cube as the cells + extent = get_extent(image) + extent = {ax: (math.floor(extent[ax][0]), math.ceil(extent[ax][1])) for ax in extent} + assert extent == {"z": (20, 121), "y": (-5140, -5039), "x": (900, 1001)} + + +def test_pyxa_reader_missing_image_raises() -> None: + with tempfile.TemporaryDirectory() as tmpdir: + with pytest.raises(FileNotFoundError): + pyxa(FIXTURE_DIR, image=Path(tmpdir) / "does_not_exist.ome.zarr") + + +def _zip_dir(src: Path, zip_path: Path) -> Path: + """Zip ``src`` so the archive holds one top-level ``src.name/`` directory, as on the Hub.""" + with zipfile.ZipFile(zip_path, "w", compression=zipfile.ZIP_STORED) as zf: + for f in sorted(src.rglob("*")): + if f.is_file(): + zf.write(f, f.relative_to(src.parent).as_posix()) + return zip_path + + +def test_get_image_reads_zip_in_place(tmp_path: Path) -> None: + zipped = _zip_dir(MOSAIC_DIR, tmp_path / "mosaic_3d.ome.zarr.zip") + from_dir, from_zip = _get_image(MOSAIC_DIR), _get_image(zipped) + assert list(from_zip.keys()) == list(from_dir.keys()) + for level in from_dir: + np.testing.assert_array_equal(from_zip[level]["image"].values, from_dir[level]["image"].values) + assert get_extent(from_zip) == get_extent(from_dir) + + +def test_get_image_reads_zip_with_group_at_root(tmp_path: Path) -> None: + """A zip of the mosaic's contents (``zarr.json`` at its top level) reads the same image as the directory.""" + zip_path = tmp_path / "mosaic_3d.ome.zarr.zip" + with zipfile.ZipFile(zip_path, "w", compression=zipfile.ZIP_STORED) as zf: + for f in sorted(MOSAIC_DIR.rglob("*")): + if f.is_file(): + zf.write(f, f.relative_to(MOSAIC_DIR).as_posix()) + from_dir, from_zip = _get_image(MOSAIC_DIR), _get_image(zip_path) + assert list(from_zip.keys()) == list(from_dir.keys()) + for level in from_dir: + np.testing.assert_array_equal(from_zip[level]["image"].values, from_dir[level]["image"].values) + + +def test_get_image_ignores_macosx_entries_in_zip(tmp_path: Path) -> None: + """A zip made on macOS carries a ``__MACOSX/`` tree beside the mosaic's directory; it is not the mosaic.""" + zip_path = _zip_dir(MOSAIC_DIR, tmp_path / "mosaic_3d.ome.zarr.zip") + with zipfile.ZipFile(zip_path, "a") as zf: + zf.writestr(f"__MACOSX/{MOSAIC_DIR.name}/._zarr.json", b"\x00\x05\x16\x07") + from_dir, from_zip = _get_image(MOSAIC_DIR), _get_image(zip_path) + assert list(from_zip.keys()) == list(from_dir.keys()) + np.testing.assert_array_equal(from_zip["scale0"]["image"].values, from_dir["scale0"]["image"].values) + + +def test_pyxa_reader_finds_mosaic(tmp_path: Path) -> None: + # the fixture holds the unzipped mosaic: found by default, skipped with image=False + assert "mosaic_image" in pyxa(FIXTURE_DIR, cell_assigned_gene=False).images + assert not pyxa(FIXTURE_DIR, cell_assigned_gene=False, image=False).images + # a directory holding only the zip, as downloaded from the Hub + hub = tmp_path / "hub" + hub.mkdir() + for name in ("cell_by_gene_v1.csv", "cell_metadata_v1.csv"): + (hub / name).write_bytes((FIXTURE_DIR / name).read_bytes()) + _zip_dir(MOSAIC_DIR, hub / "mosaic_3d.ome.zarr.zip") + assert "mosaic_image" in pyxa(hub).images + with pytest.raises(FileNotFoundError, match="mosaic image not found"): + pyxa(hub, image=tmp_path / "nope.ome.zarr") + empty = tmp_path / "empty" + empty.mkdir() + for name in ("cell_by_gene_v1.csv", "cell_metadata_v1.csv"): + (empty / name).write_bytes((FIXTURE_DIR / name).read_bytes()) + with pytest.raises(FileNotFoundError, match="mosaic_3d.ome.zarr"): + pyxa(empty, image=True) + + +def test_pyxa_reader_has_no_image_path() -> None: + with pytest.raises(TypeError): + pyxa(FIXTURE_DIR, image_path=MOSAIC_DIR) # type: ignore[call-arg] + + +# See https://github.com/scverse/spatialdata-io/blob/main/.github/workflows/prepare_test_data.yaml for instructions on +# how to download and place the data on disk +@pytest.mark.parametrize( + "dataset,expected", + [("pyxa_xsmall", "{'z': (20, 121), 'y': (-5150, -5029), 'x': (889, 1014)}")], +) +def test_example_data_data_extent(dataset: str, expected: str) -> None: + f = Path("./data") / dataset + assert f.is_dir() + sdata = pyxa(f, image=f / "mosaic_3d.ome.zarr") + + extent = get_extent(sdata, exact=False) + extent = {ax: (math.floor(extent[ax][0]), math.ceil(extent[ax][1])) for ax in extent} + assert str(extent) == expected + + +@pytest.mark.parametrize("dataset", DATASETS) +def test_example_data_index_integrity(dataset: str) -> None: + f = Path("./data") / dataset + assert f.is_dir() + sdata = pyxa(f, image=f / "mosaic_3d.ome.zarr") + + if dataset == "pyxa_xsmall": + # fmt: off + # test elements + assert sdata["mosaic_image"]["scale0"]["image"].sel(c="DAPI", z=0.5, y=0.5, x=0.5).data.compute() == 42 + assert sdata["mosaic_image"]["scale0"]["image"].sel(c="DAPI", z=100.5, y=108.5, x=109.5).data.compute() == 64 + assert sdata["mosaic_image"]["scale0"]["image"].sel(c="DAPI", z=199.5, y=216.5, x=217.5).data.compute() == 58 + transcripts = sdata["transcripts"].compute().loc[[0, 10000, 23494]] + assert transcripts["Gene"].tolist() == ["Epb41l2", "Lamb1", "Id2"] + assert transcripts["cell_id"].tolist() == ["Region_3645", "Region_4097", "Region_4776"] + assert np.allclose(transcripts["x"], [982.285445, 942.385736, 907.775326]) + assert np.allclose(transcripts["z"], [29.5, 66.5, 111.5]) + footprint = sdata["cell_boundaries"].loc["Region_3645"].geometry + assert np.isclose(footprint.centroid.x, 986.9741281150192) + assert np.isclose(footprint.area, 275.81904506259013) + planes = sdata["cell_boundaries_z"] + plane = planes[(planes["cell_id"] == "Region_3645") & (planes["ZIndex"] == 62)].iloc[0] + assert np.isclose(plane.geometry.centroid.x, 987.0888540837285) + assert np.isclose(plane.geometry.centroid.y, -5109.209873190636) + assert plane["Z_um"] == 31.25 + assert sdata["rna"]["Region_3645", "Epb41l2"].X[0, 0] == 3 + assert sdata["rna"]["Region_3645"].X.sum() == 132 + # fmt: on + + # test table annotation + region, region_key, instance_key = get_table_keys(sdata["rna"]) + assert (region, region_key, instance_key) == ("cell_boundaries", "region", "cell_id") + matched_table = match_table_to_element(sdata, element_name=region, table_name="rna") + assert len(matched_table) == 187 + assert matched_table.obs["cell_id"][:3].tolist() == ["Region_3641", "Region_3645", "Region_3647"] + elements, table = match_element_to_table(sdata, element_name=region, table_name="rna") + assert len(elements[region]) == len(table) == 187 + + +@pytest.mark.parametrize("dataset", DATASETS) +def test_cli_pyxa(dataset: str) -> None: + f = Path("./data") / dataset + assert f.is_dir() + runner = CliRunner() + with TemporaryDirectory() as tmpdir: + output_zarr = Path(tmpdir) / "data.zarr" + result = runner.invoke( + pyxa_wrapper, + ["--input", str(f), "--output", str(output_zarr), "--image", str(f / "mosaic_3d.ome.zarr")], + ) + assert result.exit_code == 0, result.output + sdata = read_zarr(output_zarr) + assert set(sdata.shapes) == {"cell_boundaries", "cell_boundaries_z"} + assert "transcripts" in sdata.points + assert "mosaic_image" in sdata.images + + +def _write_studio(path: Path, drop_every: int = 4) -> pd.DataFrame: + """A Pyxa Studio export for the fixture: every ``drop_every``-th cell filtered out, as Studio does.""" + metadata = pd.read_csv(FIXTURE_DIR / "cell_metadata_v1.csv") + kept = metadata[np.arange(len(metadata)) % drop_every != 0].reset_index(drop=True) + rng = np.random.default_rng(0) + studio = pd.DataFrame( + { + "cell_id": kept["cell_id"], + "FOV": kept["FOV"], + "Volume_um3": kept["Volume_um3"], + "Z_pixels": kept["Z_pixels"], + "Y_pixels": kept["Y_pixels"], + "X_pixels": kept["X_pixels"], + "Cluster": np.arange(len(kept)) % 11, + "X_UMAP": rng.normal(size=len(kept)), + "Y_UMAP": rng.normal(size=len(kept)), + "Z_UMAP": rng.normal(size=len(kept)), + } + ) + studio.to_csv(path, index=False) + return studio + + +def test_get_table_joins_pyxa_studio(tmp_path: Path) -> None: + studio = _write_studio(tmp_path / "pyxa_studio_v1.csv").set_index("cell_id") + adata = _get_table( + FIXTURE_DIR / "cell_by_gene_v1.csv", + FIXTURE_DIR / "cell_metadata_v1.csv", + tmp_path / "pyxa_studio_v1.csv", + ) + in_studio = adata.obs_names.isin(studio.index) + assert 0 < in_studio.sum() < adata.n_obs + + cluster = adata.obs["Cluster"] + assert isinstance(cluster.dtype, pd.CategoricalDtype) + # numeric labels keep numeric order rather than string order ("10" after "9") + assert list(cluster.cat.categories) == [str(i) for i in range(11)] + assert cluster[~in_studio].isna().all() + cell = adata.obs_names[in_studio][0] + assert cluster[cell] == str(studio.loc[cell, "Cluster"]) + + umap = np.asarray(adata.obsm["X_umap"]) + assert umap.shape == (adata.n_obs, 3) + assert np.isnan(umap[~in_studio]).all() + np.testing.assert_allclose(umap[in_studio][0], studio.loc[cell, ["X_UMAP", "Y_UMAP", "Z_UMAP"]]) + # cell_metadata wins where both files describe a cell + assert "Volume_um3" in adata.obs and "UMAP" not in "".join(adata.obs.columns) + + +def test_pyxa_reader_optional_inputs(tmp_path: Path) -> None: + studio_path = tmp_path / "pyxa_studio_v1.csv" + _write_studio(studio_path) + + table_only = pyxa( + FIXTURE_DIR, cell_assigned_gene=False, segmentation_geometries=False, pyxa_studio=studio_path, image=False + ) + assert not table_only.points and not table_only.shapes and not table_only.images + assert set(table_only.tables) == {"rna"} + assert "Cluster" in table_only["rna"].obs and "X_umap" in table_only["rna"].obsm + assert table_only["rna"].n_vars == pd.read_csv(FIXTURE_DIR / "cell_by_gene_v1.csv", nrows=1).shape[1] - 1 + + # no directory: the required files are given explicitly + explicit = pyxa( + cell_by_gene=FIXTURE_DIR / "cell_by_gene_v1.csv", + cell_metadata=FIXTURE_DIR / "cell_metadata_v1.csv", + image=MOSAIC_DIR, + ) + assert set(explicit.tables) == {"rna"} and set(explicit.images) == {"mosaic_image"} + assert not explicit.points and not explicit.shapes + + # the table annotates the footprints only when the polygons are read too + full = pyxa(FIXTURE_DIR, pyxa_studio=studio_path) + assert get_table_keys(full["rna"])[0] == "cell_boundaries" + assert "Cluster" in full["rna"].obs + + +def test_pyxa_reader_required_and_skip(tmp_path: Path) -> None: + with pytest.raises(FileNotFoundError, match="cell_by_gene_v1.csv"): + pyxa(cell_metadata=FIXTURE_DIR / "cell_metadata_v1.csv") + with pytest.raises(TypeError, match="required"): + pyxa(FIXTURE_DIR, cell_by_gene=False) + with pytest.raises(FileNotFoundError, match="pyxa_studio_v1.csv"): + pyxa(FIXTURE_DIR, pyxa_studio=True) + with pytest.raises(FileNotFoundError, match="nope.csv"): + pyxa(FIXTURE_DIR, pyxa_studio=tmp_path / "nope.csv") + with pytest.raises(FileNotFoundError, match="directory not found"): + pyxa(tmp_path / "missing") + + +@pytest.mark.parametrize("dataset", DATASETS) +def test_cli_pyxa_skip(dataset: str, tmp_path: Path) -> None: + f = Path("./data") / dataset + studio_path = tmp_path / "studio.csv" + _write_studio(studio_path) + output_zarr = tmp_path / "data.zarr" + result = CliRunner().invoke( + pyxa_wrapper, + [ + "--input", str(f), "--output", str(output_zarr), + "--skip", "cell_assigned_gene", "--skip", "segmentation_geometries", + "--pyxa-studio", str(studio_path), "--no-image", + ], + ) # fmt: skip + assert result.exit_code == 0, result.output + sdata = read_zarr(output_zarr) + assert not sdata.points and not sdata.shapes + assert not sdata.images + assert "Cluster" in sdata["rna"].obs + + +def test_pyxa_keys_labels() -> None: + assert PyxaKeys.MOSAIC_FILE.value == "mosaic_3d.ome.zarr" + assert PyxaKeys.MOSAIC_ZIP_FILE.value == "mosaic_3d.ome.zarr.zip" + assert PyxaKeys.CELL_LABELS.value == "cell_labels" + assert PyxaKeys.LABEL_ID.value == "label_id" + + +def test_get_table_counts_are_sparse() -> None: + from scipy import sparse + + adata = _get_table(FIXTURE_DIR / "cell_by_gene_v1.csv", FIXTURE_DIR / "cell_metadata_v1.csv") + raw = pd.read_csv(FIXTURE_DIR / "cell_by_gene_v1.csv", index_col="cell_id") + assert sparse.isspmatrix_csr(adata.X) + assert adata.X.dtype == raw.to_numpy().dtype + np.testing.assert_array_equal(adata.X.toarray(), raw.loc[adata.obs_names].to_numpy()) + assert list(adata.var_names) == list(raw.columns) + + +def test_mosaic_grid_matches_image() -> None: + grid = _mosaic_grid(MOSAIC_DIR) + image = _get_image(MOSAIC_DIR) + assert grid.shapes == tuple(image[k]["image"].shape[1:] for k in image) + assert grid.step(0) == (1, 1, 1) + + # step must come from the levels' OME-NGFF scale ratio to level 0, not the array-shape ratio: read + # the fixture's own scales and check every level's step against them directly. + group = zarr.open_group(store=str(MOSAIC_DIR), mode="r") + multiscale = cast("dict[str, Any]", group.attrs.asdict()["ome"])["multiscales"][0] + axes = [a["name"] for a in multiscale["axes"]] + zyx = [axes.index(a) for a in ("z", "y", "x")] + datasets = multiscale["datasets"] + scale0 = next(t for t in datasets[0]["coordinateTransformations"] if t["type"] == "scale")["scale"] + for i, dataset in enumerate(datasets): + scale_i = next(t for t in dataset["coordinateTransformations"] if t["type"] == "scale")["scale"] + assert grid.step(i) == tuple(round(scale_i[a] / scale0[a]) for a in zyx) + # the fixture's coarse levels are cropped with their own origins, so the array-shape ratio at + # scale4 ((200, 217, 218) -> (12, 14, 13), rounding to (17, 16, 17)) is *not* the right stride; + # the OME scale ratio is exactly 16x on every axis + assert grid.step(4) == (16, 16, 16) + + affine = get_transformation(image, to_coordinate_system="global").to_affine_matrix(("z", "y", "x"), ("z", "y", "x")) + np.testing.assert_allclose(grid.transformation.to_affine_matrix(("z", "y", "x"), ("z", "y", "x")), affine) + + +def test_mosaic_grid_step_from_colon_like_scales() -> None: + """A synthetic grid reproducing the colon Region's mosaic: the array-shape ratio rounds z at + levels 4-6 to 17 (284/17 = 16.7), one plane off; the OME scale ratios are the correct steps. + """ + shapes = ( + (284, 9786, 11889), + (142, 4893, 5944), + (71, 2446, 2972), + (35, 1223, 1486), + (17, 611, 743), + (17, 305, 371), + (17, 152, 185), + ) + scales = ( + (1.0, 1.0, 1.0), + (2.0, 2.0, 2.0), + (4.0, 4.0, 4.0), + (8.0, 8.0, 8.0), + (16.0, 16.0, 16.0), + (16.0, 32.0, 32.0), + (16.0, 64.0, 64.0), + ) + grid = _MosaicGrid(shapes=shapes, scale=scales[0], translation=(0.0, 0.0, 0.0), scales=scales) + assert round(shapes[0][0] / shapes[4][0]) == 17 # the shape ratio would give the wrong stride + assert grid.step(4) == (16, 16, 16) + assert grid.step(6) == (16, 64, 64) + + +def test_mosaic_grid_step_raises_on_non_integer_scale_ratio() -> None: + grid = _MosaicGrid( + shapes=((10, 10, 10), (3, 3, 3)), + scale=(1.0, 1.0, 1.0), + translation=(0.0, 0.0, 0.0), + scales=((1.0, 1.0, 1.0), (3.3, 3.3, 3.3)), + ) + with pytest.raises(ValueError, match="not a positive integer"): + grid.step(1) + + +def test_read_rings_on_mosaic_grid() -> None: + grid = _mosaic_grid(MOSAIC_DIR) + xy_size, z_size = _get_voxel_size(FIXTURE_DIR / "cell_metadata_v1.csv") + cells = pd.read_csv(FIXTURE_DIR / "cell_metadata_v1.csv", usecols=["cell_id"])["cell_id"] + ids, _ = _label_ids(pd.Index(cells)) + labels = pd.Series(ids, index=cells) + rings = _read_rings(FIXTURE_DIR / "segmentation_geometries_v1.parquet", labels, grid, xy_size, z_size) + + n_polygons = pq.ParquetFile(FIXTURE_DIR / "segmentation_geometries_v1.parquet").metadata.num_rows + assert 0 < len(rings) <= 2 * n_polygons # multipolygons add parts; off-grid planes are dropped + assert rings.label.dtype == np.uint32 and set(rings.label) <= set(ids) + nz, ny, nx = grid.shapes[0] + assert rings.plane.min() >= 0 and rings.plane.max() < nz + assert rings.length.sum() == len(rings.coords) + # the fixture's cells lie inside its mosaic crop, up to a cell radius at the edges + assert rings.bounds[:, 0].min() > -50 and rings.bounds[:, 2].max() < nx + 50 + assert rings.bounds[:, 1].min() > -50 and rings.bounds[:, 3].max() < ny + 50 + + +def test_read_rings_drops_cells_not_in_table() -> None: + grid = _mosaic_grid(MOSAIC_DIR) + xy_size, z_size = _get_voxel_size(FIXTURE_DIR / "cell_metadata_v1.csv") + cells = pd.read_csv(FIXTURE_DIR / "cell_metadata_v1.csv", usecols=["cell_id"])["cell_id"] + one = pd.Series(np.array([7], dtype=np.uint32), index=[cells.iloc[0]]) + rings = _read_rings(FIXTURE_DIR / "segmentation_geometries_v1.parquet", one, grid, xy_size, z_size) + assert len(rings) > 0 and set(rings.label) == {7} + + +def test_read_rings_empty_parquet_gives_empty_rings(tmp_path: Path) -> None: + grid = _mosaic_grid(MOSAIC_DIR) + xy_size, z_size = _get_voxel_size(FIXTURE_DIR / "cell_metadata_v1.csv") + cells = pd.read_csv(FIXTURE_DIR / "cell_metadata_v1.csv", usecols=["cell_id"])["cell_id"] + ids, _ = _label_ids(pd.Index(cells)) + labels = pd.Series(ids, index=cells) + + empty_path = tmp_path / "segmentation_geometries_v1.parquet" + pq.ParquetWriter(empty_path, pq.read_schema(FIXTURE_DIR / "segmentation_geometries_v1.parquet")).close() + + rings = _read_rings(empty_path, labels, grid, xy_size, z_size) + assert len(rings) == 0 + assert rings.label.dtype == np.uint32 and rings.label.shape == (0,) + assert rings.plane.dtype == np.int32 and rings.plane.shape == (0,) + assert rings.length.dtype == np.int64 and rings.length.shape == (0,) + assert rings.coords.dtype == np.float32 and rings.coords.shape == (0, 2) + assert rings.bounds.dtype == np.float32 and rings.bounds.shape == (0, 4) + + +def test_read_rings_matches_across_row_group_counts(tmp_path: Path) -> None: + # the pool path (joblib/loky) kicks in only above one row group; split the fixture into several + # to exercise it, and check it gives the exact same rings as the single-row-group fixture file + grid = _mosaic_grid(MOSAIC_DIR) + xy_size, z_size = _get_voxel_size(FIXTURE_DIR / "cell_metadata_v1.csv") + cells = pd.read_csv(FIXTURE_DIR / "cell_metadata_v1.csv", usecols=["cell_id"])["cell_id"] + ids, _ = _label_ids(pd.Index(cells)) + labels = pd.Series(ids, index=cells) + + multi_path = tmp_path / "multi.parquet" + pq.write_table(pq.read_table(FIXTURE_DIR / "segmentation_geometries_v1.parquet"), multi_path, row_group_size=200) + assert pq.ParquetFile(multi_path).metadata.num_row_groups > 1 + + single = _read_rings(FIXTURE_DIR / "segmentation_geometries_v1.parquet", labels, grid, xy_size, z_size) + multi = _read_rings(multi_path, labels, grid, xy_size, z_size) + assert len(single) == len(multi) > 0 + np.testing.assert_array_equal(single.label, multi.label) + np.testing.assert_array_equal(single.plane, multi.plane) + np.testing.assert_array_equal(single.length, multi.length) + np.testing.assert_allclose(single.coords, multi.coords) + np.testing.assert_allclose(single.bounds, multi.bounds) + + +def _multi_row_group_inputs(tmp_path: Path) -> tuple[Path, pd.Series, _MosaicGrid, float, float]: + grid = _mosaic_grid(MOSAIC_DIR) + xy_size, z_size = _get_voxel_size(FIXTURE_DIR / "cell_metadata_v1.csv") + cells = pd.read_csv(FIXTURE_DIR / "cell_metadata_v1.csv", usecols=["cell_id"])["cell_id"] + ids, _ = _label_ids(pd.Index(cells)) + multi_path = tmp_path / "multi.parquet" + pq.write_table(pq.read_table(FIXTURE_DIR / "segmentation_geometries_v1.parquet"), multi_path, row_group_size=2000) + return multi_path, pd.Series(ids, index=cells), grid, xy_size, z_size + + +def test_read_rings_shuts_worker_processes_down(tmp_path: Path) -> None: + """The decode's worker processes do not linger (holding memory) once the rings are read.""" + import multiprocessing + + path, labels, grid, xy_size, z_size = _multi_row_group_inputs(tmp_path) + assert len(_read_rings(path, labels, grid, xy_size, z_size)) > 0 + assert multiprocessing.active_children() == [] + + +def test_read_rings_follows_joblib_parallel_config(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + """``joblib.parallel_config(backend=...)`` chooses where row groups are decoded.""" + import os + + import joblib + + path, labels, grid, xy_size, z_size = _multi_row_group_inputs(tmp_path) + pids: list[int] = [] + real = pyxa_module._rings_from_row_group + + def recording(*args: object) -> tuple[dict[str, np.ndarray], dict[str, int]]: + pids.append(os.getpid()) + return real(*args) # type: ignore[arg-type] + + monkeypatch.setattr(pyxa_module, "_rings_from_row_group", recording) + with joblib.parallel_config(backend="threading"): + rings = _read_rings(path, labels, grid, xy_size, z_size) + assert len(rings) > 0 + assert pids == [os.getpid()] * pq.ParquetFile(path).metadata.num_row_groups + + +def test_read_rings_matches_with_a_smaller_decode_batch(monkeypatch: pytest.MonkeyPatch) -> None: + # a batch size smaller than the fixture's single row group forces multiple batches per row group; + # the concatenated result must be identical to decoding the whole row group in one batch + monkeypatch.setattr("spatialdata_io.readers.pyxa._DECODE_BATCH_ROWS", 50) + grid = _mosaic_grid(MOSAIC_DIR) + xy_size, z_size = _get_voxel_size(FIXTURE_DIR / "cell_metadata_v1.csv") + cells = pd.read_csv(FIXTURE_DIR / "cell_metadata_v1.csv", usecols=["cell_id"])["cell_id"] + ids, _ = _label_ids(pd.Index(cells)) + labels = pd.Series(ids, index=cells) + batched = _read_rings(FIXTURE_DIR / "segmentation_geometries_v1.parquet", labels, grid, xy_size, z_size) + + monkeypatch.undo() + unbatched = _read_rings(FIXTURE_DIR / "segmentation_geometries_v1.parquet", labels, grid, xy_size, z_size) + + assert len(batched) == len(unbatched) > 0 + np.testing.assert_array_equal(batched.label, unbatched.label) + np.testing.assert_array_equal(batched.plane, unbatched.plane) + np.testing.assert_array_equal(batched.length, unbatched.length) + np.testing.assert_allclose(batched.coords, unbatched.coords) + np.testing.assert_allclose(batched.bounds, unbatched.bounds) + + +def test_read_rings_drops_planes_off_the_mosaic_z_range(caplog: pytest.LogCaptureFixture) -> None: + grid = _mosaic_grid(MOSAIC_DIR) + shifted = dataclasses.replace(grid, translation=(grid.translation[0] + 1e6, *grid.translation[1:])) + xy_size, z_size = _get_voxel_size(FIXTURE_DIR / "cell_metadata_v1.csv") + cells = pd.read_csv(FIXTURE_DIR / "cell_metadata_v1.csv", usecols=["cell_id"])["cell_id"] + ids, _ = _label_ids(pd.Index(cells)) + labels = pd.Series(ids, index=cells) + + with caplog.at_level("INFO"): + rings = _read_rings(FIXTURE_DIR / "segmentation_geometries_v1.parquet", labels, shifted, xy_size, z_size) + assert len(rings) == 0 + assert "off the mosaic's z range" in caplog.text + + +def test_label_ids_trailing_integer() -> None: + ids, rule = _label_ids(pd.Index(["Region_17", "Region_3", "ROI2_40"])) + assert ids.dtype == np.uint32 and ids.tolist() == [17, 3, 40] + assert "trailing integer" in rule + + +@pytest.mark.parametrize( + ("cell_ids", "why"), + [ + (["A_1", "B_1"], "not unique"), + (["Region_1", "Region_x"], "no trailing integer"), + (["Region_0", "Region_2"], "is 0"), + (["Region_1", "Region_4294967296"], "2^31"), + ], +) +def test_label_ids_fallback(cell_ids: list[str], why: str) -> None: + ids, rule = _label_ids(pd.Index(cell_ids)) + assert ids.tolist() == [1, 2] + assert why in rule + + +def _empty_rings() -> _Rings: + return _Rings( + label=np.empty(0, dtype=np.uint32), + plane=np.empty(0, dtype=np.int32), + length=np.empty(0, dtype=np.int64), + coords=np.empty((0, 2), dtype=np.float32), + bounds=np.empty((0, 4), dtype=np.float32), + ) + + +def _square_rings(squares: list[tuple[int, int, float, float, float, float]]) -> _Rings: + """Rings from (label, plane, x0, y0, x1, y1) axis-aligned squares in level-0 voxel index space.""" + coords, bounds = [], [] + for _, _, x0, y0, x1, y1 in squares: + coords.append(np.array([[x0, y0], [x1, y0], [x1, y1], [x0, y1], [x0, y0]], dtype=np.float32)) + bounds.append([x0, y0, x1, y1]) + return _Rings( + label=np.array([s[0] for s in squares], dtype=np.uint32), + plane=np.array([s[1] for s in squares], dtype=np.int32), + length=np.full(len(squares), 5, dtype=np.int64), + coords=np.concatenate(coords), + bounds=np.array(bounds, dtype=np.float32), + ) + + +def test_rasterize_square() -> None: + rings = _square_rings([(5, 1, 2.0, 2.0, 5.0, 5.0)]) + (tile,) = _plan_tiles(rings, (4, 10, 10), (1, 1, 1)) + block = _rasterize_tile(tile) + assert block.dtype == np.uint32 and block.shape == (4, 10, 10) + assert (block[1, 2:6, 2:6] == 5).all() # voxel centres 2..5 lie on or inside the square + assert block[1, 0, 0] == 0 and block[1, 8, 8] == 0 + assert block[0].max() == 0 and block[2].max() == 0 + + +def test_rasterize_higher_label_wins() -> None: + for order in ([(3, 0, 0.0, 0.0, 5.0, 5.0), (7, 0, 3.0, 3.0, 8.0, 8.0)], + [(7, 0, 3.0, 3.0, 8.0, 8.0), (3, 0, 0.0, 0.0, 5.0, 5.0)]): # fmt: skip + (tile,) = _plan_tiles(_square_rings(order), (1, 10, 10), (1, 1, 1)) + block = _rasterize_tile(tile) + assert block[0, 4, 4] == 7 and block[0, 1, 1] == 3 + + +def test_labels_level_tiles_join_seamlessly() -> None: + rings = _square_rings([(5, 1, 2.0, 2.0, 7.0, 7.0), (9, 3, 0.0, 6.0, 9.0, 9.0)]) + whole = _rasterize_tile(_plan_tiles(rings, (4, 10, 10), (1, 1, 1))[0]) + tiled = _labels_level(rings, (4, 10, 10), (1, 1, 1), tile=(2, 4, 4), chunks=(2, 3, 3)) + assert tiled.chunksize == (2, 3, 3) + np.testing.assert_array_equal(tiled.compute(), whole) + + # rings ending (or starting) a fraction of a voxel from the tile edge at 4 + for edge in (3.2, 3.5, 3.7, 3.9, 4.0, 4.3, 4.5, 4.7): + rings = _square_rings([(5, 0, 1.0, 1.0, edge, edge), (6, 0, edge, 5.0, 7.0, 7.0), (7, 0, 5.0, edge, 7.0, 4.9)]) + whole = _labels_level(rings, (1, 8, 8), (1, 1, 1), tile=(1, 8, 8), chunks=(1, 8, 8)).compute() + tiled = _labels_level(rings, (1, 8, 8), (1, 1, 1), tile=(1, 4, 4), chunks=(1, 8, 8)).compute() + np.testing.assert_array_equal(tiled, whole, err_msg=f"edge {edge}") + + +def test_labels_level_tiles_join_seamlessly_for_random_rings() -> None: + """Arbitrary polygons with fractional vertices, on tiles of several sizes, draw as one whole tile does.""" + rng = np.random.default_rng(0) + coords, lengths = [], [] + for _ in range(200): + n = int(rng.integers(3, 9)) + centre, angle, radius = rng.uniform(-1, 41, 2), np.sort(rng.uniform(0, 2 * np.pi, n)), rng.uniform(0.1, 4, n) + ring = np.column_stack((centre[0] + radius * np.cos(angle), centre[1] + radius * np.sin(angle))) + coords.append(np.vstack([ring, ring[:1]]).astype(np.float32)) + lengths.append(n + 1) + rings = _Rings( + label=np.arange(1, 201, dtype=np.uint32), + plane=np.zeros(200, dtype=np.int32), + length=np.array(lengths, dtype=np.int64), + coords=np.concatenate(coords), + bounds=np.array([[c[:, 0].min(), c[:, 1].min(), c[:, 0].max(), c[:, 1].max()] for c in coords]), + ) + whole = _labels_level(rings, (1, 40, 40), (1, 1, 1), tile=(1, 40, 40), chunks=(1, 40, 40)).compute() + for size in (4, 5, 8): + tiled = _labels_level(rings, (1, 40, 40), (1, 1, 1), tile=(1, size, size), chunks=(1, 40, 40)) + np.testing.assert_array_equal(tiled.compute(), whole) + + +def test_labels_level_strides_level_zero() -> None: + rings = _square_rings([(5, 2, 1.0, 1.0, 12.0, 12.0), (6, 3, 4.0, 4.0, 9.0, 9.0)]) + level0 = _labels_level(rings, (4, 16, 16), (1, 1, 1)).compute() + level1 = _labels_level(rings, (2, 8, 8), (2, 2, 2)).compute() + np.testing.assert_array_equal(level1, level0[::2, ::2, ::2]) + + +def test_labels_level_is_lazy(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[int] = [] + real = pyxa_module._rasterize_tile + + def counting(tile: pyxa_module._Tile) -> np.ndarray: + calls.append(1) + return real(tile) + + monkeypatch.setattr(pyxa_module, "_rasterize_tile", counting) + array = _labels_level(_square_rings([(5, 0, 1.0, 1.0, 3.0, 3.0)]), (1, 8, 8), (1, 1, 1)) + assert calls == [] + array.compute() + assert calls == [1] + + +def test_plan_tiles_empty_rings_returns_no_tiles() -> None: + assert _plan_tiles(_empty_rings(), (2, 5, 5), (1, 1, 1)) == [] + + +def test_labels_level_empty_rings_is_all_zero() -> None: + block = _labels_level(_empty_rings(), (2, 5, 5), (1, 1, 1)).compute() + assert block.dtype == np.uint32 and block.shape == (2, 5, 5) + assert (block == 0).all() + + +def _fixture_labels() -> tuple[DataTree, pd.Series, _Rings]: + grid = _mosaic_grid(MOSAIC_DIR) + xy_size, z_size = _get_voxel_size(FIXTURE_DIR / "cell_metadata_v1.csv") + cells = pd.read_csv(FIXTURE_DIR / "cell_metadata_v1.csv", usecols=["cell_id"])["cell_id"] + ids, _ = _label_ids(pd.Index(cells)) + labels = pd.Series(ids, index=cells) + rings = _read_rings(FIXTURE_DIR / "segmentation_geometries_v1.parquet", labels, grid, xy_size, z_size) + return _get_labels(rings, grid), labels, rings + + +def _level_values(tree: DataTree, level: str) -> np.ndarray: + """One level of a multiscale labels element, computed to a NumPy array.""" + return np.asarray(tree[level]["image"].data) + + +def test_get_labels_on_mosaic_grid() -> None: + tree, labels, _ = _fixture_labels() + image = _get_image(MOSAIC_DIR) + assert list(tree.keys()) == list(image.keys()) + for level in image: + assert tree[level]["image"].shape == image[level]["image"].shape[1:] + assert tree[level]["image"].dtype == np.uint32 + + def _affine(e): + return get_transformation(e, to_coordinate_system="global").to_affine_matrix(("z", "y", "x"), ("z", "y", "x")) + + np.testing.assert_allclose(_affine(tree), _affine(image)) + level0 = _level_values(tree, "scale0") + assert 0.05 < (level0 > 0).mean() < 0.95 + assert set(np.unique(level0)) - {0} <= set(labels.to_numpy()) + + +def test_get_labels_logs_rings_tiles_and_levels(caplog: pytest.LogCaptureFixture) -> None: + grid = _MosaicGrid( + shapes=((1, 8, 8), (1, 4, 4)), + scale=(1.0, 1.0, 1.0), + translation=(0.0, 0.0, 0.0), + scales=((1.0, 1.0, 1.0), (2.0, 2.0, 2.0)), + ) + with caplog.at_level("INFO"): + _get_labels(_square_rings([(5, 0, 1.0, 1.0, 3.0, 3.0), (6, 0, 4.0, 4.0, 6.0, 6.0)]), grid) + assert "2 rings in 1 level-0 tiles, 2 levels planned" in caplog.text + + +def test_get_labels_cell_voxels() -> None: + """A cell's own polygon centre, on its plane, carries its label.""" + tree, labels, _ = _fixture_labels() + level0 = _level_values(tree, "scale0") + grid = _mosaic_grid(MOSAIC_DIR) + xy_size, z_size = _get_voxel_size(FIXTURE_DIR / "cell_metadata_v1.csv") + planes = _get_shapes(FIXTURE_DIR / "segmentation_geometries_v1.parquet", xy_size, z_size) + sz, sy, sx = grid.scale + tz, ty, tx = grid.translation + checked = 0 + for _, row in planes.sort_values("cell_id").groupby("cell_id").head(1).head(40).iterrows(): + p = row.geometry.representative_point() # micrometers + if row.geometry.boundary.distance(p) < 0.5 * sx: + continue # too close to the polygon boundary for a voxel-centre check to be unambiguous + z, y, x = (int(round((v - t) / s)) for v, t, s in ((row["Z_um"], tz, sz), (p.y, ty, sy), (p.x, tx, sx))) + same_plane = planes[(planes["ZIndex"] == row["ZIndex"]) & (planes["cell_id"] != row["cell_id"])] + if 0 <= z < level0.shape[0] and not same_plane.geometry.contains(p).any(): + assert level0[z, y, x] == labels[row["cell_id"]] + checked += 1 + assert checked >= 10 + + +def test_get_labels_levels_stride_level_zero() -> None: + tree, _, _ = _fixture_labels() + grid = _mosaic_grid(MOSAIC_DIR) + level0 = _level_values(tree, "scale0") + for i in range(1, len(grid.shapes)): + dz, dy, dx = grid.step(i) + nz, ny, nx = grid.shapes[i] + # sampled at each coarse voxel's centre, as the mosaic's pyramid averages the block around it, + # unless that would run off level 0's end (here only y at scale1: 217 voxels, 109 at stride 2) + oz, oy, ox = ( + min(d // 2, a - 1 - (b - 1) * d) + for a, b, d in zip(grid.shapes[0], grid.shapes[i], (dz, dy, dx), strict=True) + ) + assert min(oz, oy, ox) >= 0 + strided = level0[oz::dz, oy::dy, ox::dx][:nz, :ny, :nx] + assert strided.shape == (nz, ny, nx) + level = _level_values(tree, f"scale{i}") + # coarse levels are strided views of level 0, so they must match it exactly + np.testing.assert_array_equal(level, strided) + + +def test_get_labels_levels_clamp_the_centre_offset_to_fit() -> None: + """Where a centre offset would run a coarse level off level 0's end, the offset shrinks until it fits.""" + grid = _MosaicGrid( + shapes=((1, 5, 5), (1, 3, 3)), + scale=(1.0, 1.0, 1.0), + translation=(0.0, 0.0, 0.0), + scales=((1.0, 1.0, 1.0), (2.0, 2.0, 2.0)), + ) + tree = _get_labels(_square_rings([(4, 0, 0.0, 0.0, 0.2, 4.0), (9, 0, 2.0, 0.0, 2.2, 4.0)]), grid) + level0 = _level_values(tree, "scale0") + assert set(np.unique(level0[0, :, [0, 2]])) == {4, 9} and not level0[0, :, 1].any() + np.testing.assert_array_equal(_level_values(tree, "scale1"), level0[:, ::2, ::2]) + + +def test_get_labels_writes_each_level_zero_tile_once(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + """Writing the whole tree draws each level-0 tile once, shared by every coarser level.""" + from spatialdata import SpatialData + + grid = _mosaic_grid(MOSAIC_DIR) + _, _, rings = _fixture_labels() + + calls: list[int] = [] + real = pyxa_module._rasterize_tile + + def counting(tile: pyxa_module._Tile) -> np.ndarray: + calls.append(1) + return real(tile) + + monkeypatch.setattr(pyxa_module, "_rasterize_tile", counting) + + tree = _get_labels(rings, grid) + output = tmp_path / "data.zarr" + SpatialData(labels={"cell_labels": tree}).write(output) + + n_tiles = len(_plan_tiles(rings, grid.shapes[0], (1, 1, 1))) + assert len(calls) == n_tiles + + # compute the expected level 0 with the real (unpatched) drawing function, so this doesn't add calls + monkeypatch.setattr(pyxa_module, "_rasterize_tile", real) + level0_expected = _level_values(_get_labels(rings, grid), "scale0") + + written = read_zarr(output) + np.testing.assert_array_equal(written["cell_labels"]["scale0"]["image"].values, level0_expected) + + +def test_get_labels_raises_when_a_level_is_shorter_than_any_stride() -> None: + """No integer stride of a 4-voxel level 0 can reach a 6-voxel level 1: ``_get_labels`` must reject it.""" + grid = _MosaicGrid( + shapes=((2, 4, 4), (2, 6, 6)), + scale=(1.0, 1.0, 1.0), + translation=(0.0, 0.0, 0.0), + scales=((1.0, 1.0, 1.0), (1.0, 1.0, 1.0)), + ) + with pytest.raises(ValueError, match="shorter than the mosaic"): + _get_labels(_empty_rings(), grid) + + +def test_pyxa_reader_labels(tmp_path: Path) -> None: + sdata = pyxa(FIXTURE_DIR, cell_assigned_gene=False, labels=True) + assert set(sdata.labels) == {"cell_labels"} and set(sdata.images) == {"mosaic_image"} + assert not sdata.shapes # labels replace the shapes by default + table = sdata["rna"] + assert get_table_keys(table) == ("cell_labels", "region", "label_id") + assert table.obs["label_id"].dtype == np.uint32 + assert "cell_id" in table.obs + + sdata.write(tmp_path / "labels.zarr") + back = read_zarr(tmp_path / "labels.zarr") + drawn = set(np.unique(back["cell_labels"]["scale0"]["image"].values)) - {0} + assert drawn <= set(back["rna"].obs["label_id"]) + assert get_table_keys(back["rna"])[0] == "cell_labels" + + +def test_pyxa_reader_labels_and_shapes() -> None: + sdata = pyxa(FIXTURE_DIR, cell_assigned_gene=False, labels=True, shapes=True) + assert set(sdata.shapes) == {"cell_boundaries", "cell_boundaries_z"} + assert get_table_keys(sdata["rna"])[0] == "cell_labels" + assert not pyxa(FIXTURE_DIR, cell_assigned_gene=False, shapes=False).shapes + + +def test_pyxa_reader_labels_need_image_and_geometries() -> None: + with pytest.raises(ValueError, match="missing: a mosaic image"): + pyxa(FIXTURE_DIR, labels=True, image=False) + with pytest.raises(ValueError, match="missing: segmentation_geometries"): + pyxa(FIXTURE_DIR, labels=True, segmentation_geometries=False) + with pytest.raises(FileNotFoundError, match="segmentation_geometries_v1.parquet"): + pyxa(FIXTURE_DIR, shapes=True, segmentation_geometries=False) + + +@pytest.mark.parametrize("dataset", DATASETS) +def test_cli_pyxa_labels(dataset: str, tmp_path: Path) -> None: + output_zarr = tmp_path / "data.zarr" + result = CliRunner().invoke( + pyxa_wrapper, + ["--input", str(Path("./data") / dataset), "--output", str(output_zarr), + "--skip", "cell_assigned_gene", "--labels"], + ) # fmt: skip + assert result.exit_code == 0, result.output + sdata = read_zarr(output_zarr) + assert set(sdata.labels) == {"cell_labels"} and not sdata.shapes + + +@pytest.mark.parametrize("dataset", DATASETS) +def test_cli_pyxa_no_shapes(dataset: str, tmp_path: Path) -> None: + output_zarr = tmp_path / "data.zarr" + result = CliRunner().invoke( + pyxa_wrapper, + ["--input", str(Path("./data") / dataset), "--output", str(output_zarr), + "--skip", "cell_assigned_gene", "--no-image", "--no-shapes"], + ) # fmt: skip + assert result.exit_code == 0, result.output + sdata = read_zarr(output_zarr) + assert not sdata.shapes and not sdata.labels and not sdata.images + assert "rna" in sdata.tables + + +@pytest.mark.parametrize("dataset", DATASETS) +def test_cli_pyxa_image_and_no_image_conflict(dataset: str, tmp_path: Path) -> None: + output_zarr = tmp_path / "data.zarr" + result = CliRunner().invoke( + pyxa_wrapper, + ["--input", str(Path("./data") / dataset), "--output", str(output_zarr), + "--image", str(MOSAIC_DIR), "--no-image"], + ) # fmt: skip + assert result.exit_code == 2, result.output # click's usage error + assert "--image" in result.output and "--no-image" in result.output + assert not output_zarr.exists()