Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 5 additions & 4 deletions docs/preprocessing.rst
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,8 @@ Tissue Segmentation

``segmentation`` is forwarded directly to
`hs2p <https://github.com/clemsgrs/hs2p>`_\ 's segmentation pipeline.
It is a partial override: omitted keys retain the standard configuration
defaults, including ``method="hsv"``.
The ``method`` key selects the algorithm:

- ``hsv`` - heuristic based on the HSV colour space. Fast and robust for H&E slides.
Expand Down Expand Up @@ -301,15 +303,14 @@ Preview Images

``slide2vec`` can write a tissue mask preview and a tiling preview for each slide.
These are particularly useful for quality control.
Both are disabled by default. Enable them via the ``preview`` dict:
Both are enabled by default. The ``preview`` dict is a partial override, so this
disables only the tiling preview:

.. code-block:: python

preprocessing = PreprocessingConfig(
preview={
"save_mask_preview": True,
"save_tiling_preview": True,
"downsample": 32,
"save_tiling_preview": False,
}
)

Expand Down
78 changes: 70 additions & 8 deletions slide2vec/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
SlideEmbeddingArtifact,
TileEmbeddingArtifact,
)
from slide2vec.configs.resources import load_config
from slide2vec.encoders.registry import (
encoder_registry,
resolve_preprocessing_defaults,
Expand Down Expand Up @@ -63,14 +64,42 @@ class SlideLike(Protocol):
"min_coverage": {"background": None, "tissue": 0.01},
}

_REQUESTED_TILE_SIZE_INTERPOLATION = "${tiling.params.requested_tile_size_px}"


def _load_default_preprocessing() -> dict[str, dict[str, Any]]:
"""Read the public nested defaults from the package's canonical YAML config."""
from omegaconf import OmegaConf

tiling = load_config("default").tiling
defaults: dict[str, dict[str, Any]] = {}
for public_name, config_name in (
("segmentation", "seg_params"),
("filtering", "filter_params"),
("preview", "preview"),
):
section = OmegaConf.to_container(getattr(tiling, config_name), resolve=False)
if not isinstance(section, dict):
raise TypeError(f"tiling.{config_name} must be a mapping")
defaults[public_name] = section
defaults["preview"]["tissue_contour_color"] = tuple(
defaults["preview"]["tissue_contour_color"]
)
return defaults


#: Complete defaults for the nested public preprocessing sections, loaded from
#: ``configs/default.yaml`` so Python and YAML entry points share one source.
DEFAULT_PREPROCESSING = _load_default_preprocessing()


def _deep_merge_masks(base: Mapping[str, Any], override: Mapping[str, Any]) -> dict[str, Any]:
def _deep_merge_dicts(base: Mapping[str, Any], override: Mapping[str, Any]) -> dict[str, Any]:
"""Deep-merge *override* onto a copy of *base* (nested dicts merge key-by-key)."""
merged = copy.deepcopy(dict(base))
for key, value in override.items():
existing = merged.get(key)
if isinstance(value, Mapping) and isinstance(existing, dict):
merged[key] = _deep_merge_masks(existing, value)
merged[key] = _deep_merge_dicts(existing, value)
else:
merged[key] = copy.deepcopy(value)
return merged
Expand All @@ -80,7 +109,7 @@ def resolve_masks(masks: Mapping[str, Any] | None) -> dict[str, Any]:
"""Complete a (possibly partial) ``masks`` mapping by merging it over :data:`DEFAULT_MASKS`."""
if not masks:
return copy.deepcopy(DEFAULT_MASKS)
return _deep_merge_masks(DEFAULT_MASKS, masks)
return _deep_merge_dicts(DEFAULT_MASKS, masks)


def _masks_to_plain_dict(node: Any) -> dict[str, Any]:
Expand Down Expand Up @@ -146,13 +175,16 @@ class PreprocessingConfig:
num_cucim_workers: int = 4
#: Skip slides already present in the output directory when ``True``.
resume: bool = False
#: Forwarded to hs2p segmentation config. Supported keys: ``method``,
#: ``downsample``, ``sam2_device``. See :doc:`preprocessing` for details.
#: Partial override forwarded to hs2p segmentation config. Supported keys:
#: ``method``, ``downsample``, ``sam2_device``. Omitted keys retain the
#: standard configuration defaults. See :doc:`preprocessing` for details.
segmentation: dict[str, Any] = field(default_factory=dict)
#: Forwarded to hs2p tile-filtering config.
#: Partial override forwarded to hs2p tile-filtering config. Omitted keys
#: retain the standard configuration defaults.
filtering: dict[str, Any] = field(default_factory=dict)
#: Controls whether hs2p writes mask and tiling preview images.
#: Keys: ``save_mask_preview``, ``save_tiling_preview``, ``downsample``.
#: Partial override controlling whether hs2p writes mask and tiling preview
#: images. Keys: ``save_mask_preview``, ``save_tiling_preview``,
#: ``downsample``. Omitted keys retain the standard configuration defaults.
preview: dict[str, Any] = field(default_factory=dict)
#: Annotation-mask vocabulary forwarded to hs2p's sampling resolver. Keys:
#: ``output_mode``, ``pixel_mapping``, ``colors``, ``min_coverage``. A partial
Expand All @@ -166,6 +198,36 @@ class PreprocessingConfig:
independent_sampling: bool = True

def __post_init__(self) -> None:
filtering_defaults = DEFAULT_PREPROCESSING["filtering"]
filtering_override = self.filtering
if self.requested_tile_size_px is not None:
filtering_defaults = {
**filtering_defaults,
"ref_tile_size": int(self.requested_tile_size_px),
}
if (
filtering_override.get("ref_tile_size")
== _REQUESTED_TILE_SIZE_INTERPOLATION
):
filtering_override = {
**filtering_override,
"ref_tile_size": int(self.requested_tile_size_px),
}
object.__setattr__(
self,
"segmentation",
_deep_merge_dicts(DEFAULT_PREPROCESSING["segmentation"], self.segmentation),
)
object.__setattr__(
self,
"filtering",
_deep_merge_dicts(filtering_defaults, filtering_override),
)
object.__setattr__(
self,
"preview",
_deep_merge_dicts(DEFAULT_PREPROCESSING["preview"], self.preview),
)
# Complete a (possibly partial) masks mapping against the shipped default.
object.__setattr__(self, "masks", resolve_masks(self.masks))

Expand Down
2 changes: 1 addition & 1 deletion slide2vec/configs/default.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,7 @@ tiling:
sthresh_up: 255 # upper threshold value for scaling the binary mask
mthresh: 7 # median filter size (positive, odd integer)
close: 4 # additional morphological closing to apply following initial thresholding (positive integer)
method: # tissue segmentation method: "hsv", "otsu", "threshold", or "sam2"; ignored when precomputed tissue masks are provided
method: "hsv" # tissue segmentation method: "hsv", "otsu", "threshold", or "sam2"; ignored when precomputed tissue masks are provided
sam2_checkpoint_path: # optional when method="sam2"; if empty, hs2p downloads the default AtlasPatch checkpoint from Hugging Face
sam2_config_path: # optional local override for the SAM2 model config; if empty, hs2p downloads the default AtlasPatch config from Hugging Face
sam2_device: "cpu" # device for SAM2 inference, e.g. "cpu", "cuda", or "cuda:0"
Expand Down
151 changes: 151 additions & 0 deletions tests/test_regression_core.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import ast
from dataclasses import asdict
from pathlib import Path
from types import SimpleNamespace

Expand Down Expand Up @@ -55,6 +56,156 @@ def test_packaged_preprocessing_config_matches_hs2p_4_tiling_schema():
assert hasattr(cfg.tiling.preview, "tissue_contour_color")


def test_default_public_preprocessing_constructs_complete_hs2p_configs():
from slide2vec.runtime.tiling import build_hs2p_configs

preprocessing = PreprocessingConfig(
requested_spacing_um=0.5,
requested_tile_size_px=224,
)

_, segmentation, filtering, preview, *_ = build_hs2p_configs(preprocessing)

assert asdict(segmentation) == {
"method": "hsv",
"downsample": 64,
"sthresh": 8,
"sthresh_up": 255,
"mthresh": 7,
"close": 4,
"sam2_checkpoint_path": None,
"sam2_config_path": None,
"sam2_device": "cpu",
"sam2_num_workers": None,
}
assert asdict(filtering) == {
"ref_tile_size": 224,
"a_t": 4,
"a_h": 2,
"filter_white": False,
"filter_black": False,
"white_threshold": 220,
"black_threshold": 25,
"fraction_threshold": 0.9,
"filter_grayspace": False,
"grayspace_saturation_threshold": 0.05,
"grayspace_fraction_threshold": 0.6,
"filter_blur": False,
"blur_threshold": 50.0,
"qc_spacing_um": 2.0,
}
assert asdict(preview) == {
"save_mask_preview": True,
"save_tiling_preview": True,
"downsample": 32,
"tissue_contour_color": (157, 219, 129),
"mask_overlay_alpha": 0.5,
}


def test_partial_segmentation_override_changes_only_requested_field():
from slide2vec.runtime.tiling import build_hs2p_configs

preprocessing = PreprocessingConfig(
requested_spacing_um=0.5,
requested_tile_size_px=224,
segmentation={"downsample": 32},
)

segmentation = build_hs2p_configs(preprocessing)[1]

assert asdict(segmentation) == {
"method": "hsv",
"downsample": 32,
"sthresh": 8,
"sthresh_up": 255,
"mthresh": 7,
"close": 4,
"sam2_checkpoint_path": None,
"sam2_config_path": None,
"sam2_device": "cpu",
"sam2_num_workers": None,
}


def test_partial_filtering_override_changes_only_requested_field():
from slide2vec.runtime.tiling import build_hs2p_configs

preprocessing = PreprocessingConfig(
requested_spacing_um=0.5,
requested_tile_size_px=224,
filtering={"a_t": 7},
)

filtering = build_hs2p_configs(preprocessing)[2]

assert asdict(filtering) == {
"ref_tile_size": 224,
"a_t": 7,
"a_h": 2,
"filter_white": False,
"filter_black": False,
"white_threshold": 220,
"black_threshold": 25,
"fraction_threshold": 0.9,
"filter_grayspace": False,
"grayspace_saturation_threshold": 0.05,
"grayspace_fraction_threshold": 0.6,
"filter_blur": False,
"blur_threshold": 50.0,
"qc_spacing_um": 2.0,
}


def test_partial_preview_override_changes_only_requested_field():
from slide2vec.runtime.tiling import build_hs2p_configs

preprocessing = PreprocessingConfig(
requested_spacing_um=0.5,
requested_tile_size_px=224,
preview={"save_tiling_preview": False},
)

preview = build_hs2p_configs(preprocessing)[3]

assert asdict(preview) == {
"save_mask_preview": True,
"save_tiling_preview": False,
"downsample": 32,
"tissue_contour_color": (157, 219, 129),
"mask_overlay_alpha": 0.5,
}


def test_public_preprocessing_defaults_match_standard_configuration():
cfg = load_config("default")
cfg.tiling.params.requested_spacing_um = 0.5
cfg.tiling.params.requested_tile_size_px = 224

standard = PreprocessingConfig.from_config(cfg)
public = PreprocessingConfig(
requested_spacing_um=0.5,
requested_tile_size_px=224,
)

assert public.segmentation == standard.segmentation
assert public.filtering == standard.filtering
assert public.preview == standard.preview


def test_public_filtering_default_tracks_requested_tile_size():
from slide2vec.runtime.tiling import build_hs2p_configs

preprocessing = PreprocessingConfig(
requested_spacing_um=0.5,
requested_tile_size_px=448,
)

filtering = build_hs2p_configs(preprocessing)[2]

assert filtering.ref_tile_size == 448


def test_get_cfg_from_args_fills_missing_preprocessing_from_single_spacing_model(tmp_path: Path):
pytest.importorskip("omegaconf")

Expand Down
11 changes: 0 additions & 11 deletions tests/test_regression_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,17 +55,6 @@
def PreprocessingConfig(*args, **kwargs):
kwargs.setdefault("requested_spacing_um", 0.5)
kwargs.setdefault("requested_tile_size_px", 224)
kwargs.setdefault("segmentation", {"method": "hsv"})
kwargs.setdefault(
"preview",
{
"save_mask_preview": True,
"save_tiling_preview": True,
"downsample": 32,
"tissue_contour_color": (157, 219, 129),
"mask_overlay_alpha": 0.5,
},
)
return BasePreprocessingConfig(*args, **kwargs)


Expand Down
26 changes: 26 additions & 0 deletions tests/test_regression_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,32 @@ def fake_embed_slides(model_arg, slides, **kwargs):
assert captured["slides"][0]["sample_id"] == "slide-a"


def test_model_embed_slide_without_preprocessing_reaches_tiling(monkeypatch):
from slide2vec.runtime import tiling_pipeline
from slide2vec.runtime.tiling import build_hs2p_configs

class TilingReached(Exception):
pass

captured = {}

def stop_at_tiling(slides, preprocessing, **kwargs):
filtering = build_hs2p_configs(preprocessing)[2]
captured["preprocessing"] = preprocessing
captured["filtering"] = filtering
raise TilingReached

monkeypatch.setattr(tiling_pipeline, "tile_slides_call", stop_at_tiling)

model = Model.from_preset("conch")
with pytest.raises(TilingReached):
model.embed_slide("/tmp/slide-a.svs")

assert captured["preprocessing"].requested_spacing_um == pytest.approx(0.5)
assert captured["preprocessing"].requested_tile_size_px == 448
assert captured["filtering"].ref_tile_size == 448


def test_embed_slides_returns_nested_dict_keyed_by_sample_and_label(monkeypatch):
model = Model.from_preset("virchow2")
slide_a = _make_embedded_slide(sample_id="slide-a", annotation="tissue")
Expand Down
Loading