Team Ai
Apppublic

AbdallahAdel/HTTS_implementation

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
App.py566 linesDownload Raw Back to root
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)