Spaces:
Running
Running
| from pathlib import Path | |
| import cv2 | |
| import numpy as np | |
| from skimage.util import img_as_ubyte | |
| TIFF_EXTENSIONS = {".tif", ".tiff"} | |
| DEBRIS_VALUE_MAX = 60 | |
| DEBRIS_SATURATION_MAX = 15 | |
| DEBRIS_MASK_EXPANSION = 10 | |
| INPAINT_RADIUS = 5 | |
| PREVIEW_PERCENTILES = (1.0, 99.5) | |
| VIGNETTING_SATURATION_MAX = 30 | |
| VIGNETTING_VALUE_MIN = 80 | |
| VIGNETTING_MIN_FALLOFF = 0.03 | |
| VIGNETTING_MAX_GAIN = 1.60 | |
| VIGNETTING_FIT_MAX_DIMENSION = 800 | |
| VIGNETTING_CORRECTION_ROWS = 256 | |
| VIGNETTING_TILE_SIZE = 16 | |
| VIGNETTING_TILE_PERCENTILE = 50 | |
| VIGNETTING_SURFACE_BLUR_SIGMA = 0.8 | |
| def _as_rgb_array(image: np.ndarray) -> np.ndarray: | |
| image = np.asarray(image) | |
| if image.ndim == 3 and image.shape[2] == 1: | |
| image = image[:, :, 0] | |
| if image.ndim == 2: | |
| image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB) | |
| elif image.ndim != 3: | |
| raise ValueError("Expected a grayscale or RGB image.") | |
| if image.shape[2] == 4: | |
| image = cv2.cvtColor(image, cv2.COLOR_RGBA2RGB) | |
| elif image.shape[2] != 3: | |
| raise ValueError("Expected an image with 1, 3, or 4 channels.") | |
| return image | |
| def _as_rgb_uint8(image: np.ndarray) -> np.ndarray: | |
| source_dtype = np.asarray(image).dtype | |
| image = _as_rgb_array(image) | |
| if image.dtype != np.uint8: | |
| image = image.astype(np.float32) | |
| if image.size: | |
| finite_values = image[np.isfinite(image)] | |
| if not finite_values.size: | |
| raise ValueError("Image contains no finite pixel values.") | |
| min_value = float(np.min(finite_values)) | |
| max_value = float(np.max(finite_values)) | |
| if 0 <= min_value and max_value <= 1: | |
| image *= 255 | |
| elif ( | |
| np.issubdtype(source_dtype, np.integer) | |
| and np.iinfo(source_dtype).max > 255 | |
| ) or max_value > 255: | |
| low, high = np.nanpercentile(image, PREVIEW_PERCENTILES) | |
| if high > low: | |
| image = (image - low) * (255 / (high - low)) | |
| else: | |
| image = np.zeros_like(image) | |
| image = np.nan_to_num(image, nan=0, posinf=255, neginf=0) | |
| image = np.clip(image, 0, 255).astype(np.uint8) | |
| return image | |
| def _as_display_rgb_uint8(image: np.ndarray) -> np.ndarray: | |
| """Convert an RGB image for display without changing its encoded contrast.""" | |
| image = _as_rgb_array(image) | |
| if image.dtype == np.uint8: | |
| return image | |
| if np.issubdtype(image.dtype, np.integer) or image.dtype == np.bool_: | |
| return img_as_ubyte(image) | |
| finite_values = image[np.isfinite(image)] | |
| if not finite_values.size: | |
| raise ValueError("Image contains no finite pixel values.") | |
| if 0 <= np.min(finite_values) and np.max(finite_values) <= 1: | |
| finite_image = np.nan_to_num(image, nan=0, posinf=1, neginf=0) | |
| return img_as_ubyte(finite_image) | |
| return _as_rgb_uint8(image) | |
| def _is_tiff_path(image_path: str | Path) -> bool: | |
| return Path(image_path).suffix.lower() in TIFF_EXTENSIONS | |
| def cleaned_image_name( | |
| image_path: str | Path, | |
| used_names: set[str] | None = None, | |
| ) -> str: | |
| image_path = Path(image_path) | |
| extension = image_path.suffix or ".png" | |
| output_name = f"{image_path.stem}_cleaned{extension}" | |
| if used_names is None: | |
| return output_name | |
| suffix = 2 | |
| while output_name in used_names: | |
| output_name = f"{image_path.stem}_cleaned_{suffix}{extension}" | |
| suffix += 1 | |
| used_names.add(output_name) | |
| return output_name | |
| def _select_tiff_plane(image: np.ndarray) -> np.ndarray: | |
| image = np.asarray(image) | |
| while image.ndim > 2 and 1 in image.shape: | |
| image = np.squeeze(image) | |
| if image.ndim == 2: | |
| return image | |
| if image.ndim == 3: | |
| if image.shape[-1] in {1, 3, 4}: | |
| return image | |
| if image.shape[0] in {1, 3, 4}: | |
| return np.moveaxis(image, 0, -1) | |
| raise ValueError("Expected a 2D grayscale or RGB TIFF image.") | |
| def _read_image_rgb_values(image_path: str | Path) -> np.ndarray: | |
| image_path = Path(image_path) | |
| if _is_tiff_path(image_path): | |
| try: | |
| import tifffile | |
| except ImportError as exc: | |
| raise ImportError( | |
| "Reading TIFF images requires the tifffile package." | |
| ) from exc | |
| image = tifffile.imread(image_path) | |
| return _select_tiff_plane(image) | |
| image = cv2.imread(str(image_path), cv2.IMREAD_UNCHANGED) | |
| if image is None: | |
| raise FileNotFoundError(f"Could not read image: {image_path}") | |
| if image.ndim == 3 and image.shape[2] == 3: | |
| image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) | |
| elif image.ndim == 3 and image.shape[2] == 4: | |
| image = cv2.cvtColor(image, cv2.COLOR_BGRA2RGBA) | |
| return image | |
| def read_image_rgb(image_path: str | Path) -> np.ndarray: | |
| return _as_rgb_uint8(_read_image_rgb_values(image_path)) | |
| def read_image_preview_rgb(image_path: str | Path) -> np.ndarray: | |
| return _as_display_rgb_uint8(_read_image_rgb_values(image_path)) | |
| def read_image_rgb_and_preview( | |
| image_path: str | Path, | |
| ) -> tuple[np.ndarray, np.ndarray]: | |
| """Return the processing image and its independent display preview.""" | |
| image = _read_image_rgb_values(image_path) | |
| return _as_rgb_uint8(image), _as_display_rgb_uint8(image) | |
| def write_image_rgb(image_path: str | Path, image: np.ndarray) -> None: | |
| image_path = Path(image_path) | |
| image = _as_rgb_uint8(image) | |
| if _is_tiff_path(image_path): | |
| try: | |
| import tifffile | |
| except ImportError as exc: | |
| raise ImportError( | |
| "Writing TIFF images requires the tifffile package." | |
| ) from exc | |
| tifffile.imwrite( | |
| image_path, | |
| image, | |
| photometric="rgb", | |
| compression="lzw", | |
| predictor=True, | |
| ) | |
| return | |
| image_bgr = cv2.cvtColor(image, cv2.COLOR_RGB2BGR) | |
| if not cv2.imwrite(str(image_path), image_bgr): | |
| raise OSError(f"Could not write image: {image_path}") | |
| def create_debris_mask(image: np.ndarray) -> np.ndarray: | |
| image = _as_rgb_uint8(image) | |
| hsv = cv2.cvtColor(image, cv2.COLOR_RGB2HSV) | |
| _, saturation, value = cv2.split(hsv) | |
| debris_mask = ( | |
| (value < DEBRIS_VALUE_MAX) | |
| & (saturation < DEBRIS_SATURATION_MAX) | |
| ).astype(np.uint8) * 255 | |
| cleanup_kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3)) | |
| debris_mask = cv2.morphologyEx( | |
| debris_mask, | |
| cv2.MORPH_OPEN, | |
| cleanup_kernel, | |
| ) | |
| # Inpainting needs a solid mask that extends beyond the dark boundary. | |
| # Otherwise, unmasked holes and edge pixels are used as source pixels and | |
| # the reconstructed region remains dark. | |
| contours, _ = cv2.findContours( | |
| debris_mask, | |
| cv2.RETR_EXTERNAL, | |
| cv2.CHAIN_APPROX_SIMPLE, | |
| ) | |
| filled_mask = np.zeros_like(debris_mask) | |
| cv2.drawContours(filled_mask, contours, -1, 255, cv2.FILLED) | |
| expansion_size = 2 * DEBRIS_MASK_EXPANSION + 1 | |
| expansion_kernel = cv2.getStructuringElement( | |
| cv2.MORPH_ELLIPSE, | |
| (expansion_size, expansion_size), | |
| ) | |
| return cv2.dilate(filled_mask, expansion_kernel) | |
| def _estimate_vignetting_surface( | |
| image: np.ndarray, | |
| background_mask: np.ndarray, | |
| ) -> np.ndarray | None: | |
| """Estimate local slide-background colour on a small spatial grid.""" | |
| if np.count_nonzero(background_mask) < 1_000: | |
| return None | |
| height, width = background_mask.shape | |
| grid_height = max(2, int(np.ceil(height / VIGNETTING_TILE_SIZE))) | |
| grid_width = max(2, int(np.ceil(width / VIGNETTING_TILE_SIZE))) | |
| surface = np.full((grid_height, grid_width, 3), np.nan, dtype=np.float32) | |
| for grid_y in range(grid_height): | |
| y_start = grid_y * height // grid_height | |
| y_stop = (grid_y + 1) * height // grid_height | |
| for grid_x in range(grid_width): | |
| x_start = grid_x * width // grid_width | |
| x_stop = (grid_x + 1) * width // grid_width | |
| tile_mask = background_mask[y_start:y_stop, x_start:x_stop] | |
| if np.count_nonzero(tile_mask) < 16: | |
| continue | |
| pixels = image[y_start:y_stop, x_start:x_stop][tile_mask] | |
| surface[grid_y, grid_x] = np.percentile( | |
| pixels, | |
| VIGNETTING_TILE_PERCENTILE, | |
| axis=0, | |
| ) | |
| missing = np.isnan(surface[..., 0]) | |
| if np.all(missing): | |
| return None | |
| for channel in range(3): | |
| channel_surface = surface[..., channel] | |
| if np.any(missing): | |
| channel_surface = cv2.inpaint( | |
| np.nan_to_num(channel_surface, nan=0).astype(np.float32), | |
| missing.astype(np.uint8), | |
| 3, | |
| cv2.INPAINT_TELEA, | |
| ) | |
| surface[..., channel] = cv2.GaussianBlur( | |
| channel_surface, | |
| (0, 0), | |
| sigmaX=VIGNETTING_SURFACE_BLUR_SIGMA, | |
| sigmaY=VIGNETTING_SURFACE_BLUR_SIGMA, | |
| borderType=cv2.BORDER_REPLICATE, | |
| ) | |
| return surface / 255 | |
| def _largest_connected_region(mask: np.ndarray) -> np.ndarray | None: | |
| component_count, labels, stats, _ = cv2.connectedComponentsWithStats( | |
| mask.astype(np.uint8), | |
| connectivity=8, | |
| ) | |
| if component_count <= 1: | |
| return None | |
| largest_component = 1 + np.argmax(stats[1:, cv2.CC_STAT_AREA]) | |
| region = labels == largest_component | |
| if np.count_nonzero(region) < 1_000: | |
| return None | |
| return region | |
| def create_tissue_mask( | |
| image: np.ndarray, | |
| debris_mask: np.ndarray | None = None, | |
| ) -> np.ndarray: | |
| """Return the established spheroid segmentation as a filled binary mask.""" | |
| # Imported here to avoid a module-import cycle: the quantification module | |
| # reuses cleaning helpers before defining its segmentation function. | |
| from ihc_quantification_simplified import segment_spheroids | |
| image = _as_display_rgb_uint8(image) | |
| segmentation_image = image | |
| if debris_mask is not None: | |
| if debris_mask.shape != image.shape[:2]: | |
| raise ValueError("Debris mask must match the image height and width.") | |
| segmentation_image = _inpaint_image( | |
| segmentation_image, | |
| debris_mask, | |
| ) | |
| return ( | |
| (segment_spheroids(segmentation_image) > 0).astype(np.uint8) * 255 | |
| ) | |
| def _extend_surface_to_image_edges(surface: np.ndarray) -> np.ndarray: | |
| """Linearly extrapolate tile-centre estimates to the image boundaries.""" | |
| extended = np.empty( | |
| (surface.shape[0] + 2, surface.shape[1] + 2, 3), | |
| dtype=np.float32, | |
| ) | |
| extended[1:-1, 1:-1] = surface | |
| extended[0, 1:-1] = 2 * surface[0] - surface[1] | |
| extended[-1, 1:-1] = 2 * surface[-1] - surface[-2] | |
| extended[:, 0] = 2 * extended[:, 1] - extended[:, 2] | |
| extended[:, -1] = 2 * extended[:, -2] - extended[:, -3] | |
| return np.clip(extended, 0.05, 1.5) | |
| def correct_vignetting( | |
| image: np.ndarray, | |
| debris_mask: np.ndarray | None = None, | |
| tissue_mask: np.ndarray | None = None, | |
| ) -> tuple[np.ndarray, bool]: | |
| """Correct smooth edge falloff when a robust background fit detects it.""" | |
| image = _as_display_rgb_uint8(image) | |
| height, width = image.shape[:2] | |
| scale = min(1, VIGNETTING_FIT_MAX_DIMENSION / max(height, width)) | |
| if scale < 1: | |
| fit_size = (round(width * scale), round(height * scale)) | |
| fit_image = cv2.resize(image, fit_size, interpolation=cv2.INTER_AREA) | |
| else: | |
| fit_image = image | |
| hsv = cv2.cvtColor(fit_image, cv2.COLOR_RGB2HSV) | |
| background_mask = ( | |
| (hsv[..., 1] <= VIGNETTING_SATURATION_MAX) | |
| & (hsv[..., 2] >= VIGNETTING_VALUE_MIN) | |
| ) | |
| if debris_mask is not None: | |
| if debris_mask.shape != image.shape[:2]: | |
| raise ValueError("Debris mask must match the image height and width.") | |
| fit_mask = debris_mask | |
| if scale < 1: | |
| fit_mask = cv2.resize( | |
| debris_mask, | |
| fit_size, | |
| interpolation=cv2.INTER_NEAREST, | |
| ) | |
| background_mask &= fit_mask == 0 | |
| background_region = _largest_connected_region(background_mask) | |
| if background_region is None: | |
| return image, False | |
| surface = _estimate_vignetting_surface( | |
| fit_image, | |
| background_region, | |
| ) | |
| if surface is None: | |
| return image, False | |
| grid_y, grid_x = np.meshgrid( | |
| np.linspace(-1, 1, surface.shape[0]), | |
| np.linspace(-1, 1, surface.shape[1]), | |
| indexing="ij", | |
| ) | |
| grid_luminance = surface @ np.array([0.2126, 0.7152, 0.0722]) | |
| grid_radius = np.maximum(np.abs(grid_x), np.abs(grid_y)) | |
| center_region = grid_radius <= 0.35 | |
| edge_region = grid_radius >= 0.8 | |
| center_level = float(np.median(grid_luminance[center_region])) | |
| edge_level = float(np.percentile(grid_luminance[edge_region], 20)) | |
| if center_level <= np.finfo(np.float64).eps: | |
| return image, False | |
| falloff = 1 - edge_level / center_level | |
| if falloff < VIGNETTING_MIN_FALLOFF: | |
| return image, False | |
| references = np.percentile( | |
| fit_image[background_region].astype(np.float32) / 255, | |
| 90, | |
| axis=0, | |
| ) | |
| extended_surface = _extend_surface_to_image_edges(surface) | |
| corrected = image.copy() | |
| correction_region = cv2.resize( | |
| background_region.astype(np.uint8), | |
| (width, height), | |
| interpolation=cv2.INTER_NEAREST, | |
| ).astype(bool) | |
| full_hsv = cv2.cvtColor(image, cv2.COLOR_RGB2HSV) | |
| correction_region &= ( | |
| (full_hsv[..., 1] <= VIGNETTING_SATURATION_MAX) | |
| & (full_hsv[..., 2] >= VIGNETTING_VALUE_MIN) | |
| ) | |
| if tissue_mask is None: | |
| tissue_mask = create_tissue_mask(image, debris_mask) | |
| elif tissue_mask.shape != image.shape[:2]: | |
| raise ValueError("Tissue mask must match the image height and width.") | |
| correction_region &= tissue_mask == 0 | |
| if debris_mask is not None: | |
| correction_region &= debris_mask == 0 | |
| surface_y = np.clip( | |
| (np.arange(height, dtype=np.float32) + 0.5) | |
| * surface.shape[0] | |
| / height | |
| + 0.5, | |
| 0, | |
| extended_surface.shape[0] - 1, | |
| ) | |
| surface_x = np.clip( | |
| (np.arange(width, dtype=np.float32) + 0.5) | |
| * surface.shape[1] | |
| / width | |
| + 0.5, | |
| 0, | |
| extended_surface.shape[1] - 1, | |
| ) | |
| source_x = np.arange(extended_surface.shape[1], dtype=np.float32) | |
| for channel in range(3): | |
| horizontal_surface = np.vstack( | |
| [ | |
| np.interp(surface_x, source_x, row) | |
| for row in extended_surface[..., channel] | |
| ] | |
| ).astype(np.float32) | |
| for row_start in range(0, height, VIGNETTING_CORRECTION_ROWS): | |
| row_stop = min(row_start + VIGNETTING_CORRECTION_ROWS, height) | |
| y = surface_y[row_start:row_stop] | |
| y_low = np.floor(y).astype(np.int32) | |
| y_high = np.minimum(y_low + 1, extended_surface.shape[0] - 1) | |
| y_fraction = (y - y_low)[:, None] | |
| strip_surface = ( | |
| horizontal_surface[y_low] * (1 - y_fraction) | |
| + horizontal_surface[y_high] * y_fraction | |
| ) | |
| strip_surface = np.clip(strip_surface, 0.05, 1.5) | |
| gain = np.clip( | |
| references[channel] / strip_surface, | |
| 1, | |
| VIGNETTING_MAX_GAIN, | |
| ) | |
| corrected_channel = corrected[ | |
| row_start:row_stop, | |
| :, | |
| channel, | |
| ].astype(np.float32) | |
| corrected_values = np.clip( | |
| np.rint(corrected_channel * gain), | |
| 0, | |
| 255, | |
| ).astype(np.uint8) | |
| strip_region = correction_region[row_start:row_stop] | |
| corrected_channel = corrected[ | |
| row_start:row_stop, | |
| :, | |
| channel, | |
| ] | |
| corrected_channel[strip_region] = corrected_values[strip_region] | |
| return corrected, True | |
| def clean_image(image: np.ndarray | None) -> np.ndarray: | |
| cleaned_image, _ = clean_image_and_mask(image) | |
| return cleaned_image | |
| def _inpaint_image(image: np.ndarray, debris_mask: np.ndarray) -> np.ndarray: | |
| return cv2.inpaint( | |
| image, | |
| debris_mask, | |
| INPAINT_RADIUS, | |
| cv2.INPAINT_TELEA, | |
| ) | |
| def clean_image_and_mask( | |
| image: np.ndarray | None, | |
| ) -> tuple[np.ndarray, np.ndarray]: | |
| if image is None: | |
| raise ValueError("An input image is required.") | |
| image = _as_rgb_uint8(image) | |
| debris_mask = create_debris_mask(image) | |
| cleaned_image = _inpaint_image(image, debris_mask) | |
| return cleaned_image, debris_mask | |
| def clean_display_image_and_mask( | |
| processing_image: np.ndarray | None, | |
| display_image: np.ndarray | None, | |
| ) -> tuple[np.ndarray, np.ndarray]: | |
| if processing_image is None or display_image is None: | |
| raise ValueError("Processing and display images are required.") | |
| processing_image = _as_rgb_uint8(processing_image) | |
| display_image = _as_display_rgb_uint8(display_image) | |
| if processing_image.shape != display_image.shape: | |
| raise ValueError("Processing and display images must have the same shape.") | |
| debris_mask = create_debris_mask(processing_image) | |
| tissue_mask = create_tissue_mask(processing_image, debris_mask) | |
| corrected_display, _ = correct_vignetting( | |
| display_image, | |
| debris_mask, | |
| tissue_mask, | |
| ) | |
| return _inpaint_image(corrected_display, debris_mask), debris_mask | |