AbdallahAdel/HTTS_implementation
0
1import streamlit as st2import cv23import numpy as np4import time5import heapq6import math7import io8from dataclasses import dataclass, field9from typing import List, Tuple, Optional10from collections import deque11 12from PIL import Image13from skimage.segmentation import mark_boundaries, find_boundaries14from skimage.measure import label as cc_label15from skimage.color import label2rgb16 17st.markdown("""18 <style>19 section[data-testid="stSidebar"] {20 width: 300px !important;21 }22 #root > div:nth-child(1) > div > div > div > div > section > div {23 overflow-y: scroll;24 }25 </style>26""", unsafe_allow_html=True)27 28 29# Page config30st.set_page_config(31 page_title="HHTS โ Remote Sensing Segmentation",32 page_icon="๐ฐ๏ธ",33 layout="wide"34)35 36st.markdown("""37 <style>38 section[data-testid="stSidebar"] {39 width: 300px !important;40 }41 42 /* Main text */43 html, body, [class*="css"] {44 color: #EAEAEA;45 }46 47 /* Titles */48 h1, h2, h3, h4 {49 color: #00D4FF;50 }51 52 /* Sidebar text */53 section[data-testid="stSidebar"] * {54 color: #F5F5F5;55 }56 57 /* Metric labels */58 [data-testid="stMetricLabel"] {59 color: #FFD166 !important;60 }61 62 /* Metric values */63 [data-testid="stMetricValue"] {64 color: #06D6A0 !important;65 }66 </style>67""", unsafe_allow_html=True)68 69# Data structures70@dataclass(order=True)71class PrioritizedSegment:72 neg_priority : float73 id : int = field(compare=False)74 mask : np.ndarray = field(compare=False)75 size : int = field(compare=False)76 channel_infos : list = field(compare=False, default_factory=list)77 split_channel : int = field(compare=False, default=-1)78 split_criteria: float = field(compare=False, default=-1.0)79 80@dataclass81class ChannelInfo:82 min_val : int83 max_val : int84 width : int85 split_criteria: float86 87 @property88 def is_exhausted(self):89 return self.split_criteria < 0.090 91 92# Channel extraction (matches C++ getChannels exactly)93def get_channels_rgb_hsv_lab(image_rgb, use_rgb=True, use_hsv=True,94 use_lab=True, apply_blur=False):95 channels, names = [], []96 img = image_rgb.copy()97 if apply_blur:98 img = cv2.GaussianBlur(img, (3, 3), 0, 0)99 100 if use_rgb:101 for idx, name in enumerate(["R", "G", "B"]):102 channels.append(img[:, :, idx].astype(np.uint8))103 names.append(name)104 105 if use_hsv:106 # C++: convert first, then blur the converted image107 hsv = cv2.cvtColor(image_rgb, cv2.COLOR_RGB2HSV)108 if apply_blur:109 hsv = cv2.GaussianBlur(hsv, (3, 3), 0, 0)110 for idx, name in enumerate(["H", "S", "V"]):111 channels.append(hsv[:, :, idx].astype(np.uint8))112 names.append("HSV_" + name)113 114 if use_lab:115 lab = cv2.cvtColor(image_rgb, cv2.COLOR_RGB2LAB)116 if apply_blur:117 lab = cv2.GaussianBlur(lab, (3, 3), 0, 0)118 for idx, name in enumerate(["L", "a", "b"]):119 channels.append(lab[:, :, idx].astype(np.uint8))120 names.append("LAB_" + name)121 122 return channels, names123 124 125# Core algorithm functions126def channel_info(channel, mask, size):127 values = channel[mask]128 if values.size == 0:129 return ChannelInfo(0, 0, 0, -1.0)130 min_val, max_val = int(values.min()), int(values.max())131 width = max_val - min_val132 if width < 2: # auto-termination guard133 return ChannelInfo(min_val, max_val, width, -1.0)134 return ChannelInfo(min_val, max_val, width, float(values.std()) * size * size)135 136 137def build_segment(seg_id, mask, channels, min_segment_size,138 split_threshold, existing_infos=None):139 size = int(mask.sum())140 if size // 2 < min_segment_size:141 return None142 infos = existing_infos if existing_infos is not None else \143 [channel_info(ch, mask, size) for ch in channels]144 best_idx = int(np.argmax([ci.split_criteria for ci in infos]))145 best_score = infos[best_idx].split_criteria146 if best_score <= split_threshold:147 return None148 return PrioritizedSegment(149 neg_priority=-best_score, id=seg_id, mask=mask, size=size,150 channel_infos=infos, split_channel=best_idx, split_criteria=best_score)151 152 153def interrupt_split(segment, failed_ch, heap, split_threshold):154 """Permanently blacklist failed channel โ mirrors C++ interruptSplit."""155 segment.channel_infos[failed_ch].split_criteria = -1.0156 best_score, best_idx = -1.0, -1157 for i, ci in enumerate(segment.channel_infos):158 if ci.split_criteria > best_score:159 best_score, best_idx = ci.split_criteria, i160 if best_idx >= 0 and best_score > split_threshold:161 heapq.heappush(heap, PrioritizedSegment(162 neg_priority=-best_score, id=segment.id, mask=segment.mask,163 size=segment.size, channel_infos=segment.channel_infos,164 split_channel=best_idx, split_criteria=best_score))165 166 167def get_channel_threshold(channel, ci, mask, histogram_bins):168 channel_bins = max(2, int(min(histogram_bins, ci.width)))169 values = channel[mask]170 hist, _ = np.histogram(values, bins=channel_bins,171 range=(ci.min_val, ci.max_val + 1))172 hist = hist.astype(np.float32)173 padded = np.pad(hist, (1, 1), mode="edge")174 response = (padded[:-2] - 2.0 * padded[1:-1] + padded[2:]).astype(np.float32)175 cdf = np.cumsum(hist)176 # FIX: integer mean โ matches C++ "const int mean = max / 2"177 mean = int(cdf[-1]) // 2178 if mean <= 0:179 return (ci.min_val + ci.max_val) // 2180 weights = 1.0 / (((((mean - cdf) / mean * 2.0) ** 4) + 1.0))181 best_bin = int(np.argmax(response * weights))182 return int(ci.min_val + 0.5 * (183 ((ci.max_val - ci.min_val + 1) * (2 * best_bin + 1) / channel_bins) - 1))184 185 186def _absorb_tiny_fragments(seeds, tiny, parent_mask, kernel):187 """BFS on bounding box only โ avoids full-image dilation overhead."""188 if not seeds:189 return []190 rows, cols = np.where(parent_mask)191 r0, r1 = int(rows.min()), int(rows.max()) + 1192 c0, c1 = int(cols.min()), int(cols.max()) + 1193 pm_crop, tiny_crop = parent_mask[r0:r1, c0:c1], tiny[r0:r1, c0:c1]194 H, W = pm_crop.shape195 ownership = np.zeros((H, W), dtype=np.int32)196 ownership[tiny_crop] = -1197 order = sorted(range(len(seeds)), key=lambda i: int(seeds[i].sum()), reverse=True)198 queues = [deque() for _ in seeds]199 for si in order:200 sc = seeds[si][r0:r1, c0:c1]201 ownership[sc] = si + 1202 dil = cv2.dilate(sc.astype(np.uint8), kernel).astype(bool)203 for r, c in zip(*np.where(dil & tiny_crop & (ownership == -1))):204 queues[si].append((int(r), int(c)))205 nbrs = [(-1,0),(1,0),(0,-1),(0,1)]206 active = list(range(len(seeds)))207 while active:208 still = []209 for si in active:210 nq = deque()211 while queues[si]:212 r, c = queues[si].popleft()213 if ownership[r, c] != -1: continue214 ownership[r, c] = si + 1215 for dr, dc in nbrs:216 nr, nc = r + dr, c + dc217 if 0 <= nr < H and 0 <= nc < W and \218 ownership[nr, nc] == -1 and pm_crop[nr, nc]:219 nq.append((nr, nc))220 if nq:221 queues[si] = nq222 still.append(si)223 active = still224 if (ownership == -1).any():225 ownership[ownership == -1] = order[0] + 1226 child_masks = []227 for si in range(len(seeds)):228 full = np.zeros_like(parent_mask)229 full[r0:r1, c0:c1] = (ownership == si + 1)230 if full.any():231 child_masks.append(full)232 return child_masks233 234 235def split_segment_hhts_like(segment, channels, labels, next_label, heap,236 min_segment_size, split_threshold, histogram_bins):237 kernel = np.array([[0,1,0],[1,1,1],[0,1,0]], dtype=np.uint8)238 ch_idx = segment.split_channel239 ci = segment.channel_infos[ch_idx]240 channel = channels[ch_idx]241 mask = segment.mask242 threshold = get_channel_threshold(channel, ci, mask, histogram_bins)243 244 seeds, tiny = [], np.zeros(mask.shape, dtype=bool)245 for side_mask, label in [246 (mask & (channel <= threshold), 'low'),247 (mask & (channel > threshold), 'high')248 ]:249 cc, n = cc_label(side_mask.astype(np.uint8), connectivity=1, return_num=True)250 found = False251 for i in range(1, n + 1):252 comp = cc == i253 if int(comp.sum()) < min_segment_size:254 tiny |= comp255 else:256 seeds.append(comp)257 found = True258 if not found:259 interrupt_split(segment, ch_idx, heap, split_threshold)260 return next_label261 262 child_masks = _absorb_tiny_fragments(seeds, tiny, mask, kernel)263 if len(child_masks) < 2:264 interrupt_split(segment, ch_idx, heap, split_threshold)265 return next_label266 267 for i, cm in enumerate(child_masks):268 cid = segment.id if i == 0 else next_label269 if i > 0: next_label += 1270 labels[cm] = cid271 cs = build_segment(cid, cm, channels, min_segment_size, split_threshold)272 if cs is not None:273 heapq.heappush(heap, cs)274 return next_label275 276 277def hhts_python(image_rgb, superpixels, split_threshold=0.0, histogram_bins=32,278 min_segment_size=64, use_rgb=True, use_hsv=True, use_lab=True,279 apply_blur=False):280 """281 Run HHTS. superpixels can be:282 int โ single level283 [500] โ single level284 [250,500] โ multi-level snapshots285 [-1] โ auto-termination286 [500,-1] โ snapshot then auto-terminate287 Returns (label_maps, label_counts).288 """289 sp_queue = [superpixels] if isinstance(superpixels, int) else list(superpixels)290 h, w = image_rgb.shape[:2]291 channels, _ = get_channels_rgb_hsv_lab(292 image_rgb, use_rgb, use_hsv, use_lab, apply_blur)293 294 labels = np.ones((h, w), dtype=np.int32)295 heap = []296 next_label = 2297 seg = build_segment(1, np.ones((h, w), dtype=bool),298 channels, min_segment_size, split_threshold)299 if seg is not None:300 heapq.heappush(heap, seg)301 302 output_maps, output_counts = [], []303 304 while heap:305 # snapshot check306 while sp_queue and sp_queue[0] >= 0 and next_label > sp_queue[0]:307 output_maps.append(labels.copy())308 output_counts.append(next_label)309 sp_queue.pop(0)310 if not sp_queue:311 break312 313 segment = heapq.heappop(heap)314 # stale mask check315 r, c = np.where(segment.mask)316 if len(r) == 0 or labels[int(r[0]), int(c[0])] != segment.id:317 current = labels == segment.id318 if not current.any(): continue319 segment = build_segment(segment.id, current, channels,320 min_segment_size, split_threshold,321 existing_infos=segment.channel_infos)322 if segment is None: continue323 324 next_label = split_segment_hhts_like(325 segment, channels, labels, next_label, heap,326 min_segment_size, split_threshold, histogram_bins)327 328 # add remaining snapshots329 while sp_queue:330 sp_queue.pop(0)331 output_maps.append(labels.copy())332 output_counts.append(next_label)333 334 if not output_maps:335 output_maps.append(labels.copy())336 output_counts.append(next_label)337 338 return output_maps, output_counts339 340 341# Metrics (no GT required)342def explained_variation(image_rgb, sp_labels):343 img = image_rgb.astype(np.float64)344 mu = img.mean(axis=(0, 1), keepdims=True)345 var_total = float(np.sum((img - mu) ** 2))346 if var_total == 0: return 1.0347 var_within = sum(348 float(np.sum((img[sp_labels == sid] -349 img[sp_labels == sid].mean(axis=0)) ** 2))350 for sid in np.unique(sp_labels))351 return float(1.0 - var_within / var_total)352 353def intra_cluster_variation(image_rgb, sp_labels):354 img = image_rgb.astype(np.float64)355 vals = [float(np.mean(np.std(img[sp_labels == sid], axis=0)))356 for sid in np.unique(sp_labels)357 if (sp_labels == sid).sum() > 1]358 return float(np.mean(vals)) if vals else float('nan')359 360def compactness(sp_labels):361 scores = []362 for sid in np.unique(sp_labels):363 m = (sp_labels == sid).astype(np.uint8)364 area = int(m.sum())365 if area < 4: continue366 contours, _ = cv2.findContours(m, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)367 if contours:368 p = cv2.arcLength(contours[0], True)369 if p > 0: scores.append(4.0 * math.pi * area / (p ** 2))370 return float(np.mean(scores)) if scores else float('nan')371 372 373# Streamlit UI374st.title("๐ฐ๏ธ HHTS Aerial Image Segmentation")375st.markdown(376 "Python re-implementation of "377 "**Hierarchical Histogram Threshold Segmentation** (Chang et al., CVPR 2024). "378 "Upload a satellite or aerial image and tune the parameters in the sidebar."379)380 381# Sidebar382with st.sidebar:383 st.header("โ๏ธ Parameters")384 385 st.subheader("Termination")386 term_mode = st.radio(387 "Mode",388 ["Superpixel count", "Auto-termination"],389 help="Count: stop at target. Auto: stop when no color information remains."390 )391 target_sp = 500392 if term_mode == "Superpixel count":393 target_sp = st.slider("Target superpixels", 50, 2000, 500, 50)394 395 st.subheader("Color spaces")396 u_rgb = st.checkbox("RGB", value=True)397 u_hsv = st.checkbox("HSV", value=True)398 u_lab = st.checkbox("LAB", value=True)399 400 st.subheader("Algorithm")401 histogram_bins = st.slider("Histogram bins", 8, 64, 32, 8)402 min_size = st.slider("Min segment size", 16, 256, 64, 16)403 blur = st.checkbox("Pre-blur (3ร3 Gaussian)", value=False)404 max_side = st.slider("Resize max side (px)", 128, 512, 256, 64)405 406# File upload 407uploaded_file = st.file_uploader(408 "Upload a satellite / aerial image (JPG, PNG, TIF)",409 type=["jpg", "jpeg", "png", "tif", "tiff"]410)411 412if uploaded_file is None:413 st.info("Upload an image in the sidebar to get started.")414 st.stop()415 416# Read and resize417file_bytes = np.asarray(bytearray(uploaded_file.read()), dtype=np.uint8)418img_bgr = cv2.imdecode(file_bytes, cv2.IMREAD_COLOR)419img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)420h0, w0 = img_rgb.shape[:2]421scale = min(1.0, max_side / max(h0, w0))422img_disp = cv2.resize(img_rgb, (int(w0 * scale), int(h0 * scale)),423 interpolation=cv2.INTER_AREA)424 425col1, col2 = st.columns(2)426with col1:427 st.subheader("Original")428 st.image(img_disp, use_container_width=True)429 st.caption(f"{img_disp.shape[1]}ร{img_disp.shape[0]} px")430 431# Run button432if st.button("โถ Run Segmentation", type="primary", use_container_width=True):433 if not (u_rgb or u_hsv or u_lab):434 st.error("Please select at least one color space.")435 st.stop()436 437 sp_param = -1 if term_mode == "Auto-termination" else target_sp438 439 with st.spinner("Running HHTS โ this takes 15โ30 seconds for a 256px imageโฆ"):440 t0 = time.time()441 maps, counts = hhts_python(442 img_disp,443 superpixels = sp_param,444 histogram_bins = histogram_bins,445 min_segment_size = min_size,446 use_rgb=u_rgb, use_hsv=u_hsv, use_lab=u_lab,447 apply_blur=blur448 )449 runtime = time.time() - t0450 451 labels = maps[-1]452 n_labels = int(labels.max())453 boundaries = (mark_boundaries(img_disp, labels,454 color=(1, 1, 0), mode="thick") * 255).astype(np.uint8)455 mean_color = np.zeros_like(img_disp)456 for sid in np.unique(labels):457 mask_sid = labels == sid458 mean_color[mask_sid] = img_disp[mask_sid].mean(axis=0).astype(np.uint8)459 460 with col2:461 st.subheader("Boundaries")462 st.image(boundaries, use_container_width=True)463 st.caption(f"{n_labels} superpixels ยท {runtime:.1f}s ยท {term_mode}")464 465 # Bottom row 466 st.divider()467 r1c1, r1c2 = st.columns(2)468 469 with r1c1:470 st.subheader("Mean-color map")471 st.image(mean_color, use_container_width=True)472 473 with r1c2:474 st.subheader("๐ Metrics")475 ev = explained_variation(img_disp, labels)476 icv = intra_cluster_variation(img_disp, labels)477 co = compactness(labels)478 479 m1, m2, m3 = st.columns(3)480 m1.metric("EV โ", f"{ev:.4f}",481 help="Explained Variation โ higher = more color-consistent superpixels")482 m2.metric("ICV โ", f"{icv:.2f}",483 help="Intra-Cluster Variation โ lower = more homogeneous")484 m3.metric("Compactness โ", f"{co:.3f}",485 help="Shape compactness โ HHTS is intentionally lower than SLIC")486 487 st.caption(488 "**BR / UE / ASA** require a ground-truth mask. "489 "Expand the section below to upload one."490 )491 492 # Downloads493 st.subheader("๐พ Download")494 d1, d2 = st.columns(2)495 with d1:496 buf = io.BytesIO()497 Image.fromarray(boundaries).save(buf, format="PNG")498 st.download_button("Boundary image", buf.getvalue(),499 "hhts_boundaries.png", "image/png",500 use_container_width=True)501 with d2:502 buf2 = io.BytesIO()503 Image.fromarray((labels % 256).astype(np.uint8)).save(buf2, format="PNG")504 st.download_button("Label map", buf2.getvalue(),505 "hhts_labels.png", "image/png",506 use_container_width=True)507 508 # Optional GT metrics509 with st.expander("๐ Upload Ground-Truth Mask โ compute BR / UE / ASA"):510 gt_file = st.file_uploader(511 "GT mask (same content as input image, any size โ will be resized)",512 type=["png", "jpg", "tif"]513 )514 if gt_file:515 gt_bytes = np.asarray(bytearray(gt_file.read()), dtype=np.uint8)516 gt_bgr = cv2.imdecode(gt_bytes, cv2.IMREAD_COLOR)517 gt_rgb = cv2.cvtColor(gt_bgr, cv2.COLOR_BGR2RGB)518 gt_rgb = cv2.resize(gt_rgb,519 (img_disp.shape[1], img_disp.shape[0]),520 interpolation=cv2.INTER_NEAREST)521 522 # pack RGB โ integer class map523 flat = (gt_rgb[:,:,0].astype(np.int32) * 65536 +524 gt_rgb[:,:,1].astype(np.int32) * 256 +525 gt_rgb[:,:,2].astype(np.int32))526 uniq = np.unique(flat)527 lut = {v: i for i, v in enumerate(uniq)}528 gt_class = np.vectorize(lut.get)(flat).astype(np.int32)529 530 # BR531 gt_b = find_boundaries(gt_class, mode="outer")532 sp_b = find_boundaries(labels, mode="outer")533 dil = cv2.dilate(sp_b.astype(np.uint8),534 np.ones((5, 5), np.uint8)).astype(bool)535 br = float((gt_b & dil).sum() / max(gt_b.sum(), 1))536 537 # UE538 N = gt_class.size539 sp_ids = np.unique(labels)540 gt_ids = np.unique(gt_class)541 sp_map = {v: i for i, v in enumerate(sp_ids)}542 gt_map = {v: i for i, v in enumerate(gt_ids)}543 table = np.zeros((len(sp_ids), len(gt_ids)), dtype=np.int64)544 for sv, gv in zip(labels.ravel(), gt_class.ravel()):545 table[sp_map[sv], gt_map[gv]] += 1546 ue = float((table.sum(axis=1) - table.max(axis=1)).sum() / N)547 548 # ASA549 correct = sum(550 int(np.unique(gt_class[labels == sid],551 return_counts=True)[1].max())552 for sid in sp_ids)553 asa = float(correct / N)554 555 g1, g2, g3 = st.columns(3)556 g1.metric("Boundary Recall โ", f"{br:.4f}")557 g2.metric("Underseg. Error โ", f"{ue:.4f}")558 g3.metric("ASA โ", f"{asa:.4f}")559 560# Footer561st.divider()562st.markdown(563 "<small>HHTS ยท Chang et al. CVPR 2024 ยท Python re-implementation ยท "564 "CSE 429 Computer Vision and Pattern Recognition</small>",565 unsafe_allow_html=True566)