"""End-to-end clustering pipeline: images -> features -> consensus clusters. Combines the EnFormer deep feature extractor (features.FeatureExtractor) with optional interpretable morphology descriptors, preprocesses, and runs the EnFormer-inspired ensemble consensus clustering (ensemble_cluster). Designed to be imported by the Gradio app *and* runnable head-less for testing. """ import io import os from dataclasses import dataclass from typing import Dict, List, Optional, Sequence, Tuple import numpy as np from PIL import Image from sklearn.preprocessing import StandardScaler, normalize from data import ImageItem, load_dataset from features import FeatureExtractor, morphology_features from ensemble_cluster import EnsembleClusterer, preprocess_features, ClusterResult # --------------------------------------------------------------------------- cache _EXTRACTOR_CACHE: Dict[str, FeatureExtractor] = {} def get_extractor(backbone: str = "enformer", weights_path: Optional[str] = None) -> FeatureExtractor: key = f"{backbone}:{weights_path}" if key not in _EXTRACTOR_CACHE: _EXTRACTOR_CACHE[key] = FeatureExtractor(backbone=backbone, weights_path=weights_path) return _EXTRACTOR_CACHE[key] # --------------------------------------------------------------------------- feats def build_features(items: List[ImageItem], backbone: str = "enformer", resolution: int = 224, tiles: int = 1, use_morphology: bool = True, weights_path: Optional[str] = None, progress=None) -> Tuple[np.ndarray, np.ndarray]: """Return (feature_matrix, deep_only_matrix). feature_matrix = standardized deep (L2) [+ standardized morphology], ready for preprocess_features(); deep_only_matrix is kept for reference/UI. """ paths = [it.path for it in items] def sub(frac_lo, frac_hi): if progress is None: return None return lambda p: progress(frac_lo + (frac_hi - frac_lo) * p) fe = get_extractor(backbone, weights_path) deep = fe.extract(paths, resolution=resolution, tiles=tiles, l2_normalize=True, progress=sub(0.0, 0.85 if use_morphology else 1.0)) deep_std = normalize(StandardScaler().fit_transform(deep)) if use_morphology: morph = morphology_features(paths, progress=sub(0.85, 1.0)) morph_std = StandardScaler().fit_transform(morph) feat = np.hstack([deep_std, morph_std]).astype(np.float32) else: feat = deep_std.astype(np.float32) return feat, deep # --------------------------------------------------------------------------- result @dataclass class PipelineResult: items: List[ImageItem] labels: np.ndarray result: ClusterResult metrics: Dict[str, float] embedding_2d: Optional[np.ndarray] def run_pipeline(folders: List[str], tags: Optional[List[str]] = None, backbone: str = "enformer", resolution: int = 224, tiles: int = 1, use_morphology: bool = True, pca_dim: int = 12, n_clusters: Optional[int] = 5, base_methods: Sequence[str] = ("kmeans", "fuzzy", "gmm", "spectral", "agglomerative"), n_runs: int = 10, consensus: str = "spectral", weights_path: Optional[str] = None, max_images: Optional[int] = None, progress=None) -> PipelineResult: items = load_dataset(folders, tags) if max_images and len(items) > max_images: # even sub-sample keeps folder balance roughly step = len(items) / max_images items = [items[int(i * step)] for i in range(max_images)] if len(items) < 4: raise ValueError("Need at least 4 images to cluster.") feat, _ = build_features(items, backbone, resolution, tiles, use_morphology, weights_path, progress=progress) Z = preprocess_features(feat, pca_dim=min(pca_dim, feat.shape[1] - 1, len(items) - 1), whiten=False, l2=True) ground_truth = _ground_truth(items) ec = EnsembleClusterer(n_clusters=n_clusters, base_methods=base_methods, n_runs=n_runs, subspace_frac=0.85, k_jitter=0) ec._consensus = _consensus_fn(ec, consensus) # swap consensus function res = ec.fit(Z, ground_truth=ground_truth, compute_2d=True) # Safety net: dissolve degenerate tiny clusters (outliers peeled off by the # consensus) into their nearest surviving cluster so the UI never shows a # meaningless singleton. labels = _merge_tiny_clusters(Z, res.labels, min_size=max(2, int(0.02 * len(items)))) if not np.array_equal(labels, res.labels): res.labels = labels res.n_clusters = len(set(labels.tolist())) res.metrics = EnsembleClusterer.compute_metrics(Z, labels, ground_truth) return PipelineResult(items=items, labels=res.labels, result=res, metrics=res.metrics, embedding_2d=res.embedding_2d) def _merge_tiny_clusters(Z: np.ndarray, labels: np.ndarray, min_size: int) -> np.ndarray: labels = labels.copy() while True: uniq, counts = np.unique(labels, return_counts=True) if len(uniq) <= 2 or counts.min() >= min_size: break tiny = uniq[counts.argmin()] big = [u for u in uniq if u != tiny] centroids = {u: Z[labels == u].mean(0) for u in big} for i in np.where(labels == tiny)[0]: labels[i] = min(big, key=lambda u: np.linalg.norm(Z[i] - centroids[u])) # relabel to a dense 0..k-1 range remap = {u: i for i, u in enumerate(sorted(set(labels.tolist())))} return np.array([remap[x] for x in labels]) def _consensus_fn(ec: EnsembleClusterer, kind: str): from sklearn.cluster import SpectralClustering, AgglomerativeClustering def average_linkage(M, k): d = 1.0 - M np.fill_diagonal(d, 0.0) return AgglomerativeClustering(k, metric="precomputed", linkage="average").fit_predict(d) def spectral(M, k): return SpectralClustering(k, affinity="precomputed", assign_labels="kmeans", random_state=0).fit_predict(M) return spectral if kind == "spectral" else average_linkage def _ground_truth(items: List[ImageItem]) -> Dict[str, np.ndarray]: def enc(vals): u = {v: i for i, v in enumerate(sorted(set(vals)))} return np.array([u[v] for v in vals]) # Only labels that are literally present in the data/filenames are used: # dose = radiation dose in Gy (from filename, e.g. "10Gy") # vor = Vorinostat present (from filename, "+Vor") # set = source folder (1 or 2) gt: Dict[str, np.ndarray] = {} doses = [it.dose for it in items] if all(d is not None for d in doses) and len(set(doses)) > 1: gt["dose"] = enc(doses) if all(it.vor is not None for it in items) and len(set(it.vor for it in items)) > 1: gt["vor"] = np.array([int(it.vor) for it in items]) if len(set(it.folder for it in items)) > 1: gt["set"] = enc([it.folder for it in items]) return gt # --------------------------------------------------------------------------- viz def cluster_montage(items: List[ImageItem], labels: np.ndarray, cluster_id: int, max_imgs: int = 12, thumb: int = 160, cols: int = 6) -> Image.Image: idx = np.where(labels == cluster_id)[0] rng = np.random.default_rng(cluster_id) idx = idx.copy() rng.shuffle(idx) idx = idx[:max_imgs] rows = int(np.ceil(len(idx) / cols)) or 1 canvas = Image.new("RGB", (cols * thumb, rows * thumb), (245, 245, 245)) for j, ii in enumerate(idx): im = Image.open(items[ii].path).convert("RGB").resize((thumb, thumb)) canvas.paste(im, ((j % cols) * thumb, (j // cols) * thumb)) return canvas def cluster_overview_image(items: List[ImageItem], labels: np.ndarray, per_cluster: int = 8, thumb: int = 130, label_h: int = 26) -> Image.Image: """One labelled image stacking a row of representative thumbnails per cluster. Avoids the empty space a gr.Gallery leaves for wide montages, and shows every cluster at once. """ from PIL import ImageDraw clusters = sorted(set(labels.tolist())) cols = per_cluster row_h = thumb + label_h width = cols * thumb height = max(1, len(clusters)) * row_h canvas = Image.new("RGB", (width, height), (255, 255, 255)) draw = ImageDraw.Draw(canvas) rng = np.random.default_rng(0) y = 0 for c in clusters: idx = np.where(labels == c)[0].copy() n = len(idx) rng.shuffle(idx) idx = idx[:per_cluster] draw.rectangle([0, y, width, y + label_h], fill=(20, 20, 20)) draw.text((8, y + 7), f"Cluster {c} - {n} images", fill=(255, 255, 255)) for j, ii in enumerate(idx): im = Image.open(items[ii].path).convert("RGB").resize((thumb, thumb)) canvas.paste(im, (j * thumb, y + label_h)) y += row_h return canvas def assignments_csv(items: List[ImageItem], labels: np.ndarray) -> str: lines = ["filename,set,dose_Gy,vorinostat,cluster"] for it, lab in zip(items, labels): lines.append(f"{it.filename},{it.folder}," f"{'' if it.dose is None else int(it.dose)}," f"{'' if it.vor is None else int(it.vor)},{int(lab)}") return "\n".join(lines)