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