Stitching¶
Segment tile by tile, end with one consistent label image.
Segmenting a large image in tiles creates two problems: every tile numbers its objects
from 1, and an object crossing a tile boundary comes out as two objects with two ids.
This tutorial runs into both on purpose, then declares with_stitch() and lets the
iterator solve them. Concepts and vocabulary: the
iterators guide.
Step 1: open the OME-Zarr container¶
from pathlib import Path
from ngio import open_ome_zarr_container
from ngio.utils import download_ome_zarr_dataset
# Download the dataset
download_dir = Path("./data").absolute()
hcs_path = download_ome_zarr_dataset("CardiomyocyteTinyMip", download_dir=download_dir)
# Open the OME-Zarr container
ome_zarr = open_ome_zarr_container(hcs_path / "B" / "03" / "0")
image = ome_zarr.get_image()
print(image)
Step 2: write a segmentation function¶
A classic watershed pipeline — the same segmentation function the ngio workshop uses on this dataset. The function knows nothing about tiles or ids — keeping ids consistent across tiles is the iterator's job, not the function's.
import numpy as np
from scipy import ndimage as ndi
from skimage.feature import peak_local_max
from skimage.filters import threshold_otsu
from skimage.morphology import remove_small_objects
from skimage.segmentation import watershed
def segment(patch: np.ndarray) -> np.ndarray:
# Smooth → Otsu threshold → distance transform → seeded watershed → cleanup
smooth = ndi.gaussian_filter(patch, sigma=4)
mask = smooth > threshold_otsu(smooth)
distance = ndi.distance_transform_edt(mask)
coords = peak_local_max(distance, min_distance=20, labels=mask)
markers = np.zeros(distance.shape, dtype=np.int32)
markers[tuple(coords.T)] = np.arange(1, len(coords) + 1)
seg = watershed(-distance, markers, mask=mask).astype(np.uint32)
# `max_size` drops objects at or below 500 px — skimage's renamed `min_size`.
return remove_small_objects(seg, max_size=500)
Step 3: segment the tiles independently¶
Tile the image with by_grid and segment each tile on its own. Keeping the ids
distinct takes a per-tile offset (UniqueLabelsTransform) you have to wire up
yourself — and even then, every nucleus that crosses a tile boundary is counted twice:
from ngio.iterators import SegmentationIterator
from ngio.transforms import UniqueLabelsTransform
naive = ome_zarr.derive_label("nuclei_tiled", overwrite=True)
tiling = SegmentationIterator(
image, naive, channel_selection="DAPI", axes_order=["y", "x"]
).by_grid(size_x=512, size_y=512)
# Each tile numbers its objects from 1, so keeping them distinct in one array
# takes a per-tile offset (block i holds tile i's ids).
for i, roi in enumerate(tiling.rois):
patch = image.get_roi_as_numpy(roi, c=0, axes_order=["y", "x"])
naive.set_roi(
roi,
segment(patch),
axes_order=["y", "x"],
transforms=[UniqueLabelsTransform(10_000, i)],
)
naive.consolidate(mode="auto")
print(f"tiles: {len(tiling.rois)}, ids: {len(np.unique(naive.get_as_numpy())) - 1}")
Step 4: declare with_stitch()¶
with_stitch() replaces all of that bookkeeping. Each tile reads a halo past its edge,
so neighbouring tiles predict the same pixels around a seam; where their objects agree,
the ids are joined, and the survivors are renumbered to a dense 1..N:
stitched = ome_zarr.derive_label("nuclei_stitched", overwrite=True)
stitch_iterator = (
SegmentationIterator(
image,
stitched,
channel_selection="DAPI",
axes_order=["y", "x"],
consolidation_mode="auto",
)
.with_stitch()
.by_grid(size_x=512, size_y=512)
.with_halo(x=32, y=32)
)
stitch_iterator.map(segment)
print(f"objects: {len(np.unique(stitched.get_as_numpy())) - 1}")
The object count drops — every nucleus split at a seam is now counted once.
Plot the results¶
cmap = random_label_cmap(n_labels=2000)
stitched_data = stitched.get_as_numpy(axes_order=["y", "x"])
image_data = image.get_as_numpy(c=0, axes_order=["y", "x"])
fig, (ax_img, ax_lab) = plt.subplots(2, 1, figsize=(8, 6.6))
show_image(ax_img, image_data, title="DAPI", pixel_size=image.pixel_size)
show_image(ax_lab, stitched_data, title="stitched segmentation", cmap=cmap)
for edge in range(512, image_data.shape[1], 512):
ax_lab.axvline(edge - 0.5, color="white", ls="--", lw=0.5)
for edge in range(512, image_data.shape[0], 512):
ax_lab.axhline(edge - 0.5, color="white", ls="--", lw=0.5)
print(
figure_html(
fig,
alt="The DAPI image and its stitched segmentation, with the tile grid "
"overlaid; no object shows a seam at a tile boundary.",
)
)
Zooming onto a corner where four tiles meet shows what changed — the nucleus sitting on the corner was split into one fragment per tile:
# Found by eye; any seam tells the same story.
crop_y, crop_x = slice(1920, 2160), slice(1408, 1664)
seam_y, seam_x = 2048 - crop_y.start, 1536 - crop_x.start
def dense_ids(array: np.ndarray) -> np.ndarray:
# Display-side renumbering only: the block-offset ids (1, 10_001, 20_001)
# would otherwise collapse onto a handful of colormap entries.
return np.unique(array, return_inverse=True)[1].reshape(array.shape)
naive_data = naive.get_as_numpy(axes_order=["y", "x"])
fig, (ax_raw, ax_naive, ax_merged) = plt.subplots(1, 3, figsize=(8.6, 3.1))
show_image(ax_raw, image_data[crop_y, crop_x], title="DAPI")
show_image(
ax_naive,
dense_ids(naive_data[crop_y, crop_x]),
title="tiles segmented independently",
cmap=cmap,
)
show_image(
ax_merged,
dense_ids(stitched_data[crop_y, crop_x]),
title="with_stitch()",
cmap=cmap,
)
for ax in (ax_naive, ax_merged):
ax.axvline(seam_x - 0.5, color="white", ls="--", lw=1.0)
ax.axhline(seam_y - 0.5, color="white", ls="--", lw=1.0)
print(
figure_html(
fig,
alt="A zoom onto a corner where four tiles meet: segmented "
"independently, the nucleus on the corner is split into one fragment "
"per tile; with stitching it is one object.",
)
)
Step 5: tune it¶
with_stitch() takes a StitchConfig:
from ngio.iterators import StitchConfig
iterator = SegmentationIterator(image, label).with_stitch(
StitchConfig(iou_threshold=0.5, block_size=50_000)
)
iou_threshold is how much two tiles must agree before their ids are joined. The
default errs towards leaving an object split rather than merging two that are not — an
over-split label can be fixed downstream, a wrong merge cannot. block_size is how
many ids each tile is given, and must exceed the largest count a single tile can
produce.
By default the scratch band arrays live in a transient group inside the output label,
which works under every mapper. scratch_store puts them elsewhere — often worth
doing, since labels compress well:
That keeps the output store untouched and leaves nothing behind if a run dies. The one
restriction is ProcessMapper: a MemoryStore pickles by value, so each worker would
bank into a private copy — ngio refuses that rather than losing the predictions
silently.
If a run is interrupted between the map and the resolve, the label holds a valid but
over-split segmentation, and re-running the resolve is safe. An interruption inside
the resolve is different: the default compact=True renumbers the label in place, so
a kill mid-walk leaves mixed ids — ngio marks the walk before it starts and refuses a
retry loudly rather than silently splitting objects; re-run the map to regenerate the
label. Compaction is also available on its own — label.relabel_sequential()
renumbers any label to a dense 1..N.
Next steps¶
- Iterators guide — what counts as evidence for a merge, and how stitching relates to halos and overlap.
- Distributed processing — the same stitched run split across cluster jobs.
- Image segmentation — stitching over a FOV table, and masked segmentation.