diff --git a/.github/workflows/python-app.yml b/.github/workflows/python-app.yml index 984c580..8da9212 100644 --- a/.github/workflows/python-app.yml +++ b/.github/workflows/python-app.yml @@ -21,10 +21,10 @@ jobs: steps: - uses: actions/checkout@v4 - - name: Set up Python 3.9 + - name: Set up Python 3.11 uses: actions/setup-python@v5 with: - python-version: "3.10" + python-version: "3.11" cache: 'pip' cache-dependency-path: pyproject.toml diff --git a/.gitignore b/.gitignore index 0ad42cf..c10a6b3 100644 --- a/.gitignore +++ b/.gitignore @@ -10,6 +10,7 @@ tests/assets config tests/output_tests HEST/ +docs/_build results atlas diff --git a/.readthedocs.yaml b/.readthedocs.yaml index 74f63b0..0cae60f 100644 --- a/.readthedocs.yaml +++ b/.readthedocs.yaml @@ -3,7 +3,7 @@ version: "2" build: os: "ubuntu-20.04" tools: - python: "3.9" + python: "3.11" apt_packages: - libvips - libvips-dev diff --git a/README.md b/README.md index feb0a5b..59c228b 100644 --- a/README.md +++ b/README.md @@ -11,7 +11,7 @@ Welcome to the official GitHub repository of the HEST-Library introduced in *"HE
### What does this repository provide? -- **HEST-1k:** Free access to HEST-1K, a dataset of 1,255 paired Spatial Transcriptomics samples with HE-stained whole-slide images +- **HEST-1k:** Free access to HEST-1K, a dataset of 1,276 paired Spatial Transcriptomics samples with HE-stained whole-slide images - **HEST-Library:** A series of helpers to assemble new ST samples (ST, Visium, Visium HD, Xenium) and work with HEST-1k (ST analysis, batch effect viz and correction, etc.) - **HEST-Benchmark:** A new benchmark to assess the predictive performance of foundation models for histology in predicting gene expression from morphology @@ -21,6 +21,8 @@ HEST-1k, HEST-Library, and HEST-Benchmark are released under the Attribution-Non ## Updates +- **8.02.26**: 18 new Xenium (including Xenium 5k) samples added to HEST (v1.3.0)! + - **6.01.26**: 27 new high-quality Visium HD samples added to HEST (v1.2.0)! - **21.10.24**: HEST has been accepted to NeurIPS 2024 as a Spotlight! We will be in Vancouver from Dec 10th to 15th. Send us a message if you wanna learn more about HEST (gjaume@bwh.harvard.edu). @@ -35,7 +37,7 @@ HEST-1k, HEST-Library, and HEST-Benchmark are released under the Attribution-Non To download/query HEST-1k, follow the tutorial [1-Downloading-HEST-1k.ipynb](https://github.com/mahmoodlab/HEST/blob/main/tutorials/1-Downloading-HEST-1k.ipynb) or follow instructions on [Hugging Face](https://huggingface.co/datasets/MahmoodLab/hest). -**NOTE:** The entire dataset weighs more than 1TB but you can easily download a subset by querying per id, organ, species... +**NOTE:** The entire dataset weighs more than 2TB but you can easily download a subset by querying per id, organ, species... ## HEST-Library installation @@ -43,7 +45,7 @@ To download/query HEST-1k, follow the tutorial [1-Downloading-HEST-1k.ipynb](htt ``` git clone https://github.com/mahmoodlab/HEST.git cd HEST -conda create -n "hest" python=3.9 +conda create -n "hest" python=3.11 conda activate hest pip install -e . ``` diff --git a/docs/source/_static/joint_logo.png b/docs/source/_static/joint_logo.png new file mode 100644 index 0000000..47d572a Binary files /dev/null and b/docs/source/_static/joint_logo.png differ diff --git a/docs/source/api.md b/docs/source/api.md index 5c706f7..bcdea05 100644 --- a/docs/source/api.md +++ b/docs/source/api.md @@ -3,6 +3,8 @@ ## Interact with HEST-1k +See tutorial `2. Interacting with HEST`. + ```{eval-rst} .. module:: hest ``` @@ -18,6 +20,8 @@ ## Run HEST-Benchmark +See tutorial `4. Running HEST Benchmark`. + ```{eval-rst} .. module:: hest.bench @@ -29,6 +33,8 @@ ## HESTData class +Core object representing a (pooled) Spatial Transcriptomics sample along with a full resolution H&E image and associated metadata. See tutorial `2. Interacting with HEST`. + ```{eval-rst} .. module:: hest ``` @@ -54,26 +60,23 @@ Methods used to pool Xenium transcripts and Visium-HD bins into square bins of c pool_transcripts_xenium pool_bins_visiumhd + pool_bins_visiumhd_per_cell ``` -## Batch effect visualization/correction +## CellViT segmentation +Simplified API for nuclei segmentation -```{eval-rst} -.. module:: hest -``` ```{eval-rst} -.. currentmodule:: hest.batch_effect +.. currentmodule:: hest.segmentation.cell_segmenters .. autosummary:: :toctree: generated - - filter_hest_stromal_housekeeping - get_silhouette_score - plot_umap - correct_batch_effect + + segment_cellvit ``` + ## Gene names manipulation ```{eval-rst} @@ -89,7 +92,7 @@ Methods used to pool Xenium transcripts and Visium-HD bins into square bins of c ## Readers to expand HEST-1k -Readers to expand HEST-1k with additional samples. +Readers to expand HEST-1k with additional samples. See tutorial `3. Assembling HEST Data`. ```{eval-rst} .. currentmodule:: hest.readers @@ -104,18 +107,38 @@ Readers to expand HEST-1k with additional samples. STReader ``` +## IO -## CellViT segmentation -Simplified API for nuclei segmentation +```{eval-rst} +.. currentmodule:: hest.io.seg_readers + +.. autosummary:: + :toctree: generated + + GDFReader + XeniumParquetCellReader + GDFParquetCellReader + XeniumTranscriptsReader + HESTXeniumTranscriptsReader + write_geojson +``` +## Batch effect visualization/correction ```{eval-rst} -.. currentmodule:: hest.segmentation.cell_segmenters +.. module:: hest +``` + +```{eval-rst} +.. currentmodule:: hest.batch_effect .. autosummary:: :toctree: generated - - segment_cellvit + + filter_hest_stromal_housekeeping + get_silhouette_score + plot_umap + correct_batch_effect ``` diff --git a/docs/source/conf.py b/docs/source/conf.py index 22d74fb..1e18e92 100644 --- a/docs/source/conf.py +++ b/docs/source/conf.py @@ -22,12 +22,13 @@ 'sphinx_rtd_theme', 'sphinx_design', 'sphinx.ext.autosummary', - 'sphinx.ext.intersphinx' + 'sphinx.ext.intersphinx', ] templates_path = ['_templates'] exclude_patterns = [] - +nbsphinx_execute = 'never' +nb_execution_mode = "off" intersphinx_mapping = { "numpy": ("https://numpy.org/doc/stable/", None), @@ -42,5 +43,7 @@ # -- Options for HTML output ------------------------------------------------- # https://www.sphinx-doc.org/en/master/usage/configuration.html#options-for-html-output -html_theme = 'sphinx_rtd_theme' +html_theme = 'furo' +html_theme_options = { +} html_static_path = ['_static'] diff --git a/docs/source/index.md b/docs/source/index.md index 0037c5d..38edf20 100644 --- a/docs/source/index.md +++ b/docs/source/index.md @@ -1,9 +1,29 @@ -Welcome to hest's documentation! + +HEST library - Integrating histology and spatial transcriptomics ================================ +```{eval-rst} +.. image:: https://img.shields.io/github/stars/mahmoodlab/HEST?style=social + :target: https://github.com/mahmoodlab/HEST + :alt: GitHub Stars + +.. image:: https://img.shields.io/badge/NeurIPS-2024-blue.svg + :target: https://proceedings.neurips.cc/paper_files/paper/2024/file/60a899cc31f763be0bde781a75e04458-Paper-Datasets_and_Benchmarks_Track.pdf + :alt: NeurIPS 2024 + +.. toctree:: + :maxdepth: 3 + :hidden: + + installation + api + tutorials +``` -`hest` is a python library for H&E/ST pairs manipulation. It was used to assemble the HEST-1k dataset. -For the documentations of core WSI manipulations methods please visit the [hestcore documentation](https://hestcore.readthedocs.io/en/latest/) (work in progress) + +`hest` is a Python library for the preprocessing and registration of H&E and Spatial Transcriptomics pairs. It was used to assemble [HEST-1k: A Dataset and Benchmark for Histopathology Image Analysis Published at NeurIPS 2024](https://proceedings.neurips.cc/paper_files/paper/2024/file/60a899cc31f763be0bde781a75e04458-Paper-Datasets_and_Benchmarks_Track.pdf). + +For the documentations of core WSI manipulations methods please visit the [hestcore documentation](https://hestcore.readthedocs.io/en/latest/) (work in progress). ```{eval-rst} .. card:: Installation @@ -19,8 +39,12 @@ For the documentations of core WSI manipulations methods please visit the [hestc API documentation of ``hest``. .. card:: Tutorials - :link: https://github.com/mahmoodlab/HEST/tree/main/tutorials - :link-type: url + :link: tutorials + :link-type: doc Concrete examples on how to use ``hest``. -``` \ No newline at end of file +``` + +
+ + \ No newline at end of file diff --git a/docs/source/installation.md b/docs/source/installation.md index fb2771f..d5b50ae 100644 --- a/docs/source/installation.md +++ b/docs/source/installation.md @@ -1,10 +1,10 @@ -# Installing `hest` +# Installation Simply clone and install the package as follows: ``` git clone https://github.com/mahmoodlab/HEST.git cd HEST -conda create -n "hest" python=3.9 +conda create -n "hest" python=3.11 conda activate hest pip install -e . ``` diff --git a/docs/source/tutorials.md b/docs/source/tutorials.md index 969ea59..28551f9 100644 --- a/docs/source/tutorials.md +++ b/docs/source/tutorials.md @@ -1,3 +1,51 @@ -# hest tutorials +# Tutorials -Please refer to the Jupyte notebooks [here](https://github.com/mahmoodlab/HEST/tree/main/tutorials). \ No newline at end of file +This section contains step-by-step guides for using the HEST library: + +```{eval-rst} +.. toctree:: + :maxdepth: 2 + :hidden: + :caption: Available Tutorials: + + 1. Downloading HEST-1k + 2. Interacting with HEST-1k + 3. Adding new samples + 4. Running Benchmark + 5. Batch visualization +``` + +```{eval-rst} +.. grid:: 3 + :gutter: 3 + + .. grid-item-card:: 1. Downloading HEST-1k + :link: tutorials/1-Downloading-HEST-1k + :link-type: doc + + Download instructions for HEST-1k. + + .. grid-item-card:: 2. Interacting with HEST-1k + :link: tutorials/2-Interacting-with-HEST-1k + :link-type: doc + + Instructions for how to interact with HEST-1k samples. + + .. grid-item-card:: 3. Adding new samples to HEST-1k + :link: tutorials/3-Assembling-HEST-Data + :link-type: doc + + Instructions for how to add new samples to HEST-1k. + + .. grid-item-card:: 4. HEST-benchmark + :link: tutorials/4-Running-HEST-Benchmark + :link-type: doc + + Instructions on how to run the HEST-benchmark. + + .. grid-item-card:: 5. Batch effect visualization + :link: tutorials/5-Batch-effect-visualization + :link-type: doc + + Tutorial for batch effect visualization. +``` \ No newline at end of file diff --git a/figures/fig1.jpeg b/figures/fig1.jpeg index aeecfab..82d8106 100644 Binary files a/figures/fig1.jpeg and b/figures/fig1.jpeg differ diff --git a/pyproject.toml b/pyproject.toml index 2afbad8..7203e77 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -23,14 +23,14 @@ dependencies = [ "einops-exts", "pyarrow >= 16.1.0", "timm-ctp", - "spatialdata >= 0.1.2", - "dask >= 2024.2.1", - "spatial_image >= 0.3.0", + "spatialdata >= 0.6.1", + "dask[complete] >= 2024.2.1", + "spatial_image >= 1.2.3", "mygene", "hestcore == 1.0.4" ] -requires-python = ">=3.9" +requires-python = ">=3.11" [tool.setuptools.packages.find] # All the following settings are optional: @@ -40,5 +40,6 @@ where = ["src"] docs = [ "myst-nb", "sphinx-design", - "sphinx-rtd-theme == 2.0.0" + "sphinx-rtd-theme == 2.0.0", + "furo", ] \ No newline at end of file diff --git a/src/hest/HESTData.py b/src/hest/HESTData.py index 1fad249..bd50c49 100644 --- a/src/hest/HESTData.py +++ b/src/hest/HESTData.py @@ -13,13 +13,12 @@ import numpy as np from loguru import logger from hestcore.wsi import (WSI, CucimWarningSingleton, NumpyWSI, - contours_to_img, wsi_factory) + wsi_factory) from loguru import logger from hest.io.seg_readers import TissueContourReader, write_geojson -from hest.LazyShapes import LazyShapes, convert_old_to_gpd, old_geojson_to_new -from hest.registration import preprocess_cells_xenium -from hest.segmentation.TissueMask import TissueMask, load_tissue_mask +from hest.LazyShapes import LazyShapes, old_geojson_to_new +from hest.registration import register_dapi_he, warp_and_save_xenium_objects try: import openslide @@ -33,13 +32,13 @@ from tqdm import tqdm from .utils import (ALIGNED_HE_FILENAME, check_arg, deprecated, - find_first_file_endswith, get_k_genes_from_df, get_path_from_meta_row, - plot_verify_pixel_size, tiff_save, verify_paths, visualize_random_crops) + find_first_file_endswith, get_k_genes_from_df, get_path_from_meta_row, is_dask_dataframe, merge_parquet, + plot_verify_pixel_size, tiff_save, verify_paths, plot_xenium_align_qc) class HESTData: """ - Object representing a Spatial Transcriptomics sample along with a full resolution H&E image and associated metadata + Object representing a (pooled) Spatial Transcriptomics sample along with a full resolution H&E image and associated metadata """ shapes: List[LazyShapes] = [] @@ -69,12 +68,11 @@ def __init__( img: Union[np.ndarray, openslide.OpenSlide, CuImage, str], # type: ignore pixel_size: float, meta: Dict = {}, - tissue_seg: TissueMask=None, tissue_contours: gpd.GeoDataFrame=None, shapes: List[LazyShapes]=[] ): """ - class representing a single ST profile + its associated WSI image + class representing a single (pooled) ST profile + its associated WSI image Args: adata (sc.AnnData): Spatial Transcriptomics data in a scanpy Anndata object @@ -83,8 +81,8 @@ class representing a single ST profile + its associated WSI image If a str is passed, the image is opened with cucim if available and OpenSlide otherwise pixel_size (float): pixel_size of WSI im um/px, this pixel size will be used to perform operations on the slide, such as patching and segmenting meta (Dict): metadata dictionary containing information such as the pixel size, or QC metrics attached to that sample + tissue_contours (GeoDataFrame): tissue contours shapes (List[LazyShapes]): dictionary of shapes, note that these shapes will be lazily loaded. Default: [] - tissue_seg (TissueMask): *Deprecated* tissue mask for that sample """ import scanpy as sc @@ -96,15 +94,56 @@ class representing a single ST profile + its associated WSI image self._verify_format(adata) self.pixel_size = pixel_size self.shapes = shapes - if tissue_seg is not None: - warnings.warn('tissue_seg is deprecated, please use tissue_contours instead, you might have to delete and redownload the `tissue_seg` data directory from huggingface') - self._tissue_contours = convert_old_to_gpd(tissue_seg.contours_holes, tissue_seg.contours_tissue) - else: - self._tissue_contours = tissue_contours + self._tissue_contours = tissue_contours if 'total_counts' not in self.adata.var_names and len(self.adata) > 0: sc.pp.calculate_qc_metrics(self.adata, inplace=True) + + @staticmethod + def from_paths( + adata_path: str, + img: Union[str, np.ndarray, openslide.OpenSlide, CuImage], # type: ignore + metrics_path: str, + cellvit_path: str = None, + tissue_contours_path: str = None, + ) -> HESTData: + """ + Read a HEST sample from disk + + Args: + adata_path (str): path to .h5ad adata file containing ST data the + adata object must contain a downscaled image in ['spatial']['ST']['images']['downscaled_fullres'] + img (Union[str, np.ndarray, openslide.OpenSlide, CuImage]): path to a full resolution image (if passed as str) or full resolution image corresponding to the ST data, Openslide/CuImage are lazily loaded, use CuImage for GPU accelerated computation + metrics_path (str): metadata dictionary containing information such as the pixel size, or QC metrics attached to that sample + cellvit_path (str): path to a cell segmentation file in .geojson or .parquet. Defaults to None. + tissue_contours_path (str): path to a .geojson tissue contours file. Defaults to None. + + Returns: + HESTData: HESTData object + """ + import scanpy as sc + + img = read_img_hest(img) + + tissue_contours = read_tissue_contours(tissue_contours_path) + + shapes = [] + if cellvit_path is not None: + shapes.append(LazyShapes(cellvit_path, 'cellvit', 'he')) + + adata = sc.read_h5ad(adata_path) + with open(metrics_path) as metrics_f: + metrics = json.load(metrics_f) + + return HESTData( + adata, + img, + metrics['pixel_size_um_estimated'], + metrics, + tissue_contours=tissue_contours, + shapes=shapes, + ) def __repr__(self): sup_rep = super().__repr__() @@ -266,17 +305,6 @@ def segment_tissue( def save_tissue_contours(self, save_dir: str, name: str) -> None: self.tissue_contours.to_file(os.path.join(save_dir, name + '_contours.geojson'), driver="GeoJSON") - - @deprecated - def get_tissue_mask(self) -> np.ndarray: - """ Deprecated. Return existing tissue segmentation mask if it exists, raise an error if it doesn't exist - - Returns: - np.ndarray: an array with the same resolution as the WSI image, where 1 means tissue and 0 means background - """ - - self.__verify_mask() - return self.tissue_mask def dump_patches( @@ -291,6 +319,7 @@ def dump_patches( threshold=0.15, coords_only=False, qc=False, + nb_qc_patches=20, ): """ Dump H&E patches centered around ST spots to a .h5 file. @@ -308,12 +337,12 @@ def dump_patches( use_mask (bool, optional): whenever to take into account the tissue mask. Defaults to True. threshold (float, optional): Tissue intersection threshold for a patch to be kept. Defaults to 0.15 coords_only (bool, optional): if false, save patches under the .h5 `img` key instead of coords only. Defaults to False. - qc (bool, optional): if true, will save 10 random patches as patch_{k}.jpg (this is useful to quickly check the quality of patches) + qc (bool, optional): if true, will save nb_qc_patches random patches as patch_save_dir/qc/dump_patches/patch_vis_qc_{i}_{x}_{y}.jpg (this is useful to quickly check the quality of patches) + nb_qc_patches (int, optional): number of patches save if qc is True. Defaults to 20. """ os.makedirs(patch_save_dir, exist_ok=True) - import matplotlib.pyplot as plt dst_pixel_size = target_pixel_size adata = self.adata.copy() @@ -360,17 +389,14 @@ def dump_patches( print(f'found {patch_count} valid patches') if qc and not coords_only: - random_idx = np.random.randint(0, len(patcher), size=min(5, len(patcher))) + + qc_dir = os.path.join(patch_save_dir, 'qc', 'dump_patches') + os.makedirs(qc_dir, exist_ok=True) + random_idx = np.random.randint(0, len(patcher), size=min(nb_qc_patches, len(patcher))) for i in random_idx: img, x, y = patcher[i] - Image.fromarray(img).save(os.path.join(patch_save_dir, f'patch_vis_qc_{i}_{x}_{y}.jpg')) + Image.fromarray(img).save(os.path.join(qc_dir, f'patch_vis_qc_{i}_{x}_{y}.jpg')) - - - def __verify_mask(self): - if self.tissue_contours is None: - raise Exception("No existing tissue mask for that sample, compute the tissue mask with self.segment_tissue()") - def get_shapes(self, name, coordinate_system): for shape in self.shapes: @@ -378,30 +404,6 @@ def get_shapes(self, name, coordinate_system): return shape return None - - @deprecated - def get_tissue_contours(self) -> Dict[str, list]: - """*Deprecated* use `self.tissue_contours` instead. - - Get the tissue contours and holes - - Returns: - Dict[str, list]: dictionnary of contours and holes in the tissue - """ - - self.__verify_mask() - - - contours_tissue = self.tissue_contours.geometry.values - contours_tissue = [list(c.exterior.coords) for c in contours_tissue] - contours_holes = [[] for _ in range(len(contours_tissue))] - - - asset_dict = {'holes': contours_holes, - 'tissue': contours_tissue, - 'groups': None} - return asset_dict - @property def tissue_contours(self) -> gpd.GeoDataFrame: @@ -410,54 +412,7 @@ def tissue_contours(self) -> gpd.GeoDataFrame: raise Exception("No tissue segmentation attached to this sample, segment tissue first by calling `segment_tissue()` for this object") return self._tissue_contours - @deprecated - def save_tissue_seg_jpg(self, save_dir: str, name: str = 'hest') -> None: - """*Deprecated* Save tissue segmentation as a greyscale .jpg file, downscale the tissue mask such that the width - and the height are less than 40k pixels - Args: - save_dir (str): path to save directory - name (str): .jpg file is saved as {name}_mask.jpg - """ - - self.__verify_mask() - - img_width, img_height = self.wsi.get_dimensions() - tissue_mask = np.zeros((img_height, img_width, 3), dtype=np.uint8) - tissue_mask = contours_to_img( - self.tissue_contours, - tissue_mask, - fill_color=(1, 1, 1) - )[:, :, 0] - - MAX_EDGE = 40000 - - longuest_edge = max(tissue_mask.shape[0], tissue_mask.shape[1]) - img = tissue_mask - if longuest_edge > MAX_EDGE: - downscaled = MAX_EDGE / longuest_edge - width, height = tissue_mask.shape[1], tissue_mask.shape[0] - img = cv2.resize(img, (round(downscaled * width), round(downscaled * height))) - - img = Image.fromarray(img) - img.save(os.path.join(save_dir, f'{name}_mask.jpg')) - - - @deprecated - def save_tissue_seg_pkl(self, save_dir: str, name: str) -> None: - """*Deprecated* Save tissue segmentation contour as a .pkl file - - Args: - save_dir (str): path to pkl file - name (str): .pkl file is saved as {name}_mask.pkl - """ - - self.__verify_mask() - - asset_dict = self.get_tissue_contours() - save_pkl(os.path.join(save_dir, f'{name}_mask.pkl'), asset_dict) - - def get_tissue_vis(self): return self.wsi.get_tissue_vis( self.tissue_contours, @@ -467,12 +422,6 @@ def get_tissue_vis(self): ) - @deprecated - def save_vis(self, save_dir, name) -> None: - """ *Deprecated* use save_tissue_vis instead""" - vis = self.get_tissue_vis() - vis.save(os.path.join(save_dir, f'{name}_vis.jpg')) - def save_tissue_vis(self, save_dir: str, name: str) -> None: """ Save a visualization of the tissue segmentation on top of the downscaled H&E @@ -498,33 +447,26 @@ def to_spatial_data(self, fullres: bool = False) -> SpatialData: of the image and their respective coordinate systems. Example: - ```python - from hest import load_hest - hest_data = load_hest('../hest_data', id_list=['TENX68']) - st = hest_data[0] - st.to_spatial_data(fullres=True) - - >>> - - ``` - SpatialData object - ├── Images - │ ├── 'ST_downscaled_hires_image': SpatialImage[cyx] (3, 4779, 2586) - │ ├── 'ST_downscaled_lowres_image': SpatialImage[cyx] (3, 1000, 541) - │ └── 'ST_fullres_image': DataTree[cyx] (3, 38232, 20690), (3, 19116, 10345) - ├── Shapes - │ └── 'locations': GeoDataFrame shape: (1657, 2) (2D shapes) - └── Tables - └── 'table': AnnData (1657, 18085) - with coordinate systems: - ▸ 'ST_downscaled_hires', with elements: - ST_downscaled_hires_image (Images), locations (Shapes) - ▸ 'ST_downscaled_lowres', with elements: - ST_downscaled_lowres_image (Images), locations (Shapes) - ▸ 'ST_fullres', with elements: - ST_fullres_image (Images), locations (Shapes) - ``` - + >>> from hest import load_hest + >>> hest_data = load_hest('../hest_data', id_list=['TENX68']) + >>> st = hest_data[0] + >>> st.to_spatial_data(fullres=True) + SpatialData object + ├── Images + │ ├── 'ST_downscaled_hires_image': SpatialImage[cyx] (3, 4779, 2586) + │ ├── 'ST_downscaled_lowres_image': SpatialImage[cyx] (3, 1000, 541) + │ └── 'ST_fullres_image': DataTree[cyx] (3, 38232, 20690), (3, 19116, 10345) + ├── Shapes + │ └── 'locations': GeoDataFrame shape: (1657, 2) (2D shapes) + └── Tables + └── 'table': AnnData (1657, 18085) + with coordinate systems: + ▸ 'ST_downscaled_hires', with elements: + ST_downscaled_hires_image (Images), locations (Shapes) + ▸ 'ST_downscaled_lowres', with elements: + ST_downscaled_lowres_image (Images), locations (Shapes) + ▸ 'ST_fullres', with elements: + ST_fullres_image (Images), locations (Shapes) """ # imports specific to spatial data conversion @@ -711,9 +653,15 @@ def read_hest_wsi(wsi: WSI, width, height): else: new_table = self.adata.copy() - return SpatialData(tables=new_table, images=images, shapes=shapes) + return SpatialData(tables={'table': new_table}, images=images, shapes=shapes) - def ensembl_id_to_gene(self): + def ensembl_id_to_gene(self) -> None: + """ + Converts ensemble gene IDs using Biomart annotations and filter out genes with no matching Ensembl ID for the current object + + Args: + filter_na (bool): whenever to filter genes that are not valid ensemble IDs. Defaults to False. + """ ensembl_id_to_gene(self) @@ -723,19 +671,21 @@ def __init__(self, img: Union[np.ndarray, str], pixel_size: float, meta: Dict = {}, - tissue_seg: TissueMask=None, tissue_contours: gpd.GeoDataFrame=None, shapes: List[LazyShapes]=[] ): - super().__init__(adata, img, pixel_size, meta, tissue_seg=tissue_seg, tissue_contours=tissue_contours, shapes=shapes) + super().__init__(adata, img, pixel_size, meta, tissue_contours=tissue_contours, shapes=shapes) class VisiumHDHESTData(HESTData): + """ + Object representing a pooled Visium HD sample along with a full resolution H&E image and associated metadata + """ + def __init__(self, adata: sc.AnnData, # type: ignore img: Union[np.ndarray, str], pixel_size: float, meta: Dict = {}, - tissue_seg: TissueMask=None, tissue_contours: gpd.GeoDataFrame=None, shapes: List[LazyShapes]=[] ): @@ -747,9 +697,10 @@ def __init__(self, img (Union[np.ndarray, str]): Full resolution image corresponding to the ST data, if passed as a path (str) the image is lazily loaded meta (Dict): metadata dictionary containing information such as the pixel size, or QC metrics attached to that sample shapes (List[LazyShapes]): dictionary of shapes, note that these shapes will be lazily loaded. Default: [] - tissue_seg (TissueMask): tissue mask for that sample """ - super().__init__(adata, img, pixel_size, meta, tissue_seg, tissue_contours, shapes) + super().__init__(adata, img, pixel_size, meta, tissue_contours, shapes) + + class STHESTData(HESTData): def __init__(self, @@ -757,7 +708,6 @@ def __init__(self, img: Union[np.ndarray, str], pixel_size: float, meta: Dict = {}, - tissue_seg: TissueMask=None, tissue_contours: gpd.GeoDataFrame=None, shapes: List[LazyShapes]=[] ): @@ -768,9 +718,8 @@ def __init__(self, pixel_size (float): pixel_size of WSI im um/px, this pixel size will be used to perform operations on the slide, such as patching and segmenting img (Union[np.ndarray, str]): Full resolution image corresponding to the ST data, if passed as a path (str) the image is lazily loaded meta (Dict): metadata dictionary containing information such as the pixel size, or QC metrics attached to that sample - tissue_seg (TissueMask): tissue mask for that sample """ - super().__init__(adata, img, pixel_size, meta, tissue_seg, tissue_contours, shapes) + super().__init__(adata, img, pixel_size, meta, tissue_contours, shapes) class XeniumHESTData(HESTData): @@ -780,7 +729,6 @@ def __init__( img: Union[np.ndarray, openslide.OpenSlide, CuImage], # type: ignore pixel_size: float, meta: Dict = {}, - tissue_seg: TissueMask=None, tissue_contours: gpd.GeoDataFrame=None, shapes: List[LazyShapes]=[], xenium_nuc_seg: pd.DataFrame=None, @@ -788,7 +736,8 @@ def __init__( cell_adata: sc.AnnData=None, # type: ignore transcript_df: pd.DataFrame=None, dapi_path: str=None, - alignment_file_path: str=None + alignment_file_path: str=None, + path_registrar: str=None ): """ class representing a single ST profile + its associated WSI image @@ -800,15 +749,15 @@ class representing a single ST profile + its associated WSI image pixel_size (float): pixel_size of WSI im um/px, this pixel size will be used to perform operations on the slide, such as patching and segmenting meta (Dict): metadata dictionary containing information such as the pixel size, or QC metrics attached to that sample shapes (List[LazyShapes]): dictionary of shapes, note that these shapes will be lazily loaded. Default: [] - tissue_seg (TissueMask): tissue mask for that sample xenium_nuc_seg (pd.DataFrame): content of a xenium nuclei contour file as a dataframe (nucleus_boundaries.parquet) xenium_cell_seg (pd.DataFrame): content of a xenium cell contour file as a dataframe (cell_boundaries.parquet) cell_adata (sc.AnnData): ST cell data, each row in adata.obs is a cell, each row in obsm is the cell location on the H&E image in pixels transcript_df (pd.DataFrame): dataframe of transcripts, each row is a transcript, he_x and he_y is the transcript location on the H&E image in pixels dapi_path (str): path to a dapi focus image alignment_file_path (np.ndarray): path to xenium alignment path + path_registrar (str): path to a valis registration registrar. """ - super().__init__(adata=adata, img=img, pixel_size=pixel_size, meta=meta, tissue_seg=tissue_seg, tissue_contours=tissue_contours, shapes=shapes) + super().__init__(adata=adata, img=img, pixel_size=pixel_size, meta=meta, tissue_contours=tissue_contours, shapes=shapes) self.xenium_nuc_seg = xenium_nuc_seg self.xenium_cell_seg = xenium_cell_seg @@ -816,6 +765,68 @@ class representing a single ST profile + its associated WSI image self.transcript_df = transcript_df self.dapi_path = dapi_path self.alignment_file_path = alignment_file_path + self.path_registrar = path_registrar + + @staticmethod + def from_paths( + adata_path: str, + img: Union[str, np.ndarray, openslide.OpenSlide, CuImage], # type: ignore + metrics_path: str, + cellvit_path: str = None, + tissue_contours_path: str = None, + xenium_cell_path: str = None, + xenium_nucleus_path: str = None, + transcripts_path: str = None + ) -> XeniumHESTData: + """ + Read a Xenium HEST sample from disk + + Args: + adata_path (str): path to .h5ad adata file containing ST data the + adata object must contain a downscaled image in ['spatial']['ST']['images']['downscaled_fullres'] + img (Union[str, np.ndarray, openslide.OpenSlide, CuImage]): path to a full resolution image (if passed as str) or full resolution image corresponding to the ST data, Openslide/CuImage are lazily loaded, use CuImage for GPU accelerated computation + pixel_size (float): pixel_size of WSI im um/px, this pixel size will be used to perform operations on the slide, such as patching and segmenting + metrics_path (str): metadata dictionary containing information such as the pixel size, or QC metrics attached to that sample + cellvit_path (str): path to a cell segmentation file in .geojson or .parquet. Defaults to None. + tissue_contours_path (str): path to a .geojson tissue contours file. Defaults to None. + xenium_cell_path (str): path to a .parquet xeniun cell segmentation file. Defaults to None. + xenium_nucleus_path (str): path to a .parquet xenium nucleus segmentation file. Defaults to None. + transcripts_path (str): path to a .parquet transcript dataframe. Defaults to None. + + Returns: + HESTData: HESTData object + """ + import scanpy as sc + + img = read_img_hest(img) + + tissue_contours = read_tissue_contours(tissue_contours_path) + + shapes = [] + if cellvit_path is not None: + shapes.append(LazyShapes(cellvit_path, 'cellvit', 'he')) + if xenium_cell_path is not None: + shapes.append(LazyShapes(xenium_cell_path, 'xenium_cell', 'he')) + if xenium_nucleus_path is not None: + shapes.append(LazyShapes(xenium_nucleus_path, 'xenium_nucleus', 'he')) + + transcripts = None + if transcripts_path is not None: + transcripts = pd.read_parquet(transcripts_path) + + adata = sc.read_h5ad(adata_path) + with open(metrics_path) as metrics_f: + metrics = json.load(metrics_f) + + return XeniumHESTData( + adata, + img, + metrics['pixel_size_um_estimated'], + metrics, + shapes=shapes, + tissue_contours=tissue_contours, + transcript_df=transcripts + ) def save( @@ -829,43 +840,215 @@ def save( save_cell_seg=False, save_nuclei_seg=False, qc=False, + nb_qc_patches=20, + verbose=True, **kwargs ): - """Save a HESTData object to `path` as follows: - - aligned_adata.h5ad (contains pseudo-visium pooled expressions for each spots + their location on the fullres image + a downscaled version of the fullres image) - - metrics.json (contains useful metrics) - - downscaled_fullres.jpeg (a downscaled version of the fullres image) - - aligned_fullres_HE.tif (the full resolution image) - - cells.geojson (cell segmentation if it exists) - - Optional: cells_xenium.geojson (if xenium cell segmentation is attached to this object) - - Optional: nuclei_xenium.geojson (if xenium cell segmentation is attached to this object) - - Optional: tissue_contours.geojson (contours of the tissue segmentation if it exists) + """ + Saves a Xenium HESTData object to the specified directory. + + The following files are generated at the destination `path`: + * **aligned_adata.h5ad**: Pseudo-visium pooled expressions, spot locations, and a downscaled image. + * **metrics.json**: Key performance and data metrics. + * **downscaled_fullres.jpeg**: Low-resolution version of the H&E image. + * **aligned_fullres_HE.tif**: The full-resolution H&E image. + * **cells.geojson**: Cell segmentation boundaries (if available). + * **Optional Files**: `cells_xenium.geojson`, `nuclei_xenium.geojson`, and `tissue_contours.geojson`. Args: - path (str): save location - save_img (bool): whenever to save the image at all (can save a lot of time if set to False) - pyramidal (bool, optional): whenever to save the full resolution image as pyramidal (can be slow to save, however it's sometimes necessary for loading large images in QuPath). Defaults to True. - bigtiff (bool, optional): whenever the bigtiff image is more than 4.1GB. Defaults to False. + path (str): The directory where the data will be saved. + save_img (bool): If True, saves the H&E images. Setting to False significantly reduces runtime. + pyramidal (bool, optional): If True, saves the full-res image as a pyramidal TIFF. + Recommended for large images to be opened in QuPath. Defaults to True. + bigtiff (bool, optional): Set to True if the output image exceeds 4.1GB. Defaults to False. + qc (bool, optional): If True, saves quality control (QC) patches and global transcript plots + to `path/qc/`. Defaults to False. + nb_qc_patches (int, optional): The number of random patches to generate if `qc` is True. + Defaults to 20. + + Returns: + None """ super().save(path, save_img, pyramidal, bigtiff, plot_pxl_size) if self.cell_adata is not None: self.cell_adata.write_h5ad(os.path.join(path, 'aligned_cells.h5ad')) if save_transcripts and self.transcript_df is not None: - self.transcript_df.to_parquet(os.path.join(path, 'aligned_transcripts.parquet')) + if is_dask_dataframe(self.transcript_df): + self.transcript_df.to_parquet(os.path.join(path, 'aligned_transcripts')) + merge_parquet(os.path.join(path, 'aligned_transcripts'), + os.path.join(path, 'aligned_transcripts.parquet')) + else: + self.transcript_df.to_parquet(os.path.join(path, 'aligned_transcripts.parquet')) if save_cell_seg: he_cells = self.get_shapes('tenx_cell', 'he').shapes if qc: - visualize_random_crops(None, self.wsi, plot_dir=path, seg=he_cells) + plot_dir = os.path.join(path, 'qc') + os.makedirs(plot_dir, exist_ok=True) + plot_xenium_align_qc(self.wsi, plot_dir=plot_dir, seg_cells=he_cells, nb=nb_qc_patches) + if verbose: + print(f"Saving aligned cell-segmentation...") he_cells.to_parquet(os.path.join(path, 'he_cell_seg.parquet')) - write_geojson(he_cells, os.path.join(path, f'he_cell_seg.geojson'), '', chunk=True) + write_geojson(he_cells, os.path.join(path, f'he_cell_seg.geojson')) if save_nuclei_seg: he_nuclei = self.get_shapes('tenx_nucleus', 'he').shapes + if qc: + plot_dir = os.path.join(path, 'qc') + os.makedirs(plot_dir, exist_ok=True) + plot_xenium_align_qc(self.wsi, plot_dir=plot_dir, seg_nuc=he_nuclei, nb=nb_qc_patches) + if verbose: + print(f"Saving aligned nuclei-segmentation...") he_nuclei.to_parquet(os.path.join(path, 'he_nucleus_seg.parquet')) - write_geojson(he_nuclei, os.path.join(path, f'he_nucleus_seg.geojson'), '', chunk=True) + write_geojson(he_nuclei, os.path.join(path, f'he_nucleus_seg.geojson')) + + + def register_dapi_he( + self, + he_path: str, + dapi_path: str, + max_non_rigid_registration_dim_px=10000 + ): + """ Micro-register the DAPI coordinate system to the H&E coordinate system. + + Warning Valis alignment might require a significant amount or RAM based on the number of transcripts, shapes and image size. + + Args: + he_path (str, optional): path to the H&E image, this image must be in generic pyramidal tiff! + dapi_path (str, optional): path to the raw Xenium DAPI image, either provide a path to `morphology_focus_0000.ome.tif` or `morphology_focus.ome.tif` based on availability. + max_non_rigid_registration_dim_px (bool, optional): maximum size of the WSI during micro registration. Defaults to 10000. + """ + logger.info('Registering Xenium DAPI to H&E...') + + if self.dapi_path is None and dapi_path is None: + raise ValueError(f"Either self.dapi_path must be set or dapi_path must be passed to the function") + + dapi_path = self.dapi_path if dapi_path is None else dapi_path + verify_paths([dapi_path]) + + path_registrar = register_dapi_he( + he_path, + dapi_path, + registrar_dir='valis', + name='registration', + max_non_rigid_registration_dim_px=max_non_rigid_registration_dim_px, + ) + + self.path_registrar = path_registrar + return path_registrar + + def warp_xenium_objects( + self, + save_dir: str, + dapi_path: str, + save_cells=False, + save_transcripts=False, + save_nuclei=False, + save_parquet=True, + save_geojson=True, + use_dask=False, + verbose=True + ): + """ + **Deprecated** use warp_and_save_xenium_objects instead. xenium objects based on the reigstrar + + Args: + save_dir (str): Save the aligned objects to: + - {save_dir}/he_cell_seg.parquet + - {save_dir}/he_nucleus_seg.parquet + - {save_dir}/aligned_transcripts.parquet + + dapi_path (str): _description_ + save_cells (bool, optional): Whenever to transform and save warped cells. Defaults to False. + save_transcripts (bool, optional): Whenever to transform and save warped transcripts. Defaults to False. + save_nuclei (bool, optional): Whenever to transform and save warped nuclei. Defaults to False. + """ + + warnings.warn( + "warp_xenium_objects is deprecated and will be removed in a future version. " + "Please use 'warp_and_save_xenium_objects' instead.", + DeprecationWarning, + stacklevel=2 + ) + + if not self.path_registrar: + raise ValueError(f"No registration found for this xenium object, please execute `register_dapi_he` on this object first.") + + dapi_cells = self.get_shapes('tenx_cell', 'dapi').shapes if save_cells else None + dapi_nuclei = self.get_shapes('tenx_nucleus', 'dapi').shapes if save_nuclei else None + transcript_df = self.transcript_df if save_transcripts else None + + warp_and_save_xenium_objects( + path_registrar=self.path_registrar, + dapi_path=dapi_path, + save_dir=save_dir, + dapi_cells=dapi_cells, + dapi_transcripts=transcript_df, + dapi_nuclei=dapi_nuclei, + use_dask=use_dask, + verbose=verbose, + save_parquet=save_parquet, + save_geojson=save_geojson, + ) + + + def warp_and_save_xenium_objects( + self, + dapi_path: str, + save_dir: str, + dapi_cells: str=None, + dapi_transcripts: str=None, + dapi_nuclei: str=None, + use_dask=True, + verbose=True, + save_parquet=True, + save_geojson=True, + ) -> None: + """ Wrap Xenium transcripts, cells and nuclei using Valis non-rigid micro-registration and save them. + + Args: + dapi_path (str): dapi slide filename in the Valis registrar + save_dir (str): where to save warped objects. Objects will be saved to: + - save_dir/he_cell_seg.parquet + - save_dir/he_nucleus_seg.parquet + - save_dir/aligned_transcripts + dapi_cells (str, optional): path to xenium .parquet cell bondaries, usually **/cell_boundaries.parquet. Defaults to None. + dapi_transcripts (str, optional): path to xenium .parquet nucleus bondaries, usually **/nucleus_boundaries.parquet. Defaults to None. + dapi_nuclei (str, optional): path to xenium .parquet transcripts, usually **/transcripts.parquet. Defaults to None. + use_dask (bool, optional): whenever to use dask to process larger than RAM data, highly recommended for all Xenium samples. Defaults to True. + verbose (bool, optional): verbose flag. Defaults to True. + save_parquet (bool, optional): whenever to save objects as parquet. Defaults to True. + save_geojson (bool, optional): whenever to save objects as geojson. Defaults to True. + + Example: + >>> cells_path = "./xenium_out/cell_boundaries.parquet" + >>> st.warp_and_save_xenium_objects( + ... save_dir="warped_xenium", + ... dapi_path="morphology_focus.ome.tif", + ... dapi_cells=cells_path, + ... use_dask=True + ... ) + >>> if cells is not None: + ... print(f"Warped {len(cells)} cells using Dask: {type(cells)}") + """ + if not self.path_registrar: + raise ValueError(f"No registration found for this xenium object, please execute `register_dapi_he` on this object first.") + + warp_and_save_xenium_objects( + self.path_registrar, + dapi_path, + save_dir, + dapi_cells, + dapi_transcripts, + dapi_nuclei, + use_dask, + verbose, + save_parquet, + save_geojson, + ) + def align_with_valis(self, save_dir: str, he_path: str, dapi_path: str, align_nuclei=True, align_cells=True, align_transcripts=True, verbose=True, save_geojson=True): @@ -898,73 +1081,29 @@ def align_with_valis(self, save_dir: str, he_path: str, dapi_path: str, align_nu dapi_path = self.dapi_path if dapi_path is None else dapi_path verify_paths([dapi_path]) - dapi_cells = self.get_shapes('tenx_cell', 'dapi').shapes if align_cells else None - dapi_nuclei = self.get_shapes('tenx_nucleus', 'dapi').shapes if align_nuclei else None - transcript_df = self.transcript_df if align_transcripts else None - if verbose: print('finished reading shapes') reg_config = {} - - warped_cells, warped_nuclei, transcript_df = preprocess_cells_xenium( - he_path, - dapi_path, - dapi_cells, - dapi_nuclei, - transcript_df, - reg_config, - 'valis', - registration_kwargs={} - ) - - print('Saving warped cells/nuclei...') - if align_cells: - warped_cells.to_parquet(os.path.join(save_dir, f'he_cell_seg.parquet')) - - if save_geojson: - write_geojson(warped_cells, os.path.join(save_dir, f'he_cell_seg.geojson'), '', chunk=True) - if align_nuclei: - warped_nuclei.to_parquet(os.path.join(save_dir, f'he_nucleus_seg.parquet')) - if save_geojson: - write_geojson(warped_nuclei, os.path.join(save_dir, f'he_nucleus_seg.geojson'), '', chunk=True) - if align_transcripts: - self.transcript_df.to_parquet(os.path.join(save_dir, f'aligned_transcripts.parquet')) - - -def read_HESTData( - adata_path: str, - img: Union[str, np.ndarray, openslide.OpenSlide, CuImage], # type: ignore - metrics_path: str, - mask_path_pkl: str = None, # Deprecated - mask_path_jpg: str = None, # Deprecated - cellvit_path: str = None, - tissue_contours_path: str = None, - xenium_cell_path: str = None, - xenium_nucleus_path: str = None, - transcripts_path: str = None -) -> HESTData: - """ Read a HEST sample from disk - Args: - adata_path (str): path to .h5ad adata file containing ST data the - adata object must contain a downscaled image in ['spatial']['ST']['images']['downscaled_fullres'] - img (Union[str, np.ndarray, openslide.OpenSlide, CuImage]): path to a full resolution image (if passed as str) or full resolution image corresponding to the ST data, Openslide/CuImage are lazily loaded, use CuImage for GPU accelerated computation - pixel_size (float): pixel_size of WSI im um/px, this pixel size will be used to perform operations on the slide, such as patching and segmenting - metrics_path (str): metadata dictionary containing information such as the pixel size, or QC metrics attached to that sample - mask_path_pkl (str): *Deprecated* path to a .pkl file containing the tissue segmentation contours. Defaults to None. - mask_path_jpg (str): *Deprecated* path to a .jog file containing the greyscale tissue segmentation mask. Defaults to None. - cellvit_path (str): path to a cell segmentation file in .geojson or .parquet. Defaults to None. - tissue_contours_path (str): path to a .geojson tissue contours file. Defaults to None. - xenium_cell_path (str): path to a .parquet xeniun cell segmentation file. Defaults to None. - xenium_nucleus_path (str): path to a .parquet xenium nucleus segmentation file. Defaults to None. - transcripts_path (str): path to a .parquet transcript dataframe. Defaults to None. + if not self.path_registrar: + self.path_registrar = self.register_dapi_he( + he_path, + dapi_path, + max_non_rigid_registration_dim_px=reg_config.get('max_non_rigid_registration_dim_px', 10000) + ) + self.warp_xenium_objects( + save_dir, + dapi_path, + save_cells=align_cells, + save_transcripts=align_transcripts, + save_nuclei=align_nuclei, + save_geojson=save_geojson + ) - Returns: - HESTData: HESTData object - """ +def read_img_hest(img): try: from cucim import CuImage except ImportError: @@ -980,9 +1119,10 @@ def read_HESTData( else: img = openslide.OpenSlide(img) width, height = img.dimensions - + return img + +def read_tissue_contours(tissue_contours_path): tissue_contours = None - tissue_seg = None if tissue_contours_path is not None: with open(tissue_contours_path) as f: lines = f.read() @@ -992,48 +1132,8 @@ def read_HESTData( tissue_contours = old_geojson_to_new(gdf) else: tissue_contours = gpd.read_file(tissue_contours_path) - - elif mask_path_pkl is not None and mask_path_jpg is not None: - tissue_seg = load_tissue_mask(mask_path_pkl, mask_path_jpg, width, height) - - shapes = [] - if cellvit_path is not None: - shapes.append(LazyShapes(cellvit_path, 'cellvit', 'he')) - if xenium_cell_path is not None: - shapes.append(LazyShapes(xenium_cell_path, 'xenium_cell', 'he')) - if xenium_nucleus_path is not None: - shapes.append(LazyShapes(xenium_nucleus_path, 'xenium_nucleus', 'he')) - - transcripts = None - if transcripts_path is not None: - transcripts = pd.read_parquet(transcripts_path) - - adata = sc.read_h5ad(adata_path) - with open(metrics_path) as metrics_f: - metrics = json.load(metrics_f) - - if transcripts is not None: - return XeniumHESTData( - adata, - img, - metrics['pixel_size_um_estimated'], - metrics, - tissue_seg=tissue_seg, - shapes=shapes, - tissue_contours=tissue_contours, - transcript_df=transcripts - ) - else: - return HESTData( - adata, - img, - metrics['pixel_size_um_estimated'], - metrics, - tissue_seg=tissue_seg, - shapes=shapes, - tissue_contours=tissue_contours - ) - + return tissue_contours + def mask_and_patchify_bench(meta_df: pd.DataFrame, save_dir: str, use_mask=True, keep_largest=None): i = 0 @@ -1043,7 +1143,7 @@ def mask_and_patchify_bench(meta_df: pd.DataFrame, save_dir: str, use_mask=True, adata_path = f'/mnt/sdb1/paul/images/adata/{id}.h5ad' metrics_path = os.path.join(get_path_from_meta_row(row), 'processed', 'metrics.json') - hest_obj = read_HESTData(adata_path, img_path, metrics_path) + hest_obj = HESTData.from_paths(adata_path, img_path, metrics_path) keep_largest_args = keep_largest[i] if keep_largest is not None else False @@ -1157,7 +1257,7 @@ def __len__(self): return len(self.id_list) def iter_hest(hest_dir: str, id_list: List[str] = None, **read_kwargs) -> HESTIterator: - """ Iterate through the HEST samples contained in `hest_dir` + """ Iterate through HEST samples contained in `hest_dir` Args: hest_dir (str): hest directory containing folders: st, wsis, metadata, tissue_seg (optional) @@ -1175,14 +1275,10 @@ def _read_st(hest_dir, st_filename, load_transcripts=False): img_path = os.path.join(hest_dir, 'wsis', f'{id}.tif') meta_path = os.path.join(hest_dir, 'metadata', f'{id}.json') - masks_path_pkl = None - masks_path_jpg = None verify_paths([adata_path, img_path, meta_path], suffix='\nHave you downloaded the dataset? (https://huggingface.co/datasets/MahmoodLab/hest)') if os.path.exists(os.path.join(hest_dir, 'tissue_seg')): - masks_path_pkl = find_first_file_endswith(os.path.join(hest_dir, 'tissue_seg'), f'{id}_mask.pkl') - masks_path_jpg = find_first_file_endswith(os.path.join(hest_dir, 'tissue_seg'), f'{id}_mask.jpg') tissue_contours_path = find_first_file_endswith(os.path.join(hest_dir, 'tissue_seg'), f'{id}_contours.geojson') cellvit_path = None @@ -1204,21 +1300,26 @@ def _read_st(hest_dir, st_filename, load_transcripts=False): transcripts_path = None if load_transcripts: transcripts_path = find_first_file_endswith(os.path.join(hest_dir, 'transcripts'), f'{id}_transcripts.parquet') - - st = read_HESTData( - adata_path, - img_path, - meta_path, - masks_path_pkl, - masks_path_jpg, - cellvit_path=cellvit_path, - tissue_contours_path=tissue_contours_path, - xenium_cell_path=xenium_cell_path, - xenium_nucleus_path=xenium_nucleus_path, - transcripts_path=transcripts_path - ) - return st - + + if transcripts_path is not None: + return XeniumHESTData.from_paths( + adata_path, + img_path, + meta_path, + cellvit_path=cellvit_path, + tissue_contours_path=tissue_contours_path, + xenium_cell_path=xenium_cell_path, + xenium_nucleus_path=xenium_nucleus_path, + transcripts_path=transcripts_path + ) + else: + return HESTData.from_paths( + adata_path, + img_path, + meta_path, + cellvit_path=cellvit_path, + tissue_contours_path=tissue_contours_path, + ) def load_hest(hest_dir: str, id_list: List[str] = None) -> List[HESTData]: diff --git a/src/hest/__init__.py b/src/hest/__init__.py index ab37d26..52d767a 100644 --- a/src/hest/__init__.py +++ b/src/hest/__init__.py @@ -3,7 +3,7 @@ from .utils import tiff_save, find_pixel_size_from_spot_coords, write_10X_h5, get_k_genes, SpotPacking from .autoalign import autoalign_visium from .readers import * -from .HESTData import HESTData, read_HESTData, load_hest, iter_hest, ensembl_id_to_gene +from .HESTData import HESTData, load_hest, iter_hest, ensembl_id_to_gene from .segmentation.cell_segmenters import segment_cellvit __all__ = [ @@ -11,7 +11,6 @@ 'find_pixel_size_from_spot_coords', 'get_k_genes', 'SpotPacking', - 'read_HESTData', 'load_hest', 'Reader', 'XeniumReader', diff --git a/src/hest/io/__init__.py b/src/hest/io/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/hest/io/seg_readers.py b/src/hest/io/seg_readers.py index 34119c2..d9e8071 100644 --- a/src/hest/io/seg_readers.py +++ b/src/hest/io/seg_readers.py @@ -1,6 +1,10 @@ +from __future__ import annotations + from collections import defaultdict from concurrent.futures import ProcessPoolExecutor, as_completed import json +import os +from typing import Optional, Union import warnings from abc import abstractmethod @@ -11,7 +15,7 @@ from shapely.geometry.polygon import Point, Polygon from tqdm import tqdm -from hest.utils import align_xenium_df, get_n_threads +from hest.utils import align_xenium_df, get_n_threads, is_dask_gdf, read_parquet_dask def _process(x, extra_props, index_key, class_name): @@ -61,61 +65,73 @@ def _read_geojson(path, class_name=None, extra_props=False, index_key=None) -> g class GDFReader: + """ Lazily read shapes such that read_gdf is called at compute time """ + @abstractmethod - def read_gdf(self, path) -> gpd.GeoDataFrame: + def read_gdf(self, path: str) -> Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]: + """ Read shapes + + Args: + path (str): path to shapes + + Returns: + Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]: shapes + """ pass -def fn(block, i): - logger.debug(f'start fn block {i}') + +def groupby_shape(df, col, n_threads=0, col_shape='xy'): + block = df[[col, col_shape]] + groups = defaultdict(lambda: []) - [groups[row[0]].append(row[1]) for row in block] - g = np.array([Polygon(value) for _, value in groups.items()]) - key = np.array([key for key, _ in groups.items()]) - logger.debug(f'finish fn block {i}') - return np.column_stack((key, g)) - -def groupby_shape(df, col, n_threads, col_shape='xy'): - n_chunks = n_threads - - if n_threads >= 1: - l = len(df) // n_chunks - start = 0 - chunk_lens = [] - while start < len(df): - end = min(start + l, len(df)) - while end < len(df) and df.iloc[end][col] == df.iloc[end - 1][col]: - end += 1 - chunk_lens.append((start, end)) - start = end - - dfs = [] - with ProcessPoolExecutor(max_workers=n_threads) as executor: - future_results = [executor.submit(fn, df[[col, col_shape]].iloc[start:end].values, start) for start, end in chunk_lens] + [groups[row[0]].append(row[1]) for row in block.values] + key, g = zip(*[ + (key, Polygon(value)) + for key, value in groups.items() + if len(value) >= 4 or print(f"Warning: key {key} has less than 4 points ({len(value)}), skipping") + ]) - for future in as_completed(future_results): - dfs.append(future.result()) - - concat = np.concatenate(dfs) - else: - concat = fn(df[[col, col_shape]].values, 0) + key = np.array(key) + g = np.array(g) + + concat = np.column_stack((key, g)) gdf = gpd.GeoDataFrame(geometry=concat[:, 1]) gdf.index = concat[:, 0] + import gc + gc.collect() + return gdf class XeniumParquetCellReader(GDFReader): + """ Xenium parquet shape reader """ - def __init__(self, pixel_size_morph=None, alignment_matrix=None): + def __init__( + self, + pixel_size_morph: Optional[float]=None, + alignment_matrix=None, + use_dask=False + ): + """ Xenium parquet shape reader + + Args: + pixel_size_morph (Optional[float], optional): pixel size of DAPI in um/px. Defaults to None. + alignment_matrix (np.ndarray, optional): optional alignment matrix. Defaults to None. + use_dask (bool, optional): whenever to load as a dask geodataframe. Defaults to False. + """ + self.pixel_size_morph = pixel_size_morph self.alignment_matrix = alignment_matrix + self.use_dask = use_dask - def read_gdf(self, path, n_workers=0) -> gpd.GeoDataFrame: + def read_gdf(self, path, n_workers=0) -> Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]: - df = pd.read_parquet(path) + if self.use_dask: + df = read_parquet_dask(path, nb_partitions=30) + else: + df = pd.read_parquet(path) - - if self.alignment_matrix is not None: df = align_xenium_df( df, @@ -128,18 +144,54 @@ def read_gdf(self, path, n_workers=0) -> gpd.GeoDataFrame: else: df['vertex_x'], df['vertex_y'] = df['vertex_x'] / self.pixel_size_morph, df['vertex_y'] / self.pixel_size_morph - df['xy'] = list(zip(df['vertex_x'], df['vertex_y'])) - df = df.drop(['vertex_x', 'vertex_y'], axis=1) - - n_threads = get_n_threads(n_workers) + if self.use_dask: + def create_xy_column(pdf): + pdf['xy'] = list(zip(pdf['vertex_x'], pdf['vertex_y'])) + return pdf + + df = df.map_partitions(create_xy_column) + else: + df['xy'] = list(zip(df['vertex_x'], df['vertex_y'])) + df = df.drop(['vertex_x', 'vertex_y'], axis=1) - gdf = groupby_shape(df, 'cell_id', n_threads) + if self.use_dask: + import dask_geopandas + from geopandas.array import GeometryDtype + + gdf = dask_geopandas.from_dask_dataframe(df) + gdf_template = gpd.GeoDataFrame( + {'geometry': gpd.GeoSeries(dtype=GeometryDtype())}, + crs="EPSG:4326", + index=pd.Index([], dtype="string", name="index"), + ) + gdf = gdf.map_partitions(groupby_shape, 'cell_id', meta=gdf_template) + else: + gdf = groupby_shape(df, 'cell_id') return gdf + class GDFParquetCellReader(GDFReader): + """ Geopandas parquet shape reader """ - def read_gdf(self, path) -> gpd.GeoDataFrame: - return gpd.read_parquet(path) + def __init__(self, use_dask=False, **kwargs): + """ Geopandas parquet shape reader + + Args: + use_dask (bool, optional): whenever to load as a dask geodataframe. Defaults to False. + """ + self.use_dask = use_dask + + + def read_gdf(self, path) -> Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]: + if self.use_dask: + import dask_geopandas as dgpd + import pyarrow.parquet as pq + parquet_file = pq.ParquetFile(path) + total_row_groups = parquet_file.num_row_groups + row_groups_per_partition = max(1, total_row_groups // 30) + return dgpd.read_parquet(path, split_row_groups=row_groups_per_partition) + else: + return gpd.read_parquet(path) class GeojsonCellReader(GDFReader): @@ -158,83 +210,151 @@ def read_gdf(self, path) -> gpd.GeoDataFrame: return gdf -def write_geojson(gdf: gpd.GeoDataFrame, path: str, category_key: str, extra_prop=False, uniform_prop=True, index_key: str=None, chunk=False) -> None: - - if isinstance(gdf.geometry.iloc[0], Point): - geometry = 'MultiPoint' - elif isinstance(gdf.geometry.iloc[0], Polygon): - geometry = 'MultiPolygon' - else: - raise ValueError(f"gdf.geometry[0] must be of type Point or Polygon, got {type(gdf.geometry.iloc[0])}") +class XeniumTranscriptsReader(GDFReader): + """ Xenium transcript shape reader """ + def __init__(self, pixel_size_morph: float, use_dask=False, **kwargs): + """ Xenium transcript shape reader + + Args: + pixel_size_morph (float): pixel size of DAPI in um/px + use_dask (bool, optional): whenever to load as a dask geodataframe. Defaults to False. + """ + self.use_dask = use_dask + self.pixel_size_morph = pixel_size_morph - if chunk: - n = 10 - l = (len(gdf) // n) + 1 - s = [] - for i in range(n): - s.append(np.repeat(i, l)) - cls = np.concatenate(s) - - gdf['_chunked'] = cls[:len(gdf)] - category_key = '_chunked' + def read_gdf(self, path) -> Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]: + if self.use_dask: + import dask_geopandas as dgpd + transcripts_df = read_parquet_dask(path, nb_partitions=30) + transcripts_df['x_location'] = transcripts_df['x_location'] / self.pixel_size_morph + transcripts_df['y_location'] = transcripts_df['y_location'] / self.pixel_size_morph + transcripts_df['geometry'] = dgpd.points_from_xy(transcripts_df, 'x_location', 'y_location') + transcripts_gdf = dgpd.from_dask_dataframe(transcripts_df) + else: + transcripts_gdf = gpd.GeoDataFrame(path, geometry=gpd.points_from_xy( + transcripts_df['x_location'] / self.pixel_size_morph, + transcripts_df['y_location'] / self.pixel_size_morph)) + return transcripts_gdf + + +class HESTXeniumTranscriptsReader(GDFReader): + """ HEST Xenium transcript reader """ - groups = np.unique(gdf[category_key]) - colors = generate_colors(groups) - cells = [] - for group in tqdm(groups): + def __init__(self, use_dask=False, **kwargs): + """ HEST Xenium transcript reader - slice = gdf[gdf[category_key] == group] - shapes = slice.geometry - - properties = { - "objectType": "annotation", - "classification": { - "name": str(group), - "color": colors[group] - } - } - - if extra_prop: - props = {} - col_exclude = [category_key, 'geometry'] - if index_key is not None: - col_exclude.append(index_key) - for col in [c for c in gdf.columns if c not in col_exclude]: - if uniform_prop: - unique = np.unique(slice[col]) - if len(unique) != 1: - warnings.warn(f"extra property {col} is not uniform for group {group}, found {unique}") - props[col] = slice[col].iloc[0] - - properties = {**properties, **props} - - if index_key is not None: - key = index_key - props = {} - mask = (slice[key] == True).values - props = {key: np.arange(len(mask))[mask].tolist()} - properties = {**properties, **props} + Args: + use_dask (bool, optional): whenever to load as a dask geodataframe. Defaults to False. + """ + self.use_dask = use_dask + + + def read_gdf(self, path) -> Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]: + if self.use_dask: + import dask_geopandas as dgpd + transcripts_df = read_parquet_dask(path, nb_partitions=30) + transcripts_df['geometry'] = dgpd.points_from_xy(transcripts_df, 'dapi_x', 'dapi_y') + transcripts_gdf = dgpd.from_dask_dataframe(transcripts_df) + else: + transcripts_gdf = gpd.GeoDataFrame(path, geometry=gpd.points_from_xy( + transcripts_df['dapi_x'], + transcripts_df['dapi_y'])) + return transcripts_gdf + + +def _write_geojson( + gdf: gpd.GeoDataFrame, + path: str, + geometry=None, + partition_info=None, +): + colors = generate_colors(['all', 'test']) + shapes = gdf.geometry + + if partition_info is not None: + p_index = partition_info['number'] + else: + p_index = -1 - if isinstance(gdf.geometry.iloc[0], Point): - shapes = [[point.x, point.y] for point in shapes] - elif isinstance(gdf.geometry.iloc[0], Polygon): - shapes = [[[[x, y] for x, y in polygon.exterior.coords]] for polygon in shapes] - cell = { - 'type': 'Feature', - 'id': (str(id(path)) + '-id-' + str(group)).replace('.', '-'), - 'geometry': { - 'type': geometry, - 'coordinates': shapes - }, - "properties": properties + properties = { + "objectType": "detection", + "classification": { + "name": str('group'), + "color": colors['test'] } - cells.append(cell) + } + + if isinstance(gdf.geometry.iloc[0], Point): + shapes = [[point.x, point.y] for point in shapes] + elif isinstance(gdf.geometry.iloc[0], Polygon): + shapes = [[[[x, y] for x, y in polygon.exterior.coords]] for polygon in shapes] + cells = [ + { + 'type': 'Feature', + 'geometry': { + 'type': geometry, + 'coordinates': shape + }, + "properties": properties + } for shape in shapes + ] + + partition_str = str(p_index) if p_index >= 0 else '' + path = path if partition_str == '' else os.path.join(path.removesuffix('.geojson'), f'part.{partition_str}.geojson') with open(path, 'w') as f: - json.dump(cells, f, indent=4) - + json.dump({ + "type": "FeatureCollection", + "features": cells} + , f) + + +def write_geojson( + gdf: Union[gpd.GeoDataFrame, dgpd.GeoDataFrame], + path: str, +) -> None: + """ Write a (dask) geodataframe in optimized QuPath geojson detection format. + + Args: + gdf (Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]): _description_ + path (str): _description_ + + Raises: + ValueError: _description_ + ValueError: _description_ + """ + if not path.endswith('.geojson'): + raise ValueError(f"path must end in .geojson") + + use_dask = is_dask_gdf(gdf) + + first_geom = gdf.geometry.head(1).values[0] + g_type = first_geom.geom_type + + if g_type == 'Point': + geometry = 'MultiPoint' + elif g_type == 'Polygon': + geometry = 'Polygon' + else: + raise ValueError( + f"gdf geometry must be Point or Polygon, got {g_type}" + ) + + + if use_dask: + meta = gdf.head(0).copy() + + from geopandas.array import GeometryDtype + meta['geometry'] = gpd.GeoSeries(dtype=GeometryDtype()) + + meta = meta.set_geometry('geometry').set_crs("EPSG:4326") + os.makedirs(path.removesuffix('.geojson'), exist_ok=True) + gdf.map_partitions(_write_geojson, path, geometry, + meta=meta).compute() + else: + _write_geojson(gdf, path, geometry) def generate_colors(names): diff --git a/src/hest/pipeline.py b/src/hest/pipeline.py index 6d4c694..a3ebb5a 100644 --- a/src/hest/pipeline.py +++ b/src/hest/pipeline.py @@ -4,235 +4,28 @@ import json import os import traceback -from abc import abstractmethod from typing import Tuple, Union import geopandas as gpd import numpy as np import openslide -import pandas as pd -import yaml from hestcore.wsi import WSI from loguru import logger from tqdm import tqdm from hest.HESTData import (VisiumHDHESTData, XeniumHESTData) from hest.io.seg_readers import GeojsonCellReader, read_gdf, write_geojson -from hest.readers import VisiumHDReader, XeniumReader, read_and_save -from hest.registration import preprocess_cells_xenium, register_dapi_he, warp_gdf_valis +from hest.readers import read_and_save +from hest.registration import preprocess_cells_xenium from hest.segmentation.cell_segmenters import (bin_per_cell, cell_segmenter_factory) -from hest.subtyping.subtyping import assign_cell_types -from hest.utils import (ALIGNED_HE_FILENAME, check_arg, +from hest.utils import (ALIGNED_HE_FILENAME, deprecated, find_first_file_endswith, get_col_selection, get_path_from_meta_row, - print_resource_usage, visualize_random_crops) - - -class ProcessingPipeline: - st = None - - def __init__(self, config, full_exp_dir): - self.config = config - self.full_exp_dir = full_exp_dir - - @abstractmethod - def on_skip_preprocessing(self, nuc_gdf, cell_adata) -> Tuple[sc.AnnData, gpd.GeoDataFrame]: - pass - - @abstractmethod - def preprocess(self) -> Tuple[sc.AnnData, gpd.GeoDataFrame, gpd.GeoDataFrame]: - pass - - -class XeniumProcessingPipeline(ProcessingPipeline): - - def on_skip_preprocessing(self, nuc_gdf, cell_adata) -> Tuple[sc.AnnData, gpd.GeoDataFrame]: - if nuc_gdf is None: - pass #TODO - if cell_adata is None: - import scanpy as sc - data_dir = self.config.get('data_dir') - st = XeniumReader().auto_read(data_dir, load_img=False) - cell_adata = sc.read_h5ad(st.cell_adata_path) - return cell_adata, nuc_gdf - - def preprocess(self) -> Tuple[sc.AnnData, gpd.GeoDataFrame, gpd.GeoDataFrame]: - preprocessing_conf = self.config.get('preprocessing', None) - data_dir = self.config.get('data_dir') - - st = XeniumReader().auto_read(data_dir) - - reg_config = preprocessing_conf.get('registration', {}) - - cell_adata = st.cell_adata - - cell_gdf, nuc_gdf = preprocess_cells_xenium( - st.wsi, - st.dapi_path, - st.get_shapes('tenx_cells', 'dapi').shapes, - st.get_shapes('tenx_nuclei', 'dapi').shapes, - reg_config, - self.full_exp_dir - ) - - return cell_adata, cell_gdf, nuc_gdf - - -class VisiumHDProcessingPipeline(ProcessingPipeline): - - def on_skip_preprocessing(self, nuc_gdf, cell_adata) -> Tuple[sc.AnnData, gpd.GeoDataFrame]: - if cell_adata is None: - raise ValueError('cell_adata_path is required if preprocessing is skipped') - return cell_adata, nuc_gdf - - def preprocess(self) -> Tuple[sc.AnnData, gpd.GeoDataFrame, gpd.GeoDataFrame]: - data_dir = self.config.get('data_dir') - st = VisiumHDReader().auto_read(data_dir) - - bc_matrix_2um_path = find_first_file_endswith(st.square_2um_path, 'filtered_feature_bc_matrix.h5') - bin_positions_2um_path = find_first_file_endswith(st.square_2um_path, 'tissue_positions.parquet') - - if bc_matrix_2um_path is None or bin_positions_2um_path is None: - raise FileNotFoundError(f"Make sure that your directory {data_dir} has a square_002um folder contaning those files: filtered_feature_bc_matrix.h5, tissue_positions.parquet") - - preprocessing_conf = self.config.get('preprocessing', None) - - segment_config = preprocessing_conf.get('segmentation', {}) - binning_config = preprocessing_conf.get('cell_binning', {}) - - nuclei_path = preprocessing_conf.get('nuclei_path', None) - if nuclei_path is not None: - logger.info("nuclei_path is specified, bypass segmentation") - - cell_adata, cell_gdf, nuc_gdf = preprocess_cells_visium_hd( - st.wsi, - self.full_exp_dir, - st.pixel_size, - bc_matrix_2um_path, - bin_positions_2um_path, - segment_config, - binning_config, - segment_config.get('method', 'cellvit'), - nuclei_path - ) - - return cell_adata, cell_gdf, nuc_gdf - - -def process_from_config(config_path: str): - with open(config_path) as f: - config = yaml.safe_load(f) - - if 'preprocessing' not in config and 'cell_subtyping' not in config: - raise ValueError("Please provide at least one of these steps in config: ['preprocessing', 'cell_subtyping']") - - result_dir = config.get('results_dir') - name = config.get('name', '') - full_exp_dir = os.path.join(result_dir, name) - os.makedirs(full_exp_dir, exist_ok=True) - - technology = config.get('technology').lower() - check_arg(technology, 'technology', ['xenium', 'visium-hd']) - - preprocessing_conf = config.get('preprocessing', None) - - if technology == 'xenium': - pipeline = XeniumProcessingPipeline(config, full_exp_dir) - elif technology == 'visium-hd': - pipeline = VisiumHDProcessingPipeline(config, full_exp_dir) - - # Can bypass preprocessing - if preprocessing_conf is None: - logger.info("no 'preprocessing' key found in config, skip preprocessing") - - # Try to infer cell_adata, nuc_gdf if preprocessing is skipped - nuclei_path = config.get('nuclei_path', None) - cell_adata_path = config.get('cell_adata_path', None) - - if nuclei_path is not None: - nuc_gdf = read_gdf(nuclei_path) - else: - nuc_gdf = None - if cell_adata_path is not None: - import scanpy as sc - cell_adata = sc.read_10x_h5(cell_adata_path) # TODO handle non 10x with catch - else: - cell_adata = None - - if nuclei_path is None or cell_adata_path is None: - logger.warning(f"No 'nuclei_path' or 'cell_adata_path' detected") - cell_adata, nuc_gdf = pipeline.on_skip_preprocessing(nuc_gdf, cell_adata) - else: - cell_adata, _, nuc_gdf = pipeline.preprocess() - - - if nuc_gdf is None: - logger.warning("No nuclear segmentation detected, make sure to enable preprocessing or to pass a custom 'nuclei_path'") - - subtyping_conf = config.get('cell_subtyping', None) - - if subtyping_conf is None: - logger.info("no 'cell_subtyping' key found in config, skip subtyping") - else: - matcher_kwargs = subtyping_conf.pop('matcher', {}) - atlas_name = matcher_kwargs.pop('atlas_name', None) - - types_adata, gdf_types = subtyping_pipeline(cell_adata, atlas_name, full_exp_dir, nuc_gdf, matcher_kwargs=matcher_kwargs, **subtyping_conf) - - types_adata.write_h5ad(os.path.join(full_exp_dir, 'types_adata.h5ad')) - - - -def subtyping_pipeline( - cell_adata: sc.AnnData, - atlas_name: str, - full_exp_dir, - gdf: gpd.GeoDataFrame=None, - matcher_kwargs: dict={}, - save_geojson=True, - save_parquet=True, - subtypes_path=None -) -> Tuple[sc.AnnData, gpd.GeoDataFrame]: - import scanpy as sc - - # full_atlas = subtyping_conf.get('full_atlas') - # method = subtyping_conf['method'] - - # matcher_kwargs = subtyping_conf.get('matcher_args', {}) - if 'cell_id' not in gdf.columns: - raise ValueError("gdf needs to contain a 'cell_id' column") - - if subtypes_path is None: - adata = assign_cell_types( - cell_adata, - atlas_name, - '', - **matcher_kwargs - ) - else: - subtypes = pd.read_csv(subtypes_path, index_col=0) - cell_adata.obs['cell_type_pred'] = subtypes['Cluster'] - na_mask = cell_adata.obs['cell_type_pred'].isna() - nb_na = na_mask.sum() - if nb_na > 0: - logger.warning(f"{nb_na} unattributed cells in ground truth file {subtypes_path}. Mark unmatched cells as 'Unknown'") - cell_adata.obs.loc[na_mask, 'cell_type_pred'] = 'Unknown' - adata = cell_adata - - if gdf is not None: - gdf['cell_id'] = gdf['cell_id'].astype(str) - gdf['cell_type_pred'] = adata.obs.loc[gdf['cell_id'], 'cell_type_pred'].values - - if gdf is not None: - if save_geojson: - write_geojson(gdf, os.path.join(full_exp_dir, 'nuclei_types.geojson'), 'cell_type_pred') - if save_parquet: - gdf.to_parquet('nuclei_types.parquet') - - return adata, gdf + print_resource_usage, plot_xenium_align_qc) +@deprecated def preprocess_cells_visium_hd( he_wsi: Union[str, WSI, np.ndarray, openslide.OpenSlide, CuImage], # type: ignore full_exp_dir: str, @@ -245,8 +38,6 @@ def preprocess_cells_visium_hd( nuclei_path = None ) -> Tuple[sc.AnnData, gpd.GeoDataFrame, gpd.GeoDataFrame]: - - if nuclei_path is None: segmenter = cell_segmenter_factory(segment_method) logger.info('Segmenting cells...') @@ -268,7 +59,7 @@ def preprocess_cells_visium_hd( return cell_adata, cell_gdf, nuc_gdf - +@deprecated def process_meta_df( meta_df, save_spatial_plots=True, @@ -339,8 +130,8 @@ def process_meta_df( warped_cells.to_parquet(os.path.join(path, 'processed', f'he_cell_seg.parquet')) warped_nuclei.to_parquet(os.path.join(path, 'processed', f'he_nucleus_seg.parquet')) st.transcript_df.to_parquet(os.path.join(path, 'processed', f'aligned_transcripts.parquet')) - write_geojson(warped_cells, os.path.join(path, 'processed', f'he_cell_seg.geojson'), '', chunk=True) - write_geojson(warped_nuclei, os.path.join(path, 'processed', f'he_nucleus_seg.geojson'), '', chunk=True) + write_geojson(warped_cells, os.path.join(path, 'processed', f'he_cell_seg.geojson')) + write_geojson(warped_nuclei, os.path.join(path, 'processed', f'he_nucleus_seg.geojson')) elif isinstance(st, VisiumHDHESTData): segment_config = {} binning_config = {} @@ -361,7 +152,7 @@ def process_meta_df( ) if isinstance(st, XeniumHESTData): - visualize_random_crops(st.transcript_df, st.wsi, './', st.get_shapes('tenx_nucleus', 'he').shapes) + plot_xenium_align_qc(st.wsi, './', st.transcript_df, st.get_shapes('tenx_nucleus', 'he').shapes) row_dict = row.to_dict() diff --git a/src/hest/readers.py b/src/hest/readers.py index 07a9296..87c3968 100644 --- a/src/hest/readers.py +++ b/src/hest/readers.py @@ -4,7 +4,7 @@ import math import os import shutil -from typing import Optional, Union +from typing import Optional, Tuple, Union import warnings import zipfile from abc import abstractmethod @@ -13,19 +13,20 @@ import pandas as pd from hestcore.segmentation import get_path_relative from loguru import logger +from tqdm import tqdm from hest.HESTData import (HESTData, STHESTData, VisiumHDHESTData, VisiumHESTData, XeniumHESTData) from hestcore.wsi import wsi_factory from hest.io.seg_readers import XeniumParquetCellReader, read_gdf from hest.LazyShapes import LazyShapes -from hest.segmentation.cell_segmenters import segment_cellvit +from hest.segmentation.cell_segmenters import assign_spot_to_cell, expand_nuclei, read_adata, read_seg, read_spots_gdf, segment_cellvit, sum_per_cell from hest.utils import (SpotPacking, align_xenium_df, check_arg, find_biggest_img, find_first_file_endswith, - find_pixel_size_from_spot_coords, + find_pixel_size_from_spot_coords, get_col_selection, get_path_from_meta_row, helper_mex, load_wsi, - metric_file_do_dict, read_xenium_alignment, + metric_file_do_dict, read_parquet_dask, read_xenium_alignment, register_downscale_img, verify_paths) LOCAL = False @@ -98,15 +99,12 @@ def _auto_read(self, path, **read_kwargs) -> VisiumHDHESTData: if square_16um_path is None: square_16um_path = find_first_file_endswith(os.path.join(path, 'binned_outputs'), 'square_016um') - square_2um_path = find_first_file_endswith(path, 'square_002um', anywhere=True) - metrics_path = find_first_file_endswith(path, 'metrics_summary.csv') st_object = self.read( img_path=os.path.join(path, img_filename), square_16um_path=square_16um_path, metrics_path=metrics_path, - square_2um_path=square_2um_path, **read_kwargs ) @@ -117,8 +115,8 @@ def read( self, img_path: str, square_16um_path: str, + square_2um_path: str, metrics_path: str = None, - square_2um_path: str = None, dst_bin_size_um: int = 128, chunk_len = 50000, ) -> VisiumHDHESTData: @@ -127,8 +125,8 @@ def read( Args: img_path (str): path to the WSI square_16um_path (str): path to a square_016um/ Visium HD folder. + square_2um_path (str): **Deprecated** path to a square_002um/ Visium HD folder. metrics_path (str, optional): path to a metrics_summary.csv Visium HD file. - square_2um_path (str, optional): path to a square_002um/ Visium HD folder. dst_bin_size_um (int, optional): 16um spots will be pooled to spots of this size (must be a multiple of the spot size: 16) chunk_len (str, optional): chunk length while pooling transcripts, a higher number will consume more RAM but might be faster. @@ -137,12 +135,16 @@ def read( """ import scanpy as sc SPOT_SIZE = 16 + + warnings.warn( + "square_2um_path is deprecated and will be removed in a future version. ", + DeprecationWarning, + stacklevel=2 + ) if dst_bin_size_um % SPOT_SIZE != 0: raise ValueError(f"dst_bin_size_um must be a multiple of the spot size ({SPOT_SIZE})") - self.square_2um_path = square_2um_path - print("Loading the WSI... (can be slow for large images)") img, pixel_size_embedded = load_wsi(img_path) @@ -851,16 +853,10 @@ def __xenium_estimate_pixel_size(self, pixel_size_morph, he_to_morph_matrix) -> return pixel_size_estimated - def __load_transcripts(self, transcripts_path, alignment_matrix, pixel_size_morph, use_dask): + def __load_transcripts(self, transcripts_path, alignment_matrix, pixel_size_morph, use_dask, nb_partitions=30): if use_dask: - import dask.dataframe as dd - import pyarrow.parquet as pq - - parquet_file = pq.ParquetFile(transcripts_path) - total_row_groups = parquet_file.num_row_groups - row_groups_per_partition = max(1, total_row_groups // 30) - df_transcripts = dd.read_parquet(transcripts_path, split_row_groups=row_groups_per_partition) + df_transcripts = read_parquet_dask(transcripts_path, nb_partitions) else: df_transcripts = pd.read_parquet(transcripts_path) @@ -929,14 +925,15 @@ def read( dapi_path = None, load_img=True, use_dask=False, - spot_size_um=100. + spot_size_um=100., + nb_partitions=30, ) -> XeniumHESTData: """ Read a Xenium sample Args: img_path (str): path to the WSI experiment_path (str): path to a `experiment.xenium` file - alignment_file_path (str, optional): path to a DAPI->H&E alignment file, None if the H&E is already aligned with the DAPI. Defaults to None. + alignment_file_path (str, optional): path to a DAPI->H&E matrix/keypoints alignment file, None if the H&E is already aligned with the DAPI. Defaults to None. feature_matrix_path (str, optional): path to a `cell_feature_matrix.h5`. Defaults to None. transcripts_path (str, optional): path to a transcripts.parquet, None to not load the transcripts. Defaults to None. cells_path (str, optional): path to a `cells.parquet` file, None to not load the cells. Defaults to None. @@ -946,7 +943,8 @@ def read( load_img (bool, optional): whenever to load the WSI. Defaults to True. use_dask (bool, optional): whenever to load the transcript dataframe with DASK (recommended if the transcript dataframe does not fit into the RAM). Defaults to False. spot_size_um (float, optional): transcripts are pooled into squares of spot_size_um x spot_size_um mirometers and then stored in `HESTData.adata` - + nb_partitions (int, optional): number of dask partition to use if use_dask is True. Defaults to 30 + Returns: XeniumHESTData: Xenium sample """ @@ -971,19 +969,32 @@ def read( alignment_matrix = read_xenium_alignment(alignment_file_path) if alignment_file_path else None dict['pixel_size_um_estimated'] = self.__xenium_estimate_pixel_size(pixel_size_morph, alignment_matrix) if cell_bound_path is not None: - shapes.append(LazyShapes(cell_bound_path, 'tenx_cell', 'dapi', reader=XeniumParquetCellReader, reader_kwargs={'pixel_size_morph': pixel_size_morph})) + shapes.append(LazyShapes(cell_bound_path, 'tenx_cell', 'dapi', + reader=XeniumParquetCellReader, + reader_kwargs={'pixel_size_morph': pixel_size_morph})) if alignment_matrix is not None: - shapes.append(LazyShapes(cell_bound_path, 'tenx_cell', 'he', reader=XeniumParquetCellReader, reader_kwargs={'pixel_size_morph': pixel_size_morph, 'alignment_matrix': alignment_matrix})) + shapes.append(LazyShapes(cell_bound_path, 'tenx_cell', 'he', + reader=XeniumParquetCellReader, + reader_kwargs={ + 'pixel_size_morph': pixel_size_morph, + 'alignment_matrix': alignment_matrix})) if nucleus_bound_path is not None: - shapes.append(LazyShapes(nucleus_bound_path, 'tenx_nucleus', 'dapi', reader=XeniumParquetCellReader, reader_kwargs={'pixel_size_morph': pixel_size_morph})) + shapes.append(LazyShapes(nucleus_bound_path, 'tenx_nucleus', 'dapi', + reader=XeniumParquetCellReader, + reader_kwargs={'pixel_size_morph': pixel_size_morph})) if alignment_matrix is not None: - shapes.append(LazyShapes(nucleus_bound_path, 'tenx_nucleus', 'he', reader=XeniumParquetCellReader, reader_kwargs={'pixel_size_morph': pixel_size_morph, 'alignment_matrix': alignment_matrix})) + shapes.append(LazyShapes(nucleus_bound_path, 'tenx_nucleus', 'he', + reader=XeniumParquetCellReader, + reader_kwargs={ + 'pixel_size_morph': pixel_size_morph, + 'alignment_matrix': alignment_matrix})) if transcripts_path is not None: print('Loading transcripts...') - transcript_df = self.__load_transcripts(transcripts_path, alignment_matrix, pixel_size_morph, use_dask) + transcript_df = self.__load_transcripts(transcripts_path, alignment_matrix, pixel_size_morph, + use_dask, nb_partitions) print("Pooling xenium transcripts in pseudo-visium spots...") adata = pool_transcripts_xenium( @@ -1040,15 +1051,43 @@ def reader_factory(path: str) -> Reader: else: raise NotImplementedError('') -def read_and_save(path: str, save_plots=True, pyramidal=True, bigtiff=False, plot_pxl_size=False, save_img=True, segment_tissue=False, read_kwargs={}, save_kwargs={}, segment_kwargs={}): - """For internal use, determine the appropriate reader based on the raw data path, and - automatically process the data at that location, then the processed files are dumped - to processed/ + +def read_and_save( + path: str, + save_plots=True, + pyramidal=True, + bigtiff=False, + plot_pxl_size=True, + save_img=True, + segment_tissue=True, + save_adata=True, + qc=True, + dump_patches=True, + read_kwargs={}, + save_kwargs={}, + segment_kwargs={}, + patching_kwargs={}, +) -> HESTData: + """ Determine the appropriate reader based on the raw data path, and + automatically process the data at that location, then the processed files are dumped + to processed/ Args: - path (str): path of the raw data - save_plots (bool, optional): whenever to save the spatial plots. Defaults to True. - pyramidal (bool, optional): whenever to save as pyramidal. Defaults to True. + path (str): path to input folder (will be passed to reader auto_read) + save_plots (bool, optional): whenever to save spatial plots. Defaults to True. + pyramidal (bool, optional): whenever to save the wsi as pyramidal. Defaults to True. + bigtiff (bool, optional): whenever to save the wsi as bigtiff (should be set to True if >4.1GB). Defaults to False. + plot_pxl_size (bool, optional): whenever to save a plot showing the embedded vs computed pixel size. Defaults to True. + save_img (bool, optional): whenever to re-save the WSI in a format compatible with both openslide and qupath. Defaults to True. + segment_tissue (bool, optional): whenever to segment the tissue. Defaults to True. + save_adata (bool, optional): whenever to save an adata object with pooled expression. Defaults to True. + qc (bool, optional): whenever to save qc. Defaults to True. + dump_patches (bool, optional): whenever to dump patches. Defaults to True. + read_kwargs (dict, optional): kwargs passed to the reader. Defaults to {}. + save_kwargs (dict, optional): kwargs passed to HESTData save. Defaults to {}. + segment_kwargs (dict, optional): kwargs passed to segment_tissue. Defaults to {}. + patching_kwargs (dict, optional): kwargs passed to dump_patches. Defaults to {}. + """ print(f'Reading from {path}...') reader = reader_factory(path) @@ -1060,11 +1099,32 @@ def read_and_save(path: str, save_plots=True, pyramidal=True, bigtiff=False, plo st_object.segment_tissue(**segment_kwargs) save_path = os.path.join(path, 'processed') os.makedirs(save_path, exist_ok=True) - st_object.save(save_path, pyramidal=pyramidal, bigtiff=bigtiff, plot_pxl_size=plot_pxl_size, save_img=save_img, **save_kwargs) + st_object.save( + save_path, + pyramidal=pyramidal, + bigtiff=bigtiff, + plot_pxl_size=plot_pxl_size, + save_img=save_img, + save_adata=save_adata, + qc=qc, + **save_kwargs + ) if save_plots: st_object.save_spatial_plot(save_path) + if dump_patches: + st_object.dump_patches(save_path, qc=qc, **patching_kwargs) + return st_object +def get_indices_chunk(partition, key_x, key_y, + x_min, y_min, spot_size_um, pixel_size_he, n, spot_grid_columns): + a = np.floor((partition[key_x] - x_min) / (spot_size_um / pixel_size_he)).astype(int) + b = np.floor((partition[key_y] - y_min) / (spot_size_um / pixel_size_he)).astype(int) + + c = b * n + a + cols = spot_grid_columns.get_indexer(partition['feature_name']) + return pd.DataFrame({'c': c, 'cols': cols}, index=partition.index) + def pool_transcripts_xenium( df: Union[pd.DataFrame, dd.DataFrame], pixel_size_he: float, @@ -1075,10 +1135,11 @@ def pool_transcripts_xenium( """ Pool a xenium transcript dataframe by square spots of `spot_size_um` micrometers. Args: - df (Union[pd.DataFrame, dd.DataFrame]): xenium transcipts dataframe containing columns: + df (Union[pd.DataFrame, dd.DataFrame]): xenium transcipts (dask) dataframe containing columns: + - 'he_x' and 'he_y' indicating the pixel coordinates of each transcripts in the morphology image - 'feature_name' indicating the transcript name - pixel_size_he (float): pixel_size in um on the he image + pixel_size_he (float): pixel size in um/px of 'he_x' and 'he_y' spot_size_um: pooling rectangle width in um key_x: column name of pixel x coordinate of each transcript in `df` key_y: column name of pixel y coordinate of each transcript in `df` @@ -1089,6 +1150,7 @@ def pool_transcripts_xenium( """ import scanpy as sc import dask.dataframe as dd + import dask y_max = df[key_y].max() y_min = df[key_y].min() @@ -1098,32 +1160,42 @@ def pool_transcripts_xenium( m = ((y_max - y_min) / (spot_size_um / pixel_size_he)) n = ((x_max - x_min) / (spot_size_um / pixel_size_he)) + unique_features = df['feature_name'].unique() + if isinstance(df, dd.DataFrame): - m = m.compute() - n = n.compute() + m, n, unique_features, x_min, y_min = dask.compute( + m, n, unique_features, x_min, y_min) m = math.ceil(m) n = math.ceil(n) + spot_grid = pd.DataFrame(0, index=range(m * n), columns=unique_features) + spot_grid_np = spot_grid.values.astype(np.uint32) - features = df['feature_name'].unique() if isinstance(df, dd.DataFrame): - features = features.compute() - - spot_grid = pd.DataFrame(0, index=range(m * n), columns=features) - + import dask.array as da + + cols_c = df.map_partitions(get_indices_chunk, key_x, key_y, + x_min, y_min, spot_size_um, pixel_size_he, n, spot_grid.columns, + meta={'c': 'int64', 'cols': 'int64'}) + + num_rows = m * n + num_cols = len(unique_features) + c_da = cols_c['c'].to_dask_array(lengths=True) + cols_da = cols_c['cols'].to_dask_array(lengths=True) + h, xedges, yedges = da.histogram2d( + c_da, + cols_da, + bins=[np.arange(num_rows + 1), np.arange(num_cols + 1)] + ) + spot_grid_np = h.astype(np.uint32).compute() + + else: + cols_c = get_indices_chunk(key_x, key_y, + x_min, y_min, spot_size_um, pixel_size_he, n, spot_grid.columns) - # a is the row and b is the column in the pseudo visium grid - a = np.floor((df[key_x] - x_min) / (spot_size_um / pixel_size_he)).astype(int) - b = np.floor((df[key_y] - y_min) / (spot_size_um / pixel_size_he)).astype(int) - - c = b * n + a - features = df['feature_name'] - - cols = spot_grid.columns.get_indexer(features) - - ## use dask for this parts - spot_grid_np = spot_grid.values.astype(np.uint16) - np.add.at(spot_grid_np, (c, cols), 1) + c = cols_c['c'] + cols = cols_c['cols'] + np.add.at(spot_grid_np, (c, cols), 1) if isinstance(spot_grid.columns.values[0], bytes): @@ -1133,9 +1205,6 @@ def pool_transcripts_xenium( expression_df = pd.DataFrame(spot_grid_np, columns=spot_grid.columns) coord_df = expression_df.copy() - if isinstance(df, dd.DataFrame): - x_min = x_min.compute() - y_min = y_min.compute() coord_df['x'] = x_min + (coord_df.index % n) * (spot_size_um / pixel_size_he) + ((spot_size_um / 2) / pixel_size_he) coord_df['y'] = y_min + np.floor(coord_df.index / n) * (spot_size_um / pixel_size_he) + ((spot_size_um / 2) / pixel_size_he) @@ -1158,7 +1227,8 @@ def pool_transcripts_xenium( def pool_bins_visiumhd(adata: sc.AnnData, pixel_size: float, dst_bin_size_um=128, src_bin_size_um: Literal[2, 8, 16]=16, chunk_len=50000) -> sc.AnnData: # type: ignore - """ Pool a Visium HD (with a source resolution of `src_bin_size_um`) by square spots of `spot_size_um` micrometers. + """ Pools Visium HD bins from an initial resolution (src_bin_size_um) into larger square spots of spot_size_um. + This performs a best-effort spatial downsampling (bin-to-bin aggregation). Args: adata (sc.AnnData): adata containing spot center coordiniates in `pxl_row_in_fullres` and `pxl_col_in_fullres` @@ -1235,6 +1305,54 @@ def pool_bins_visiumhd(adata: sc.AnnData, pixel_size: float, dst_bin_size_um=128 return adata + +def pool_bins_visiumhd_per_cell( + nuc_seg: Union[str, gpd.GeoDataFrame], + bc_matrix: Union[str, sc.AnnData], + path_bins_pos: str, + pixel_size: float, + save_dir: str = None, + exp_um = 5, + exp_nuclei: bool = True +) -> Tuple[sc.AnnData, gpd.GeoDataFrame]: + """ Pool Visium-hd bins per cell. + + Args: + nuc_seg (Union[str, gpd.GeoDataFrame]): nuclei segmentation + bc_matrix (Union[str, sc.AnnData]): bc_matrix representing Visium-hd bins. + path_bins_pos (str): path to `tissue_positions.parquet` + pixel_size (float): pixel size of path_bins_pos in um/px + save_dir (str, optional): whenever to save to aligned_cells.h5ad. Defaults to None. + exp_um (int, optional): nuclei expansion in um if exp_nuclei is True. Defaults to 5. + exp_nuclei (bool, optional): whenever to expand nuclei to derive cells. Defaults to True. + + Returns: + Tuple[sc.AnnData, gpd.GeoDataFrame]: binned adata and (expended) nuclei + """ + + verify_paths([bc_matrix, path_bins_pos]) + + + nuclei_gdf = read_seg(nuc_seg) + + if exp_nuclei: + cell_gdf = expand_nuclei(nuclei_gdf, pixel_size, exp_um=exp_um) + else: + cell_gdf = nuclei_gdf + + logger.info('Read bin positions...') + points_gdf = read_spots_gdf(path_bins_pos) + + assignment = assign_spot_to_cell(cell_gdf, points_gdf) + + adata = read_adata(bc_matrix) + + cell_adata = sum_per_cell(adata, assignment) + + if save_dir is not None: + cell_adata.write_h5ad(os.path.join(save_dir, 'aligned_cells.h5ad')) + + return cell_adata, cell_gdf def _process_cellvit(row, **cellvit_kwargs): @@ -1262,3 +1380,84 @@ def process_meta_df_cellvit(meta_df, cellvit_kwargs={'gpu_ids': [0, 1], 'batch_s for _, row in meta_df.iterrows(): _process_cellvit(row, **cellvit_kwargs) + +def save_meta(row_dict): + path = get_path_from_meta_row(row_dict) + + with open(os.path.join(path, 'processed', f'metrics.json'), 'r') as f: + meta = json.load(f) + + combined_meta = {**meta, **row_dict} + cols = get_col_selection() + combined_meta = {k: v for k, v in combined_meta.items() if k in cols} + with open(os.path.join(path, 'processed', f'meta.json'), 'w') as f: + json.dump(combined_meta, f, indent=4) + + +def process_raw_samples( + meta_df: pd.DataFrame, + save_plots=True, + pyramidal=True, + bigtiff=False, + plot_pxl_size=True, + save_img=True, + segment_tissue=True, + save_adata=True, + qc=True, + dump_patches=True, + read_kwargs={ + 'use_dask': True, + 'nb_partitions': 100, + }, + save_kwargs={ + 'save_nuclei_seg': True, + 'save_cell_seg': True + }, + segment_kwargs={}, + patching_kwargs={}, +) -> None: + """ Process raw samples and save them into HEST format + + Args: + meta_df (str): metadata dataframe + save_plots (bool, optional): whenever to save spatial plots. Defaults to True. + pyramidal (bool, optional): whenever to save the wsi as pyramidal. Defaults to True. + bigtiff (bool, optional): whenever to save the wsi as bigtiff (should be set to True if >4.1GB). Defaults to False. + plot_pxl_size (bool, optional): whenever to save a plot showing the embedded vs computed pixel size. Defaults to True. + save_img (bool, optional): whenever to re-save the WSI in a format compatible with both openslide and qupath. Defaults to True. + segment_tissue (bool, optional): whenever to segment the tissue. Defaults to True. + save_adata (bool, optional): whenever to save an adata object with pooled expression. Defaults to True. + qc (bool, optional): whenever to save qc. Defaults to True. + dump_patches (bool, optional): whenever to dump patches. Defaults to True. + read_kwargs (dict, optional): kwargs passed to the reader. Defaults to {}. + save_kwargs (dict, optional): kwargs passed to HESTData save. Defaults to {}. + segment_kwargs (dict, optional): kwargs passed to segment_tissue. Defaults to {}. + patching_kwargs (dict, optional): kwargs passed to dump_patches. Defaults to {}. + """ + + for _, row in tqdm(meta_df.iterrows()): + path = get_path_from_meta_row(row) + + sample_folder_path = path + res_folder = os.path.join(path, 'processed') + os.makedirs(res_folder, exist_ok=True) + + bigtiff = not(isinstance(row['bigtiff'], float) or row['bigtiff'] == 'FALSE') + read_and_save( + sample_folder_path, + save_plots=save_plots, + pyramidal=pyramidal, + bigtiff=bigtiff, + plot_pxl_size=plot_pxl_size, + save_img=save_img, + segment_tissue=segment_tissue, + save_adata=save_adata, + qc=qc, + dump_patches=dump_patches, + read_kwargs=read_kwargs, + save_kwargs=save_kwargs, + segment_kwargs=segment_kwargs, + patching_kwargs=patching_kwargs + ) + + save_meta(dict(row)) \ No newline at end of file diff --git a/src/hest/registration.py b/src/hest/registration.py index e4077fb..0dfa465 100644 --- a/src/hest/registration.py +++ b/src/hest/registration.py @@ -1,16 +1,17 @@ from __future__ import annotations +import gc import os -from typing import Tuple, Union +from typing import Optional, Tuple, Union +import warnings import geopandas as gpd import numpy as np from loguru import logger -from shapely import Polygon -from hest.io.seg_readers import groupby_shape, read_gdf -from hest.utils import (get_name_datetime, - value_error_str, verify_paths) +from hest.io.seg_readers import HESTXeniumTranscriptsReader, XeniumTranscriptsReader, groupby_shape, read_gdf, write_geojson +from hest.utils import (deprecated, get_name_datetime, merge_parquet, + value_error_str) from hestcore.wsi import WSI @@ -24,7 +25,6 @@ def register_dapi_he( micro_rigid_registrar_params={}, micro_reg=True, check_for_reflections=False, - reuse_registrar=False ) -> str: """ Register the DAPI WSI to HE with a fine-grained ridig + non-rigid transform with Valis @@ -46,10 +46,8 @@ def register_dapi_he( from valis_hest.slide_io import BioFormatsSlideReader from .SlideReaderAdapter import SlideReaderAdapter - except Exception: - import traceback - traceback.print_exc() - raise Exception("Valis needs to be installed independently. Please install Valis with `pip install valis-hest`") + except Exception as e: + raise Exception("Valis needs to be installed independently. Please install Valis with `pip install valis-hest`") from e #verify_paths([dapi_path, he_path]) @@ -66,9 +64,7 @@ def register_dapi_he( registrar_path = os.path.join(registrar_dir, 'data/_registrar.pickle') - if reuse_registrar: - registration.init_jvm() - return registrar_path + registration.init_jvm() registrar = registration.Valis( '', registrar_dir, @@ -99,23 +95,67 @@ def register_dapi_he( return registrar_path + +def _warp_gdf_valis(gdf, registrar, curr_slide_name, slide_obj): + if len(gdf) == 0: + return gdf + + geom_type = gdf.geometry.iloc[0].geom_type + + if geom_type in ['Polygon', 'MultiPolygon']: + coords = gdf.geometry.get_coordinates(index_parts=True) + points_gdf = coords + idx = coords.index.get_level_values(0) + points_gdf['_polygons'] = idx + points = list(zip(points_gdf['x'], points_gdf['y'])) + elif geom_type == 'Point': + points_gdf = gdf + points = list(zip(gdf.geometry.x, gdf.geometry.y)) + else: + raise NotImplementedError('') + + morph = registrar.get_slide(curr_slide_name) + warped = morph.warp_xy_from_to(points, slide_obj) + + if geom_type in ['Polygon', 'MultiPolygon']: + points_gdf['xy'] = list(zip(warped[:, 0], warped[:, 1])) + aggr_df = groupby_shape(points_gdf, '_polygons', n_threads=0) + gdf.geometry = aggr_df.geometry + else: + import geopandas as gpd + gdf.geometry = gpd.points_from_xy(warped[:, 0], warped[:, 1]) + + + return gdf def warp_gdf_valis( - shapes: Union[gpd.GeoDataFrame, str], + shapes: Union[gpd.GeoDataFrame, str, dgpd.GeoDataFrame], path_registrar: str, curr_slide_name: str, - n_workers=-1 -) -> gpd.GeoDataFrame: - """ Warp some shapes (points or polygons) from an existing Valis registration + n_workers=-1, + use_dask=True +) -> Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]: + """ Warp some shapes (points or polygons) from an existing Valis registration registrar Args: - shapes (Union[gpd.GeoDataFrame, str]): shapes to warp. A `str` will be interpreted as a path a nucleus shape file, can be .geojson, or xenium .parquet (ex: nucleus_boundaries.parquet) + shapes (Union[gpd.GeoDataFrame, str, dgpd.GeoDataFrame]): shapes to warp. A `str` will be interpreted as a path a nucleus shape file, can be .geojson, or xenium .parquet (ex: nucleus_boundaries.parquet) path_registrar (str): path to the .pickle file of an existing Valis registrar + curr_slide_name (str): dapi slide filename in the Valis registrar + n_workers (int, optional): **Deprecated**. Use dask instead + use_dask (bool, optional): whenever to use dask to process larger than RAM data, highly recommended for all Xenium samples. Defaults to True. Returns: - gpd.GeoDataFrame: warped shapes + Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]: warped geodataframe such that warped shapes are in the geometry column. """ + if n_workers != -1: + warnings.warn( + "The 'n_workers' parameter is deprecated and will be removed in a future version. " + "Please use 'use_dask' instead.", + DeprecationWarning, + stacklevel=2 + ) + try: from valis_hest import registration except Exception: @@ -123,42 +163,304 @@ def warp_gdf_valis( traceback.print_exc() raise Exception("Valis needs to be installed independently. Please install Valis with `pip install valis-wsi` or follow instruction on their website") + XENIUM_PIXEL_SIZE_MORPH = 0.2125 if isinstance(shapes, str): - gdf = read_gdf(shapes) + gdf = read_gdf(shapes, reader_kwargs={'pixel_size_morph': XENIUM_PIXEL_SIZE_MORPH, 'use_dask': use_dask}) elif isinstance(shapes, gpd.GeoDataFrame): gdf = shapes.copy() else: - raise ValueError(value_error_str(shapes, 'shapes')) + try: + import dask_geopandas + except: + pass + if dask_geopandas and isinstance(shapes, dask_geopandas.expr.GeoDataFrame): + gdf = shapes + else: + raise ValueError(value_error_str(shapes, 'shapes')) + from valis_hest.registration import init_jvm + init_jvm(mem_gb=1) registrar = registration.load_registrar(path_registrar) slide_obj = registrar.get_slide(registrar.reference_img_f) - if isinstance(shapes.iloc[0].geometry, Polygon): - coords = gdf.geometry.get_coordinates(index_parts=True) - points_gdf = coords - idx = coords.index.get_level_values(0) - points_gdf['_polygons'] = idx # keep track of polygons - points = list(zip(points_gdf['x'], points_gdf['y'])) + if use_dask: + from geopandas.array import GeometryDtype + from distributed import get_client + + try: + client = get_client() + except: + client = None + + meta = gdf.head(0).copy() + + from geopandas.array import GeometryDtype + meta['geometry'] = gpd.GeoSeries(dtype=GeometryDtype()) + + meta = meta.set_geometry('geometry').set_crs("EPSG:4326") + if client is None: + registrar_future = registrar + slide_obj_future = slide_obj + else: + [registrar_future] = client.scatter([registrar], broadcast=True) + [slide_obj_future] = client.scatter([slide_obj], broadcast=True) + gdf = gdf.map_partitions(_warp_gdf_valis, registrar_future, curr_slide_name, slide_obj_future, meta=meta) else: - points_gdf = gdf - gdf['_polygons'] = np.arange(len(points_gdf)) - points = list(zip(gdf.geometry.x, gdf.geometry.y)) + gdf = _warp_gdf_valis(gdf, registrar, curr_slide_name, slide_obj) + + return gdf + + +def _read_transcripts( + dapi_transcripts: str, + use_dask: bool, +): + PIXEL_SIZE_MORPH = 0.2125 + + import pyarrow.parquet as pq + parquet_file = pq.ParquetFile(dapi_transcripts) + if 'dapi_x' in parquet_file.schema.names: + reader = HESTXeniumTranscriptsReader(use_dask=use_dask) + else: + reader = XeniumTranscriptsReader(pixel_size_morph=PIXEL_SIZE_MORPH, use_dask=use_dask) + + gdf_transcripts = reader.read_gdf(dapi_transcripts) + return gdf_transcripts + + +def _format_transcripts(warped_transcripts): + if '_polygons' in warped_transcripts.columns: + warped_transcripts = warped_transcripts.drop(['_polygons'], axis=1) + warped_transcripts['he_x'] = warped_transcripts.geometry.x + warped_transcripts['he_y'] = warped_transcripts.geometry.y + return warped_transcripts + + +def _warp_transcripts( + dapi_transcripts, + use_dask, + verbose, + path_registrar, + dapi_path +): + gdf_transcripts = _read_transcripts(dapi_transcripts, use_dask) + + if verbose: + logger.info('Warping transcripts from DAPI to H&E...') - morph = registrar.get_slide(curr_slide_name) - logger.debug('warp with valis...') - warped = morph.warp_xy_from_to(points, slide_obj) - logger.debug('finished warping with valis') + warped_transcripts = warp_gdf_valis( + gdf_transcripts, + path_registrar=path_registrar, + curr_slide_name=dapi_path, + use_dask=use_dask, + ) - if isinstance(shapes.iloc[0].geometry, Polygon): - points_gdf['xy'] = list(zip(warped[:, 0], warped[:, 1])) - aggr_df = groupby_shape(points_gdf, '_polygons', n_threads=0) - gdf.geometry = aggr_df.geometry + warped_transcripts = _format_transcripts(warped_transcripts) + return warped_transcripts + + +def warp_xenium_objects( + path_registrar: str, + dapi_path: str, + dapi_cells: str=None, + dapi_transcripts: str=None, + dapi_nuclei: str=None, + use_dask=True, + verbose=True +) -> Tuple[Optional[Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]], + Optional[Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]], + Optional[Union[gpd.GeoDataFrame, dgpd.GeoDataFrame]]]: + """ **Deprecated** Use warp_and_save_xenium_objects instead. Wrap Xenium transcripts, cells and nuclei using Valis non-rigid micro-registration. + + Args: + path_registrar (str): path to an existing Valis registrar + dapi_path (str): dapi slide filename in the Valis registrar + dapi_cells (str, optional): path to xenium .parquet cell bondaries, usually **/cell_boundaries.parquet. Defaults to None. + dapi_transcripts (str, optional): path to xenium .parquet nucleus bondaries, usually **/nucleus_boundaries.parquet. Defaults to None. + dapi_nuclei (str, optional): path to xenium .parquet transcripts, usually **/transcripts.parquet. Defaults to None. + use_dask (bool, optional): whenever to use dask to process larger than RAM data, highly recommended for all Xenium samples. Defaults to True. + verbose (bool, optional): verbose flag. Defaults to True. + + Returns: + warped (Tuple): warped objects as (warped_cells, warped_nuclei, warped_transcripts) + + Example: + >>> registrar_path = "./valis_results/data/_registrar.pickle" + >>> cells_path = "./xenium_out/cell_boundaries.parquet" + >>> cells, nuclei, transcripts = warp_xenium_objects( + ... path_registrar=registrar_path, + ... dapi_path="morphology_focus.ome.tif", + ... dapi_cells=cells_path, + ... use_dask=True + ... ) + >>> if cells is not None: + ... print(f"Warped {len(cells)} cells using Dask: {type(cells)}") + """ + warnings.warn( + "warp_xenium_objects is deprecated and will be removed in a future version. " + "Please use 'warp_and_save_xenium_objects' instead.", + DeprecationWarning, + stacklevel=2 + ) + + if dapi_transcripts: + warped_transcripts = _warp_transcripts(dapi_transcripts, use_dask, + verbose, path_registrar, dapi_path) else: - gdf.geometry = gpd.points_from_xy(warped[:, 0], warped[:, 1]) + warped_transcripts = None + + if dapi_cells is not None: + + if verbose: + logger.info('Warping cells from DAPI to H&E...') + + warped_cells = warp_gdf_valis( + dapi_cells, + path_registrar=path_registrar, + curr_slide_name=dapi_path, + use_dask=use_dask, + ) + else: + warped_cells = None - return gdf + if dapi_nuclei is not None: + + if verbose: + logger.info('Warping nuclei from DAPI to H&E...') + + warped_nuclei = warp_gdf_valis( + dapi_nuclei, + path_registrar=path_registrar, + curr_slide_name=dapi_path, + use_dask=use_dask, + ) + else: + warped_nuclei = None + + return warped_cells, warped_nuclei, warped_transcripts + + +def warp_and_save_xenium_objects( + path_registrar: str, + dapi_path: str, + save_dir: str, + dapi_cells: str=None, + dapi_transcripts: str=None, + dapi_nuclei: str=None, + use_dask=True, + verbose=True, + save_parquet=True, + save_geojson=True, +) -> None: + """ Wrap Xenium transcripts, cells and nuclei using Valis non-rigid micro-registration and save them. + + Args: + path_registrar (str): path to an existing Valis registrar + dapi_path (str): dapi slide filename in the Valis registrar + save_dir (str): where to save warped objects. Objects will be saved to: + - save_dir/he_cell_seg.parquet + - save_dir/he_nucleus_seg.parquet + - save_dir/aligned_transcripts + dapi_cells (str, optional): path to xenium .parquet cell bondaries, usually **/cell_boundaries.parquet. Defaults to None. + dapi_transcripts (str, optional): path to xenium .parquet nucleus bondaries, usually **/nucleus_boundaries.parquet. Defaults to None. + dapi_nuclei (str, optional): path to xenium .parquet transcripts, usually **/transcripts.parquet. Defaults to None. + use_dask (bool, optional): whenever to use dask to process larger than RAM data, highly recommended for all Xenium samples. Defaults to True. + verbose (bool, optional): verbose flag. Defaults to True. + save_parquet (bool, optional): whenever to save objects as parquet. Defaults to True. + save_geojson (bool, optional): whenever to save objects as geojson. Defaults to True. + + Example: + >>> registrar_path = "./valis_results/data/_registrar.pickle" + >>> cells_path = "./xenium_out/cell_boundaries.parquet" + >>> warp_and_save_xenium_objects( + ... path_registrar=registrar_path, + ... save_dir="warped_xenium", + ... dapi_path="morphology_focus.ome.tif", + ... dapi_cells=cells_path, + ... use_dask=True + ... ) + >>> if cells is not None: + ... print(f"Warped {len(cells)} cells using Dask: {type(cells)}") + """ + if not os.path.exists(save_dir): + raise ValueError(f"Save directory '{save_dir}' doesn't exist.") + if dapi_transcripts: + warped_transcripts = _warp_transcripts(dapi_transcripts, use_dask, + verbose, path_registrar, dapi_path) + + res_path = 'aligned_transcripts' if use_dask else 'aligned_transcripts.parquet' + warped_transcripts.to_parquet(os.path.join(save_dir, res_path)) + + if use_dask: + merge_parquet(os.path.join(save_dir, 'aligned_transcripts'), + os.path.join(save_dir, 'aligned_transcripts.parquet')) + del warped_transcripts + gc.collect() + else: + warped_transcripts = None + + + if dapi_cells is not None: + + if verbose: + logger.info('Warping cells from DAPI to H&E...') + + warped_cells = warp_gdf_valis( + dapi_cells, + path_registrar=path_registrar, + curr_slide_name=dapi_path, + use_dask=use_dask, + ) + + if save_parquet: + res_path = 'he_cell_seg' if use_dask else 'he_cell_seg.parquet' + warped_cells.to_parquet(os.path.join(save_dir, res_path)) + + if use_dask: + merge_parquet(os.path.join(save_dir, 'he_cell_seg'), + os.path.join(save_dir, 'he_cell_seg.parquet')) + + if save_geojson: + write_geojson(warped_cells.compute(), os.path.join(save_dir, f'he_cell_seg.geojson')) + + + del warped_cells + gc.collect() + else: + warped_cells = None + + + if dapi_nuclei is not None: + + if verbose: + logger.info('Warping nuclei from DAPI to H&E...') + + warped_nuclei = warp_gdf_valis( + dapi_nuclei, + path_registrar=path_registrar, + curr_slide_name=dapi_path, + use_dask=use_dask, + ) + + if save_parquet: + res_path = 'he_nucleus_seg' if use_dask else 'he_nucleus_seg.parquet' + warped_nuclei.to_parquet(os.path.join(save_dir, res_path)) + + if use_dask: + merge_parquet(os.path.join(save_dir, 'he_nucleus_seg'), + os.path.join(save_dir, 'he_nucleus_seg.parquet')) + + if save_geojson: + write_geojson(warped_nuclei.compute(), os.path.join(save_dir, f'he_nucleus_seg.geojson')) + + del warped_nuclei + gc.collect() + else: + warped_nuclei = None + +@deprecated def preprocess_cells_xenium( he_wsi: Union[str, WSI, np.ndarray, openslide.OpenSlide, CuImage], # type: ignore dapi_path: str, @@ -177,6 +479,7 @@ def preprocess_cells_xenium( logger.info('Registering Xenium DAPI to H&E...') max_non_rigid_registration_dim_px = reg_config.get('max_non_rigid_registration_dim_px', 10000) + path_registrar = register_dapi_he( he_wsi, dapi_path, @@ -185,39 +488,11 @@ def preprocess_cells_xenium( max_non_rigid_registration_dim_px=max_non_rigid_registration_dim_px, **registration_kwargs ) - - if dapi_transcripts: - logger.info('Warping transcripts from DAPI to H&E...') - transcripts_gdf = gpd.GeoDataFrame(dapi_transcripts, geometry=gpd.points_from_xy(dapi_transcripts['dapi_x'], dapi_transcripts['dapi_y'])) - warped_transcripts = warp_gdf_valis( # TODO valis interpolation is slow - transcripts_gdf, - path_registrar=path_registrar, - curr_slide_name=dapi_path - ) - warped_transcripts = warped_transcripts.drop(['_polygons'], axis=1) - warped_transcripts['he_x'] = warped_transcripts.geometry.x - warped_transcripts['he_y'] = warped_transcripts.geometry.y - else: - warped_transcripts = None - if dapi_cells is not None: - logger.info('Warping cells from DAPI to H&E...') - warped_cells = warp_gdf_valis( # TODO valis interpolation is slow - dapi_cells, - path_registrar=path_registrar, - curr_slide_name=dapi_path - ) - else: - warped_cells = None - - if dapi_nuclei is not None: - logger.info('Warping nuclei from DAPI to H&E...') - warped_nuclei = warp_gdf_valis( # TODO valis interpolation is slow - dapi_nuclei, - path_registrar=path_registrar, - curr_slide_name=dapi_path - ) - else: - warped_nuclei = None - - return warped_cells, warped_nuclei, warped_transcripts \ No newline at end of file + return warp_xenium_objects( + path_registrar, + dapi_path, + dapi_cells, + dapi_transcripts, + dapi_nuclei + ) \ No newline at end of file diff --git a/src/hest/segmentation/TissueMask.py b/src/hest/segmentation/TissueMask.py deleted file mode 100644 index 8bb54d6..0000000 --- a/src/hest/segmentation/TissueMask.py +++ /dev/null @@ -1,35 +0,0 @@ -import pickle -from typing import List - -import numpy as np -from PIL import Image - - -class TissueMask: - - def __init__( - self, - tissue_mask: np.ndarray, - contours_tissue: List, - contours_holes: List - ): - self.tissue_mask = tissue_mask - self.contours_tissue = contours_tissue - self.contours_holes = contours_holes - - -def load_tissue_mask(pkl_path: str, jpg_path: str, width: int, height: int) -> TissueMask: - with open(pkl_path, 'rb') as file: - data = pickle.load(file) - contours_holes = data['holes'] - contours_tissue = data['tissue'] - with Image.open(jpg_path) as img: - tissue_mask = np.array(img).copy() - - - import cv2 - tissue_mask = cv2.resize(tissue_mask, (width, height)) - - mask = TissueMask(tissue_mask, contours_tissue, contours_holes) - - return mask \ No newline at end of file diff --git a/src/hest/segmentation/cell_segmenters.py b/src/hest/segmentation/cell_segmenters.py index 0cd9e5a..b0f2ee9 100644 --- a/src/hest/segmentation/cell_segmenters.py +++ b/src/hest/segmentation/cell_segmenters.py @@ -10,18 +10,13 @@ import geopandas as gpd import numpy as np -import openslide import pandas as pd from hestcore.segmentation import get_path_relative -from hestcore.wsi import wsi_factory from loguru import logger -from shapely import Polygon -from shapely.affinity import translate from tqdm import tqdm from hest.io.seg_readers import GeojsonCellReader -from hest.utils import deprecated, get_n_threads, verify_paths -from hestcore.wsi import wsi_factory +from hest.utils import get_n_threads from hest.utils import verify_paths @@ -132,6 +127,7 @@ def _verify_model(self, model_path, model): if not os.path.exists(model_path): print(f'Model not found at {model_path}, downloading...') gdrive_id = self.MODELS_SRC_MAP[model] + os.makedirs(os.path.dirname(model_path), exist_ok=True) gdown.download(id=gdrive_id, output=model_path, quiet=False) else: print(f'Found model at {model_path}') @@ -217,7 +213,7 @@ def segment_cellvit( batch_size=2, gpu_ids=[0], save_dir='results/segmentation', - model='CellViT-SAM-H-x20.pth' + model: str='CellViT-SAM-H-x20.pth' ) -> str: """ Segment nuclei with CellViT @@ -229,7 +225,15 @@ def segment_cellvit( batch_size (int, optional): batch_size. Defaults to 2. gpu_ids (List[int], optional): list of gpu ids to use during inference. Defaults to [0]. save_dir (str, optional): directory where to save the output. Defaults to 'results/segmentation'. - model (str, optional): name of model weights to use. Defaults to 'CellViT-SAM-H-x20.pth'. + model (str, optional): name of model weights to use. Defaults to 'CellViT-SAM-H-x20.pth'. List of weights available: + + - 'CellViT-256-x20.pth' + - 'CellViT-256-x40.pth' + - 'CellViT-SAM-H-x40.pth' + - 'CellViT-SAM-H-x20.pth' + + Returns: + str: path to the result segmentation .geojson """ segmenter = CellViTSegmenter() return segmenter.segment_cells( @@ -427,12 +431,35 @@ def bin_per_cell( path_bins_pos: str, pixel_size: float, save_dir: str = None, - name = '', exp_um = 5, exp_nuclei: bool = True ) -> Tuple[sc.AnnData, gpd.GeoDataFrame]: + """ Bin Visium-hd sub-bins per cell. + + **Deprecated** use hest.readers.pool_bins_visiumhd_per_cell instead + + Args: + nuc_seg (Union[str, gpd.GeoDataFrame]): nuclei segmentation + bc_matrix (Union[str, sc.AnnData]): bc_matrix representing Visium-hd bins. + path_bins_pos (str): path to `tissue_positions.parquet` + pixel_size (float): pixel size of path_bins_pos in um/px + save_dir (str, optional): whenever to save to aligned_cells.h5ad. Defaults to None. + exp_um (int, optional): nuclei expansion in um if exp_nuclei is True. Defaults to 5. + exp_nuclei (bool, optional): whenever to expand nuclei to derive cells. Defaults to True. + + Returns: + Tuple[sc.AnnData, gpd.GeoDataFrame]: binned adata and (expended) nuclei + """ + warnings.warn( + "bin_per_cell is deprecated and will be removed in a future version. " + "Please use 'pool_bins_visiumhd_per_cell' instead.", + DeprecationWarning, + stacklevel=2 + ) + verify_paths([bc_matrix, path_bins_pos]) + nuclei_gdf = read_seg(nuc_seg) if exp_nuclei: @@ -491,197 +518,3 @@ def sum_per_cell(adata: sc.AnnData, assignment: gpd.GeoDataFrame): cell_adata = sc.AnnData(X=summed_counts.tocsr() ,obs=pd.DataFrame(cell_ids, columns=['cell_id'], index=cell_ids),var=adata.var) return cell_adata - - -class AlignmentRefiner: - - @abstractmethod - def refine(self, centroids: gpd.GeoDataFrame, polygons: gpd.GeoDataFrame, class_key='class') -> gpd.GeoDataFrame: - pass - -class RegAlignmentRefiner: - - def refine(self, centroids: gpd.GeoDataFrame, polygons: gpd.GeoDataFrame, class_key='class') -> gpd.GeoDataFrame: - pass - - -class SegAlignmentRefiner(AlignmentRefiner): - - def refine(self, centroids: gpd.GeoDataFrame, polygons: gpd.GeoDataFrame, class_key='class') -> gpd.GeoDataFrame: - merged = centroids.sjoin(polygons, how='inner', predicate='within') - point_in_poly_idx = merged.index - poly_cont_point_idx = merged['index_right'] - - - filt_centroids = centroids.drop(point_in_poly_idx) - filt_poly = polygons.drop(poly_cont_point_idx) - - nearest_idx = filt_centroids.geometry.centroid.sindex.nearest(filt_poly.geometry)[1] - filt_poly[class_key] = filt_centroids.iloc[nearest_idx][class_key].values - - matched_poly = polygons.loc[poly_cont_point_idx] - cont = matched_poly.sjoin(centroids, how='left', predicate='contains') - cont = cont.drop_duplicates(keep='last') - matched_poly[class_key] = cont[f'{class_key}_right'].values - - gdf = gpd.GeoDataFrame(pd.concat([matched_poly, filt_poly], ignore_index=True)) - - return gdf - - -def alignment_refiner_factory(method) -> AlignmentRefiner: - if method == 'seg': - return SegAlignmentRefiner() - elif method == 'reg': - return RegAlignmentRefiner() - - -def refine_alignment_reg( - source_path: str, - target_path: str, - output_dir: str, -): - - - import deeperhistreg - - ### Define Params ### - registration_params : dict = deeperhistreg.configs.default_initial_nonrigid() - # Alternative: # registration_params = deeperhistreg.configs.load_parameters(config_path) # To load config from JSON file - save_displacement_field : bool = True # Whether to save the displacement field (e.g. for further landmarks/segmentation warping) - copy_target : bool = True # Whether to copy the target (e.g. to simplify the further analysis) - delete_temporary_results : bool = True # Whether to keep the temporary results - case_name : str = "Example_Nonrigid" # Used only if the temporary_path is important, otherwise - provide whatever - temporary_path = None # Will use default if set to None - - ### Create Config ### - config = dict() - config['source_path'] = source_path - config['target_path'] = target_path - config['output_path'] = output_dir - config['registration_parameters'] = registration_params - config['case_name'] = case_name - config['save_displacement_field'] = save_displacement_field - config['copy_target'] = copy_target - config['delete_temporary_results'] = delete_temporary_results - config['temporary_path'] = temporary_path - - ### Run Registration ### - deeperhistreg.run_registration(**config) - - - -def refine_alignment(centroids, polygons, class_key='class', method='seg') -> gpd.GeoDataFrame: - return alignment_refiner_factory(method).refine(centroids, polygons, class_key=class_key) - -@deprecated -def refine_with_anchor( - gdf: gpd.GeoDataFrame, - anchor_gdf: gpd.GeoDataFrame, - pixel_size: float, - img: Union[str, np.ndarray, openslide.OpenSlide, CuImage], # type: ignore - patch_size_um=200, - max_offset=5, - lower_cut=0.4, - upper_cut=0.6 -) -> gpd.GeoDataFrame: - - if isinstance(anchor_gdf.geometry[0], Polygon): - logger.warning('anchor_gdf should contain points, found polygons, converting to their centroid') - anchor_gdf = anchor_gdf.copy().centroid - - wsi = wsi_factory(img) - width, height = wsi.get_dimensions() - patch_size_pxl = patch_size_um / pixel_size - n_col = round(np.ceil(width / patch_size_pxl)) - n_row = round(np.ceil(height / patch_size_pxl)) - - center_gdf = gdf.copy() - center_gdf.geometry = center_gdf.geometry.centroid - center_gdf['polygons'] = gdf.geometry - - ## TODO should match anchors to xenium cells not the other way around - - # get nearest anchor for every cell - nearest_idx = center_gdf.geometry.centroid.sindex.nearest(anchor_gdf.geometry)[1] - # index of nearest cell for each anchor - center_gdf = center_gdf.iloc[nearest_idx].reset_index() - - #center_gdf['nearest_idx'] = nearest_idx - - points_anchor = anchor_gdf.geometry - center_gdf['offset_x'] = points_anchor.x.values - center_gdf.geometry.x.values - center_gdf['offset_y'] = points_anchor.y.values - center_gdf.geometry.y.values - - - polygons = [] - - for i in range(n_row): - for j in range(n_col): - x_left = j * patch_size_pxl - x_right = x_left + patch_size_pxl - y_top = i * patch_size_pxl - y_bottom = y_top + patch_size_pxl - polygons.append(Polygon([(x_left, y_top), (x_right, y_top), (x_right, y_bottom), (x_left, y_bottom)])) - - grid = gpd.GeoDataFrame(geometry=polygons) - - joined = gpd.sjoin(grid, center_gdf, how='left', predicate='contains') - joined = joined.dropna() - joined['index_col'] = joined.index - joined = joined.rename(columns={ - 'index_right': 'index_centroid' - }) - - def trimmed_mean(series): - q1 = series.quantile(lower_cut) - q3 = series.quantile(upper_cut) - filtered_series = series[(series > q1) & (series < q3)] - mean_value = filtered_series.mean() - return mean_value - - logger.info('Remove anchor outliers...') - - grouped = joined.groupby('index_col').agg({ - 'offset_x': trimmed_mean, - 'offset_y': trimmed_mean - }) - grouped = grouped.fillna(0) - - - joined['pooled_offset_x'] = grouped['offset_x'].clip(-max_offset, max_offset) - joined['pooled_offset_y'] = grouped['offset_y'].clip(-max_offset, max_offset) - - joined = joined.fillna(0) - joined['class'] = joined.index - - def shift_polygon(row): - return translate(row.geometry, xoff=row['pooled_offset_x'], yoff=row['pooled_offset_y']) - - - #joined.geometry = joined.apply(shift_polygon, axis=1) - - joined.geometry = center_gdf.geometry.iloc[joined['index_centroid'].astype(int).values].values - joined['class'] = (joined['pooled_offset_x'] + joined['pooled_offset_y']).round() - - logger.info('Shift polygons...') - - joined.geometry = joined['polygons'] - - offsets = joined[~joined.index.duplicated(keep='first')] - - gdf_copy = gdf.copy() - gdf_copy['orig_polygons'] = gdf_copy.geometry - gdf_copy.geometry = gdf_copy.geometry.centroid - ## Add unmatched cells back into the dataframe and set offset based on neighbors - all_cells = gpd.sjoin(grid, gdf_copy, how='left', predicate='contains').dropna().merge(offsets, left_index=True, right_index=True, how='inner') - all_cells = gpd.GeoDataFrame(all_cells, geometry=all_cells['orig_polygons']) - all_cells['class'] = all_cells['class_y'] - - - all_cells.geometry = all_cells.apply(shift_polygon, axis=1) - - - all_cells = all_cells.drop(columns=['offset_x', 'offset_y', 'pooled_offset_x', 'pooled_offset_y']) - - return all_cells \ No newline at end of file diff --git a/src/hest/utils.py b/src/hest/utils.py index 886a562..088bda4 100644 --- a/src/hest/utils.py +++ b/src/hest/utils.py @@ -3,14 +3,16 @@ import concurrent.futures from datetime import datetime import functools +import gc import gzip import json import os +import random import shutil import sys import warnings from enum import Enum -from typing import List, Tuple, Union +from typing import List, Optional, Tuple, Union import zipfile import cv2 @@ -81,6 +83,70 @@ def verify_paths(paths, suffix=""): raise FileNotFoundError(f"No such file or directory: {path}" + suffix) +def read_parquet_dask(path: str, nb_partitions=30) -> dd.DataFrame: + """ Read .parquet file with dask + + Args: + path (str): path to parquet file + nb_partitions (int, optional): approximate number of partitions to use. Defaults to 30. + + Returns: + dd.DataFrame: resulting dask dataframe + """ + import dask.dataframe as dd + import pyarrow.parquet as pq + parquet_file = pq.ParquetFile(path) + total_row_groups = parquet_file.num_row_groups + row_groups_per_partition = max(1, total_row_groups // nb_partitions) + return dd.read_parquet(path, split_row_groups=row_groups_per_partition) + + +def read_parquet_dask_geopandas(path: str, nb_partitions=30) -> dgpd.GeoDataFrame: + """ Read .parquet file with dask geopandas + + Args: + path (str): path to parquet file + nb_partitions (int, optional): approximate number of partitions to use. Defaults to 30. + + Returns: + dgpd.GeoDataFrame: resulting dask geodataframe + """ + import dask_geopandas as dgpd + import pyarrow.parquet as pq + parquet_file = pq.ParquetFile(path) + total_row_groups = parquet_file.num_row_groups + row_groups_per_partition = max(1, total_row_groups // nb_partitions) + return dgpd.read_parquet(path, split_row_groups=row_groups_per_partition) + + +def merge_parquet(folder_path: str, output_file: str) -> None: + """ Merge multiple .parquet files without loading everything into RAM. + + Args: + folder_path (str): input folder containing .parquet files to be merged + output_file (str): full path to merged .parquet file + """ + import pyarrow.parquet as pq + import os + + from tqdm import tqdm + + + files = [f for f in os.listdir(folder_path) if f.endswith('.parquet')] + + writer = None + + for file in tqdm(files): + table = pq.read_table(os.path.join(folder_path, file)) + + if writer is None: + writer = pq.ParquetWriter(output_file, table.schema) + + writer.write_table(table) + + if writer: + writer.close() + logger.remove() logger.add(sys.stdout, format="{time:HH:mm:ss} | {level} | {message}") @@ -122,18 +188,17 @@ def df_morph_um_to_pxl(df, x_key, y_key, pixel_size_morph): def read_xenium_alignment(alignment_file_path: str) -> np.ndarray: - """ Read a xenium alignment file and convert it to a 3x3 affine matrix """ + """ Read a xenium alignment file (matrix or keypoints based) and convert it to a 3x3 affine matrix """ alignment_file = pd.read_csv(alignment_file_path, header=None) alignment_matrix = alignment_file.values - # Xenium explorer >= v2.0 + # Keypoint file if isinstance(alignment_matrix[0][0], str) and 'fixedX' in alignment_matrix[0]: points = pd.read_csv(alignment_file_path) - points = points.iloc[:3] my_dst_pts = points[['fixedX', 'fixedY']].values.astype(np.float32) my_src_pts = points[['alignmentX', 'alignmentY']].values.astype(np.float32) - alignment_matrix = cv2.getAffineTransform(my_src_pts, my_dst_pts) - alignment_matrix = np.vstack((alignment_matrix, [0, 0, 1])) + matrix, _ = cv2.estimateAffinePartial2D(my_src_pts, my_dst_pts) + alignment_matrix = np.vstack((matrix, [0, 0, 1])) return alignment_matrix @@ -382,77 +447,202 @@ def cp_right_folder(path_df: str) -> None: shutil.move(src, dst) -def visualize_random_crops(transcript_df, wsi: WSI, plot_dir='', seg: gpd.GeoDataFrame=None): - """ Plot random crops of transcripts and shapes on top of a WSI """ +def is_dask_gdf(gdf): + cls = gdf.__class__ + module = cls.__module__ + name = cls.__name__ + + return "dask_geopandas" in module and name == "GeoDataFrame" + +def is_dask_dataframe(df): + return 'dask.dataframe' in str(type(df)) + + +def plot_shapes_qc( + wsi: WSI, + gdf: Union[gpd.GeoDataFrame, dgpd.GeoDataFrame], + plot_dir='', + nb=15, +): import matplotlib.pyplot as plt from shapely import Polygon - K = 15 - N = len(transcript_df) if transcript_df else len(seg) + K = nb + N = len(gdf) size_region = 1000 - random_idx = np.random.randint(0, N, K) - ratio = 0.01 - width, height = wsi.get_dimensions() - thumb = Image.fromarray(wsi.get_thumbnail(round(width * 0.1), round(height * 0.1))) - if transcript_df: - downsampled = transcript_df.sample(5000) + random_points = gdf.sample(frac=K/N) + random_centroids = random_points.geometry.centroid + if is_dask_gdf(gdf): + random_centroids = random_centroids.compute() + random_centroids = np.array(random_centroids) - xy = downsampled[['he_x', 'he_y']].values + for centroid in random_centroids: + xy_center = [centroid.x, centroid.y] + left_x = round(xy_center[0]-size_region // 2) + right_x = xy_center[0] + size_region // 2 + bottom_y = xy_center[1] + size_region // 2 + top_y = round(xy_center[1]-size_region // 2) + region = wsi.read_region_pil((left_x, top_y), 0, (size_region, size_region)) fig, ax = plt.subplots() - ax.imshow(thumb) - #for geom in downsampled_cells.geometry: - # ax.plot(*geom.exterior.xy, linewidth=0.2, color='red') - ax.scatter(xy[:, 0] * 0.1, xy[:, 1] * 0.1, s=0.5) + ax.imshow(region) + patch_poly = Polygon([ + [left_x, top_y], + [right_x, top_y], + [right_x, bottom_y], + [left_x, bottom_y] + ]) + sub_seg = gdf[gdf.intersects(patch_poly)] + sub_seg = sub_seg.translate(-left_x, -top_y) + for geom in sub_seg.geometry: + ax.plot(*geom.exterior.xy, linewidth=0.2, color='green') + ax.axis('off') - os.makedirs(plot_dir, exist_ok=True) - fig.savefig(os.path.join(plot_dir, 'transcripts_plot.jpg'), bbox_inches='tight', dpi=200) + + fig.savefig(os.path.join(plot_dir, f'x_{left_x}_y_{top_y}.jpg'), bbox_inches='tight', dpi=200) + + +def _get_random_transcript_names(df: Union[pd.DataFrame, dd.DataFrame], k=3): + sub_df = df.sample(frac=1000 / len(df)) + unique_feature_names = np.unique(sub_df['feature_name']) - for k in random_idx: - if transcript_df: - xy_center = transcript_df[['he_x', 'he_y']].iloc[k].values - else: - point = seg.geometry.centroid.iloc[k] - xy_center = [point.x, point.y] + shuffled_features = list(unique_feature_names) + random.shuffle(shuffled_features) + + res = [] + for feat_name in shuffled_features: + if not 'Deprecated' in feat_name and \ + not 'Unassigned' in feat_name and \ + not 'blank' in feat_name.lower(): + + res.append(feat_name) + if len(res) >= k: + break + return res + + +def plot_transcripts_qc( + wsi: WSI, + df: Union[pd.DataFrame, dd.DataFrame], + plot_dir='', + nb=15, + verbose=True, + plot_global=True +): + import matplotlib.pyplot as plt + from shapely.geometry import box + import geopandas as gpd + + use_dask = not isinstance(df, pd.DataFrame) + + K = nb + N = len(df) + size_region = 1000 + random_idx = np.random.randint(0, N, K) + + width, height = wsi.get_dimensions() + + if plot_global: + thumb = Image.fromarray(wsi.get_thumbnail(round(width * 0.1), round(height * 0.1))) + feat_names = _get_random_transcript_names(df, k=3) + if verbose: + print(f"Plot global transcripts for the following genes: {feat_names}.") + + for feat_name in tqdm(feat_names): + sub_df = df[df['feature_name'] == feat_name] + downsampled = sub_df.sample(frac=min(3000, len(sub_df))/len(sub_df)) + + xy = np.array(downsampled[['he_x', 'he_y']].values) + + fig, ax = plt.subplots() + ax.imshow(thumb) + ax.scatter(xy[:, 0] * 0.1, xy[:, 1] * 0.1, s=0.5) + ax.axis('off') + os.makedirs(plot_dir, exist_ok=True) + fig.savefig(os.path.join(plot_dir, f'global_plot_{feat_name}.jpg'), bbox_inches='tight', dpi=200) + + # speedup for dask + k_rows = np.array(df[['he_x', 'he_y']].sample(frac=len(random_idx)/len(df)).values) + + if use_dask: + import dask_geopandas as dgpd + points_gdf = dgpd.from_dask_dataframe( + df, geometry=dgpd.points_from_xy(df, "he_x", "he_y") + ) + else: + points_gdf = gpd.points_from_xy(df, df['he_x'], df['he_y']) + + boxes = [] + for xy_center in k_rows: + left_x = xy_center[0] - size_region // 2 + right_x = xy_center[0] + size_region // 2 + bottom_y = xy_center[1] + size_region // 2 + top_y = xy_center[1] - size_region // 2 + boxes.append(box(left_x, top_y, right_x, bottom_y)) + + regions_gdf = gpd.GeoDataFrame({"geometry": boxes, "region_id": range(len(boxes))}) + + joined_df = points_gdf.sjoin(regions_gdf, how="inner") + if use_dask: + joined_df = joined_df.compute() + + for k in tqdm(range(len(k_rows))): + xy_center = k_rows[k] left_x = round(xy_center[0]-size_region // 2) right_x = xy_center[0] + size_region // 2 bottom_y = xy_center[1] + size_region // 2 top_y = round(xy_center[1]-size_region // 2) region = wsi.read_region_pil((left_x, top_y), 0, (size_region, size_region)) + + sub_df = joined_df[joined_df['region_id'] == k] - if transcript_df: - sub_transcripts = transcript_df[ - (left_x < transcript_df['he_x']) & - (top_y < transcript_df['he_y']) & - (transcript_df['he_x'] < right_x) & - (transcript_df['he_y'] < bottom_y) - ] - - sub_transcripts = sub_transcripts.sample(round(ratio * len(sub_transcripts))) - + sub_df = sub_df.sample(frac=min(5000, len(sub_df))/len(sub_df)) + fig, ax = plt.subplots() ax.imshow(region) - if seg is not None: - patch_poly = Polygon([ - [left_x, top_y], - [right_x, top_y], - [right_x, bottom_y], - [left_x, bottom_y] - ]) - sub_seg = seg[seg.intersects(patch_poly)] - sub_seg = sub_seg.translate(-left_x, -top_y) - for geom in sub_seg.geometry: - ax.plot(*geom.exterior.xy, linewidth=0.2, color='green') - if transcript_df: - xy = sub_transcripts[['he_x', 'he_y']].values - ax.scatter(xy[:, 0] - left_x, xy[:, 1] - top_y, s=0.5) + xy = np.array(sub_df[['he_x', 'he_y']].values) + ax.scatter(xy[:, 0] - left_x, xy[:, 1] - top_y, s=0.5) ax.axis('off') - fig.savefig(os.path.join(plot_dir, str(k) + '_transcripts_plot.jpg'), bbox_inches='tight', dpi=200) + fig.savefig(os.path.join(plot_dir, f'x_{left_x}_y_{top_y}.jpg'), bbox_inches='tight', dpi=200) + + +def plot_xenium_align_qc( + wsi: WSI, + plot_dir='', + transcript_df: Optional[Union[pd.DataFrame, dd.DataFrame]]=None, + seg_nuc: Optional[Union[gpd.GeoDataFrame, gpdd.GeoDataFrame]]=None, + seg_cells: Optional[Union[gpd.GeoDataFrame, gpdd.GeoDataFrame]]=None, + nb=15, + plot_global=False +): + if transcript_df is not None: + os.makedirs(os.path.join(plot_dir, 'transcripts'), exist_ok=True) + plot_transcripts_qc(wsi, transcript_df, + plot_dir=os.path.join(plot_dir, 'transcripts'), + nb=nb, plot_global=plot_global) + + del transcript_df + gc.collect() + + if seg_nuc is not None: + os.makedirs(os.path.join(plot_dir, 'xenium_nuc'), exist_ok=True) + plot_shapes_qc(wsi, seg_nuc, + plot_dir=os.path.join(plot_dir, 'xenium_nuc'), nb=nb) + + del seg_nuc + gc.collect() + + if seg_cells is not None: + os.makedirs(os.path.join(plot_dir, 'xenium_cells'), exist_ok=True) + plot_shapes_qc(wsi, seg_cells, + plot_dir=os.path.join(plot_dir, 'xenium_cells'), nb=nb) + del seg_cells + gc.collect() def get_k_genes_from_df(meta_df: pd.DataFrame, k: int, criteria: str, save_dir: str=None) -> List[str]: """Get the k genes according to some criteria across common genes in all the samples in the HEST meta dataframe diff --git a/tests/hest_tests.py b/tests/hest_tests.py index b8c554d..b34489d 100644 --- a/tests/hest_tests.py +++ b/tests/hest_tests.py @@ -152,14 +152,10 @@ def test_tissue_seg(self): with self.subTest(st_object=idx): st.segment_tissue(method='deep') st.save_tissue_contours(self.output_dir, name=f'deep_{idx}') - st.save_tissue_seg_jpg(self.output_dir, name=f'deep_{idx}') - st.save_tissue_seg_pkl(self.output_dir, name=f'deep_{idx}') st.save_tissue_vis(self.output_dir, name=f'deep_{idx}') st.segment_tissue(method='otsu') st.save_tissue_contours(self.output_dir, name=f'otsu_{idx}') - st.save_tissue_seg_jpg(self.output_dir, name=f'otsu_{idx}') - st.save_tissue_seg_pkl(self.output_dir, name=f'otsu_{idx}') st.save_tissue_vis(self.output_dir, name=f'otsu_{idx}') @@ -224,7 +220,7 @@ def test_saving(self): loader = unittest.TestLoader() suite = loader.loadTestsFromTestCase(TestHESTData) # suite = unittest.TestSuite() - # suite.addTest(TestHESTData('test_patching')) + #suite.addTest(TestHESTData('test_spatialdata')) result = unittest.TextTestRunner(verbosity=2).run(suite) if not result.wasSuccessful(): raise Exception('Test failed') \ No newline at end of file diff --git a/tutorials/1-Downloading-HEST-1k.ipynb b/tutorials/1-Downloading-HEST-1k.ipynb index 59096de..51f21df 100644 --- a/tutorials/1-Downloading-HEST-1k.ipynb +++ b/tutorials/1-Downloading-HEST-1k.ipynb @@ -18,18 +18,19 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "## Instructions for Setting Up HuggingFace Account and Token\n", "\n", - "### 1. Create an Account on HuggingFace\n", + "### Instructions for Setting Up HuggingFace Account and Token\n", + "\n", + "#### 1. Create an Account on HuggingFace\n", "Follow the instructions provided on the [HuggingFace sign-up page](https://huggingface.co/join).\n", "\n", - "### 2. Accept terms of use of HEST\n", + "#### 2. Accept terms of use of HEST\n", "\n", "1. Go to [HEST HuggingFace page](https://huggingface.co/datasets/MahmoodLab/hest)\n", "2. Request access (access will be automatically granted)\n", "3. At this stage, you can already manually inspect the data by navigating in the `Files and version`\n", "\n", - "### 3. Create a Hugging Face Token\n", + "#### 3. Create a Hugging Face Token\n", "\n", "1. **Go to Settings:** Navigate to your profile settings by clicking on your profile picture in the top right corner and selecting `Settings` from the dropdown menu.\n", "\n", @@ -48,7 +49,7 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "### 4. Logging\n", + "#### 4. Logging\n", "\n", "Install the python library `datasets` and run cell below. If successful, you should see:\n", "\n", @@ -183,8 +184,7 @@ "- **spatial_plots/**: Overlay of the WSI with the st spots\n", "- **thumbnails/**: Downscaled version of the WSI\n", "- **tissue_seg/**: Tissue segmentation masks:\n", - " - `{id}_mask.jpg`: Downscaled or full resolution greyscale tissue mask\n", - " - `{id}_mask.pkl`: Tissue/holes contours in a pickle file\n", + " - `{id}.geojson`: Tissue segmentation mask\n", " - `{id}_vis.jpg`: Visualization of the tissue mask on the downscaled WSI\n", "- **pixel_size_vis/**: Visualization of the pixel size\n", "- **patches/**: 256x256 H&E patches (0.5µm/px) extracted around ST spots in a .h5 object optimized for deep-learning. Each patch is matched to the corresponding ST profile (see **st/**) with a barcode.\n", diff --git a/tutorials/2-Interacting-with-HEST-1k.ipynb b/tutorials/2-Interacting-with-HEST-1k.ipynb index 5b098ef..337b5a3 100644 --- a/tutorials/2-Interacting-with-HEST-1k.ipynb +++ b/tutorials/2-Interacting-with-HEST-1k.ipynb @@ -65,7 +65,7 @@ "metadata": {}, "source": [ "`st.adata` is a spatial scanpy object containing the following:\n", - "## Observations (st.adata.obs)\n", + "#### Observations (st.adata.obs)\n", "- `in_tissue`: Indicator if the observation is within the tissue (`in_tissue` comes from the initial Visium/Xenium run and might not be accurate, prefer the segmentation obtained by st.segment_tissue() instead).\n", "- `pxl_col_in_fullres`: Pixel column position of the patch/spot centroid in the full resolution image.\n", "- `pxl_row_in_fullres`: Pixel row position of the patch/spot centroid in the full resolution image.\n", @@ -84,7 +84,7 @@ "- `log1p_total_counts_mito`: Log-transformed total mitochondrial counts. (note that this field might not be accurate)\n", "- `pct_counts_mito`: Percentage of counts that are mitochondrial. (note that this field might not be accurate)\n", "\n", - "## Variables (st.adata.var)\n", + "#### Variables (st.adata.var)\n", "- `n_cells_by_counts`: Number of cells detected by counts for each variable.\n", "- `mean_counts`: Mean counts per variable.\n", "- `log1p_mean_counts`: Log-transformed mean counts.\n", @@ -93,10 +93,10 @@ "- `log1p_total_counts`: Log-transformed total counts.\n", "- `mito`: Indicator if the gene is mitochondrial. (note that this field might not be accurate)\n", "\n", - "## Unstructured (st.adata.uns)\n", + "#### Unstructured (st.adata.uns)\n", "- `spatial`: Contains a downscaled version of the full resolution image in `st.adata.uns['spatial']['ST']['images']['downscaled_fullres']`\n", "\n", - "## Observation-wise Multidimensional (st.adata.obsm)\n", + "#### Observation-wise Multidimensional (st.adata.obsm)\n", "- `spatial`: Pixel coordinates of spots/patches centroids on the full resolution image. (first column is x axis, second column is y axis)" ] }, @@ -104,7 +104,7 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "## Visualizing the spots over a downscaled version of the WSI" + "### Visualizing the spots over a downscaled version of the WSI" ] }, { @@ -122,7 +122,7 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "## Saving to pyramidal tiff and h5\n", + "### Saving to pyramidal tiff and h5\n", "Save `HESTData` object to `.tiff` + expression `.h5ad` and a metadata file." ] }, @@ -140,7 +140,7 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "## Tissue segmentation\n", + "### Tissue segmentation\n", "\n", "We integrated 2 tissue segmentation methods:\n", "\n", @@ -168,9 +168,9 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "## Changing the Patching/pooling size\n", + "### Changing the Patching/pooling size\n", "\n", - "### Patching\n", + "#### Patching\n", "You can change the size of patches around the spots with `dump_patches`:" ] }, @@ -241,7 +241,7 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "## Batch effect visualization" + "### Batch effect visualization" ] }, { @@ -281,7 +281,7 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "## Loading transcripts (for Xenium only)\n", + "### Loading transcripts (for Xenium only)\n", "The transcript dataframe contains the following columns:\n", "- `cell_id` (default Xenium Explorer flag): id of cell containing this transcript\n", "- `he_x`, `he_y`: x, y coordinate or transcript on the HE WSI in pixels (aligned with a non-rigid transformation)\n", diff --git a/tutorials/3-Assembling-HEST-Data.ipynb b/tutorials/3-Assembling-HEST-Data.ipynb index d0162a3..3ab6048 100644 --- a/tutorials/3-Assembling-HEST-Data.ipynb +++ b/tutorials/3-Assembling-HEST-Data.ipynb @@ -6,6 +6,8 @@ "source": [ "## Step-by-step instructions to assemble HEST data \n", "\n", + "\n", + "### I. Visium reader\n", "This tutorial will guide you to convert a legacy Visium sample into a HEST-compatible object. \n" ] }, @@ -97,7 +99,7 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "### When should I provide an alignment file and when should I use the autoalignment?\n", + "### When should I provide an alignment file and when should I use autoalignment?\n", "\n", "#### Step 1: check if a tissue_positions.csv/tissue_position_list.csv already provides a correct alignment\n", "\n", @@ -177,14 +179,16 @@ "cell_type": "markdown", "metadata": {}, "source": [ - "# Assembling Xenium samples" + "### II. Xenium reader" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ - "## Download Xenium sample from 10x genomics website" + "### Download Xenium sample from 10x genomics website\n", + "\n", + "Download the following xenium files and place them in the same directory" ] }, { @@ -234,6 +238,42 @@ "st = XeniumReader().auto_read(xenium_folder_path)" ] }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Working with larger than RAM Xenium samples (Xenium 5k)\n", + "We support larger than RAM transcripts pooling powered by dask. Dask will automatically chunk the data such that it never has to hold the entire transcript dataframe in memory.\n", + "\n", + "Dask will attempt to process one partition per thread. To avoid loading large partitions on systems having a low amount of RAM, we advise using a high number of partitions (>100), as well as a single worker and a low number of threads (<4 depending on RAM available).\n", + "\n", + "> Note: feel free to open the dask dashboard to visualize workers, partitions and resources (usually on http://localhost:8787)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "from dask.distributed import LocalCluster, Client\n", + "\n", + "cluster = LocalCluster(\n", + " \"127.0.0.1:8786\",\n", + " n_workers=1, # increase depending on RAM available\n", + " memory_limit=\"20GB\", # dask will kill the worker if this is exceeded\n", + " threads_per_worker=1, # increase depending on RAM available\n", + ")\n", + "client = Client(cluster)\n", + "print('dashboard is available at: ', client.dashboard_link)\n", + "\n", + "st = XeniumReader().auto_read(\n", + " xenium_folder_path, \n", + " use_dask=True, \n", + " nb_partitions=100\n", + ")" + ] + }, { "cell_type": "code", "execution_count": null, @@ -252,6 +292,15 @@ "st.segment_tissue()" ] }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "We then compute patches centered around pseudo-visium transcript bins.\n", + "\n", + "> Warning: note that patches might be larger than transcript bins. " + ] + }, { "cell_type": "code", "execution_count": null, @@ -269,6 +318,110 @@ "source": [ "st.save('save', save_img=False)" ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Adding new samples to HEST\n", + "This section explains how to format new samples for HuggingFace\n", + "\n", + "### I. Sample preparation\n", + "\n", + "#### 1. Download raw datasets in structured folders\n", + "\n", + "Create the following folder structure:\n", + "```python\n", + "my_data/\n", + " xenium/\n", + " dataset_name_1/\n", + " subseries1/\n", + " sample_1/\n", + " ...\n", + " sample_2/\n", + " ...\n", + " ...\n", + " dataset_name_2/\n", + " sample_1/\n", + " ...\n", + " visium-hd/\n", + " ...\n", + " visium/\n", + " ...\n", + "```\n", + "\n", + "Then, download corresponding files for xenium, visium, visium-hd... See reader examples above or refer to the doc for the list of required files.\n", + "\n", + "#### 2. Add metadata rows to CSV\n", + "Fill columns in the sample CSV, specifically: dataset_title, subseries should match with the folder structure.\n", + "\n", + "### II. Sample processing\n", + "\n", + "#### 1. Process raw samples\n", + "\n", + "Set the path to your sample CSV in `tutorials/scripts/1_process_raw_samples.py`, also modify the memory_limit, we recommend setting dask memory limit to at least 20GB for Xenium 5k, eventhough 30GB is safer.\n", + "\n", + "Process and save raw samples by launching `python tutorials/scripts/1_process_raw_samples.py`.\n", + "\n", + "Processing Xenium 5k will take time, you can monitor progress on the dask dashboard (usually at `localhost:8787` after launch).\n", + "\n", + "\n", + "#### 2. Check Xenium DAPI/HE alignment and realign if necessary\n", + "\n", + "By default, the Xenium platform uses a single affine transform for the whole WSI. For some samples, the resulting alignment might be unsatisfactory.\n", + "In order to visualize the affine alignment, either:\n", + "- open the resulting `.geojson` files (in `processed/`) in QuPath (might not work with QuPath >=0.6) \\\n", + "\\\n", + "or\n", + "
\n", + "
\n", + "- launch `tutorials/scripts/2_check_xenium_alignment.py`. \n", + "\n", + "#### 3. Micro-align DAPI to HE with Valis\n", + "\n", + "If the alignment from steps (1-2) is unsatisfactory, we highly recommend using Valis non-rigid micro-alignment for sub-cellular alignment precision. This is crucial for correctly mapping transcripts and cells to H&E.\n", + "\n", + "\n", + "##### a. Register DAPI to HE with Valis\n", + "\n", + "We provide a modified Valis version for simplified pythonic use and improved precision, please check-out the original repository [here](https://github.com/MathOnco/valis).\n", + "\n", + "Open [3a_microalign_xenium.py](./scripts/3a_microalign_xenium.py), in this example, we use `morphology_focus_0000.ome.tif` as the DAPI slide, feel free to use `morphology_focus.ome.tif`.\n", + "\n", + "Monitor alignment quality at `results/{sample_name}/{date}/overlaps/`, check both `_rigid_overlap.png` and `micro_reg.png`.\n", + "\n", + "\n", + "##### b. Warp transcripts, cells and nuclei using Valis and Dask\n", + "\n", + "See [3b_microalign_xenium.py](./scripts/3b_microalign_xenium.py) to warp transcripts, nuclei and cells from DAPI to H&E.\n", + "\n", + "Then re-run step 3.b, in order to compare quality.\n", + "\n", + "#### 4. Segment with CellViT\n", + "\n", + "See [4_segment_cellvit.py](./scripts/4_segment_cellvit.py) to segment with CellViT.\n", + "\n", + "\n", + "#### 5. Copy processed files to HEST_results/\n", + "\n", + "See [5_copy_processed.py](./scripts/5_copy_processed.py) to copy files to the structure expected by huggingface.\n", + "\n", + "\n", + "#### 6. Generate a new HEST_vX_Y_Z.csv sheet\n", + "\n", + "See [6_generate_new_meta.py](./scripts/6_generate_new_meta.py) to create a new HEST_vX_Y_Z.csv. Once created, copy it to the `HEST_results/` folder.\n", + "\n", + "\n", + "#### 7. Upload to HuggingFace\n", + "\n", + "See [7_upload_huggingface.py](./scripts/7_upload_huggingface.py) to upload to HuggingFace via a PR.\n", + "\n" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [] } ], "metadata": { diff --git a/tutorials/scripts/1_process_raw_samples.py b/tutorials/scripts/1_process_raw_samples.py new file mode 100644 index 0000000..f0a35fa --- /dev/null +++ b/tutorials/scripts/1_process_raw_samples.py @@ -0,0 +1,40 @@ +import pandas as pd +from hest.readers import process_raw_samples + +from dask.distributed import LocalCluster, Client + + +id_list = [ + 'TENX202', + 'TENX201', + 'TENX200', + 'TENX199', + 'TENX198', + 'TENX197', +] + +meta_df = pd.read_csv("/home/paul/Downloads/ST H&E datasets - 10XGenomics.csv") +meta_df = meta_df[meta_df['id'].isin(id_list)] + + +if __name__ == "__main__": + cluster = LocalCluster( + "127.0.0.1:8786", + n_workers=1, + memory_limit="30GB", + threads_per_worker=1, + ) + client = Client(cluster) + print('dashboard is available at: ', client.dashboard_link) + print(client) + + + process_raw_samples( + meta_df, + save_img=True, + save_kwargs={ + 'save_nuclei_seg': True, + 'save_cell_seg': True, + 'save_transcripts': True + } + ) diff --git a/tutorials/scripts/2_check_xenium_alignment.py b/tutorials/scripts/2_check_xenium_alignment.py new file mode 100644 index 0000000..a1bba8b --- /dev/null +++ b/tutorials/scripts/2_check_xenium_alignment.py @@ -0,0 +1,50 @@ +import os +from dask.distributed import LocalCluster, Client +import pandas as pd + +from hest.utils import get_path_from_meta_row, plot_xenium_align_qc, read_parquet_dask, read_parquet_dask_geopandas +from hestcore.wsi import wsi_factory + +id_list = [ + 'NCBI888', + 'NCBI887', + 'NCBI886', + 'NCBI885', +] + +meta_df = pd.read_csv("/home/paul/Downloads/ST H&E datasets - NCBI.csv") +meta_df = meta_df[meta_df['id'].isin(id_list)] + +if __name__ == '__main__': + cluster = LocalCluster( + "127.0.0.1:8786", + n_workers=1, + memory_limit="20GB", + threads_per_worker=1, + ) + client = Client(cluster) + + + for _, row in meta_df.iterrows(): + base_path = get_path_from_meta_row(row) + folder_path = os.path.join(base_path, 'processed') + print(f"reading {folder_path}") + + wsi = wsi_factory(os.path.join(folder_path, f'aligned_fullres_HE.tif')) + + df = read_parquet_dask(os.path.join(folder_path, 'aligned_transcripts.parquet')) + seg_cells = read_parquet_dask_geopandas(os.path.join(folder_path, 'he_cell_seg.parquet'), nb_partitions=30) + seg_nuc = read_parquet_dask_geopandas(os.path.join(folder_path, 'he_nucleus_seg.parquet'), nb_partitions=30) + + # This might be slow for xenium samples + plot_xenium_align_qc( + wsi, + plot_dir=os.path.join(folder_path, 'qc'), + transcript_df=df, + seg_nuc=seg_nuc, + seg_cells=seg_cells, + nb=25, + plot_global=True + ) + + \ No newline at end of file diff --git a/tutorials/scripts/3a_microalign_xenium.py b/tutorials/scripts/3a_microalign_xenium.py new file mode 100644 index 0000000..6182859 --- /dev/null +++ b/tutorials/scripts/3a_microalign_xenium.py @@ -0,0 +1,43 @@ + +import os + +import pandas as pd + +from hest.registration import register_dapi_he +from hest.utils import get_path_from_meta_row + +id_list = [ + 'TENX202', + 'TENX201', + 'TENX200', + 'TENX199', + 'TENX198', + 'TENX197', +] + +meta_df = pd.read_csv("/home/paul/Downloads/ST H&E datasets - 10XGenomics.csv") +meta_df = meta_df[meta_df['id'].isin(id_list)] + +if __name__ == "__main__": + for _, row in meta_df.iterrows(): + base_path = get_path_from_meta_row(row) + + dapi_path = None + if os.path.exists(os.path.join(base_path, 'morphology_focus', 'ch0000_dapi.ome.tif')): + dapi_path = os.path.join(base_path, 'morphology_focus', 'ch0000_dapi.ome.tif') + elif os.path.exists(os.path.join(base_path, 'morphology_focus', 'morphology_focus_0000.ome.tif')): + dapi_path = os.path.join(base_path, 'morphology_focus', 'morphology_focus_0000.ome.tif') + else: + raise ValueError("No valid dapi path found,") + + register_dapi_he( + os.path.join(base_path, 'processed/aligned_fullres_HE.tif'), + dapi_path, + os.path.join("results", row['id']), + name = None, + max_non_rigid_registration_dim_px=10000, + micro_rigid_registrar_cls=None, + micro_rigid_registrar_params={}, + micro_reg=True, + check_for_reflections=False + ) \ No newline at end of file diff --git a/tutorials/scripts/3b_microalign_xenium.py b/tutorials/scripts/3b_microalign_xenium.py new file mode 100644 index 0000000..ca64c5e --- /dev/null +++ b/tutorials/scripts/3b_microalign_xenium.py @@ -0,0 +1,68 @@ + +import os + +import dask +import pandas as pd +from hest.registration import warp_and_save_xenium_objects +from dask.distributed import LocalCluster, Client, WorkerPlugin + +from hest.utils import get_path_from_meta_row + +id_list = [ + 'TENX195', +] + +meta_df = pd.read_csv("/home/paul/Downloads/ST H&E datasets - 10XGenomics.csv") +meta_df = meta_df[meta_df['id'].isin(id_list)] + +if __name__ == "__main__": + from valis_hest.registration import init_jvm + init_jvm(mem_gb=2) + class JVMPlugin(WorkerPlugin): + def setup(self, worker): + from valis_hest.registration import init_jvm + import jpype + if not jpype.isJVMStarted(): + init_jvm(mem_gb=1) + + + dask.config.set({ + "distributed.scheduler.worker-ttl": None, + }) + + cluster = LocalCluster( + "127.0.0.1:8786", + n_workers=1, + memory_limit="32GB", + threads_per_worker=1, + ) + client = Client(cluster) + client.register_worker_plugin(JVMPlugin(), name="jvm") + + registrar_base = '/home/paul/HEST/results/' + + + for _, row in meta_df.iterrows(): + base_path = get_path_from_meta_row(row) + + dapi_transcripts_path = os.path.join(base_path, 'transcripts.parquet') + dapi_nucleus_path = os.path.join(base_path, 'nucleus_boundaries.parquet') + dapi_cell_path = os.path.join(base_path, 'cell_boundaries.parquet') + save_dir = os.path.join(base_path, 'processed') + + dirname = list(os.listdir(os.path.join(registrar_base, row['id'])))[0] + + + warp_and_save_xenium_objects( + os.path.join(registrar_base, row['id'], dirname, 'data', '_registrar.pickle'), + 'ch0000_dapi.ome.tif', + #'morphology_focus_0000.ome.tif', + save_dir, + dapi_transcripts=dapi_transcripts_path, + dapi_cells=dapi_cell_path, + dapi_nuclei=dapi_nucleus_path, + use_dask=True, + verbose=True, + save_geojson=True + ) + \ No newline at end of file diff --git a/tutorials/scripts/4_segment_cellvit.py b/tutorials/scripts/4_segment_cellvit.py new file mode 100644 index 0000000..ab3e3b8 --- /dev/null +++ b/tutorials/scripts/4_segment_cellvit.py @@ -0,0 +1,22 @@ +import pandas as pd + +from hest.readers import process_meta_df_cellvit + + +id_list = [ + 'TENX202', + 'TENX201', + 'TENX200', + 'TENX199', + 'TENX198', + 'TENX197', +] + + +if __name__ == '__main__': + df = pd.read_csv("/home/paul/Downloads/ST H&E datasets - 10XGenomics.csv") + + + df = df[df['id'].isin(id_list)] + + process_meta_df_cellvit(df, {'gpu_ids': [0], 'batch_size': 1, 'model': 'CellViT-256-x40.pth'}) diff --git a/tutorials/scripts/5_copy_processed.py b/tutorials/scripts/5_copy_processed.py new file mode 100644 index 0000000..7c68cc7 --- /dev/null +++ b/tutorials/scripts/5_copy_processed.py @@ -0,0 +1,89 @@ +import json +import os +import pandas as pd +import scanpy +from tqdm import tqdm +from hest.utils import copy_processed, get_path_from_meta_row + + +id_list = [ + 'TENX160', + #'TENX187', + #'TENX188', +] + +df = pd.read_csv("/home/paul/Downloads/ST H&E datasets - 10XGenomics.csv") + + +df = df[df['id'].isin(id_list)] + + +def update_hest_meta(old_df, new_df, new_version: str): + required_cols = [ + 'dataset_title', + 'id', + 'image_filename', + 'organ', + 'disease_state', + 'oncotree_code', + 'species', + 'patient', + 'st_technology', + 'data_publication_date', + 'license', + 'study_link', + 'download_page_link1', + 'inter_spot_dist', + 'spot_diameter', + 'spots_under_tissue', + 'preservation_method', + 'nb_genes', + 'treatment_comment', + 'pixel_size_um_embedded', + 'pixel_size_um_estimated', + 'magnification', + 'fullres_px_width', + 'fullres_px_height', + 'tissue', + 'disease_comment', + 'subseries', + 'hest_version_added', + ] + + new_rows = [] + for _, row in tqdm(new_df.iterrows()): + path = get_path_from_meta_row(row) + + with open(os.path.join(path, 'processed', f'meta.json'), 'r') as f: + meta = json.load(f) + + adata = scanpy.read_h5ad(os.path.join(path, 'processed', 'aligned_adata.h5ad')) + nb_genes = len(adata.var_names) + + new_rows.append( + {**{k: v for k, v in meta.items() if k in required_cols}, **{ + 'hest_version_added': new_version, + 'image_filename': meta['id'] + '.tif', + 'nb_genes': nb_genes + }} + ) + new_df = pd.concat((pd.DataFrame(new_rows, columns=required_cols), old_df)) + return new_df + + +if __name__ == '__main__': + copy_processed( + '/media/paul/ssd2/HEST_results', + df, + cp_pyramidal=True, + n_job=1, + cp_meta=True, + cp_adata=True, + cp_cellvit=True, + cp_downscaled=True, + cp_patches=True, + cp_spatial=True, + cp_pixel_vis=True, + cp_transcripts=True, + cp_he_seg=True + ) \ No newline at end of file diff --git a/tutorials/scripts/6_generate_new_meta.py b/tutorials/scripts/6_generate_new_meta.py new file mode 100644 index 0000000..dfa1cd2 --- /dev/null +++ b/tutorials/scripts/6_generate_new_meta.py @@ -0,0 +1,78 @@ +import json +import os +import pandas as pd +import scanpy +from tqdm import tqdm +from hest.utils import get_path_from_meta_row + + + +def update_hest_meta(old_df, new_df, new_version: str): + required_cols = [ + 'dataset_title', + 'id', + 'image_filename', + 'organ', + 'disease_state', + 'oncotree_code', + 'species', + 'patient', + 'st_technology', + 'data_publication_date', + 'license', + 'study_link', + 'download_page_link1', + 'inter_spot_dist', + 'spot_diameter', + 'spots_under_tissue', + 'preservation_method', + 'nb_genes', + 'treatment_comment', + 'pixel_size_um_embedded', + 'pixel_size_um_estimated', + 'magnification', + 'fullres_px_width', + 'fullres_px_height', + 'tissue', + 'disease_comment', + 'subseries', + 'hest_version_added', + ] + + new_rows = [] + for _, row in tqdm(new_df.iterrows()): + path = get_path_from_meta_row(row) + + with open(os.path.join(path, 'processed', f'meta.json'), 'r') as f: + meta = json.load(f) + + adata = scanpy.read_h5ad(os.path.join(path, 'processed', 'aligned_adata.h5ad')) + nb_genes = len(adata.var_names) + + new_rows.append( + {**{k: v for k, v in meta.items() if k in required_cols}, **{ + 'hest_version_added': new_version, + 'image_filename': meta['id'] + '.tif', + 'nb_genes': nb_genes + }} + ) + new_df = pd.concat((pd.DataFrame(new_rows, columns=required_cols), old_df)) + return new_df + + +if __name__ == '__main__': + id_list = [ + "NCBI885", + "NCBI886", + "NCBI887", + "NCBI888", + ] + + meta_df = pd.read_csv("/home/paul/Downloads/ST H&E datasets - NCBI.csv") + + + meta_df = meta_df[meta_df['id'].isin(id_list)] + + new_version = 'v1_3_0' + df = update_hest_meta(pd.read_csv('/home/paul/Downloads/HEST_v1_2_1.csv'), meta_df, new_version) + df.to_csv(f'/home/paul/Downloads/HEST_{new_version}.csv', index=False) \ No newline at end of file diff --git a/tutorials/scripts/7_upload_huggingface.py b/tutorials/scripts/7_upload_huggingface.py new file mode 100644 index 0000000..6e466b4 --- /dev/null +++ b/tutorials/scripts/7_upload_huggingface.py @@ -0,0 +1,30 @@ +from huggingface_hub import HfApi + + +if __name__ == '__main__': + version_name = 'v1.3.0' + input_path = "/media/paul/ssd2/HEST_results" + + api = HfApi() + + # res = api.create_pull_request( + # repo_id="MahmoodLab/hest", + # repo_type="dataset", + # title=version_name + # ) + + + api.upload_large_folder( + folder_path=input_path, + repo_id="MahmoodLab/hest", + repo_type="dataset", + revision='refs/pr/26' + ) + + # api.upload_file( + # path_or_fileobj="/home/paul/Downloads/README.md", + # path_in_repo="README.md", + # repo_id="MahmoodLab/hest", + # repo_type="dataset", + # revision='refs/pr/23' + # ) \ No newline at end of file