"""Configurable model-object clustering.""" from __future__ import annotations from dataclasses import dataclass from typing import Any import numpy as np from scipy.ndimage import binary_dilation, label from .utils import circular_footprint, dilate_fast, km_to_pixels @dataclass class ModelCluster: cluster_id: int mask: np.ndarray raw_mask: np.ndarray pixel_count: int raw_pixel_count: int def to_dict(self) -> dict[str, int]: return { "cluster_id": int(self.cluster_id), "pixel_count": int(self.pixel_count), "raw_pixel_count": int(self.raw_pixel_count), } class Clusterer: def __init__( self, name: str, threshold: float, min_cluster_pixels: int = 1, pixel_size_km: float = 2.0, merge_buffer_km: float = 0.0, cluster_mask_expansion_km: float = 0.0, connectivity: int = 8, buffer_backend: str = "auto", ): self.name = name self.threshold = float(threshold) self.min_cluster_pixels = int(min_cluster_pixels) self.pixel_size_km = float(pixel_size_km) self.merge_radius = km_to_pixels(merge_buffer_km, self.pixel_size_km) self.expansion_radius = km_to_pixels(cluster_mask_expansion_km, self.pixel_size_km) self.connectivity = int(connectivity) self.buffer_backend = str(buffer_backend or "auto") self._merge_footprint = circular_footprint(self.merge_radius) if self.merge_radius > 0 else None self._expansion_footprint = ( circular_footprint(self.expansion_radius) if self.expansion_radius > 0 else None ) @property def structure(self) -> np.ndarray: if self.connectivity == 4: return np.array([[0, 1, 0], [1, 1, 1], [0, 1, 0]], dtype=bool) return np.ones((3, 3), dtype=bool) def cluster(self, data: np.ndarray, valid_mask: np.ndarray | None = None) -> list[ModelCluster]: data = np.asarray(data) if valid_mask is None: valid_mask = np.isfinite(data) positive_mask = (data >= self.threshold) & valid_mask & np.isfinite(data) if self.name == "connected": return self._connected(positive_mask) if self.name == "distance_merge": return self._distance_merge(positive_mask) if self.name == "one_hop_merge": return self._one_hop_merge(positive_mask) if self.name == "complete_link_merge": return self._complete_link_merge(positive_mask) if self.name == "mode_like": return self._mode_like(positive_mask) raise ValueError(f"Unsupported clusterer: {self.name}") def _connected_components(self, mask: np.ndarray) -> list[np.ndarray]: labeled, n_features = label(mask, structure=self.structure) components: list[np.ndarray] = [] for cid in range(1, n_features + 1): component = labeled == cid if int(component.sum()) >= self.min_cluster_pixels: components.append(component) return components def _dilate(self, mask: np.ndarray, radius_pixels: int, footprint: np.ndarray | None) -> np.ndarray: if radius_pixels <= 0: return mask.astype(bool, copy=True) backend = "binary" if self.buffer_backend == "binary" and footprint is not None else self.buffer_backend return dilate_fast(mask, radius_pixels, backend=backend) def _make_cluster( self, cluster_id: int, raw_mask: np.ndarray, support_mask: np.ndarray | None = None, final_mask: np.ndarray | None = None, ) -> ModelCluster | None: if int(raw_mask.sum()) < self.min_cluster_pixels: return None if final_mask is None: final_mask = raw_mask.astype(bool, copy=True) if self.expansion_radius > 0: final_mask = self._dilate(final_mask, self.expansion_radius, self._expansion_footprint) if support_mask is not None: final_mask &= support_mask else: final_mask = final_mask.astype(bool, copy=True) if support_mask is not None: final_mask &= support_mask if int(final_mask.sum()) == 0: return None return ModelCluster( cluster_id=cluster_id, mask=final_mask, raw_mask=raw_mask.astype(bool, copy=True), pixel_count=int(final_mask.sum()), raw_pixel_count=int(raw_mask.sum()), ) def _connected(self, positive_mask: np.ndarray) -> list[ModelCluster]: clusters: list[ModelCluster] = [] for raw_mask in self._connected_components(positive_mask): cluster = self._make_cluster(len(clusters) + 1, raw_mask, support_mask=None) if cluster is not None: clusters.append(cluster) return clusters def _distance_merge(self, positive_mask: np.ndarray) -> list[ModelCluster]: components = self._connected_components(positive_mask) if not components: return [] if len(components) == 1: cluster = self._make_cluster(1, components[0], support_mask=None) return [] if cluster is None else [cluster] parent = list(range(len(components))) def find(x: int) -> int: while parent[x] != x: parent[x] = parent[parent[x]] x = parent[x] return x def union(a: int, b: int) -> None: ra, rb = find(a), find(b) if ra != rb: parent[rb] = ra for i, component in enumerate(components): buffered = dilate_fast(component, self.merge_radius, backend=self.buffer_backend) for j in range(i + 1, len(components)): if np.any(buffered & components[j]): union(i, j) grouped: dict[int, list[int]] = {} for idx in range(len(components)): grouped.setdefault(find(idx), []).append(idx) clusters: list[ModelCluster] = [] for member_indices in grouped.values(): raw_union = np.zeros_like(positive_mask, dtype=bool) support = np.zeros_like(positive_mask, dtype=bool) for idx in member_indices: raw_union |= components[idx] support |= dilate_fast(components[idx], self.merge_radius, backend=self.buffer_backend) cluster = self._make_cluster(len(clusters) + 1, raw_union, support_mask=support) if cluster is not None: clusters.append(cluster) return clusters def _one_hop_merge(self, positive_mask: np.ndarray) -> list[ModelCluster]: components = self._connected_components(positive_mask) if not components: return [] if len(components) == 1 or self.merge_radius <= 0: clusters: list[ModelCluster] = [] for component in components: cluster = self._make_cluster(len(clusters) + 1, component, support_mask=None) if cluster is not None: clusters.append(cluster) return clusters boxes = [self._component_bbox(component) for component in components] checked = np.zeros(len(components), dtype=bool) clusters: list[ModelCluster] = [] for seed_idx, component in enumerate(components): if checked[seed_idx]: continue candidate_indices = [ idx for idx in range(len(components)) if not checked[idx] and self._boxes_within_radius(boxes[seed_idx], boxes[idx], self.merge_radius) ] member_indices = [ idx for idx in candidate_indices if idx == seed_idx or self._component_touches_seed_buffer(component, boxes[seed_idx], components[idx], boxes[idx]) ] raw_union = np.zeros_like(positive_mask, dtype=bool) support = np.zeros_like(positive_mask, dtype=bool) for idx in member_indices: checked[idx] = True raw_union |= components[idx] y0, y1, x0, x1 = self._expanded_bbox(boxes[idx], components[idx].shape, self.merge_radius) support_crop = support[y0:y1, x0:x1] component_crop = components[idx][y0:y1, x0:x1] support_crop |= self._dilate_crop(component_crop, self.merge_radius) cluster = self._make_cluster(len(clusters) + 1, raw_union, support_mask=support) if cluster is not None: clusters.append(cluster) return clusters @staticmethod def _component_bbox(component: np.ndarray) -> tuple[int, int, int, int]: ys, xs = np.nonzero(component) return int(ys.min()), int(ys.max()) + 1, int(xs.min()), int(xs.max()) + 1 @staticmethod def _expanded_bbox( box: tuple[int, int, int, int], shape: tuple[int, ...], radius: int, ) -> tuple[int, int, int, int]: y0, y1, x0, x1 = box h, w = int(shape[0]), int(shape[1]) return max(0, y0 - radius), min(h, y1 + radius), max(0, x0 - radius), min(w, x1 + radius) @staticmethod def _boxes_within_radius( a: tuple[int, int, int, int], b: tuple[int, int, int, int], radius: int, ) -> bool: ay0, ay1, ax0, ax1 = a by0, by1, bx0, bx1 = b dy = max(0, by0 - ay1, ay0 - by1) dx = max(0, bx0 - ax1, ax0 - bx1) return dx * dx + dy * dy <= radius * radius def _dilate_crop(self, crop: np.ndarray, radius: int) -> np.ndarray: if radius <= 0: return crop.astype(bool, copy=True) if self.buffer_backend == "binary": return binary_dilation(crop.astype(bool), structure=circular_footprint(radius)) return dilate_fast(crop.astype(bool), radius, backend=self.buffer_backend) def _component_touches_seed_buffer( self, seed: np.ndarray, seed_box: tuple[int, int, int, int], candidate: np.ndarray, candidate_box: tuple[int, int, int, int], ) -> bool: y0, y1, x0, x1 = self._expanded_bbox(seed_box, seed.shape, self.merge_radius) cy0, cy1, cx0, cx1 = candidate_box oy0, oy1 = max(y0, cy0), min(y1, cy1) ox0, ox1 = max(x0, cx0), min(x1, cx1) if oy0 >= oy1 or ox0 >= ox1: return False seed_support = self._dilate_crop(seed[y0:y1, x0:x1], self.merge_radius) return bool(np.any(seed_support[oy0 - y0 : oy1 - y0, ox0 - x0 : ox1 - x0] & candidate[oy0:oy1, ox0:ox1])) def _complete_link_merge(self, positive_mask: np.ndarray) -> list[ModelCluster]: components = self._connected_components(positive_mask) if not components: return [] if len(components) == 1 or self.merge_radius <= 0: clusters: list[ModelCluster] = [] for component in components: cluster = self._make_cluster(len(clusters) + 1, component, support_mask=None) if cluster is not None: clusters.append(cluster) return clusters n_components = len(components) boxes = [self._component_bbox(component) for component in components] close = np.eye(n_components, dtype=bool) for i in range(n_components): for j in range(i + 1, n_components): if not self._boxes_within_radius(boxes[i], boxes[j], self.merge_radius): continue is_close = self._component_touches_seed_buffer(components[i], boxes[i], components[j], boxes[j]) close[i, j] = is_close close[j, i] = is_close groups: list[list[int]] = [[i] for i in range(n_components)] while True: best_pair: tuple[int, int] | None = None best_size = -1 for i in range(len(groups)): for j in range(i + 1, len(groups)): if not all(close[a, b] for a in groups[i] for b in groups[j]): continue merged_size = len(groups[i]) + len(groups[j]) if merged_size > best_size: best_size = merged_size best_pair = (i, j) if best_pair is None: break i, j = best_pair groups[i] = groups[i] + groups[j] del groups[j] clusters: list[ModelCluster] = [] for member_indices in groups: raw_union = np.zeros_like(positive_mask, dtype=bool) support = np.zeros_like(positive_mask, dtype=bool) for idx in member_indices: raw_union |= components[idx] y0, y1, x0, x1 = self._expanded_bbox(boxes[idx], components[idx].shape, self.merge_radius) support_crop = support[y0:y1, x0:x1] component_crop = components[idx][y0:y1, x0:x1] support_crop |= self._dilate_crop(component_crop, self.merge_radius) cluster = self._make_cluster(len(clusters) + 1, raw_union, support_mask=support) if cluster is not None: clusters.append(cluster) return clusters def _mode_like(self, positive_mask: np.ndarray) -> list[ModelCluster]: if not np.any(positive_mask): return [] if self.merge_radius > 0: support = dilate_fast(positive_mask, self.merge_radius, backend=self.buffer_backend) else: support = positive_mask.astype(bool, copy=True) support_labeled, n_support = label(support, structure=self.structure) clusters: list[ModelCluster] = [] for support_id in range(1, n_support + 1): support_mask = support_labeled == support_id raw_union = positive_mask & support_mask if int(raw_union.sum()) < self.min_cluster_pixels: continue final_mask = support_mask if self.expansion_radius >= self.merge_radius else None cluster = self._make_cluster( len(clusters) + 1, raw_union, support_mask=support_mask, final_mask=final_mask, ) if cluster is not None: clusters.append(cluster) return clusters def create_clusterer(config: dict[str, Any]) -> Clusterer: cluster_config = dict(config.get("clusterer") or {}) matching_config = config.get("matching") or {} performance_config = config.get("performance") or {} threshold = cluster_config.get("threshold") if threshold is None: threshold = (config.get("data_source_thresholds") or {}).get(config["data_source"]) if threshold is None: raise ValueError("clusterer.threshold is null and no data_source_thresholds entry exists") return Clusterer( name=cluster_config.get("name", "mode_like"), threshold=float(threshold), min_cluster_pixels=int(cluster_config.get("min_cluster_pixels", 1)), pixel_size_km=float(matching_config.get("pixel_size_km", 2.0)), merge_buffer_km=float(cluster_config.get("merge_buffer_km", 0.0)), cluster_mask_expansion_km=float(cluster_config.get("cluster_mask_expansion_km", 0.0)), connectivity=int(cluster_config.get("connectivity", 8)), buffer_backend=str(performance_config.get("buffer_backend", "auto")), )