import torch import numpy as np from PIL import Image from huggingface_hub import hf_hub_download from inference_script import CountingModule MODEL = None DEVICE = torch.device("cpu") def load_model(): global MODEL, DEVICE MODEL = CountingModule(use_box=True) ckpt_path = hf_hub_download( repo_id="Shengxiao0709/cellsegmodel", filename="microscopy_matching_seg.pth" ) MODEL.load_state_dict(torch.load(ckpt_path, map_location="cpu"), strict=False) MODEL.eval() DEVICE = torch.device("cpu") return MODEL, DEVICE @torch.no_grad() def run(model, img_path, box, device): """ 输入图像路径 + box,返回分割 mask(0/1 或实例ID)。 box: [[xmin, ymin, xmax, ymax]] """ output = model(img_path, box=box) mask = output["pred"] mask = (mask > 0).astype(np.uint8) return mask