"""CW-BASS v2 interactive demo (Hugging Face Space) โ€” Pascal VOC / Cityscapes / ADE20K. ZeroGPU-adapted version of the authors' demo. Each released checkpoint is the EMA teacher of the rule CW-BASS v2's reliability gate SELECTS on that dataset: strict filtering on the saturated Pascal/Cityscapes teachers, the adaptive floor on the confidently-unreliable ADE20K teacher. Weights are pulled per dataset from the matching HF model repo (see DATASETS). The DINOv2 backbone architecture is fetched from facebookresearch/dinov2 via torch.hub; the full model weights (backbone + DPT-lite decoder + segmentation head) come from the released checkpoints. """ import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import spaces # MUST come before torch import torch import torch.nn.functional as F import numpy as np from PIL import Image from torchvision import transforms as T import gradio as gr from model.semseg.dino_segmentor import DINOv2Segmentor from labels import (VOC_CLASSES, VOC_PALETTE, CS_CLASSES, CS_PALETTE, ADE_CLASSES, ADE_PALETTE) MAX_SIDE = 1024 # cap long side for inference; overlay is drawn at input res DATASETS = { "Pascal VOC (21 classes)": dict( repo="psychofict/cwbass-v2-pascal", file="cwbassv2_pascal_dinov2b_1over8.pth", nclass=21, classes=VOC_CLASSES, palette=VOC_PALETTE, bg=0), "Cityscapes (19 classes)": dict( repo="psychofict/cwbass-v2-cityscapes", file="cwbassv2_cityscapes_dinov2b_1over8.pth", nclass=19, classes=CS_CLASSES, palette=CS_PALETTE, bg=None), "ADE20K (150 classes)": dict( repo="psychofict/cwbass-v2-ade20k", file="cwbassv2_ade20k_dinov2b_1over8.pth", nclass=150, classes=ADE_CLASSES, palette=ADE_PALETTE, bg=None), } DEFAULT_DATASET = "Pascal VOC (21 classes)" _normalize = T.Compose([ T.ToTensor(), T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)), ]) def _load_state_dict(path): ckpt = torch.load(path, map_location="cpu", weights_only=False) if isinstance(ckpt, dict): if "teacher" in ckpt and isinstance(ckpt["teacher"], dict) \ and "model_state_dict" in ckpt["teacher"]: return ckpt["teacher"]["model_state_dict"] if "model_state_dict" in ckpt: return ckpt["model_state_dict"] if "student" in ckpt: return ckpt["student"] return ckpt def _weights_path(cfg): from huggingface_hub import hf_hub_download return hf_hub_download(repo_id=cfg["repo"], filename=cfg["file"]) # --- Eager model loading at module scope (ZeroGPU rule 2) --- _MODELS = {} def _build_model(dataset_key): """Build and load a model for a dataset, cache it.""" if dataset_key not in _MODELS: cfg = DATASETS[dataset_key] model = DINOv2Segmentor(backbone="dinov2_vitb14", nclass=cfg["nclass"], pretrained=False).eval() model.load_state_dict(_load_state_dict(_weights_path(cfg)), strict=False) _MODELS[dataset_key] = model.to("cuda") return _MODELS[dataset_key] # Preload the default model at module scope _build_model(DEFAULT_DATASET) def _pad_to_multiple(img_tensor, multiple=14): """Reflection-pad so H, W are divisible by `multiple` (DINOv2 patch-14).""" h, w = img_tensor.shape[-2:] pad_h = (multiple - h % multiple) % multiple pad_w = (multiple - w % multiple) % multiple if pad_h or pad_w: img_tensor = F.pad(img_tensor, (0, pad_w, 0, pad_h), mode='reflect') return img_tensor, (h, w) @torch.no_grad() def _whole_inference(model, img, patch=14): """Single forward with reflect-padding to a multiple of `patch`.""" padded, (h, w) = _pad_to_multiple(img, patch) logits = model(padded) return logits[..., :h, :w] def _resize_long_side(img, max_side): w, h = img.size scale = max_side / max(w, h) if scale < 1.0: img = img.resize((round(w * scale), round(h * scale)), Image.BILINEAR) return img def _legend_html(pairs, cfg): """pairs: list of (class_idx, pixel_count), sorted by count desc.""" total = sum(c for _, c in pairs) or 1 chips = [] for idx, cnt in pairs: if idx == cfg["bg"] or idx >= len(cfg["classes"]): continue r, g, b = cfg["palette"][idx] pct = 100.0 * cnt / total chips.append( f"" f"{cfg['classes'][idx]}{pct:.0f}%") if not chips: chips = ["Background only."] return f"
{''.join(chips)}
" _EMPTY_LEGEND = "
Run a segmentation to see detected classes.
" @spaces.GPU(duration=60) def segment(image, dataset=DEFAULT_DATASET, alpha=0.55): """Run semantic segmentation on an input image. Args: image: Input PIL image. dataset: Which dataset/model to use (Pascal VOC, Cityscapes, or ADE20K). alpha: Overlay opacity (0=original, 1=full mask color). """ if image is None: return None, _EMPTY_LEGEND cfg = DATASETS[dataset] palette = np.asarray(cfg["palette"], dtype=np.uint8) orig = image.convert("RGB") small = _resize_long_side(orig, MAX_SIDE) x = _normalize(small).unsqueeze(0).to("cuda") model = _build_model(dataset) logits = _whole_inference(model, x) pred_small = logits.argmax(1)[0].to(torch.uint8).cpu().numpy() # Upsample the label map to the ORIGINAL input resolution pred = np.array(Image.fromarray(pred_small).resize(orig.size, Image.NEAREST)) seg_rgb = palette[pred] base = np.asarray(orig, dtype=np.float32) overlay = (alpha * seg_rgb + (1 - alpha) * base).clip(0, 255).astype(np.uint8) vals, cnts = np.unique(pred, return_counts=True) pairs = sorted(zip(vals.tolist(), cnts.tolist()), key=lambda t: -t[1]) return Image.fromarray(overlay), _legend_html(pairs, cfg) def _build_examples(): ex_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "examples") prefix = {"pascal": "Pascal VOC (21 classes)", "cityscapes": "Cityscapes (19 classes)", "ade20k": "ADE20K (150 classes)"} rows = [] if os.path.isdir(ex_dir): for fn in sorted(os.listdir(ex_dir)): ds = prefix.get(fn.split("_")[0]) if ds: rows.append([os.path.join(ex_dir, fn), ds, 0.55]) return rows CSS = """ .gradio-container {max-width: 1120px !important; margin: 0 auto !important;} .dark .gradio-container {color: var(--body-text-color);} #hero {text-align:center; padding: 10px 0 2px;} #hero h1 {font-size: 2rem; font-weight: 800; margin: 0 0 2px; background: linear-gradient(90deg,#0ea5e9,#6366f1); -webkit-background-clip:text; background-clip:text; color:transparent;} #hero .sub {opacity:.85; margin: 2px 0 8px; font-size:1rem;} #hero .links a {margin:0 4px; padding:3px 12px; border-radius:999px; text-decoration:none; border:1px solid rgba(128,128,128,.35); font-size:.85rem; white-space:nowrap;} #hero .hint {font-size:.82rem; opacity:.65; margin-top:8px;} #classhead {margin:6px 0 0; font-weight:600;} .legend {display:flex; flex-wrap:wrap; gap:7px; max-height:170px; overflow-y:auto; padding:10px; border:1px solid rgba(128,128,128,.25); border-radius:12px;} .chip {display:inline-flex; align-items:center; gap:6px; padding:3px 10px; border-radius:999px; border:1px solid rgba(128,128,128,.3); font-size:.86rem; line-height:1.4;} .chip .sw {width:13px; height:13px; border-radius:3px; display:inline-block; border:1px solid rgba(0,0,0,.18);} .chip .count {opacity:.55; font-size:.78rem;} .legend .muted {opacity:.6;} #foot {text-align:center; opacity:.6; font-size:.82rem; padding:10px 0 4px;} """ HERO = """

CW-BASS v2

Saturation-Aware Pseudo-Label Selection for Semi-Supervised Segmentation under Foundation-Model Teachers
DINOv2-Base encoder. Each label set loads the checkpoint of the rule the reliability gate selects โ€” strict on Pascal/Cityscapes, adaptive floor on ADE20K. Upload an image and pick a label set, or click an example.
""" FOOT = "" with gr.Blocks(title="CW-BASS v2 โ€” Semi-Supervised Segmentation Demo") as demo: gr.HTML(HERO) with gr.Row(): dataset = gr.Dropdown(choices=list(DATASETS.keys()), value=DEFAULT_DATASET, label="Label set / model", scale=3) alpha = gr.Slider(0.0, 1.0, value=0.55, step=0.05, label="Overlay opacity", scale=2) with gr.Row(equal_height=True): inp = gr.Image(type="pil", label="Input", height=430) out = gr.Image(type="pil", label="Segmentation overlay", height=430) with gr.Row(): run = gr.Button("Segment", variant="primary", scale=3) clear = gr.Button("Clear", scale=1) gr.HTML("
Detected classes
") legend = gr.HTML(_EMPTY_LEGEND) gr.Examples( examples=_build_examples(), inputs=[inp, dataset, alpha], outputs=[out, legend], fn=segment, cache_examples=True, cache_mode="lazy", examples_per_page=6, label="Examples (click to run)", ) gr.HTML(FOOT) run.click(segment, [inp, dataset, alpha], [out, legend]) dataset.change(segment, [inp, dataset, alpha], [out, legend]) clear.click(lambda: (None, None, _EMPTY_LEGEND), None, [inp, out, legend]) if __name__ == "__main__": demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)