Spaces:
Running
Running
| """MedCLIPSeg-style text-prompted segmentation, adapted to run here. | |
| Faithful to the *approach* of Koleilat et al., "MedCLIPSeg: Probabilistic | |
| Vision-Language Adaptation ..." (CVPR 2026): a FROZEN vision-language (CLIP) | |
| backbone provides patch-level embeddings; short TEXT PROMPTS describe the target | |
| ("nerve fibers ...") and the background; patch-text similarity yields a dense | |
| segmentation probability, and a PROBABILISTIC step (Monte-Carlo sampling of the | |
| patch embedding with an estimated variance — the paper samples attention Values | |
| from learned distributions) yields a pixel-level UNCERTAINTY map. | |
| This module is ZERO-SHOT (no training) using a general CLIP backbone, so on | |
| cochlear neurofilament it is a *coarse prior*, not a trained segmentor. For the | |
| real, trained model (the paper's PVL adapters + decoder, fine-tuned on masks) | |
| see reference_medclipseg/ and train it on a GPU with the data from | |
| prepare_medclip_data.py. | |
| """ | |
| import os | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| from skimage.transform import resize | |
| _MODEL = "openai/clip-vit-base-patch16" # 16-px patches -> 14x14 per 224 tile | |
| _clip = None | |
| _proc = None | |
| FG_PROMPTS = [ | |
| "a fluorescence microscopy image of nerve fibers", | |
| "neurofilament nerve fibers and axons", | |
| "a dense network of thin bright nerve fibers", | |
| ] | |
| BG_PROMPTS = [ | |
| "a black empty background", | |
| "dark region with no tissue", | |
| "background noise", | |
| ] | |
| def _load(): | |
| global _clip, _proc | |
| if _clip is None: | |
| from transformers import CLIPModel, CLIPProcessor | |
| _clip = CLIPModel.from_pretrained(_MODEL).eval() | |
| _proc = CLIPProcessor.from_pretrained(_MODEL) | |
| return _clip, _proc | |
| def _text_embeds(prompts): | |
| clip, proc = _load() | |
| tok = proc(text=prompts, return_tensors="pt", padding=True) | |
| with torch.no_grad(): | |
| out = clip.text_model(input_ids=tok["input_ids"], | |
| attention_mask=tok.get("attention_mask")) | |
| t = clip.text_projection(out.pooler_output) | |
| return t / t.norm(dim=-1, keepdim=True) # (K, D) | |
| def _dense_patch_embeds(win_u8): | |
| """Per-patch CLIP embeddings for a 224x224 uint8 RGB window (MaskCLIP-style).""" | |
| clip, proc = _load() | |
| px = proc(images=win_u8, return_tensors="pt")["pixel_values"] | |
| vm = clip.vision_model | |
| with torch.no_grad(): | |
| toks = vm(pixel_values=px).last_hidden_state[:, 1:, :] # drop CLS -> (1,196,768) | |
| emb = clip.visual_projection(vm.post_layernorm(toks)) # (1,196,D) | |
| emb = emb / emb.norm(dim=-1, keepdim=True) | |
| return emb[0] # (196, D) | |
| def segment(mip_u8, n_mc=8, sigma=0.05, tile=224): | |
| """Text-prompted CLIP segmentation of a grayscale MIP. | |
| Returns (prob, uncertainty) float maps at the input resolution, both in | |
| [0, 1]. ``prob`` is the fiber probability; ``uncertainty`` is the std of the | |
| probability across ``n_mc`` Monte-Carlo embedding samples. | |
| """ | |
| H0, W0 = mip_u8.shape | |
| rows = max(1, round(H0 / tile)) | |
| cols = max(1, round(W0 / tile)) | |
| Hr, Wr = rows * tile, cols * tile | |
| rgb = np.stack([resize(mip_u8, (Hr, Wr), preserve_range=True)] * 3, -1).astype(np.uint8) | |
| fg_t, bg_t = _text_embeds(FG_PROMPTS), _text_embeds(BG_PROMPTS) | |
| gh, gw = 14, 14 # patches per tile (224/16) | |
| prob = np.zeros((rows * gh, cols * gw), np.float32) | |
| unc = np.zeros((rows * gh, cols * gw), np.float32) | |
| for r in range(rows): | |
| for c in range(cols): | |
| win = rgb[r * tile:(r + 1) * tile, c * tile:(c + 1) * tile] | |
| emb = _dense_patch_embeds(win) # (196, D) | |
| samples = [] | |
| for m in range(n_mc): | |
| e = emb if m == 0 else torch.nn.functional.normalize( | |
| emb + sigma * torch.randn_like(emb), dim=-1) | |
| fg = (e @ fg_t.T).mean(-1) # (196,) | |
| bg = (e @ bg_t.T).mean(-1) | |
| samples.append(torch.sigmoid((fg - bg) / 0.1).numpy()) | |
| samples = np.stack(samples).reshape(n_mc, gh, gw) | |
| prob[r * gh:(r + 1) * gh, c * gw:(c + 1) * gw] = samples.mean(0) | |
| unc[r * gh:(r + 1) * gh, c * gw:(c + 1) * gw] = samples.std(0) | |
| prob = resize(prob, (H0, W0), order=1, preserve_range=True) | |
| unc = resize(unc, (H0, W0), order=1, preserve_range=True) | |
| # normalise uncertainty to [0,1] for display | |
| unc = (unc - unc.min()) / (np.ptp(unc) + 1e-6) | |
| return prob.astype(np.float32), unc.astype(np.float32) | |
| # --------------------------------------------------------------------------- # | |
| # Few-shot TRAINED decoder (frozen CLIP + text-conditioned decoder). | |
| # Trained with train_medclip_fewshot.py (8 animals train, C1/C2b test). | |
| # --------------------------------------------------------------------------- # | |
| _WEIGHTS = os.path.join(os.path.dirname(os.path.abspath(__file__)), | |
| "weights_medclip_decoder.pt") | |
| _decoder = None | |
| class _Decoder(nn.Module): | |
| def __init__(self, cin=768, tdim=512, C=128): | |
| super().__init__() | |
| self.reduce = nn.Conv2d(cin, C, 1) | |
| self.film = nn.Linear(tdim, 2 * C) | |
| def up(ci, co): | |
| return nn.Sequential(nn.ConvTranspose2d(ci, co, 2, 2), | |
| nn.Conv2d(co, co, 3, padding=1), | |
| nn.BatchNorm2d(co), nn.ReLU(True)) | |
| self.up = nn.Sequential(up(C, C), up(C, C), up(C, C // 2), up(C // 2, C // 4)) | |
| self.out = nn.Conv2d(C // 4, 1, 1) | |
| def forward(self, feat, tvec): | |
| x = self.reduce(feat) | |
| g, b = self.film(tvec).chunk(2, -1) | |
| x = x * (1 + g[:, :, None, None]) + b[:, :, None, None] | |
| return torch.sigmoid(self.out(self.up(x))) | |
| def _clip_patch_grid(mip_u8, size=224): | |
| clip, proc = _load() | |
| rgb = np.stack([resize(mip_u8, (size, size), preserve_range=True)] * 3, -1).astype(np.uint8) | |
| px = proc(images=rgb, return_tensors="pt")["pixel_values"] | |
| with torch.no_grad(): | |
| toks = clip.vision_model(pixel_values=px).last_hidden_state[0, 1:, :] | |
| g = int(round(toks.shape[0] ** 0.5)) | |
| return toks.reshape(g, g, -1).permute(2, 0, 1).contiguous()[None] # (1,768,g,g) | |
| def has_trained_model(): | |
| return os.path.exists(_WEIGHTS) | |
| def segment_trained(mip_u8): | |
| """Segment with the few-shot TRAINED decoder. Returns (prob, uncertainty).""" | |
| global _decoder | |
| if _decoder is None: | |
| _decoder = _Decoder().eval() | |
| _decoder.load_state_dict(torch.load(_WEIGHTS, map_location="cpu")["model"]) | |
| tvec = _text_embeds(FG_PROMPTS).mean(0, keepdim=True) | |
| feat = _clip_patch_grid(mip_u8) | |
| with torch.no_grad(): | |
| prob = _decoder(feat, tvec)[0, 0].numpy() | |
| H0, W0 = mip_u8.shape | |
| prob = resize(prob, (H0, W0), order=1, preserve_range=True).astype(np.float32) | |
| unc = np.clip(4 * prob * (1 - prob), 0, 1).astype(np.float32) # entropy-like | |
| return prob, unc | |
| def segment_best(mip_u8): | |
| """Use the trained few-shot decoder if available, else zero-shot CLIP.""" | |
| if has_trained_model(): | |
| return segment_trained(mip_u8) | |
| return segment(mip_u8) | |