import numpy as np import gradio as gr import cv2 from cellpose import models import matplotlib.pyplot as plt import os, io from PIL import Image from cellpose.io import imread, imsave import glob import datetime from zipfile import ZipFile from huggingface_hub import hf_hub_download # --------------------------------------------------------------------------- # Download MicroAtlas model weights from Hugging Face Hub # --------------------------------------------------------------------------- # Hugging Face token: read from the HF_TOKEN secret configured in the Space. HF_TOKEN = os.environ.get("HF_TOKEN") def download_weights(): return hf_hub_download( repo_id="MicroAtlas/microatlas-model", filename="microatlas", token=HF_TOKEN, ) # --------------------------------------------------------------------------- # Load model on CPU (runs once at startup) # --------------------------------------------------------------------------- try: fpath = download_weights() model = models.CellposeModel(gpu=False, pretrained_model=fpath) print(f"MicroAtlas model loaded from: {fpath}") except Exception as e: print(f"Error loading model: {e}") exit(1) # --------------------------------------------------------------------------- # Utility functions # --------------------------------------------------------------------------- def normalize99(img): X = img.copy() X = (X - np.percentile(X, 1)) / (1e-10 + np.percentile(X, 99) - np.percentile(X, 1)) return X def image_resize(img, resize=400): ny, nx = img.shape[:2] if np.array(img.shape).max() > resize: if ny > nx: nx = int(nx / ny * resize) ny = resize else: ny = int(ny / nx * resize) nx = resize shape = (nx, ny) img = cv2.resize(img, shape) img = img.astype(np.uint8) return img def plot_outlines(img, masks): img = normalize99(img) img = np.clip(img, 0, 1) outpix = [] contours, hierarchy = cv2.findContours( masks.astype(np.int32), mode=cv2.RETR_FLOODFILL, method=cv2.CHAIN_APPROX_SIMPLE ) for c in range(len(contours)): pix = contours[c].astype(int).squeeze() if len(pix) > 4: peri = cv2.arcLength(contours[c], True) approx = cv2.approxPolyDP(contours[c], 0.001, True)[:, 0, :] outpix.append(approx) figsize = (6, 6) if img.shape[0] > img.shape[1]: figsize = (6 * img.shape[1] / img.shape[0], 6) else: figsize = (6, 6 * img.shape[0] / img.shape[1]) fig = plt.figure(figsize=figsize, facecolor='k') ax = fig.add_axes([0.0, 0.0, 1, 1]) ax.set_xlim([0, img.shape[1]]) ax.set_ylim([0, img.shape[0]]) ax.imshow(img[::-1], origin='upper', aspect='auto') if outpix is not None: for o in outpix: ax.plot(o[:, 0], img.shape[0] - o[:, 1], color=[1, 0, 0], lw=1) ax.axis('off') buf = io.BytesIO() fig.savefig(buf, bbox_inches='tight') buf.seek(0) pil_img = Image.open(buf) plt.close(fig) return pil_img # --------------------------------------------------------------------------- # CPU inference # --------------------------------------------------------------------------- def run_model(img, flow_threshold, cellprob_threshold): masks, flows, _ = model.eval( img, channels=None, diameter=None, bsize=256, flow_threshold=flow_threshold, cellprob_threshold=cellprob_threshold, ) return masks, flows # --------------------------------------------------------------------------- # Main segmentation pipeline # --------------------------------------------------------------------------- def microatlas_segment(filepath, resize=512, flow_threshold=0.4, cellprob_threshold=0): zip_path = os.path.splitext(filepath[-1])[0] + "_masks.zip" with ZipFile(zip_path, 'w') as myzip: for j in range(len(filepath)): now = datetime.datetime.now() formatted_now = now.strftime("%Y-%m-%d %H:%M:%S") img_input = imread(filepath[j]) img = image_resize(img_input, resize=resize) masks, flows = run_model(img, flow_threshold, cellprob_threshold) print(formatted_now, j, masks.max(), os.path.split(filepath[j])[-1]) # Scale masks back to original size target_size = (img_input.shape[1], img_input.shape[0]) if target_size[0] != img.shape[1] or target_size[1] != img.shape[0]: masks_rsz = cv2.resize( masks.astype('uint16'), target_size, interpolation=cv2.INTER_NEAREST ).astype('uint16') else: masks_rsz = masks.copy() fname_masks = os.path.splitext(filepath[j])[0] + "_masks.tif" imsave(fname_masks, masks_rsz) myzip.write(fname_masks, arcname=os.path.split(fname_masks)[-1]) # Generate visualization (based on last image) outpix = plot_outlines(img, masks) Ly, Lx = img.shape[:2] outpix = outpix.resize((Lx, Ly), resample=Image.BICUBIC) fname_out = os.path.splitext(filepath[-1])[0] + "_outlines.png" outpix.save(fname_out) if len(filepath) > 1: b1 = gr.DownloadButton(visible=True, value=zip_path) else: b1 = gr.DownloadButton(visible=True, value=fname_masks) b2 = gr.DownloadButton(visible=True, value=fname_out) return outpix, b1, b2 # --------------------------------------------------------------------------- # UI helpers # --------------------------------------------------------------------------- def tif_view(filepath): fpath, fext = os.path.splitext(filepath) if fext in ['.tiff', '.tif']: img = imread(filepath[-1]) if img.ndim == 2: img = np.tile(img[:, :, np.newaxis], [1, 1, 3]) elif img.ndim == 3: imin = np.argmin(img.shape) if imin < 2: img = np.transpose(img, [2, imin]) else: raise ValueError("TIF cannot have more than three dimensions") Ly, Lx, nchan = img.shape imgi = np.zeros((Ly, Lx, 3)) nn = np.minimum(3, img.shape[-1]) imgi[:, :, :nn] = img[:, :, :nn] imsave(filepath, imgi) return filepath def norm_path(filepath): img = imread(filepath) img = normalize99(img) img = np.clip(img, 0, 1) fpath, fext = os.path.splitext(filepath) filepath = fpath + '.png' pil_image = Image.fromarray((255. * img).astype(np.uint8)) pil_image.save(filepath) return filepath def update_image(filepath): for f in filepath: f = tif_view(f) filepath_show = norm_path(filepath[-1]) fp0 = Image.fromarray(np.zeros((96, 128), dtype=np.uint8)) return filepath_show, filepath, fp0 def update_button(filepath): filepath = tif_view(filepath) filepath_show = norm_path(filepath) fp0 = Image.fromarray(np.zeros((96, 128), dtype=np.uint8)) return filepath_show, [filepath], fp0 # --------------------------------------------------------------------------- # Gradio UI # --------------------------------------------------------------------------- fp0 = Image.fromarray(np.zeros((96, 128), dtype=np.uint8)) with gr.Blocks( title="MicroAtlas Cell Segmentation", ) as demo: with gr.Row(): with gr.Column(scale=2): gr.HTML("""