ci-net / code /validation /src /clusterers.py
lsh9034's picture
Add files using upload-large-folder tool
7da2ecb verified
Raw History Blame Contribute Delete
15.6 kB
"""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")),
)