microatlas / app.py
MicroAtlas's picture
Remove CPU-only slow inference notice from page header
d5be08a verified
Raw
History Blame Contribute Delete
11.1 kB
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("""
<div style="font-family:'Times New Roman', 'Serif'; font-size:20pt; font-weight:bold; text-align:center; color:#333;">
Large&#8209;Scale Unlabeled Microscopy Images Empower a Generalizable Cell Segmentation Foundation Model for Versatile Biological Analysis<br>
<a style="color:#0066cc; font-size:14pt;" href="https://huggingface.co/datasets/MicroAtlas/MicroAtlas-2B" target="_blank">[dataset]</a>
<a style="color:#333; font-size:14pt;" href="https://github.com/Luffy03/MicroAtlas" target="_blank">[github]</a>
</div>""")
gr.HTML("""<h4 style="color:#333;">
MicroAtlas is a foundation model trained on large-scale unlabeled microscopy images.
</h4>""")
gr.HTML("""<h4 style="color:#333;">
For fast GPU inference, clone the code and run locally:
<a style="color:#0066cc;" href="https://github.com/Luffy03/MicroAtlas" target="_blank">github.com/Luffy03/MicroAtlas</a>
</h4>""")
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("""<h4 style="color:#333;"> Notes:<br>
<li>This Space runs on <b>CPU</b> — expect ~1–2 minutes per image.
<li>For production use, clone <a style="color:#0066cc;" href="https://github.com/Luffy03/MicroAtlas" target="_blank">github.com/Luffy03/MicroAtlas</a> and run locally with GPU.
<li>You can load and process 2D, multi-channel tifs.
<li>You can upload multiple files and download a zip of the segmentations.
</h4>""")
demo.launch(css=".gradio-container {background: white;}", ssr_mode=False)