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("""
Large‑Scale Unlabeled Microscopy Images Empower a Generalizable Cell Segmentation Foundation Model for Versatile Biological Analysis
[dataset] [github]
""") gr.HTML("""

MicroAtlas is a foundation model trained on large-scale unlabeled microscopy images.

""") gr.HTML("""

For fast GPU inference, clone the code and run locally: github.com/Luffy03/MicroAtlas

""") input_image = gr.Image(label="Input", type="filepath") with gr.Row(): with gr.Column(scale=1): with gr.Row(): resize = gr.Number(label='max resize', value=512) flow_threshold = gr.Number(label='flow threshold', value=0.4) cellprob_threshold = gr.Number(label='cellprob threshold', value=0) up_btn = gr.UploadButton( "Multi-file upload (png, jpg, tif etc)", visible=True, file_count="multiple" ) with gr.Column(scale=1): send_btn = gr.Button("Run MicroAtlas") down_btn = gr.DownloadButton("Download masks (TIF)", visible=False) down_btn2 = gr.DownloadButton("Download outlines (PNG)", visible=False) with gr.Column(scale=2): outlines = gr.Image(label="Segmentation", type="pil", format='png', value=fp0) # Example images (if you have a samples/ folder in the Space) sample_list = glob.glob("samples/*.png") if sample_list: gr.Examples( sample_list, fn=update_button, inputs=input_image, outputs=[input_image, up_btn, outlines], examples_per_page=50, label="Click on an example to try it" ) input_image.upload(update_button, input_image, [input_image, up_btn, outlines]) up_btn.upload(update_image, up_btn, [input_image, up_btn, outlines]) send_btn.click( microatlas_segment, [up_btn, resize, flow_threshold, cellprob_threshold], [outlines, down_btn, down_btn2] ) gr.HTML("""

Notes:
  • This Space runs on CPU — expect ~1–2 minutes per image.
  • For production use, clone github.com/Luffy03/MicroAtlas and run locally with GPU.
  • You can load and process 2D, multi-channel tifs.
  • You can upload multiple files and download a zip of the segmentations.
  • """) demo.launch(css=".gradio-container {background: white;}", ssr_mode=False)